garvitsachdeva commited on
Commit
1bc6b3d
·
1 Parent(s): fe39ff6

Unify OpenEnv and benchmark episode scoring

Browse files
src/benchmark.py CHANGED
@@ -7,8 +7,8 @@ import random
7
  from typing import Any
8
 
9
  from src.models import Action, DispatchAction
 
10
  from src.openenv_environment import OpenEnvEnvironment
11
- from src.rewards import TaskGrader
12
  from src.tasks.registry import TaskRegistry
13
 
14
 
@@ -70,9 +70,8 @@ async def _run_episode_async(task_id: str, seed: int) -> tuple[float, list[float
70
  metadata={},
71
  )
72
 
73
- # Score episodes the same way as the OpenEnv evaluation path:
74
- # a normalized aggregate of per-step rewards.
75
- final_score = TaskGrader().grade_episode(rewards, task_id=task_id)
76
  return final_score, rewards
77
 
78
 
 
7
  from typing import Any
8
 
9
  from src.models import Action, DispatchAction
10
+ from src.grading import grade_episode
11
  from src.openenv_environment import OpenEnvEnvironment
 
12
  from src.tasks.registry import TaskRegistry
13
 
14
 
 
70
  metadata={},
71
  )
72
 
73
+ # Score episodes the same way as the OpenEnv evaluation path.
74
+ final_score = grade_episode(task_id=task_id, state=final_state, rewards=rewards)
 
75
  return final_score, rewards
76
 
77
 
src/grading.py ADDED
@@ -0,0 +1,49 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Episode grading utilities.
2
+
3
+ This module centralizes "final score" computation so benchmark runs and
4
+ OpenEnv runs report the same episode score.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from src.models import State
10
+
11
+
12
+ def grade_episode(task_id: str, state: State | None, rewards: list[float]) -> float:
13
+ """Compute a final episode score in [0.0, 1.0].
14
+
15
+ Args:
16
+ task_id: Task identifier.
17
+ state: Final (or current) state.
18
+ rewards: Per-step rewards.
19
+
20
+ Returns:
21
+ Normalized score in [0.0, 1.0].
22
+ """
23
+
24
+ if not rewards:
25
+ return 0.0
26
+
27
+ # Lazy imports avoid circular dependencies (task graders import src.rewards).
28
+ if task_id == "single_incident":
29
+ from src.tasks.single_incident import SingleIncidentGrader
30
+
31
+ return float(SingleIncidentGrader().grade(state, rewards) if state is not None else 0.0)
32
+
33
+ if task_id == "multi_incident":
34
+ from src.tasks.multi_incident import MultiIncidentGrader
35
+
36
+ return float(MultiIncidentGrader().grade(state, rewards) if state is not None else 0.0)
37
+
38
+ if task_id == "mass_casualty":
39
+ from src.tasks.mass_casualty import MassCasualtyGrader
40
+
41
+ return float(MassCasualtyGrader().grade(state, rewards) if state is not None else 0.0)
42
+
43
+ if task_id == "shift_surge":
44
+ from src.tasks.shift_surge import ShiftSurgeGrader
45
+
46
+ return float(ShiftSurgeGrader().grade(state, rewards) if state is not None else 0.0)
47
+
48
+ # Fallback: mean of rewards (legacy behavior).
49
+ return float(sum(rewards) / max(len(rewards), 1))
src/openenv_environment.py CHANGED
@@ -3,6 +3,7 @@
3
  import uuid
4
 
5
  from src.city_schema import CitySchemaLoader
 
6
  from src.models import Action, Observation, State
7
  from src.state_machine import DispatchStateMachine
8
 
@@ -20,6 +21,8 @@ class OpenEnvEnvironment:
20
  episode_id = str(uuid.uuid4())
21
  self._state = self._machine.reset(task_id=self.task_id, episode_id=episode_id)
22
  self._state.metadata["cumulative_reward"] = 0.0
 
 
23
  self._last_observation = Observation(
24
  result="dispatch center online",
25
  score=0.0,
@@ -40,12 +43,29 @@ class OpenEnvEnvironment:
40
  raise RuntimeError("Environment not initialized. Call reset() first.")
41
  state, obs = self._machine.step(self._state, action)
42
  self._state = state
43
- self._last_observation = obs
 
 
 
 
 
 
 
 
 
 
44
  cumulative = float(self._state.metadata.get("cumulative_reward", 0.0))
45
- cumulative += float(obs.score)
46
- self._state.metadata["cumulative_reward"] = cumulative
 
 
 
 
 
47
  done = self._machine.is_terminal(state)
48
- return obs, obs.score, done
 
 
49
 
50
  def state(self) -> State:
51
  if self._state is None:
 
3
  import uuid
4
 
5
  from src.city_schema import CitySchemaLoader
6
+ from src.grading import grade_episode
7
  from src.models import Action, Observation, State
8
  from src.state_machine import DispatchStateMachine
9
 
 
21
  episode_id = str(uuid.uuid4())
22
  self._state = self._machine.reset(task_id=self.task_id, episode_id=episode_id)
23
  self._state.metadata["cumulative_reward"] = 0.0
24
+ self._state.metadata["episode_rewards"] = []
25
+ self._state.metadata["episode_score"] = 0.0
26
  self._last_observation = Observation(
27
  result="dispatch center online",
28
  score=0.0,
 
43
  raise RuntimeError("Environment not initialized. Call reset() first.")
44
  state, obs = self._machine.step(self._state, action)
45
  self._state = state
46
+
47
+ # `DispatchStateMachine.step()` sets `obs.score` to the per-step reward.
48
+ # OpenEnv consumers often interpret `observation.score` as an episode score,
49
+ # so we keep the per-step reward in `reward` and publish the episode score
50
+ # into `observation.score`.
51
+ step_reward = float(obs.score)
52
+
53
+ rewards: list[float] = list(self._state.metadata.get("episode_rewards", []))
54
+ rewards.append(step_reward)
55
+ self._state.metadata["episode_rewards"] = rewards
56
+
57
  cumulative = float(self._state.metadata.get("cumulative_reward", 0.0))
58
+ self._state.metadata["cumulative_reward"] = cumulative + step_reward
59
+
60
+ # Episode score is derived from the same grading logic as benchmark runs.
61
+ episode_score = grade_episode(task_id=self.task_id, state=self._state, rewards=rewards)
62
+ episode_score = max(0.0, min(1.0, float(episode_score)))
63
+ self._state.metadata["episode_score"] = episode_score
64
+
65
  done = self._machine.is_terminal(state)
66
+ obs = obs.model_copy(update={"score": episode_score})
67
+ self._last_observation = obs
68
+ return obs, step_reward, done
69
 
70
  def state(self) -> State:
71
  if self._state is None:
tests/test_benchmark_integration.py CHANGED
@@ -2,7 +2,11 @@
2
 
3
  from __future__ import annotations
4
 
 
 
 
5
  from src.benchmark import list_tasks, run_all, run_task
 
6
 
7
 
8
  def test_list_tasks_has_four() -> None:
@@ -16,13 +20,34 @@ def test_run_task_score_in_range() -> None:
16
  result = run_task("single_incident", seed=42)
17
  assert 0.0 <= result["score"] <= 1.0
18
  assert result["task_id"] == "single_incident"
19
- # Benchmark scoring must match the OpenEnv evaluation path: mean step reward.
20
- rewards = result["rewards"]
21
- if rewards:
22
- expected = sum(rewards) / len(rewards)
23
- else:
24
- expected = 0.0
25
- assert abs(result["score"] - expected) < 1e-9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
26
 
27
 
28
  def test_run_all_scores_in_range() -> None:
 
2
 
3
  from __future__ import annotations
4
 
5
+ import asyncio
6
+
7
+ from src.models import Action, DispatchAction
8
  from src.benchmark import list_tasks, run_all, run_task
9
+ from src.openenv_environment import OpenEnvEnvironment
10
 
11
 
12
  def test_list_tasks_has_four() -> None:
 
20
  result = run_task("single_incident", seed=42)
21
  assert 0.0 <= result["score"] <= 1.0
22
  assert result["task_id"] == "single_incident"
23
+
24
+
25
+ def test_benchmark_and_openenv_use_same_episode_grader(monkeypatch) -> None:
26
+ from src.tasks.single_incident import SingleIncidentGrader
27
+
28
+ expected_score = 0.777
29
+ monkeypatch.setattr(SingleIncidentGrader, "grade", lambda self, state, rewards: expected_score)
30
+
31
+ # Benchmark path.
32
+ result = run_task("single_incident", seed=42)
33
+ assert abs(result["score"] - expected_score) < 1e-9
34
+
35
+ # OpenEnv path.
36
+ env = OpenEnvEnvironment(task_id="single_incident", seed=42)
37
+ asyncio.run(env.reset())
38
+ obs, reward, done = asyncio.run(
39
+ env.step(
40
+ Action(
41
+ action_type=DispatchAction.DISPATCH,
42
+ unit_id="MED-1",
43
+ incident_id="INC-001",
44
+ )
45
+ )
46
+ )
47
+ assert isinstance(reward, float)
48
+ assert isinstance(done, bool)
49
+ assert abs(float(obs.score) - expected_score) < 1e-9
50
+ env.close()
51
 
52
 
53
  def test_run_all_scores_in_range() -> None: