math-openenv / src /math_env /server /environment.py
emrekuruu's picture
Upload src/math_env/server/environment.py with huggingface_hub
af304a5 verified
Raw History Blame Contribute Delete
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,
)
@property
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,
)