FastestAI/CodeBert_Redundant_Detection_Task
0
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)