Team Ai
Apppublic

ashkanhsn/phi2-python-bugfixer-api

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
app.py432 linesDownload Raw Back to root
1"""2HuggingFace Space - Python Bug Fixer API3FAST API backend for fine-tuned Phi-24 5FIXED VERSION: Better parsing of model output format6"""7 8import ast9import re10from contextlib import nullcontext11from typing import Optional, List, Dict12 13import torch14from fastapi import FastAPI, HTTPException15from fastapi.middleware.cors import CORSMiddleware16from pydantic import BaseModel, Field, ConfigDict17from transformers import AutoModelForCausalLM, AutoTokenizer18from peft import PeftModel19import uvicorn20 21 22app = FastAPI(23    title="Python Bug Fixer API",24    description="AI-powered Python bug detection and fixing",25    version="2.3.0"  # Updated version26)27 28# Enable CORS for VS Code extension29app.add_middleware(30    CORSMiddleware,31    allow_origins=["*"],32    allow_credentials=True,33    allow_methods=["*"],34    allow_headers=["*"],35)36 37# Global variables for models38base_model = None39finetuned_model = None40tokenizer = None41device = None42 43try:44    device = "cuda" if torch.cuda.is_available() else "cpu"45    print(f"Using device: {device}")46 47    # tokenizer48    print("Load tokenizer")49    tokenizer = AutoTokenizer.from_pretrained(50        "microsoft/phi-2",51        trust_remote_code=True52    )53    tokenizer.pad_token = tokenizer.eos_token54    print("Tokenizer loaded!")55 56    # load one base model (will be shared)57    print("Loading Phi-2 base model (fp16)")58    base_model = AutoModelForCausalLM.from_pretrained(59        "microsoft/phi-2",60        trust_remote_code=True,61        torch_dtype=torch.float16 if device == "cuda" else torch.float32,62        device_map="auto" if device == "cuda" else None,63    )64    base_model.eval()65    print("Base model loaded!")66 67    # load PEFT adapter on top of base model68    print("Load fine-tuned PEFT adapter")69    finetuned_model = PeftModel.from_pretrained(70        base_model,71        "ashkanhsn/phi2-python-bugfixer"72    )73    finetuned_model.eval()74    print("Fine-tuned model ready!")75    76    print("Both models are ready!")77    78except Exception as e:79    print(f"Error loading models: {e}")80    import traceback81    traceback.print_exc()82    raise83 84 85class CodeRequest(BaseModel):86    model_config = ConfigDict(populate_by_name=True)87    code: str88    task: str = 'fix'89    temperature: Optional[float] = None90    # Accept both snake_case and camelCase from clients91    use_base_model: bool = Field(False, alias='useBaseModel')92 93 94class CodeResponse(BaseModel):95    text: str96    detectedBugs: list[str]97    fixedCode: Optional[str] = None98    explanation: Optional[str] = None99    whyFix: Optional[str] = None100    tests: Optional[str] = None101    error: Optional[str] = None102 103 104def generate_response(code: str, task: str, use_base_model: bool = False, temperature_override: Optional[float] = None) -> str:105    106    # Build prompt based on task107    if task == "fix":108        prompt = f"""### instruction:109Fix this buggy Python code:110 111{code}112 113### output:114"""115        temp = temperature_override if temperature_override is not None else 0.2116        max_tokens = 300117        do_sample = temp > 0.0118 119    elif task == "analyze":120        prompt = f"""### instruction:121Fix this buggy Python code.122 123{code}124 125Respond with ### headers: Explanation, Fixed Code, Why the fix works.126 127### output:128### Fixed Code:129```python130"""131        temp = temperature_override if temperature_override is not None else 0.6132        max_tokens = 1024133        do_sample = True134    135    elif task == "chat":136        prompt = f"""### instruction:137{code}138 139### output:140"""141        temp = temperature_override if temperature_override is not None else 0.2142        max_tokens = 200143        do_sample = True144    145    else:146        return generate_response(code, "fix", use_base_model, temperature_override)147    148    # use ONE model object (PEFT) and toggle adapter off for base inference149    model_name = 'BASE' if use_base_model else 'FINE-TUNED'150    print(f'Using {model_name} model')151 152    inputs = tokenizer(prompt, return_tensors='pt')153    inputs = {k: v.to(finetuned_model.device) for k, v in inputs.items()}154 155    # generate response156    ctx = finetuned_model.disable_adapter() if use_base_model else nullcontext()157    with torch.no_grad():158        with ctx:159            if do_sample:160                outputs = finetuned_model.generate(161                    **inputs,162                    max_new_tokens=max_tokens,163                    temperature=temp,164                    do_sample=True,165                    top_p=0.95,166                    pad_token_id=tokenizer.eos_token_id,167                    eos_token_id=tokenizer.eos_token_id168                )169            else:170                outputs = finetuned_model.generate(171                    **inputs,172                    max_new_tokens=max_tokens,173                    do_sample=False,174                    pad_token_id=tokenizer.eos_token_id,175                    eos_token_id=tokenizer.eos_token_id176                )177    178    response = tokenizer.decode(outputs[0], skip_special_tokens=True)179    if "### output:" in response:180        response = response.split("### output:")[-1].strip()181    182    return response183 184# Regex pattern used in "fix" task to split code from tests185TOPLEVEL_TEST_RE = re.compile(r"(?m)^(assert\s+|print\()")186 187def longest_valid_python_prefix(code: str) -> str:188    """Keep the longest prefix that parses as Python (prevents explanation text leaking into code)."""189    lines = code.splitlines()190    best = ""191    buf = []192    for line in lines:193        buf.append(line)194        candidate = "\n".join(buf).strip() + "\n"195        try:196            ast.parse(candidate)197            best = "\n".join(buf)198        except SyntaxError:199            pass200    return best.strip() if best.strip() else code.strip()201 202 203def parse_response(text: str, task: str) -> CodeResponse:204    """Parse model response based on actual output format"""205    raw = text.strip()206 207    if task == "fix":208        # extract only the function/class definition209        code_part = TOPLEVEL_TEST_RE.split(raw, maxsplit=1)[0].strip()210        fixed_code = longest_valid_python_prefix(code_part)211        212        # Remove trailing comments and empty lines213        lines = fixed_code.split('\n')214        while lines:215            last_line = lines[-1].strip()216            if not last_line or last_line.startswith('#'):217                lines.pop()218            else:219                break220        221        fixed_code = '\n'.join(lines)222        223        return CodeResponse(text=raw, detectedBugs=[], fixedCode=fixed_code)224 225    if task == "analyze":226        lines = raw.split('\n')227        228        # fixed code229        fixed_code = None230        code_lines = []  231 232        fixed_code_match = re.search(233            r'###?\s*fixed\s+code.*?:?\s*\n+```(?:python)?\s*\n(.*?)\n```',234            raw,235            re.IGNORECASE | re.DOTALL236        )237        if fixed_code_match:238            fixed_code = fixed_code_match.group(1).strip()239        240        # Fallback: extract function/class definition from raw text241        if not fixed_code:242            in_function = False243            244            for line in lines:245                if line.startswith('def ') or line.startswith('class '):246                    in_function = True247                    code_lines.append(line)248                elif in_function and (line.startswith('    ') or line.startswith('\t') or 249                                      line.strip().startswith('"""') or line.strip().startswith("'''")):250                    code_lines.append(line)251                elif in_function and line.strip() and not line.startswith(' '):252                    break253                elif in_function:254                    code_lines.append(line)255            256            fixed_code = '\n'.join(code_lines).strip() if code_lines else None257        258        # extract explanation259        explanation = None260        261        # "# Explanation:" comment section262        # captures everything from "# Explanation:" until "# Test" or end263        comment_explanation_match = re.search(264            r'#\s*Explanation[^:\n]*:?\s*\n((?:#[^\n]*\n?)+)',265            raw,266            re.IGNORECASE267        )268        if comment_explanation_match:269            explanation_text = comment_explanation_match.group(1)270            # remove # prefixes and join lines271            cleaned_lines = []272            for line in explanation_text.split('\n'):273                # remove leading # and whitespace274                clean = line.strip().lstrip('#').strip()275                # stop if we hit test-related content276                if clean.lower().startswith('test'):277                    break278                if clean:279                    cleaned_lines.append(clean)280            explanation = ' '.join(cleaned_lines) if cleaned_lines else None281        282        # Try markdown style ### Explanation:283        if not explanation:284            markdown_explanation_match = re.search(285                r'###?\s*(?:bug\s+)?explanation[^\n]*:?\s*\n+(.*?)(?=\n###?\s*(?:test|why|fixed)|#\s*Test|$)',286                raw,287                re.IGNORECASE | re.DOTALL288            )289            if markdown_explanation_match:290                explanation = markdown_explanation_match.group(1).strip()291                # Clean up any # prefixes292                cleaned_lines = []293                for line in explanation.split('\n'):294                    clean = line.strip().lstrip('#').strip()295                    if clean:296                        cleaned_lines.append(clean)297                explanation = ' '.join(cleaned_lines) if cleaned_lines else None298        299        # extract test assertion300        test_lines = []301        302        test_start_match = re.search(303            r'(?:#\s*Test(?:s|ing)?[^\n]*|###?\s*Test(?:s|ing)?[^\n]*)\n',304            raw,305            re.IGNORECASE306        )307        308        if test_start_match:309            test_section = raw[test_start_match.end():]310            311            # extract all assert statements312            for line in test_section.split('\n'):313                stripped = line.strip()314                if stripped.startswith('assert '):315                    # remove inline comments for cleaner output316                    clean_assert = re.sub(r'\s*#.*$', '', stripped)317                    test_lines.append(clean_assert)318        319        # Fallback: just find all assert statements in the raw text320        if not test_lines:321            for line in lines:322                stripped = line.strip()323                if stripped.startswith('assert '):324                    clean_assert = re.sub(r'\s*#.*$', '', stripped)325                    test_lines.append(clean_assert)326        327        tests = '\n'.join(test_lines) if test_lines else None328        329        # extract "WHY THE FIX WORKS" (mainly for base model)330        whyFix = None331        332        # Comment style: # Why the fix works:333        why_comment_match = re.search(334            r'#\s*Why\s+(?:the\s+)?fix\s+works:?\s*\n((?:#[^\n]*\n?)+)',335            raw,336            re.IGNORECASE337        )338        if why_comment_match:339            why_text = why_comment_match.group(1)340            cleaned_lines = []341            for line in why_text.split('\n'):342                clean = line.strip().lstrip('#').strip()343                if clean.lower().startswith('test'):344                    break345                if clean:346                    cleaned_lines.append(clean)347            whyFix = ' '.join(cleaned_lines) if cleaned_lines else None348        349        # Markdown style: ### Why the fix works:350        if not whyFix:351            why_markdown_match = re.search(352                r'###?\s*why\s+(?:the\s+)?fix\s+works[^\n]*:?\s*\n+(.*?)(?=\n###?\s*test|#\s*Test|Testing|$)',353                raw,354                re.IGNORECASE | re.DOTALL355            )356            if why_markdown_match:357                whyFix = why_markdown_match.group(1).strip()358                # Clean up359                cleaned_lines = []360                for line in whyFix.split('\n'):361                    clean = line.strip().lstrip('#').strip()362                    if clean:363                        cleaned_lines.append(clean)364                whyFix = ' '.join(cleaned_lines) if cleaned_lines else None365        366        #bug summary367        detected = []368        if explanation:369            # Take first sentence for the summary370            sentences = explanation.split('.')371            for s in sentences:372                s = s.strip()373                if len(s) > 20:374                    detected = [s + "."]375                    break376        377        if not detected:378            detected = ["Potential bug detected - see explanation for details"]379        380        return CodeResponse(381            text=raw,382            detectedBugs=detected,383            fixedCode=fixed_code,384            explanation=explanation,385            whyFix=whyFix,386            tests=tests,387        )388 389    # For chat or unknown tasks390    return CodeResponse(text=raw, detectedBugs=[], fixedCode=None)391 392 393@app.get("/")394def read_root():395    return {396        "name": "Python Bug Fixer API",397        "status": "running",398        "models": {399            "base": "microsoft/phi-2",400            "finetuned": "ashkanhsn/phi2-python-bugfixer (PEFT adapter)"401        },402        "version": "2.3.0",403        "info": "Shared base model approach - efficient memory usage"404    }405 406@app.get("/health")407def health_check():408    return {409        "status": "healthy",410        "base_model_loaded": base_model is not None,411        "finetuned_model_loaded": finetuned_model is not None,412        "device": str(device)413    }414 415 416@app.post("/analyze", response_model=CodeResponse)417async def analyze_code(request: CodeRequest):    418    try:419        print(f"Request: task={request.task}, use_base_model={request.use_base_model}")420        response_text = generate_response(421            request.code, 422            request.task, 423            request.use_base_model,424            request.temperature425        )426        return parse_response(response_text, request.task)427    except Exception as e:428        print(f"Error: {str(e)}")429        raise HTTPException(status_code=500, detail=f"Error processing request: {str(e)}")430 431if __name__ == "__main__":432    uvicorn.run(app, host="0.0.0.0", port=7860)