CoolFace
Apppublic

AMD21/codefixerenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.py160 linesDownload Raw Back to env
1import os2import sys3sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))4 5from env.models import Observation, Action, Reward6from tasks.tasks import TASKS7 8class CodeFixerEnv:9    MAX_STEPS: int = 510 11    def __init__(self) -> None:12        self._task        = None13        self._grader      = None14        self._difficulty  = None15        self._history     = []16        self._step_number = 017        self._done        = False18        self._best_score  = 0.019        self._cumulative_reward = 0.020 21    def reset(self, difficulty: str = "easy") -> Observation:22        if difficulty not in TASKS:23            raise ValueError(24                f"Unknown difficulty '{difficulty}'. Valid choices: {sorted(TASKS.keys())}"25            )26 27        task_def, grader = TASKS[difficulty]28        self._task        = task_def29        self._grader      = grader30        self._difficulty  = difficulty31        self._history     = []32        self._step_number = 033        self._done        = False34        self._best_score  = 0.035        self._cumulative_reward = 0.036 37        return Observation(38            input=task_def["buggy_code"],39            context=task_def["context"],40            history=[],41            step_number=0,42            max_steps=self.MAX_STEPS,43        )44 45    def step(self, action: Action) -> tuple[Observation, Reward, bool, dict]:46        if self._done:47            raise RuntimeError("Episode is over. Call reset() to start a new episode.")48        if self._task is None:49            raise RuntimeError("Environment not initialised. Call reset() first.")50 51        self._step_number += 152        info = {53            "task_id":    self._task["id"],54            "difficulty": self._difficulty,55            "step":       self._step_number,56        }57 58        raw_value       = 0.059        grader_score    = None60        improvement     = None61        efficiency_bonus = 0.062        penalty         = 0.063        feedback        = ""64 65        if action.type == "give_up":66            penalty   = -0.1067            raw_value = penalty68            self._done = True69            feedback  = "Agent chose to give up."70            info["outcome"] = "gave_up"71 72        elif action.type == "explain":73            penalty   = -0.0574            raw_value = penalty75            feedback  = (76                f"Hint: re-read the loop bounds carefully. "77                f"Task goal: {self._task['context']}"78            )79            info["outcome"] = "explained"80 81        elif action.type == "fix":82            grader_score, grade_detail = self._grader(action.content)83            info["grade_detail"]  = grade_detail84            info["grader_score"]  = grader_score85 86            improvement = grader_score - self._best_score87 88            if improvement > 0:89                raw_value = improvement90                self._best_score = grader_score91            else:92                penalty   = -0.0593                raw_value = penalty94 95            if grader_score >= 1.0 and self._step_number <= 2:96                efficiency_bonus = 0.2097                raw_value += efficiency_bonus98                info["efficiency_bonus"] = True99 100            if grader_score >= 1.0:101                self._done = True102                feedback   = "All test cases passed — perfect fix!"103                info["outcome"] = "solved"104            else:105                passed = grade_detail.get("tests_passed", "?")106                total  = grade_detail.get("tests_total",  "?")107                feedback = (108                    f"Score {grader_score:.2f} — "109                    f"{passed}/{total} tests passed. Keep refining."110                )111                info["outcome"] = "partial"112 113        raw_value = round(raw_value, 4)114        self._cumulative_reward = round(self._cumulative_reward + raw_value, 4)115 116        if self._step_number >= self.MAX_STEPS and not self._done:117            self._done = True118            info["outcome"] = info.get("outcome", "timeout")119 120        reward = Reward(121            value            = raw_value,122            grader_score     = grader_score,123            improvement      = round(improvement, 4) if improvement is not None else None,124            efficiency_bonus = efficiency_bonus,125            penalty          = penalty,126            reason           = feedback,127        )128 129        self._history.append({130            "step":           self._step_number,131            "action_type":    action.type,132            "action_content": (133                action.content[:120] + "..."134                if len(action.content) > 120 else action.content135            ),136            "reward":   raw_value,137            "feedback": feedback,138        })139 140        next_obs = Observation(141            input       = self._task["buggy_code"],142            context     = self._task["context"],143            history     = list(self._history),144            step_number = self._step_number,145            max_steps   = self.MAX_STEPS,146        )147        return next_obs, reward, self._done, info148 149    def state(self) -> dict:150        return {151            "task_id":           self._task["id"] if self._task else None,152            "difficulty":        self._difficulty,153            "step_number":       self._step_number,154            "max_steps":         self.MAX_STEPS,155            "done":              self._done,156            "best_score":        self._best_score,157            "cumulative_reward": self._cumulative_reward,158            "history_length":    len(self._history),159        }160