Team Ai
Apppublic

Hannounaaa/phi2-python-bugfixer-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
app.py433 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 the following 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 125Provide:1261. Bug explanation1272. Fixed code1283. Why the fix works129 130### output:131"""132        temp = temperature_override if temperature_override is not None else 0.7133        max_tokens = 1024134        do_sample = True135    136    elif task == "chat":137        prompt = f"""### instruction:138{code}139 140### output:141"""142        temp = temperature_override if temperature_override is not None else 0.2143        max_tokens = 200144        do_sample = True145    146    else:147        return generate_response(code, "fix", use_base_model, temperature_override)148    149    # use ONE model object (PEFT) and toggle adapter off for base inference150    model_name = 'BASE' if use_base_model else 'FINE-TUNED'151    print(f'Using {model_name} model')152 153    inputs = tokenizer(prompt, return_tensors='pt')154    inputs = {k: v.to(finetuned_model.device) for k, v in inputs.items()}155 156    # generate response157    ctx = finetuned_model.disable_adapter() if use_base_model else nullcontext()158    with torch.no_grad():159        with ctx:160            if do_sample:161                outputs = finetuned_model.generate(162                    **inputs,163                    max_new_tokens=max_tokens,164                    temperature=temp,165                    do_sample=True,166                    top_p=0.95,167                    pad_token_id=tokenizer.eos_token_id,168                    eos_token_id=tokenizer.eos_token_id169                )170            else:171                outputs = finetuned_model.generate(172                    **inputs,173                    max_new_tokens=max_tokens,174                    do_sample=False,175                    pad_token_id=tokenizer.eos_token_id,176                    eos_token_id=tokenizer.eos_token_id177                )178    179    response = tokenizer.decode(outputs[0], skip_special_tokens=True)180    if "### output:" in response:181        response = response.split("### output:")[-1].strip()182    183    return response184 185# Regex pattern used in "fix" task to split code from tests186TOPLEVEL_TEST_RE = re.compile(r"(?m)^(assert\s+|print\()")187 188def longest_valid_python_prefix(code: str) -> str:189    """Keep the longest prefix that parses as Python (prevents explanation text leaking into code)."""190    lines = code.splitlines()191    best = ""192    buf = []193    for line in lines:194        buf.append(line)195        candidate = "\n".join(buf).strip() + "\n"196        try:197            ast.parse(candidate)198            best = "\n".join(buf)199        except SyntaxError:200            pass201    return best.strip() if best.strip() else code.strip()202 203 204def parse_response(text: str, task: str) -> CodeResponse:205    """Parse model response based on actual output format"""206    raw = text.strip()207 208    if task == "fix":209        # extract only the function/class definition210        code_part = TOPLEVEL_TEST_RE.split(raw, maxsplit=1)[0].strip()211        fixed_code = longest_valid_python_prefix(code_part)212        213        # Remove trailing comments and empty lines214        lines = fixed_code.split('\n')215        while lines:216            last_line = lines[-1].strip()217            if not last_line or last_line.startswith('#'):218                lines.pop()219            else:220                break221        222        fixed_code = '\n'.join(lines)223        224        return CodeResponse(text=raw, detectedBugs=[], fixedCode=fixed_code)225 226    if task == "analyze":227        lines = raw.split('\n')228        229        # fixed code230        fixed_code = None231        code_lines = []  232 233        fixed_code_match = re.search(234            r'###?\s*fixed\s+code.*?:?\s*\n+```(?:python)?\s*\n(.*?)\n```',235            raw,236            re.IGNORECASE | re.DOTALL237        )238        if fixed_code_match:239            fixed_code = fixed_code_match.group(1).strip()240        241        # Fallback: extract function/class definition from raw text242        if not fixed_code:243            in_function = False244            245            for line in lines:246                if line.startswith('def ') or line.startswith('class '):247                    in_function = True248                    code_lines.append(line)249                elif in_function and (line.startswith('    ') or line.startswith('\t') or 250                                      line.strip().startswith('"""') or line.strip().startswith("'''")):251                    code_lines.append(line)252                elif in_function and line.strip() and not line.startswith(' '):253                    break254                elif in_function:255                    code_lines.append(line)256            257            fixed_code = '\n'.join(code_lines).strip() if code_lines else None258        259        # extract explanation260        explanation = None261        262        # "# Explanation:" comment section263        # captures everything from "# Explanation:" until "# Test" or end264        comment_explanation_match = re.search(265            r'#\s*Explanation[^:\n]*:?\s*\n((?:#[^\n]*\n?)+)',266            raw,267            re.IGNORECASE268        )269        if comment_explanation_match:270            explanation_text = comment_explanation_match.group(1)271            # remove # prefixes and join lines272            cleaned_lines = []273            for line in explanation_text.split('\n'):274                # remove leading # and whitespace275                clean = line.strip().lstrip('#').strip()276                # stop if we hit test-related content277                if clean.lower().startswith('test'):278                    break279                if clean:280                    cleaned_lines.append(clean)281            explanation = ' '.join(cleaned_lines) if cleaned_lines else None282        283        # Try markdown style ### Explanation:284        if not explanation:285            markdown_explanation_match = re.search(286                r'###?\s*(?:bug\s+)?explanation[^\n]*:?\s*\n+(.*?)(?=\n###?\s*(?:test|why|fixed)|#\s*Test|$)',287                raw,288                re.IGNORECASE | re.DOTALL289            )290            if markdown_explanation_match:291                explanation = markdown_explanation_match.group(1).strip()292                # Clean up any # prefixes293                cleaned_lines = []294                for line in explanation.split('\n'):295                    clean = line.strip().lstrip('#').strip()296                    if clean:297                        cleaned_lines.append(clean)298                explanation = ' '.join(cleaned_lines) if cleaned_lines else None299        300        # extract test assertion301        test_lines = []302        303        test_start_match = re.search(304            r'(?:#\s*Test(?:s|ing)?[^\n]*|###?\s*Test(?:s|ing)?[^\n]*)\n',305            raw,306            re.IGNORECASE307        )308        309        if test_start_match:310            test_section = raw[test_start_match.end():]311            312            # extract all assert statements313            for line in test_section.split('\n'):314                stripped = line.strip()315                if stripped.startswith('assert '):316                    # remove inline comments for cleaner output317                    clean_assert = re.sub(r'\s*#.*$', '', stripped)318                    test_lines.append(clean_assert)319        320        # Fallback: just find all assert statements in the raw text321        if not test_lines:322            for line in lines:323                stripped = line.strip()324                if stripped.startswith('assert '):325                    clean_assert = re.sub(r'\s*#.*$', '', stripped)326                    test_lines.append(clean_assert)327        328        tests = '\n'.join(test_lines) if test_lines else None329        330        # extract "WHY THE FIX WORKS" (mainly for base model)331        whyFix = None332        333        # Comment style: # Why the fix works:334        why_comment_match = re.search(335            r'#\s*Why\s+(?:the\s+)?fix\s+works:?\s*\n((?:#[^\n]*\n?)+)',336            raw,337            re.IGNORECASE338        )339        if why_comment_match:340            why_text = why_comment_match.group(1)341            cleaned_lines = []342            for line in why_text.split('\n'):343                clean = line.strip().lstrip('#').strip()344                if clean.lower().startswith('test'):345                    break346                if clean:347                    cleaned_lines.append(clean)348            whyFix = ' '.join(cleaned_lines) if cleaned_lines else None349        350        # Markdown style: ### Why the fix works:351        if not whyFix:352            why_markdown_match = re.search(353                r'###?\s*why\s+(?:the\s+)?fix\s+works[^\n]*:?\s*\n+(.*?)(?=\n###?\s*test|#\s*Test|Testing|$)',354                raw,355                re.IGNORECASE | re.DOTALL356            )357            if why_markdown_match:358                whyFix = why_markdown_match.group(1).strip()359                # Clean up360                cleaned_lines = []361                for line in whyFix.split('\n'):362                    clean = line.strip().lstrip('#').strip()363                    if clean:364                        cleaned_lines.append(clean)365                whyFix = ' '.join(cleaned_lines) if cleaned_lines else None366        367        #bug summary368        detected = []369        if explanation:370            # Take first sentence for the summary371            sentences = explanation.split('.')372            for s in sentences:373                s = s.strip()374                if len(s) > 20:375                    detected = [s + "."]376                    break377        378        if not detected:379            detected = ["Potential bug detected - see explanation for details"]380        381        return CodeResponse(382            text=raw,383            detectedBugs=detected,384            fixedCode=fixed_code,385            explanation=explanation,386            whyFix=whyFix,387            tests=tests,388        )389 390    # For chat or unknown tasks391    return CodeResponse(text=raw, detectedBugs=[], fixedCode=None)392 393 394@app.get("/")395def read_root():396    return {397        "name": "Python Bug Fixer API",398        "status": "running",399        "models": {400            "base": "microsoft/phi-2",401            "finetuned": "ashkanhsn/phi2-python-bugfixer (PEFT adapter)"402        },403        "version": "2.3.0",404        "info": "Shared base model approach - efficient memory usage"405    }406 407@app.get("/health")408def health_check():409    return {410        "status": "healthy",411        "base_model_loaded": base_model is not None,412        "finetuned_model_loaded": finetuned_model is not None,413        "device": str(device)414    }415 416 417@app.post("/analyze", response_model=CodeResponse)418async def analyze_code(request: CodeRequest):    419    try:420        print(f"Request: task={request.task}, use_base_model={request.use_base_model}")421        response_text = generate_response(422            request.code, 423            request.task, 424            request.use_base_model,425            request.temperature426        )427        return parse_response(response_text, request.task)428    except Exception as e:429        print(f"Error: {str(e)}")430        raise HTTPException(status_code=500, detail=f"Error processing request: {str(e)}")431 432if __name__ == "__main__":433    uvicorn.run(app, host="0.0.0.0", port=7860)