Spaces:
Paused
Paused
Download src/math_env/server/environment.py from emrekuruu/math-openenv: direct link, hf CLI and curl.
- Browser
- Download file 5.87 kB
-
https://huggingface.co/spaces/emrekuruu/math-openenv/resolve/main/src/math_env/server/environment.py
- Command line
-
hf download hf://spaces/emrekuruu/math-openenv/src/math_env/server/environment.py
-
curl -L -o environment.py https://huggingface.co/spaces/emrekuruu/math-openenv/resolve/main/src/math_env/server/environment.py
5.87 kB
| """Genuine OpenEnv adapter around one MATH core instance.""" | |
| from __future__ import annotations | |
| from collections.abc import Callable | |
| from typing import Any | |
| from uuid import uuid4 | |
| from openenv.core.env_server.interfaces import Environment | |
| from openenv.core.env_server.types import EnvironmentMetadata | |
| from math_env.catalog import get_split | |
| from math_env.core import CoreObservation, MathCore | |
| from math_env.interface import environment_description | |
| from math_env.models import ( | |
| MathAction, | |
| MathObservation, | |
| MathState, | |
| ) | |
| _CoreFactory = Callable[[], MathCore] | |
| class MathEnvironment( | |
| Environment[MathAction, MathObservation, MathState] | |
| ): | |
| """One connection-owned, typed MATH episode.""" | |
| SUPPORTS_CONCURRENT_SESSIONS = True | |
| REQUIRES_SINGLE_THREAD_EXECUTOR = False | |
| def __init__( | |
| self, | |
| core_factory: _CoreFactory = MathCore, | |
| *, | |
| default_task_id: int | None = None, | |
| ) -> None: | |
| super().__init__() | |
| self._core_factory = core_factory | |
| self._core: MathCore | None = None | |
| self._default_task_id = default_task_id | |
| self._episode_id: str | None = None | |
| self._closed = False | |
| def _get_core(self) -> MathCore: | |
| if self._closed: | |
| raise RuntimeError("MATH environment is closed") | |
| if self._core is None: | |
| self._core = self._core_factory() | |
| return self._core | |
| def get_metadata(self) -> EnvironmentMetadata: | |
| return EnvironmentMetadata( | |
| name="MATH", | |
| description=environment_description(), | |
| version="0.1.0", | |
| documentation_url="https://huggingface.co/datasets/HuggingFaceH4/MATH", | |
| ) | |
| def list_splits(self) -> list[dict[str, Any]]: | |
| return [ | |
| { | |
| "name": name, | |
| "role": get_split(name).role, | |
| "num_tasks": len(get_split(name).tasks), | |
| "source_revision": get_split(name).source_revision, | |
| "sha256": get_split(name).sha256, | |
| } | |
| for name in ("development", "test") | |
| ] | |
| def list_tasks(self, split: str) -> list[dict[str, Any]]: | |
| task_split = get_split(split) | |
| return [ | |
| { | |
| "id": task.task_id, | |
| "source_id": task.source_id, | |
| "instruction": task.instruction, | |
| "task_type": task.source_class, | |
| "source_class": task.source_class, | |
| "level": task.level, | |
| "split": task_split.role, | |
| "source_revision": task_split.source_revision, | |
| "catalog_sha256": task_split.sha256, | |
| } | |
| for task in task_split.tasks | |
| ] | |
| def num_tasks(self, split: str) -> int: | |
| return len(get_split(split).tasks) | |
| def get_task(self, split: str, index: int) -> dict[str, Any]: | |
| tasks = self.list_tasks(split) | |
| try: | |
| return tasks[index] | |
| except IndexError as error: | |
| raise IndexError( | |
| f"MATH {split} task index out of range: {index}" | |
| ) from error | |
| def get_task_range( | |
| self, | |
| split: str, | |
| start: int | None = None, | |
| stop: int | None = None, | |
| ) -> list[dict[str, Any]]: | |
| return self.list_tasks(split)[slice(start, stop)] | |
| def reset( | |
| self, | |
| seed: int | None = None, | |
| episode_id: str | None = None, | |
| *, | |
| task_id: int | None = None, | |
| **kwargs: Any, | |
| ) -> MathObservation: | |
| del seed, kwargs | |
| resolved = self._default_task_id if task_id is None else int(task_id) | |
| if resolved is None: | |
| raise ValueError("MATH reset requires task_id") | |
| core_observation = self._get_core().reset(resolved) | |
| self._episode_id = episode_id or str(uuid4()) | |
| return self._observation(core_observation, reward=0.0, terminated=False) | |
| def step( | |
| self, | |
| action: MathAction, | |
| timeout_s: float | None = None, | |
| **kwargs: Any, | |
| ) -> MathObservation: | |
| del timeout_s, kwargs | |
| transition = self._get_core().step(action.answer) | |
| return self._observation( | |
| transition.observation, | |
| reward=transition.reward, | |
| terminated=transition.terminated, | |
| ) | |
| def state(self) -> MathState: | |
| if self._core is None: | |
| return MathState( | |
| episode_id=self._episode_id, | |
| step_count=0, | |
| closed=self._closed, | |
| ) | |
| state = self._core.state() | |
| return MathState( | |
| episode_id=self._episode_id, | |
| step_count=state.step_count, | |
| task_id=state.task_id, | |
| goal=state.goal, | |
| source_class=state.source_class, | |
| level=state.level, | |
| reward=state.reward, | |
| success=state.success, | |
| progress=state.progress, | |
| terminated=state.terminated, | |
| truncated=False, | |
| step_limit=state.step_limit, | |
| diagnostics=state.diagnostics, | |
| closed=state.closed, | |
| ) | |
| def close(self) -> None: | |
| if self._closed: | |
| return | |
| if self._core is not None: | |
| self._core.close() | |
| self._closed = True | |
| def _observation( | |
| self, | |
| observation: CoreObservation, | |
| *, | |
| reward: float, | |
| terminated: bool, | |
| ) -> MathObservation: | |
| return MathObservation( | |
| task_id=observation.task_id, | |
| text=observation.text, | |
| goal=observation.goal, | |
| success=observation.success, | |
| progress=observation.progress, | |
| terminated=terminated, | |
| truncated=False, | |
| done=terminated, | |
| reward=reward, | |
| step_limit=observation.step_limit, | |
| ) |