cactus183/patchbench-dev
0
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 