Team Ai
Apppublic

Sarra22/phi2-python-bugfixer-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
app.py380 linesDownload Raw Back to root
1"""2HuggingFace Space - Python Bug Fixer API3FAST API backend for fine-tuned Phi-24"""5 6import ast7import re8from contextlib import nullcontext9from typing import Optional, List, Dict10 11import torch12from fastapi import FastAPI, HTTPException13from fastapi.middleware.cors import CORSMiddleware14from pydantic import BaseModel, Field, ConfigDict15from transformers import AutoModelForCausalLM, AutoTokenizer16from peft import PeftModel17import uvicorn18 19 20app = FastAPI(21    title="Python Bug Fixer API",22    description="AI-powered Python bug detection and fixing",23    version="2.2.0"24)25 26# Enable CORS for VS Code extension27app.add_middleware(28    CORSMiddleware,29    allow_origins=["*"],30    allow_credentials=True,31    allow_methods=["*"],32    allow_headers=["*"],33)34 35# Global variables for models36base_model = None37finetuned_model = None38tokenizer = None39device = None40 41try:42    device = "cuda" if torch.cuda.is_available() else "cpu"43    print(f"Using device: {device}")44 45    # tokenizer46    print("Load tokenizer")47    tokenizer = AutoTokenizer.from_pretrained(48        "microsoft/phi-2",49        trust_remote_code=True50    )51    tokenizer.pad_token = tokenizer.eos_token52    print("Tokenizer loaded!")53 54    # load one base model (will be shared)55    print("Loading Phi-2 base model (fp16)")56    base_model = AutoModelForCausalLM.from_pretrained(57        "microsoft/phi-2",58        trust_remote_code=True,59        torch_dtype=torch.float16 if device == "cuda" else torch.float32,60        device_map="auto" if device == "cuda" else None,61    )62    base_model.eval()63    print("Base model loaded!")64 65    # load PEFT adapter on top of base model66    print("Load fine-tuned PEFT adapter")67    finetuned_model = PeftModel.from_pretrained(68        base_model,69        "ashkanhsn/phi2-python-bugfixer"70    )71    finetuned_model.eval()72    print("Fine-tuned model ready!")73    74    print("Both models are ready!")75    76except Exception as e:77    print(f"Error loading models: {e}")78    import traceback79    traceback.print_exc()80    raise81 82 83class CodeRequest(BaseModel):84    model_config = ConfigDict(populate_by_name=True)85    code: str86    task: str = 'fix'87    temperature: Optional[float] = None88    # Accept both snake_case and camelCase from clients89    use_base_model: bool = Field(False, alias='useBaseModel')90 91 92class CodeResponse(BaseModel):93    text: str94    detectedBugs: list[str]95    fixedCode: Optional[str] = None96    explanation: Optional[str] = None97    whyFix: Optional[str] = None98    tests: Optional[str] = None99    error: Optional[str] = None100 101 102def generate_response(code: str, task: str, use_base_model: bool = False, temperature_override: Optional[float] = None) -> str:103    104    # Build prompt based on task105    if task == "fix":106        prompt = f"""### instruction:107Fix the following buggy Python code:108 109{code}110 111### output:112"""113        temp = temperature_override if temperature_override is not None else 0.2114        max_tokens = 300115        do_sample = temp > 0.0116 117    elif task == "analyze":118        prompt = f"""### instruction:119Fix this buggy Python code.120 121{code}122 123Provide:1241. Bug explanation1252. Fixed code1263. Why the fix works127 128### output:129"""130        temp = temperature_override if temperature_override is not None else 0.7131        max_tokens = 1024132        do_sample = True133    134    elif task == "chat":135        prompt = f"""### instruction:136{code}137 138### output:139"""140        temp = temperature_override if temperature_override is not None else 0.2141        max_tokens = 200142        do_sample = True143    144    else:145        return generate_response(code, "fix", use_base_model, temperature_override)146    147    # use ONE model object (PEFT) and toggle adapter off for base inference148    model_name = 'BASE' if use_base_model else 'FINE-TUNED'149    print(f'Using {model_name} model')150 151    inputs = tokenizer(prompt, return_tensors='pt')152    inputs = {k: v.to(finetuned_model.device) for k, v in inputs.items()}153 154    # generate response155    ctx = finetuned_model.disable_adapter() if use_base_model else nullcontext()156    with torch.no_grad():157        with ctx:158            if do_sample:159                outputs = finetuned_model.generate(160                    **inputs,161                    max_new_tokens=max_tokens,162                    temperature=temp,163                    do_sample=True,164                    top_p=0.95,165                    pad_token_id=tokenizer.eos_token_id,166                    eos_token_id=tokenizer.eos_token_id167                )168            else:169                outputs = finetuned_model.generate(170                    **inputs,171                    max_new_tokens=max_tokens,172                    do_sample=False,173                    pad_token_id=tokenizer.eos_token_id,174                    eos_token_id=tokenizer.eos_token_id175                )176    177    response = tokenizer.decode(outputs[0], skip_special_tokens=True)178    if "### output:" in response:179        response = response.split("### output:")[-1].strip()180    181    return response182 183STOP_HEADING_RE = re.compile(184    r"""(?im)^\s*(?:###\s*(?:explanation|why.*works|testing|tests?)\s*:?185            |#\s*Explanation(?:\s+of\s+the\s+fix)?\s*:?186            |Explanation\s+of\s+the\s+fix\s*:?187            )\s*$""",188    re.VERBOSE,189)190TOPLEVEL_TEST_RE = re.compile(r"(?m)^(assert\s+|print\()")191TEST_START_RE = re.compile(r"(?m)^(assert\s+|print\()")192FENCED_CODE_RE = re.compile(r"```(?:python)?\s*\n([\s\S]*?)\n```", re.IGNORECASE)193 194def extract_first_fenced_code(text: str) -> Optional[str]:195    m = FENCED_CODE_RE.search(text)196    return m.group(1).strip() if m else None197 198def longest_valid_python_prefix(code: str) -> str:199    """Keep the longest prefix that parses as Python (prevents explanation text leaking into code)."""200    lines = code.splitlines()201    best = ""202    buf = []203    for line in lines:204        buf.append(line)205        candidate = "\n".join(buf).strip() + "\n"206        try:207            ast.parse(candidate)208            best = "\n".join(buf)209        except SyntaxError:210            pass211    return best.strip() if best.strip() else code.strip()212 213def extract_section(text: str, heading_patterns: List[str]) -> Optional[str]:214    """extract section content following a heading, stopping at next heading or top-level tests"""215    for hp in heading_patterns:216        m = re.search(hp, text, flags=re.IGNORECASE | re.MULTILINE)217        if not m:218            continue219        start = m.end()220        tail = text[start:]221 222        # stop at next heading or tests223        stop_m = re.search(r"(?im)^\s*###\s*[A-Za-z].*?:?\s*$|(?m)^(assert\s+|print\()", tail)224        chunk = tail[: stop_m.start()] if stop_m else tail225        chunk = chunk.strip()226        if len(chunk) >= 5:227            return chunk228    return None229 230def extract_tests(text: str) -> Optional[str]:231    # prefer explicit "Tests" section232    t = extract_section(text, [233        r"(?im)^\s*###\s*Tests?\s*:?\s*$",234        r"(?im)^\s*###\s*Testing\s*:?\s*$",235        r"(?im)^\s*#\s*Tests?\s*:?\s*$",236        r"(?im)^\s*#\s*Testing\s*:?\s*$",237    ])238    if t:239        # keep only assert/print lines if the section contains mixed content240        lines = [ln.rstrip() for ln in t.splitlines() if ln.strip()]241        # If there are asserts, show from first assert onward242        for i, ln in enumerate(lines):243            if ln.lstrip().startswith("assert"):244                return "\n".join(lines[i:]).strip()245        return "\n".join(lines).strip()246 247    # Fallback: from first top-level assert to end248    m = re.search(r"(?m)^assert\s+.*$", text)249    if m:250        return text[m.start():].strip()251    return None252 253def guess_bug_summaries(explanation: Optional[str]) -> List[str]:254    # if not explanation:255    #     return ["Code analyzed and fixed - check explanation for details"]256    # # Take first 1-2 decent sentences257    # cleaned = " ".join([ln.strip().lstrip("#").strip() for ln in explanation.splitlines() if ln.strip()])258    # sentences = re.split(r"[.!?]\s+", cleaned)259    # out = []260    # for s in sentences:261    #     s = s.strip()262    #     if len(s) >= 18:263    #         out.append(s + ".")264    #     if len(out) == 2:265    #         break266    # return out if out else ["Code analyzed and fixed - check explanation for details"]267    return []268 269def parse_response(text: str, task: str) -> CodeResponse:270    """parse model response"""271    raw = text.strip()272 273    if task == "fix":274        # extract only the function/class definition275        code_part = TOPLEVEL_TEST_RE.split(raw, maxsplit=1)[0].strip()276        fixed_code = longest_valid_python_prefix(code_part)277        return CodeResponse(text=raw, detectedBugs=[], fixedCode=fixed_code)278 279    if task == "analyze":280        lines = raw.split('\n')281        282        # extract the fixed code (function/class definition only)283        code_lines = []284        in_function = False285        286        for line in lines:287            if line.startswith('def ') or line.startswith('class '):288                in_function = True289                code_lines.append(line)290            elif in_function and (line.startswith('    ') or line.startswith('\t') or line.strip().startswith('"""') or line.strip().startswith("'''")):291                code_lines.append(line)292            elif in_function and line.strip() and not line.startswith(' '):293                break294            elif in_function:295                code_lines.append(line)296        297        fixed_code = '\n'.join(code_lines).strip()298        299        # extract explanatory comments (lines starting with #)300        explanation_lines = []301        for line in lines:302            stripped = line.strip()303            # Get comment lines that are NOT part of the function definition304            if stripped.startswith('#') and line not in code_lines:305                explanation_lines.append(stripped.lstrip('#').strip())306        307        explanation = '\n'.join(explanation_lines).strip() if explanation_lines else None308        309        # extract test assertions310        test_lines = []311        for line in lines:312            stripped = line.strip()313            if stripped.startswith('assert '):314                test_lines.append(stripped)315        316        tests = '\n'.join(test_lines).strip() if test_lines else None317        318        # create bug summary from explanation319        detected = []320        if explanation:321            cleaned = explanation.replace("Explanation:", "").replace("\n", " ").strip()322            sentences = cleaned.split('.')323            if sentences and len(sentences[0].strip()) > 20:324                detected = [sentences[0].strip() + "."]325        326        if not detected:327            detected = ["Potential bug detected - see explanation for details"]328        329        return CodeResponse(330            text=raw,331            detectedBugs=detected,332            fixedCode=fixed_code,333            explanation=explanation,334            whyFix=None,  # Your model doesn't separate this335            tests=tests,336        )337 338    # For chat or unknown tasks339    return CodeResponse(text=raw, detectedBugs=[], fixedCode=None)340 341@app.get("/")342def read_root():343    return {344        "name": "Python Bug Fixer API",345        "status": "running",346        "models": {347            "base": "microsoft/phi-2",348            "finetuned": "ashkanhsn/phi2-python-bugfixer (PEFT adapter)"349        },350        "version": "2.1.0",351        "info": "Shared base model approach - efficient memory usage"352    }353 354@app.get("/health")355def health_check():356    return {357        "status": "healthy",358        "base_model_loaded": base_model is not None,359        "finetuned_model_loaded": finetuned_model is not None,360        "device": str(device)361    }362 363 364@app.post("/analyze", response_model=CodeResponse)365async def analyze_code(request: CodeRequest):    366    try:367        print(f"Request: task={request.task}, use_base_model={request.use_base_model}")368        response_text = generate_response(369            request.code, 370            request.task, 371            request.use_base_model,372            request.temperature373        )374        return parse_response(response_text, request.task)375    except Exception as e:376        print(f"Error: {str(e)}")377        raise HTTPException(status_code=500, detail=f"Error processing request: {str(e)}")378 379if __name__ == "__main__":380    uvicorn.run(app, host="0.0.0.0", port=7860)