RohitChandramouli6618 commited on
Commit
9b98195
Β·
1 Parent(s): 14e1e76

Add /grade endpoint, use real grader scores in evaluator

Browse files
Files changed (3) hide show
  1. baseline/evaluator.py +36 -9
  2. server/app.py +15 -1
  3. server/environment.py +17 -0
baseline/evaluator.py CHANGED
@@ -20,7 +20,9 @@ from core.trajectory import EpisodicMemory
20
  from core.reward import normalise_score
21
  from core.policy_update import compute_advantage, update_memory
22
 
23
- N_ROLLOUTS = 3
 
 
24
 
25
 
26
  # ── Prompt Builder With Memory ────────────────────────────────────────────────
@@ -46,11 +48,11 @@ def build_prompt_with_memory(obs: CityObservation, memory: EpisodicMemory) -> st
46
 
47
  def run_rollout(
48
  env: Any,
49
- task_name: str,
50
  client: OpenAI,
51
  memory: EpisodicMemory,
52
  verbose: bool = True,
53
- ) -> Tuple[float, int, List[dict]]:
54
  """Run one complete episode using memory-augmented prompts."""
55
  result = env.reset(task_name=task_name)
56
  obs = result.observation
@@ -64,7 +66,15 @@ def run_rollout(
64
  response = call_llm(prompt, client)
65
  action = parse_action(response, len(obs.districts))
66
 
67
- result = env.step(action)
 
 
 
 
 
 
 
 
68
  next_obs = result.observation
69
  reward = result.reward or 0.0
70
  done = result.done
@@ -97,6 +107,7 @@ def run_task_grpo(
97
  env: Any,
98
  task_name: str,
99
  client: OpenAI,
 
100
  verbose: bool = True,
101
  ) -> float:
102
  """GRPO-style simulated learning loop for one task."""
@@ -113,11 +124,27 @@ def run_task_grpo(
113
  print(f"\n Rollout {i+1}/{N_ROLLOUTS} [{label}]")
114
 
115
  total_reward, steps, trajectory = run_rollout(env, task_name, client, memory, verbose)
116
- score = normalise_score(total_reward, steps)
117
- rollouts.append((total_reward, steps, score))
118
 
119
- if verbose:
120
- print(f" β†’ Reward: {total_reward:+.4f} | Score: {score:.4f}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
121
 
122
  # GRPO advantage computation
123
  completed_rewards = [r[0] for r in rollouts]
@@ -164,7 +191,7 @@ def run_evaluation(
164
  with CascadeContainmentEnv(base_url=base_url).sync() as env:
165
  for task_name in ["easy", "medium", "hard"]:
166
  try:
167
- score = run_task_grpo(env, task_name, client, verbose)
168
  scores[task_name] = score
169
  if verbose:
170
  print(f"\n βœ“ {task_name.upper()} final score: {score:.4f}")
 
20
  from core.reward import normalise_score
21
  from core.policy_update import compute_advantage, update_memory
22
 
23
+ import requests as http_requests
24
+
25
+ N_ROLLOUTS = 4
26
 
27
 
28
  # ── Prompt Builder With Memory ────────────────────────────────────────────────
 
48
 
49
  def run_rollout(
50
  env: Any,
51
+ task_name: str,
52
  client: OpenAI,
53
  memory: EpisodicMemory,
54
  verbose: bool = True,
55
+ ) -> Tuple[float, int, List[dict]]:
56
  """Run one complete episode using memory-augmented prompts."""
57
  result = env.reset(task_name=task_name)
58
  obs = result.observation
 
66
  response = call_llm(prompt, client)
67
  action = parse_action(response, len(obs.districts))
68
 
69
+ try:
70
+ result = env.step(action)
71
+ except Exception as e:
72
+ if "close frame" in str(e).lower() or "websocket" in str(e).lower():
73
+ if verbose:
74
+ print(f" ⚠ WebSocket dropped at step {step+1}, ending rollout early")
75
+ break
76
+ raise
77
+
78
  next_obs = result.observation
79
  reward = result.reward or 0.0
80
  done = result.done
 
107
  env: Any,
108
  task_name: str,
109
  client: OpenAI,
110
+ base_url: str,
111
  verbose: bool = True,
112
  ) -> float:
113
  """GRPO-style simulated learning loop for one task."""
 
124
  print(f"\n Rollout {i+1}/{N_ROLLOUTS} [{label}]")
125
 
126
  total_reward, steps, trajectory = run_rollout(env, task_name, client, memory, verbose)
 
 
127
 
128
+ # Use real grader score from server
129
+ num_districts = {"easy": 2, "medium": 4, "hard": 6}.get(task_name, 2)
130
+ try:
131
+ grade_resp = http_requests.get(
132
+ base_url.rstrip('/') + '/grade', timeout=10
133
+ )
134
+ if grade_resp.status_code == 200:
135
+ data = grade_resp.json()
136
+ score = data["final_score"]
137
+ if verbose:
138
+ print(
139
+ f" β†’ Grader: containment={data['containment_score']:.3f} "
140
+ f"hospital={data['hospital_score']:.3f} "
141
+ f"efficiency={data['efficiency_score']:.3f} "
142
+ f"speed={data['speed_score']:.3f}"
143
+ )
144
+ else:
145
+ score = normalise_score(total_reward, steps, num_districts)
146
+ except Exception:
147
+ score = normalise_score(total_reward, steps, num_districts)
148
 
149
  # GRPO advantage computation
150
  completed_rewards = [r[0] for r in rollouts]
 
191
  with CascadeContainmentEnv(base_url=base_url).sync() as env:
192
  for task_name in ["easy", "medium", "hard"]:
193
  try:
194
+ score = run_task_grpo(env, task_name, client, base_url, verbose)
195
  scores[task_name] = score
196
  if verbose:
197
  print(f"\n βœ“ {task_name.upper()} final score: {score:.4f}")
server/app.py CHANGED
@@ -10,9 +10,10 @@ import os
10
  sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))
11
 
12
  from openenv.core.env_server import create_app
 
13
  from server.environment import EpidemicContainmentEnv
14
  from models import ContainmentAction, CityObservation
15
-
16
 
17
  app = create_app(
18
  EpidemicContainmentEnv,
@@ -20,6 +21,19 @@ app = create_app(
20
  CityObservation,
21
  )
22
 
 
 
 
 
 
 
 
 
 
 
 
 
 
23
  from fastapi.responses import HTMLResponse
24
 
25
  @app.get("/", response_class=HTMLResponse)
 
10
  sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))
11
 
12
  from openenv.core.env_server import create_app
13
+ from fastapi.responses import JSONResponse
14
  from server.environment import EpidemicContainmentEnv
15
  from models import ContainmentAction, CityObservation
16
+ import server.environment as env_module
17
 
18
  app = create_app(
19
  EpidemicContainmentEnv,
 
21
  CityObservation,
22
  )
23
 
24
+ @app.get("/grade")
25
+ async def grade_last_episode():
26
+ """
27
+ Returns the deterministic grader score for the most recently completed episode.
28
+ Computed automatically when an episode ends via the WebSocket session.
29
+ """
30
+ if not env_module._last_grade:
31
+ return JSONResponse(
32
+ {"error": "No completed episode yet β€” run a full episode first"},
33
+ status_code=400
34
+ )
35
+ return JSONResponse(env_module._last_grade)
36
+
37
  from fastapi.responses import HTMLResponse
38
 
39
  @app.get("/", response_class=HTMLResponse)
server/environment.py CHANGED
@@ -52,6 +52,8 @@ from server.utils import (
52
  )
53
  from server.tasks.registry import get_task
54
 
 
 
55
 
56
  class EpidemicContainmentEnv(Environment):
57
  """
@@ -159,6 +161,21 @@ class EpidemicContainmentEnv(Environment):
159
  # ── 10. Check terminal conditions ─────────────────────────────────────
160
  done, terminal_message = self._check_terminal()
161
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
162
  # ── 11. Build and return observation ──────────────────────────────────
163
  final_message = terminal_message if terminal_message else message
164
 
 
52
  )
53
  from server.tasks.registry import get_task
54
 
55
+ from server.grader import grade_trajectory
56
+ _last_grade: dict = {}
57
 
58
  class EpidemicContainmentEnv(Environment):
59
  """
 
161
  # ── 10. Check terminal conditions ─────────────────────────────────────
162
  done, terminal_message = self._check_terminal()
163
 
164
+ if done and self._trajectory:
165
+ import server.environment as _self_module
166
+ result = grade_trajectory(self._trajectory, self._task_name)
167
+ _self_module._last_grade = {
168
+ "final_score": result.final_score,
169
+ "containment_score": result.containment_score,
170
+ "hospital_score": result.hospital_score,
171
+ "efficiency_score": result.efficiency_score,
172
+ "speed_score": result.speed_score,
173
+ "hospital_breached": result.hospital_breached,
174
+ "districts_contained": result.districts_contained,
175
+ "total_steps": result.total_steps,
176
+ "task_name": self._task_name,
177
+ }
178
+
179
  # ── 11. Build and return observation ──────────────────────────────────
180
  final_message = terminal_message if terminal_message else message
181