Spaces:
Sleeping
Sleeping
| from typing import Dict | |
| from env import IndianTrafficEnv | |
| from models import BaselineOutput, LANES, TrafficAction, TrafficState | |
| def _lane_total(state: TrafficState, lane: str) -> int: | |
| return state.lane_queues[lane].total | |
| def baseline_policy(state: TrafficState) -> TrafficAction: | |
| """Simple fixed-cycle policy with a queue override and basic emergency handling.""" | |
| if state.emergency_vehicle.present: | |
| return TrafficAction.EMERGENCY_OVERRIDE | |
| if state.pedestrian_count >= 16 and state.pedestrian_wait_time > 18: | |
| return TrafficAction.PEDESTRIAN_CROSS | |
| if state.time_since_last_phase_switch < 3 and state.current_signal_phase in ( | |
| TrafficAction.NS_GREEN, | |
| TrafficAction.EW_GREEN, | |
| TrafficAction.LEFT_PRIORITY, | |
| ): | |
| return TrafficAction.EXTEND_GREEN | |
| ns_queue = _lane_total(state, "N") + _lane_total(state, "S") | |
| ew_queue = _lane_total(state, "E") + _lane_total(state, "W") | |
| if abs(ns_queue - ew_queue) >= 12: | |
| return TrafficAction.NS_GREEN if ns_queue > ew_queue else TrafficAction.EW_GREEN | |
| cycle = (state.tick // 8) % 4 | |
| return [ | |
| TrafficAction.NS_GREEN, | |
| TrafficAction.EW_GREEN, | |
| TrafficAction.LEFT_PRIORITY, | |
| TrafficAction.PEDESTRIAN_CROSS, | |
| ][cycle] | |
| def run_baseline(task_id: str = "single_intersection", seed: int = 42) -> BaselineOutput: | |
| from grader import grade_rollout | |
| env = IndianTrafficEnv(task_id=task_id) | |
| env.reset(seed=seed, task_id=task_id) | |
| total_reward = 0.0 | |
| steps_taken = 0 | |
| for _ in range(int(env.task.constraints["max_steps"])): | |
| _, reward, done, _ = env.step(baseline_policy(env.get_state())) | |
| total_reward += reward | |
| steps_taken += 1 | |
| if done: | |
| break | |
| grader = grade_rollout(task_id=task_id, seed=seed, actions=None) | |
| return BaselineOutput( | |
| task_id=task_id, | |
| seed=seed, | |
| score=grader.score, | |
| total_reward=round(total_reward, 4), | |
| steps=steps_taken, | |
| grader=grader, | |
| ) | |
| def baseline_scores(seed: int = 42) -> Dict[str, BaselineOutput]: | |
| return { | |
| task_id: run_baseline(task_id=task_id, seed=seed) | |
| for task_id in ("single_intersection", "rush_hour", "emergency_priority") | |
| } | |