Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse files- client.py +16 -3
- models.py +4 -0
- server/smart_emergency_environment.py +8 -6
- train_sft_grpo.py +3 -1
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 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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({
|