emrekuruu's picture
Upload src/math_env/core.py with huggingface_hub
8611d85 verified
Raw History Blame Contribute Delete
4.28 kB
"""Framework-neutral MATH single-attempt episode lifecycle."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from math_env.catalog import MathTask, get_task
from math_env.grader import compare_answers, verification_report
from math_env.interface import build_task_prompt, extract_answer
@dataclass(frozen=True, slots=True)
class CoreObservation:
task_id: int
text: str
goal: str
success: bool
progress: float
step_limit: int = 1
@dataclass(frozen=True, slots=True)
class CoreStep:
observation: CoreObservation
reward: float
terminated: bool
@dataclass(frozen=True, slots=True)
class CoreState:
task_id: int | None
goal: str | None
source_class: str | None
level: str | None
reward: float
success: bool
progress: float
terminated: bool
step_count: int
step_limit: int
diagnostics: dict[str, str]
closed: bool
_TaskLookup = Callable[[int], MathTask]
class MathCore:
"""Own one isolated, single-attempt MATH episode.
MATH has no simulator, sandbox, or code execution: reset exposes the problem and step
grades a submitted answer against the hidden reference. This is the whole runtime.
"""
def __init__(
self,
*,
task_lookup: _TaskLookup = get_task,
) -> None:
self._task_lookup = task_lookup
self._task: MathTask | None = None
self._reward = 0.0
self._progress = 0.0
self._terminated = False
self._step_count = 0
self._diagnostics: dict[str, str] = {}
self._closed = False
def reset(self, task_id: int) -> CoreObservation:
self._ensure_open()
task = self._task_lookup(int(task_id))
self._task = task
self._reward = 0.0
self._progress = 0.0
self._terminated = False
self._step_count = 0
self._diagnostics = {}
prompt = build_task_prompt(task.instruction)
return self._observation(prompt)
def step(self, action_answer: str) -> CoreStep:
self._ensure_open()
if self._task is None:
raise RuntimeError("MATH core must be reset before step")
if self._terminated:
raise RuntimeError("MATH episode has terminated")
answer = extract_answer(action_answer)
matched = compare_answers(answer, self._task.answer)
self._step_count = 1
self._terminated = True
self._progress = 1.0 if matched else 0.0
self._reward = 1.0 if matched else 0.0
self._diagnostics = {
"verification": verification_report(answer, self._task.answer),
"fail_reason": "" if matched else "answer-mismatch",
}
summary = "Answer graded: PASS" if matched else "Answer graded: FAIL"
observation = self._observation(summary)
return CoreStep(
observation=observation,
reward=self._reward,
terminated=True,
)
def state(self) -> CoreState:
return CoreState(
task_id=None if self._task is None else self._task.task_id,
goal=None if self._task is None else self._task.instruction,
source_class=None if self._task is None else self._task.source_class,
level=None if self._task is None else self._task.level,
reward=self._reward,
success=self._terminated and self._reward == 1.0,
progress=self._progress,
terminated=self._terminated,
step_count=self._step_count,
step_limit=1,
diagnostics=dict(self._diagnostics),
closed=self._closed,
)
def close(self) -> None:
self._closed = True
def _observation(self, text: str) -> CoreObservation:
if self._task is None:
raise RuntimeError("MATH core has no current task")
return CoreObservation(
task_id=self._task.task_id,
text=text,
goal=self._task.instruction,
success=self._terminated and self._reward == 1.0,
progress=self._progress,
)
def _ensure_open(self) -> None:
if self._closed:
raise RuntimeError("MATH core is closed")