CoolFace
Apppublic

cactus183/patchbench-dev

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py172 linesDownload Raw Back to root
1"""2PatchBench baseline inference script.3 4Required env vars:5    API_BASE_URL  — LLM endpoint (default: HuggingFace Router)6    MODEL_NAME    — model identifier (default: Qwen/Qwen2.5-72B-Instruct)7    HF_TOKEN      — API key8 9Output format (strict, single-line per log):10    [START] task=<task_name> env=<benchmark> model=<model_name>11    [STEP] step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>12    [END] success=<true|false> steps=<n> score=<score> rewards=<r1,r2,...,rn>13"""14import os15import sys16import traceback17from typing import Optional18 19from openai import OpenAI20 21from patchbench.environment import PatchBenchEnv22from patchbench.models import PatchBenchAction23 24API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")25MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct")26API_KEY = os.getenv("HF_TOKEN") or os.getenv("API_KEY") or ""27BENCHMARK = "patchbench"28 29TASK_IDS = ["easy_01", "medium_01", "hard_01"]30 31SYSTEM_PROMPT = (32    "You are an expert Python developer. You will be given a buggy Python file "33    "and a summary of its failing tests. Return ONLY the corrected full Python "34    "source code for the file — no explanations, no markdown fences, no prose. "35    "Your output must be a complete drop-in replacement for the file."36)37 38 39def _sanitize_action(action_str: str) -> str:40    snippet = action_str[:80]41    snippet = snippet.replace("\r", "").replace("\n", "\\n")42    return snippet43 44 45def _sanitize_error(error: Optional[str]) -> str:46    if error is None:47        return "null"48    cleaned = str(error).replace("\n", " ").replace("\r", " ")49    return cleaned[:120]50 51 52def _extract_code(response_text: str) -> str:53    text = (response_text or "").strip()54    if text.startswith("```"):55        lines = text.split("\n")56        lines = lines[1:]57        if lines and lines[-1].strip().startswith("```"):58            lines = lines[:-1]59        text = "\n".join(lines)60    return text61 62 63def _call_llm(64    client: OpenAI, buggy_code: str, task_description: str, failing_tests: str65) -> tuple[str, Optional[str]]:66    user_msg = (67        f"Task: {task_description}\n\n"68        f"Current buggy code:\n```python\n{buggy_code}\n```\n\n"69        f"Test status:\n{failing_tests}\n\n"70        "Return the corrected full Python file now."71    )72    try:73        completion = client.chat.completions.create(74            model=MODEL_NAME,75            messages=[76                {"role": "system", "content": SYSTEM_PROMPT},77                {"role": "user", "content": user_msg},78            ],79            temperature=0.0,80            max_tokens=2048,81        )82        content = completion.choices[0].message.content or ""83        return (_extract_code(content), None)84    except Exception as exc:85        return ("", f"llm_error:{type(exc).__name__}:{exc}")86 87 88def run_episode(env: PatchBenchEnv, client: OpenAI, task_id: str) -> float:89    print(f"[START] task={task_id} env={BENCHMARK} model={MODEL_NAME}", flush=True)90 91    rewards: list[float] = []92    steps_taken = 093    final_done = False94    final_reward = 0.095 96    try:97        obs = env.reset(task_id=task_id)98        current_code = obs.buggy_code99 100        for step_idx in range(1, obs.max_steps + 1):101            steps_taken = step_idx102 103            patch, llm_error = _call_llm(104                client=client,105                buggy_code=current_code,106                task_description=obs.task_description,107                failing_tests=obs.failing_tests,108            )109            if not patch:110                patch = current_code111 112            try:113                obs = env.step(PatchBenchAction(patched_code=patch))114                reward = obs.reward115                done = obs.done116                step_error = llm_error117            except Exception as exc:118                reward = 0.0119                done = True120                step_error = f"env_error:{type(exc).__name__}:{exc}"121 122            rewards.append(reward)123            final_reward = reward124            final_done = done125 126            action_str = _sanitize_action(patch)127            error_str = _sanitize_error(step_error)128            done_str = "true" if done else "false"129            print(130                f"[STEP] step={step_idx} action={action_str!r} reward={reward:.2f} done={done_str} error={error_str}",131                flush=True,132            )133 134            if done:135                break136            current_code = obs.buggy_code137 138    except Exception:139        traceback.print_exc()140    finally:141        score = (sum(rewards) / len(rewards)) if rewards else 0.0142        score = max(0.0, min(1.0, score))143        success = final_done and final_reward >= 0.8144        success_str = "true" if success else "false"145        rewards_str = ",".join(f"{r:.2f}" for r in rewards) if rewards else "0.00"146        print(147            f"[END] success={success_str} steps={steps_taken} score={score:.2f} rewards={rewards_str}",148            flush=True,149        )150    return score151 152 153def main() -> int:154    if not API_KEY:155        print(156            "WARNING: no HF_TOKEN / API_KEY set; LLM calls will fail but script will still run.",157            file=sys.stderr,158            flush=True,159        )160 161    client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY or "placeholder")162    env = PatchBenchEnv()163 164    for task_id in TASK_IDS:165        run_episode(env, client, task_id)166 167    return 0168 169 170if __name__ == "__main__":171    sys.exit(main())172