Team Ai
Apppublic

PRANAV05092003/autonomous-code-refactoring-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
openenv_interface.py139 linesDownload Raw Back to root
1from __future__ import annotations2 3import os4import sys5from typing import Any, Dict, Optional, Tuple6 7PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))8if PROJECT_ROOT not in sys.path:9    sys.path.insert(0, PROJECT_ROOT)10 11try:12    from openenv.env import Env as OpenEnvBase13except Exception:  # pragma: no cover14    class OpenEnvBase:15        def __init__(self, *args: Any, **kwargs: Any) -> None:16            return None17 18from acre.datasets.code_samples import CodeSample, CodeSampleDataset19from acre.env.refactor_env import RefactorEnv20from acre.tasks.task_registry import TaskRegistry21from models import ActionModel, ObservationModel, RewardModel, StateResponse22 23 24class OpenEnvRefactorEnv(OpenEnvBase):25    """26    Canonical OpenEnv interface for ACRE.27 28    This wrapper keeps the strict hackathon contract:29    - reset() -> ObservationModel30    - step(action) -> (ObservationModel, RewardModel, done, info)31    - state() -> StateResponse32    """33 34    def __init__(35        self,36        *,37        env: Optional[RefactorEnv] = None,38        registry: Optional[TaskRegistry] = None,39    ) -> None:40        super().__init__(41            name="ACRE",42            state_space="ObservationModel",43            action_space="ActionModel",44            episode_max_length=RefactorEnv.MAX_STEPS,45        )46        self._env = env or RefactorEnv()47        self._registry = registry or TaskRegistry()48        self._task_id: Optional[str] = None49        self._last_reset_info: Dict[str, Any] = {}50 51    @property52    def action_meanings(self) -> Dict[int, str]:53        return self._env.ACTION_MEANINGS54 55    @property56    def last_reset_info(self) -> Dict[str, Any]:57        return dict(self._last_reset_info)58 59    def _load_episode_source(self, *, task_id: Optional[str], code: Optional[str]) -> None:60        initial_code = code61        if initial_code is None and task_id:62            task = self._registry.get_task(task_id)63            if task is None:64                raise ValueError(f"Task '{task_id}' not found")65            # Load a multi-sample dataset for this task. Sample selection is66            # deterministic given the `seed` passed to `reset()`.67            samples = list(getattr(task, "samples", []) or [])68            if not samples:69                initial_code = task.initial_code70            else:71                self._env.dataset = CodeSampleDataset(72                    [73                        CodeSample(74                            id=f"{task_id}:{i}",75                            language="python",76                            code=str(src),77                        )78                        for i, src in enumerate(samples)79                    ]80                )81                return None82 83        if initial_code is None:84            return None85 86        self._env.dataset = CodeSampleDataset(87            [88                CodeSample(89                    id=task_id or "custom",90                    language="python",91                    code=initial_code,92                )93            ]94        )95        return None96 97    def reset(98        self,99        *,100        seed: Optional[int] = None,101        task_id: Optional[str] = None,102        code: Optional[str] = None,103    ) -> ObservationModel:104        self._task_id = task_id105        self._load_episode_source(task_id=task_id, code=code)106        observation, info = self._env.reset(seed=seed)107        self._last_reset_info = dict(info)108        return ObservationModel.from_vector(observation.tolist())109 110    def step(self, action: int | ActionModel) -> Tuple[ObservationModel, RewardModel, bool, Dict[str, Any]]:111        action_value = action.action if isinstance(action, ActionModel) else int(action)112        observation, raw_reward, terminated, truncated, info = self._env.step(action_value)113        reward = RewardModel(114            raw=float(raw_reward),115            normalized=float(info.get("normalized_reward", 0.0)),116            components=dict(info.get("reward_components", {})),117        )118        done = bool(terminated or truncated)119        return ObservationModel.from_vector(observation.tolist()), reward, done, dict(info)120 121    def state(self) -> StateResponse:122        raw_state = self._env.state()123        observation_vector = list(raw_state.get("observation", [0.0, 0.0, 0.0, 0.0]))124        observation = ObservationModel.from_vector(observation_vector)125        return StateResponse(126            current_code=str(raw_state.get("current_code", "")),127            episode_steps=int(raw_state.get("episode_steps", 0)),128            max_steps=int(raw_state.get("max_steps", RefactorEnv.MAX_STEPS)),129            complexity=float(raw_state.get("complexity", 0.0)),130            last_runtime=float(raw_state.get("last_runtime", 0.0)),131            last_error=bool(raw_state.get("last_error", False)),132            sample_id=raw_state.get("sample_id"),133            language=raw_state.get("language"),134            task_id=self._task_id,135            observation=observation,136            observation_vector=observation.to_vector(),137            action_meanings=dict(raw_state.get("action_meanings", {})),138        )139