Spaces:
Sleeping
Sleeping
| from typing import Iterable, Optional | |
| from baseline import baseline_policy | |
| from env import IndianTrafficEnv | |
| from models import GraderOutput, TrafficAction | |
| MIN_SCORE = 0.001 | |
| MAX_SCORE = 0.999 | |
| def _clamp(value: float) -> float: | |
| """Clamp grader outputs into the validator-safe open interval (0, 1).""" | |
| return max(MIN_SCORE, min(MAX_SCORE, float(value))) | |
| def grade_rollout( | |
| task_id: str = "single_intersection", | |
| seed: int = 42, | |
| actions: Optional[Iterable[TrafficAction]] = None, | |
| max_steps: Optional[int] = None, | |
| ) -> GraderOutput: | |
| env = IndianTrafficEnv(task_id=task_id) | |
| env.reset(seed=seed, task_id=task_id) | |
| limit = max_steps or int(env.task.constraints["max_steps"]) | |
| provided_actions = list(actions) if actions is not None else None | |
| total_reward = 0.0 | |
| for step_idx in range(limit): | |
| if provided_actions is None: | |
| action = baseline_policy(env.get_state()) | |
| elif step_idx < len(provided_actions): | |
| action = provided_actions[step_idx] | |
| else: | |
| action = TrafficAction.ALL_RED | |
| _, reward, done, _ = env.step(action) | |
| total_reward += reward | |
| if done: | |
| break | |
| metrics = env.metrics | |
| average_waiting_time = ( | |
| metrics.total_wait_observations / metrics.wait_samples if metrics.wait_samples else 0.0 | |
| ) | |
| emergency_efficiency = ( | |
| metrics.emergency_cleared_fast / metrics.emergency_seen if metrics.emergency_seen else 1.0 | |
| ) | |
| wait_score = 1.0 - min(1.0, average_waiting_time / 950.0) | |
| queue_score = 1.0 - min(1.0, metrics.max_queue_length / float(env.task.constraints["max_queue_before_failure"])) | |
| clearance_score = min(1.0, metrics.total_vehicles_cleared / float(env.task.termination["target_cleared"])) | |
| safety_score = 1.0 - min(1.0, metrics.unsafe_switches / 18.0) | |
| score = ( | |
| 0.30 * wait_score | |
| + 0.22 * queue_score | |
| + 0.25 * clearance_score | |
| + 0.18 * emergency_efficiency | |
| + 0.05 * safety_score | |
| ) | |
| return GraderOutput( | |
| score=round(_clamp(score), 4), | |
| average_waiting_time=round(average_waiting_time, 3), | |
| max_queue_length=metrics.max_queue_length, | |
| total_vehicles_cleared=metrics.total_vehicles_cleared, | |
| emergency_handling_efficiency=round(emergency_efficiency, 4), | |
| details={ | |
| "task_id": task_id, | |
| "seed": seed, | |
| "total_reward": round(total_reward, 4), | |
| "unsafe_switches": metrics.unsafe_switches, | |
| "emergency_seen": metrics.emergency_seen, | |
| "emergency_cleared_fast": metrics.emergency_cleared_fast, | |
| "full_clearances": metrics.full_clearances, | |
| }, | |
| ) | |