Team Ai
Apppublic

PRANAV05092003/autonomous-code-refactoring-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
inference.py395 linesDownload Raw Back to root
1"""2ACRE inference script for OpenEnv submission evaluation.3 4Environment variables:5    - API_BASE_URL: LLM API endpoint injected by evaluator6  - MODEL_NAME: model identifier (default allowed)7    - API_KEY: API token for the OpenAI-compatible proxy endpoint8  - ENV_URL: running ACRE server base URL (required)9  - LOCAL_IMAGE_NAME: present for evaluator compatibility (optional)10    - USE_LLM: set to "0" to disable LLM action selection11 12STRICT stdout format (do not change):13    [START] task=<task_id>14    [STEP] action=<action_int>15    [END] task=<task_id> score=<score_float>16"""17from __future__ import annotations18 19import json20import os21import re22import sys23import time24from typing import Dict, List, Optional, Tuple25 26import requests27from openai import OpenAI28 29MODEL_NAME = os.getenv("MODEL_NAME") or "gpt-4o-mini"30# Phase-2 validator expects API_KEY through provided proxy.31API_KEY = os.getenv("API_KEY")32ENV_URL: str = os.getenv("ENV_URL", "http://localhost:7860")33LOCAL_IMAGE_NAME: str | None = os.getenv("LOCAL_IMAGE_NAME")34 35TASKS: List[str] = ["rename_variables", "remove_dead_code", "full_refactor"]36 37ACTION_MEANINGS: Dict[int, str] = {38    0: "rename_variable",39    1: "remove_dead_code",40    2: "simplify_loop",41    3: "optimize_condition",42    4: "inline_function",43}44 45SYSTEM_PROMPT = """\46You are an RL agent that refactors Python code. Choose one action per step.47 48Actions:49  0 rename_variable   - rename generic names (x, tmp, i) to descriptive ones50  1 remove_dead_code  - remove unreachable stmts, if False blocks, unused vars51  2 simplify_loop     - convert append-loops to list comprehensions52  3 optimize_condition- simplify 'not not x', 'if True/False', 'x==True'53  4 inline_function   - inline simple single-return module-level functions54 55Respond ONLY with valid JSON (no markdown):56{"action": <0-4>, "reason": "<one sentence>"}"""57 58SAFE_FALLBACK_SCORES: Dict[str, float] = {59    "easy": 0.0,60    "medium": 0.0,61    "hard": 0.0,62    "final": 0.0,63}64 65 66def _safe_scores() -> Dict[str, float]:67    return dict(SAFE_FALLBACK_SCORES)68 69 70def _env_url() -> str:71    # Never crash due to missing env var.72    return str(ENV_URL or "http://localhost:7860").rstrip("/")73 74 75def _post(path: str, payload: dict | None = None) -> dict:76    try:77        response = requests.post(f"{_env_url()}{path}", json=payload or {}, timeout=5)78        response.raise_for_status()79        return response.json()80    except Exception:81        print("Warning: Could not reach environment", file=sys.stderr)82        return {}83 84 85def _get(path: str) -> dict:86    try:87        response = requests.get(f"{_env_url()}{path}", timeout=5)88        response.raise_for_status()89        return response.json()90    except Exception:91        print("Warning: Could not reach environment", file=sys.stderr)92        return {}93 94 95def reset_env(task_id: str) -> dict:96    return _post("/reset", {"task_id": task_id})97 98 99def step_env(action: int) -> dict:100    return _post("/step", {"action": action})101 102 103def get_state() -> dict:104    return _get("/state")105 106 107def grade(task_id: str, code: str) -> float:108    try:109        response = requests.post(110            f"{_env_url()}/tasks/{task_id}/grade",111            json={"code": code},112            timeout=5,113        )114        response.raise_for_status()115        return float(response.json().get("score", 0.0))116    except Exception:117        print("Warning: Could not reach environment", file=sys.stderr)118        return 0.0119 120 121def choose_action(client: Optional[OpenAI], state: dict, task_id: str) -> Tuple[int, str]:122    def heuristic_action() -> Tuple[int, str]:123        code = str(state.get("current_code", ""))124        step_i = int(state.get("episode_steps", 0))125 126        has_generic = re.search(r"\b(x|tmp|i)\b", code) is not None127        has_if_false = re.search(r"\bif\s+False\b", code) is not None128        has_if_true = re.search(r"\bif\s+True\b", code) is not None129        has_append_loop = ".append(" in code and "for " in code130        has_double_not = "not not" in code131        has_add_call = "add(" in code132 133        if task_id == "rename_variables":134            if has_generic:135                return 0, "heuristic: remove generic names first"136            if has_if_false or "unused" in code:137                return 1, "heuristic: remove dead code"138            if has_append_loop:139                return 2, "heuristic: simplify loop"140            if has_if_true or has_double_not:141                return 3, "heuristic: optimize conditions"142            return 4, "heuristic: inline simple function"143 144        if task_id == "remove_dead_code":145            if has_if_false or "unused" in code:146                return 1, "heuristic: remove dead code patterns"147            if has_append_loop:148                return 2, "heuristic: convert append-loop"149            if has_if_true or has_double_not:150                return 3, "heuristic: simplify conditions"151            if has_generic:152                return 0, "heuristic: clean generic names"153            return 4, "heuristic: inline helper"154 155        if has_generic:156            return 0, "heuristic: rename generic variables"157        if has_append_loop:158            return 2, "heuristic: simplify loop into listcomp"159        if has_if_false or has_if_true or has_double_not:160            return 3, "heuristic: optimize boolean branches"161        if has_add_call:162            return 4, "heuristic: inline add() call"163        if step_i >= 2:164            return 1, "heuristic: remove remaining dead code"165        return 3, "heuristic: condition optimization as safe default"166 167    # Enable LLM by default when credentials are present.168    use_llm = bool(API_KEY) and os.getenv("USE_LLM", "1") == "1"169    if (not use_llm) or client is None:170        return heuristic_action()171 172    messages = [173        {"role": "system", "content": SYSTEM_PROMPT},174        {175            "role": "user",176            "content": (177                f"Task: {task_id}\n"178                f"Steps remaining: {state.get('max_steps', 5) - state.get('episode_steps', 0)}\n"179                f"Complexity: {state.get('complexity', 0)}\n\n"180                f"Current code:\n```python\n{state.get('current_code', '')}\n```\n\n"181                "Choose the best action."182            ),183        },184    ]185    try:186        response = client.chat.completions.create(187            model=MODEL_NAME,188            messages=messages,189            temperature=0.0,190            max_tokens=120,191        )192        raw = (response.choices[0].message.content or "").strip()193        json_blob = raw194 195        if "{" not in json_blob or "}" not in json_blob:196            return heuristic_action()197 198        match = re.search(r"\{.*\}", json_blob, flags=re.DOTALL)199        if match:200            json_blob = match.group(0)201 202        parsed = json.loads(json_blob)203        action = int(parsed.get("action", -1))204        reason = str(parsed.get("reason", ""))205        if 0 <= action <= 4:206            return action, reason or "llm-selected action"207        return heuristic_action()208    except Exception:209        return heuristic_action()210 211 212def _build_openai_client() -> Optional[OpenAI]:213    """214    Build OpenAI-compatible client using hackathon-required proxy env vars.215    Falls back safely when vars are absent in local runs.216    """217    base_url = os.getenv("API_BASE_URL")218    api_key = os.getenv("API_KEY")219 220    if not base_url or not api_key:221        return None222 223    try:224        return OpenAI(base_url=base_url, api_key=api_key)225    except Exception:226        return None227 228 229def _touch_proxy(client: Optional[OpenAI]) -> None:230    """231    Ensure at least one request is sent through the provided proxy in Phase-2.232    """233    if client is None:234        return None235    try:236        client.chat.completions.create(237            model=MODEL_NAME,238            messages=[{"role": "user", "content": "Return exactly: ok"}],239            temperature=0.0,240            max_tokens=2,241        )242    except Exception:243        # Keep inference resilient even if proxy is temporarily unavailable.244        return None245    return None246 247 248def run_episode(client: Optional[OpenAI], task_id: str, episode_num: int) -> float:249    reset_env(task_id)250    state = get_state()251 252    # STRICT logging format required by evaluator.253    print(f"[START] task={task_id}", flush=True)254 255    cumulative_reward = 0.0256 257    for step_num in range(1, 6):258        action, reason = choose_action(client, state, task_id)259        result = step_env(action)260        state = get_state()261 262        reward_payload = result.get("reward", {})263        raw_reward = float(reward_payload.get("raw", 0.0))264        norm_reward = float(reward_payload.get("normalized", (raw_reward + 32) / 52))265        cumulative_reward += raw_reward266 267        # STRICT logging format required by evaluator.268        print(f"[STEP] action={int(action)}", flush=True)269 270        if result.get("done") or result.get("terminated") or result.get("truncated"):271            break272 273    final_state = get_state()274    task_score = grade(task_id, final_state.get("current_code", ""))275 276    # STRICT logging format required by evaluator.277    print(f"[END] task={task_id} score={task_score:.4f}", flush=True)278 279    return task_score280 281 282def run_all_tasks() -> Dict[str, float]:283    """284    Run all three tasks and return deterministic scores.285 286    This is used by the FastAPI server to show live demo results on the Space.287    """288    try:289        # Prefer local in-process execution when running inside the server (no ENV_URL needed).290        try:291            from acre.tasks.task_registry import TaskRegistry292            from openenv_interface import OpenEnvRefactorEnv293        except Exception:294            TaskRegistry = None  # type: ignore[assignment]295            OpenEnvRefactorEnv = None  # type: ignore[assignment]296 297        registry = TaskRegistry() if TaskRegistry is not None else None298        env = OpenEnvRefactorEnv(registry=registry) if OpenEnvRefactorEnv is not None else None299 300        client = _build_openai_client()301        _touch_proxy(client)302 303        task_plan = [304            "rename_variables",305            "remove_dead_code",306            "full_refactor",307        ]308 309        results: Dict[str, float] = _safe_scores()310        scores: List[float] = []311 312        # If we have a local env, use it. Otherwise fall back to HTTP.313        if env is None or registry is None:314            # Network safety: quick health probe before running.315            try:316                r = requests.get(f"{_env_url()}/health", timeout=5)317                r.raise_for_status()318            except Exception:319                print("Warning: Could not reach environment", file=sys.stderr)320                return _safe_scores()321 322            for task_id in task_plan:323                print(f"[START] task={task_id}", flush=True)324                reset_env(task_id)325                for _ in range(5):326                    state = get_state()327                    action, _reason = choose_action(client, state, task_id)328                    print(f"[STEP] action={int(action)}", flush=True)329                    step_env(action)330                final_state = get_state()331                score = float(grade(task_id, final_state.get("current_code", "")))332                print(f"[END] task={task_id} score={float(score):.4f}", flush=True)333                scores.append(score)334                if task_id == "rename_variables":335                    results["easy"] = score336                elif task_id == "remove_dead_code":337                    results["medium"] = score338                else:339                    results["hard"] = score340 341            results["final"] = float(sum(scores) / len(scores)) if scores else 0.0342            return results343 344        else:345            # Local in-process execution (fast + no network recursion).346            for task_id in task_plan:347                print(f"[START] task={task_id}", flush=True)348                env.reset(seed=0, task_id=task_id)349                for _ in range(5):350                    st = env.state()351                    state_payload = {352                        "current_code": str(st.current_code),353                        "episode_steps": int(st.episode_steps),354                        "max_steps": int(st.max_steps),355                        "complexity": float(st.complexity),356                    }357                    action, _reason = choose_action(client, state_payload, task_id)358                    action = int(action)359                    print(f"[STEP] action={int(action)}", flush=True)360                    env.step(action)361                st = env.state()362                task = registry.get_task(task_id)363                score = float(task.grade_against_expected(st.current_code)) if task is not None else 0.0364                print(f"[END] task={task_id} score={float(score):.4f}", flush=True)365                scores.append(score)366                if task_id == "rename_variables":367                    results["easy"] = score368                elif task_id == "remove_dead_code":369                    results["medium"] = score370                else:371                    results["hard"] = score372 373        results["final"] = float(sum(scores) / len(scores)) if scores else 0.0374        return results375    except Exception as e:376        print(f"ERROR: {str(e)}", file=sys.stderr)377        return _safe_scores()378 379 380def main() -> None:381    # Never crash. Always produce output.382    result = run_all_tasks()383    print(f"Easy: {float(result.get('easy', 0.0)):.4f}", file=sys.stderr)384    print(f"Medium: {float(result.get('medium', 0.0)):.4f}", file=sys.stderr)385    print(f"Hard: {float(result.get('hard', 0.0)):.4f}", file=sys.stderr)386    print(f"Final: {float(result.get('final', 0.0)):.4f}", file=sys.stderr)387    return None388 389 390if __name__ == "__main__":391    try:392        run_all_tasks()393    except Exception as e:394        print(f"Fatal error: {e}", file=sys.stderr)395