Team Ai
Apppublic

FastestAI/CodeBert_Redundant_Detection_Task

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py394 linesDownload Raw Back to root
1import os2import logging3import torch4import torch.nn.functional as F5from fastapi import FastAPI, HTTPException6from fastapi.middleware.cors import CORSMiddleware7from pydantic import BaseModel8from typing import List9import uvicorn10from datetime import datetime11from transformers import AutoTokenizer, AutoModel12import requests13import re14import tempfile15 16# Set up logging17logging.basicConfig(18    level=logging.INFO,19    format='%(asctime)s - %(levelname)s - %(message)s',20    handlers=[logging.StreamHandler()]21)22logger = logging.getLogger(__name__)23 24# System information - with your current values25DEPLOYMENT_DATE = "2025-06-22 22:15:13"26DEPLOYED_BY = "FASTESTAI"27 28# Get device29device = torch.device("cuda" if torch.cuda.is_available() else "cpu")30logger.info(f"Using device: {device}")31 32# HuggingFace model repository path just for weights file33REPO_ID = "FastestAI/Redundant_Model"34MODEL_WEIGHTS_URL = f"https://huggingface.co/{REPO_ID}/resolve/main/pytorch_model.bin"35 36# Initialize FastAPI app37app = FastAPI(38    title="Test Similarity Analyzer API",39    description="API for analyzing similarity between test cases. Deployed by " + DEPLOYED_BY,40    version="1.0.0",41    docs_url="/",42)43 44# Add CORS middleware45app.add_middleware(46    CORSMiddleware,47    allow_origins=["*"],48    allow_credentials=True,49    allow_methods=["*"],50    allow_headers=["*"],51)52 53# Define label to class mapping54label_to_class = {0: "Duplicate", 1: "Redundant", 2: "Distinct"}55 56# Define input models for API57class SourceCode(BaseModel):58    class_name: str59    code: str60 61class TestCase(BaseModel):62    id: str63    test_fixture: str64    name: str65    code: str66    target_class: str67    target_method: List[str]68 69class SimilarityInput(BaseModel):70    pair_id: str71    source_code: SourceCode72    test_case_1: TestCase73    test_case_2: TestCase74 75# Define the model class76class CodeSimilarityClassifier(torch.nn.Module):77    def __init__(self, model_name="microsoft/codebert-base", num_labels=3):78        super().__init__()79        self.encoder = AutoModel.from_pretrained(model_name)80        self.dropout = torch.nn.Dropout(0.1)81 82        # Create a more powerful classification head83        hidden_size = self.encoder.config.hidden_size84 85        self.classifier = torch.nn.Sequential(86            torch.nn.Linear(hidden_size, hidden_size),87            torch.nn.LayerNorm(hidden_size),88            torch.nn.GELU(),89            torch.nn.Dropout(0.1),90            torch.nn.Linear(hidden_size, 512),91            torch.nn.LayerNorm(512),92            torch.nn.GELU(),93            torch.nn.Dropout(0.1),94            torch.nn.Linear(512, num_labels)95        )96 97    def forward(self, input_ids, attention_mask):98        outputs = self.encoder(99            input_ids=input_ids,100            attention_mask=attention_mask,101            return_dict=True102        )103 104        pooled_output = outputs.pooler_output105        logits = self.classifier(pooled_output)106 107        return logits108 109def extract_features(source_code, test_code_1, test_code_2):110    """Extract specific features to help the model identify similarities"""111    112    # Extract test fixtures113    fixture1 = re.search(r'TEST(?:_F)?\s*\(\s*(\w+)', test_code_1)114    fixture1 = fixture1.group(1) if fixture1 else ""115 116    fixture2 = re.search(r'TEST(?:_F)?\s*\(\s*(\w+)', test_code_2)117    fixture2 = fixture2.group(1) if fixture2 else ""118 119    # Extract test names120    name1 = re.search(r'TEST(?:_F)?\s*\(\s*\w+\s*,\s*(\w+)', test_code_1)121    name1 = name1.group(1) if name1 else ""122 123    name2 = re.search(r'TEST(?:_F)?\s*\(\s*\w+\s*,\s*(\w+)', test_code_2)124    name2 = name2.group(1) if name2 else ""125 126    # Extract assertions127    assertions1 = re.findall(r'(EXPECT_|ASSERT_)(\w+)', test_code_1)128    assertions2 = re.findall(r'(EXPECT_|ASSERT_)(\w+)', test_code_2)129 130    # Extract function/method calls131    calls1 = re.findall(r'(\w+)\s*\(', test_code_1)132    calls2 = re.findall(r'(\w+)\s*\(', test_code_2)133 134    # Create explicit feature section135    same_fixture = "SAME_FIXTURE" if fixture1 == fixture2 else "DIFFERENT_FIXTURE"136    common_assertions = set([a[0] + a[1] for a in assertions1]).intersection(set([a[0] + a[1] for a in assertions2]))137    common_calls = set(calls1).intersection(set(calls2))138    139    # Calculate assertion ratio with safety check for zero140    assertion_ratio = 0141    if assertions1 and assertions2:142        total_assertions = len(assertions1) + len(assertions2)143        if total_assertions > 0:144            assertion_ratio = len(common_assertions) / total_assertions145 146    features = (147        f"METADATA: {same_fixture} | "148        f"FIXTURE1: {fixture1} | FIXTURE2: {fixture2} | "149        f"NAME1: {name1} | NAME2: {name2} | "150        f"COMMON_ASSERTIONS: {len(common_assertions)} | "151        f"COMMON_CALLS: {len(common_calls)} | "152        f"ASSERTION_RATIO: {assertion_ratio}"153    )154 155    return features156 157# Global variables for model and tokenizer158tokenizer = None159model = None160 161def download_model_weights(url, save_path):162    """Download model weights from URL to a local file"""163    try:164        logger.info(f"Downloading model weights from {url}...")165        response = requests.get(url, stream=True)166        if response.status_code != 200:167            logger.error(f"Failed to download: HTTP {response.status_code}")168            return False169            170        with open(save_path, 'wb') as f:171            for chunk in response.iter_content(chunk_size=8192):172                if chunk:173                    f.write(chunk)174        logger.info(f"Successfully downloaded model weights to {save_path}")175        return True176    except Exception as e:177        logger.error(f"Error downloading model weights: {e}")178        return False179 180# Load model and tokenizer on startup181@app.on_event("startup")182async def startup_event():183    global tokenizer, model184    185    try:186        logger.info("=== Starting model loading process ===")187        188        # Step 1: Load the tokenizer from the base model189        logger.info(f"Loading tokenizer from microsoft/codebert-base...")190        try:191            tokenizer = AutoTokenizer.from_pretrained("microsoft/codebert-base")192            logger.info("✅ Base tokenizer loaded successfully")193        except Exception as e:194            logger.error(f"❌ Failed to load tokenizer: {str(e)}")195            raise196        197        # Step 2: Create model with base architecture198        logger.info("Creating model architecture...")199        try:200            # Initialize with base CodeBERT201            model = CodeSimilarityClassifier(model_name="microsoft/codebert-base")202            logger.info("✅ Model architecture created successfully")203        except Exception as e:204            logger.error(f"❌ Failed to create model architecture: {str(e)}")205            raise206        207        # Step 3: Download and load weights208        model_path = "pytorch_model.bin"209        210        # First check if the file already exists211        if not os.path.exists(model_path):212            # Try downloading213            if not download_model_weights(MODEL_WEIGHTS_URL, model_path):214                logger.error("❌ Failed to download model weights")215                raise RuntimeError("Failed to download model weights")216        217        # Try to load the model weights218        try:219            # Check if the weights are a state dict or the whole model220            logger.info(f"Loading weights from {model_path}...")221            checkpoint = torch.load(model_path, map_location=device)222            223            if isinstance(checkpoint, dict):224                # If it's a state dict directly225                if "state_dict" in checkpoint:226                    logger.info("Loading from checkpoint['state_dict']")227                    model.load_state_dict(checkpoint["state_dict"])228                elif "model_state_dict" in checkpoint:229                    logger.info("Loading from checkpoint['model_state_dict']")230                    model.load_state_dict(checkpoint["model_state_dict"])231                else:232                    logger.info("Loading from checkpoint directly")233                    model.load_state_dict(checkpoint)234            else:235                logger.error("❌ Unsupported model format")236                raise RuntimeError("Unsupported model format")237                238            logger.info("✅ Model weights loaded successfully")239        except Exception as e:240            logger.error(f"❌ Error loading model weights: {str(e)}")241            raise242        243        # Move model to device and set to evaluation mode244        model.to(device)245        model.eval()246        logger.info(f"✅ Model moved to {device} and set to evaluation mode")247        logger.info("=== Model loading process complete ===")248        249    except Exception as e:250        logger.error(f"❌ CRITICAL ERROR in startup: {str(e)}")251        import traceback252        logger.error(traceback.format_exc())253        model = None254        tokenizer = None255 256@app.get("/health")257async def health_check():258    """Health check endpoint that also returns deployment information"""259    model_status = model is not None260    tokenizer_status = tokenizer is not None261    status = "ok" if (model_status and tokenizer_status) else "error"262    263    return {264        "status": status, 265        "model_loaded": model_status,266        "tokenizer_loaded": tokenizer_status,267        "model": REPO_ID, 268        "device": str(device),269        "deployment_date": DEPLOYMENT_DATE,270        "deployed_by": DEPLOYED_BY,271        "current_time": datetime.utcnow().strftime("%Y-%m-%d %H:%M:%S")272    }273 274@app.post("/predict")275async def predict(data: SimilarityInput):276    """277    Predict similarity class between two test cases for a given source class.278    """279    if model is None or tokenizer is None:280        raise HTTPException(status_code=500, detail="Model not loaded correctly")281    282    try:283        # Apply heuristics for method and class differences284        class_1 = data.test_case_1.target_class285        class_2 = data.test_case_2.target_class286        method_1 = data.test_case_1.target_method287        method_2 = data.test_case_2.target_method288        289        # Check if we can determine similarity without using the model290        if class_1 and class_2 and class_1 != class_2:291            logger.info(f"Heuristic detection: Different target classes - Distinct")292            model_prediction = 2  # Distinct293            probs = [0.0, 0.0, 1.0]  # 100% confidence in Distinct294        elif method_1 and method_2 and not set(method_1).intersection(set(method_2)):295            logger.info(f"Heuristic detection: Different target methods - Distinct")296            model_prediction = 2  # Distinct297            probs = [0.0, 0.0, 1.0]  # 100% confidence in Distinct298        else:299            # No clear heuristic match, use the model300            # Extract features to help with classification301            features = extract_features(data.source_code.code, data.test_case_1.code, data.test_case_2.code)302            303            # Format the input text with clear section markers as done during training304            formatted_text = (305                f"{features}\n\n"306                f"SOURCE CODE:\n{data.source_code.code.strip()}\n\n"307                f"TEST CASE 1:\n{data.test_case_1.code.strip()}\n\n"308                f"TEST CASE 2:\n{data.test_case_2.code.strip()}"309            )310 311            # Tokenize input312            inputs = tokenizer(313                formatted_text, 314                return_tensors="pt", 315                padding="max_length", 316                truncation=True, 317                max_length=512318            ).to(device)319 320            # Model inference321            with torch.no_grad():322                logits = model(323                    input_ids=inputs["input_ids"],324                    attention_mask=inputs["attention_mask"]325                )326 327            # Process results328            probs = F.softmax(logits, dim=-1)[0].cpu().tolist()329            model_prediction = torch.argmax(logits, dim=-1).item()330            logger.info(f"Model prediction: {label_to_class[model_prediction]}")331        332        # Map prediction to class name333        classification = label_to_class.get(model_prediction, "Unknown")334        335        # For API compatibility, map the model outputs (0,1,2) to API scores (1,2,3)336        api_score = model_prediction + 1337        338        return {339            "pair_id": data.pair_id,340            "test_case_1_name": data.test_case_1.name,341            "test_case_2_name": data.test_case_2.name,342            "similarity": {343                "score": api_score,344                "classification": classification,345            },346            "probabilities": probs347        }348    349    except Exception as e:350        import traceback351        error_trace = traceback.format_exc()352        logger.error(f"Prediction error: {str(e)}")353        logger.error(error_trace)354        raise HTTPException(status_code=500, detail=f"Prediction error: {str(e)}")355 356# Root and example endpoints357@app.get("/")358async def root():359    return {360        "message": "Test Similarity Analyzer API",361        "documentation": "/docs",362        "deployment_date": DEPLOYMENT_DATE,363        "deployed_by": DEPLOYED_BY364    }365 366@app.get("/example", response_model=SimilarityInput)367async def get_example():368    """Get an example input to test the API"""369    return SimilarityInput(370        pair_id="example-1",371        source_code=SourceCode(372            class_name="Calculator",373            code="class Calculator {\n    public int add(int a, int b) {\n        return a + b;\n    }\n}"374        ),375        test_case_1=TestCase(376            id="test-1",377            test_fixture="CalculatorTest",378            name="testAddsTwoPositiveNumbers",379            code="TEST(CalculatorTest, AddsTwoPositiveNumbers) {\n    Calculator calc;\n    EXPECT_EQ(5, calc.add(2, 3));\n}",380            target_class="Calculator",381            target_method=["add"]382        ),383        test_case_2=TestCase(384            id="test-2",385            test_fixture="CalculatorTest",386            name="testAddsTwoPositiveIntegers",387            code="TEST(CalculatorTest, AddsTwoPositiveIntegers) {\n    Calculator calc;\n    EXPECT_EQ(5, calc.add(2, 3));\n}",388            target_class="Calculator",389            target_method=["add"]390        )391    )392 393if __name__ == "__main__":394    uvicorn.run("app:app", host="0.0.0.0", port=7860, reload=True)