Harsh-Gupta-07 commited on
Commit
bb543f9
·
verified ·
1 Parent(s): 86e83a8

Upload folder using huggingface_hub

Browse files
client.py CHANGED
@@ -54,14 +54,26 @@ class SmartEmergencyEnv(
54
  payload["reroute"] = {
55
  "vehicle_to_reroute": action.reroute.vehicle_to_reroute,
56
  "from_event_id": action.reroute.from_event_id,
57
- "to_new_event": action.reroute.to_new_event,
58
  "replacement_vehicle_id": action.reroute.replacement_vehicle_id,
59
  }
60
  return payload
61
 
62
  def _parse_result(self, payload: Dict) -> StepResult[SmartEmergencyObservation]:
63
- """Parse server response into StepResult."""
 
 
 
 
 
 
64
  obs_data = payload.get("observation", {})
 
 
 
 
 
 
 
65
  observation = SmartEmergencyObservation(
66
  prompt=obs_data.get("prompt", ""),
67
  step=obs_data.get("step", 0),
@@ -71,7 +83,8 @@ class SmartEmergencyEnv(
71
  fleet_utilisation=obs_data.get("fleet_utilisation", 0.0),
72
  done=payload.get("done", False),
73
  reward=payload.get("reward"),
74
- metadata=obs_data.get("metadata", {}),
 
75
  )
76
  return StepResult(
77
  observation=observation,
 
54
  payload["reroute"] = {
55
  "vehicle_to_reroute": action.reroute.vehicle_to_reroute,
56
  "from_event_id": action.reroute.from_event_id,
 
57
  "replacement_vehicle_id": action.reroute.replacement_vehicle_id,
58
  }
59
  return payload
60
 
61
  def _parse_result(self, payload: Dict) -> StepResult[SmartEmergencyObservation]:
62
+ """Parse server response into StepResult.
63
+
64
+ Note: OpenEnv's serialize_observation() intentionally strips 'metadata',
65
+ 'done', and 'reward' from the nested observation dict and promotes them
66
+ to the top level. ground_truth is now a first-class field on the
67
+ observation model so it survives serialization.
68
+ """
69
  obs_data = payload.get("observation", {})
70
+ # metadata is stripped by the framework; ground_truth is now a dedicated field
71
+ metadata = payload.get("metadata", obs_data.get("metadata", {}))
72
+ # Support both the new dedicated ground_truth field and the legacy metadata path
73
+ gt = obs_data.get("ground_truth") or metadata.get("ground_truth", {})
74
+ if gt:
75
+ metadata = dict(metadata)
76
+ metadata["ground_truth"] = gt
77
  observation = SmartEmergencyObservation(
78
  prompt=obs_data.get("prompt", ""),
79
  step=obs_data.get("step", 0),
 
83
  fleet_utilisation=obs_data.get("fleet_utilisation", 0.0),
84
  done=payload.get("done", False),
85
  reward=payload.get("reward"),
86
+ ground_truth=gt or {},
87
+ metadata=metadata,
88
  )
89
  return StepResult(
90
  observation=observation,
models.py CHANGED
@@ -86,3 +86,7 @@ class SmartEmergencyObservation(Observation):
86
  fleet_utilisation: float = Field(
87
  default=0.0, description="Fraction of fleet currently busy"
88
  )
 
 
 
 
 
86
  fleet_utilisation: float = Field(
87
  default=0.0, description="Fraction of fleet currently busy"
88
  )
89
+ ground_truth: Dict = Field(
90
+ default_factory=dict,
91
+ description="Hidden ground truth for the current call (populated after step)",
92
+ )
server/smart_emergency_environment.py CHANGED
@@ -154,6 +154,12 @@ class SmartEmergencyEnvironment(Environment):
154
  )
155
  obs_text = self._build_observation() if not done else "Episode complete."
156
 
 
 
 
 
 
 
157
  return SmartEmergencyObservation(
158
  prompt=obs_text,
159
  step=self._state.step_count,
@@ -163,13 +169,9 @@ class SmartEmergencyEnvironment(Environment):
163
  fleet_utilisation=self._fleet_util(),
164
  done=done,
165
  reward=breakdown.get("total", 0.0),
 
166
  metadata={
167
- "ground_truth": {
168
- "severity": call.severity,
169
- "emergency_type": call.emergency_type,
170
- "is_duplicate": call.is_duplicate_of is not None,
171
- "required_vehicle_type": call.required_vehicle_type,
172
- },
173
  "city_seed": self._seed,
174
  },
175
  )
 
154
  )
155
  obs_text = self._build_observation() if not done else "Episode complete."
156
 
157
+ gt = {
158
+ "severity": call.severity,
159
+ "emergency_type": call.emergency_type,
160
+ "is_duplicate": call.is_duplicate_of is not None,
161
+ "required_vehicle_type": call.required_vehicle_type,
162
+ }
163
  return SmartEmergencyObservation(
164
  prompt=obs_text,
165
  step=self._state.step_count,
 
169
  fleet_utilisation=self._fleet_util(),
170
  done=done,
171
  reward=breakdown.get("total", 0.0),
172
+ ground_truth=gt,
173
  metadata={
174
+ "ground_truth": gt,
 
 
 
 
 
175
  "city_seed": self._seed,
176
  },
177
  )
train_sft_grpo.py CHANGED
@@ -298,7 +298,9 @@ def generate_sft_data(env, num_episodes=60):
298
  )
299
 
300
  result = env.step(action)
301
- gt = result.observation.metadata.get("ground_truth")
 
 
302
  if gt:
303
  ideal = build_ideal_action(gt, prev_obs)
304
  examples.append({
 
298
  )
299
 
300
  result = env.step(action)
301
+ # ground_truth is now a first-class field on the observation;
302
+ # fall back to metadata for backward compatibility with older servers.
303
+ gt = result.observation.ground_truth or result.observation.metadata.get("ground_truth")
304
  if gt:
305
  ideal = build_ideal_action(gt, prev_obs)
306
  examples.append({