Spaces:
Running on Zero
Running on Zero
Download simulation/runner.py from BJyotibrat/BanTrel-Inference: direct link, hf CLI and curl.
- Browser
- Download file 5.74 kB
-
https://huggingface.co/spaces/BJyotibrat/BanTrel-Inference/resolve/main/simulation/runner.py
- Command line
-
hf download hf://spaces/BJyotibrat/BanTrel-Inference/simulation/runner.py
-
curl -L -o runner.py https://huggingface.co/spaces/BJyotibrat/BanTrel-Inference/resolve/main/simulation/runner.py
5.74 kB
| """ | |
| Orchestrates one "simulate" job: builds demand -> runs the requested | |
| policy/policies -> aggregates metrics -> (optionally) produces a | |
| PPO-vs-Fixed-Cycle verdict table, using the exact metric definitions and | |
| tie-thresholds from BanTRel.py Cell 17. | |
| """ | |
| import shutil | |
| import tempfile | |
| from typing import Literal | |
| import numpy as np | |
| from . import demand | |
| from .baseline import FixedCyclePolicy | |
| from .env import BangaloreSumoEnv | |
| from .model import PPOPolicy | |
| ALL_METRICS_DEF = [ | |
| # (key, label, higher_is_better, unit) | |
| ("system_total_stopped", "Total Stopped Vehicles", False, "vehicles"), | |
| ("system_total_waiting_time", "Total Waiting Time", False, "seconds"), | |
| ("system_mean_waiting_time", "Mean Waiting Time", False, "s / vehicle"), | |
| ("system_mean_speed", "Mean Speed", True, "m/s"), | |
| ("avg_waiting_time", "Avg Waiting Time", False, "s / vehicle"), | |
| ("avg_travel_time", "Avg Travel Time", False, "seconds"), | |
| ("queue_length", "Queue Length", False, "vehicles"), | |
| ("throughput", "Throughput", True, "veh / step"), | |
| ("delay", "Delay vs Free-Flow", False, "seconds"), | |
| ] | |
| TIE_THRESHOLD = { | |
| "system_total_stopped": 0.05, "system_total_waiting_time": 1.0, | |
| "system_mean_waiting_time": 0.5, "system_mean_speed": 0.005, | |
| "avg_waiting_time": 0.5, "avg_travel_time": 1.0, | |
| "queue_length": 0.05, "throughput": 0.01, "delay": 0.5, | |
| } | |
| def _run_policy_episodes(policy, cfg_path: str, sumo_config_dir: str, n_runs: int) -> dict: | |
| """Ported from Cell 13's run_evaluation().""" | |
| buffers = {k: [] for k, *_ in ALL_METRICS_DEF} | |
| for _run in range(n_runs): | |
| env = BangaloreSumoEnv(cfg_path, sumo_config_dir) | |
| state = env.reset() | |
| if hasattr(policy, "reset"): | |
| policy.reset() | |
| done = False | |
| try: | |
| while not done: | |
| action = policy.choose_action(state) | |
| state, _reward, done = env.step(action) | |
| m = env.get_metrics() | |
| for k, *_ in ALL_METRICS_DEF: | |
| if k in m and len(m[k]): | |
| buffers[k].append(m[k]) | |
| finally: | |
| env.close() | |
| averaged = {} | |
| for k, runs in buffers.items(): | |
| if not runs: | |
| continue | |
| min_len = min(len(r) for r in runs) | |
| averaged[k] = np.mean([r[:min_len] for r in runs], axis=0) | |
| return averaged | |
| def _summarize(metrics: dict) -> dict: | |
| """Scalar table: mean of each per-step array (what Cell 17 prints).""" | |
| out = {} | |
| for key, label, higher_better, unit in ALL_METRICS_DEF: | |
| arr = metrics.get(key) | |
| if arr is None or not len(arr): | |
| continue | |
| out[key] = { | |
| "label": label, "unit": unit, "higher_is_better": higher_better, | |
| "value": round(float(np.mean(arr)), 3), | |
| } | |
| return out | |
| def _verdict(ppo_summary: dict, fc_summary: dict) -> dict: | |
| wins_ppo = wins_fixed = ties = 0 | |
| rows = [] | |
| for key, label, higher_better, unit in ALL_METRICS_DEF: | |
| if key not in ppo_summary or key not in fc_summary: | |
| continue | |
| pm, fcm = ppo_summary[key]["value"], fc_summary[key]["value"] | |
| delta = pm - fcm | |
| threshold = TIE_THRESHOLD.get(key, 0.01) | |
| if abs(delta) < threshold: | |
| outcome = "tie"; ties += 1 | |
| elif (delta > 0) if higher_better else (delta < 0): | |
| outcome = "ppo"; wins_ppo += 1 | |
| else: | |
| outcome = "fixed_cycle"; wins_fixed += 1 | |
| rows.append({"key": key, "label": label, "unit": unit, | |
| "ppo": pm, "fixed_cycle": fcm, "delta": round(delta, 3), | |
| "outcome": outcome}) | |
| overall = ("ppo" if wins_ppo > wins_fixed | |
| else "fixed_cycle" if wins_fixed > wins_ppo else "draw") | |
| return {"rows": rows, "wins_ppo": wins_ppo, "wins_fixed_cycle": wins_fixed, | |
| "ties": ties, "overall": overall} | |
| def run_simulation( | |
| *, | |
| net_path: str, | |
| sumo_config_dir: str, | |
| checkpoint_path: str, | |
| policy: Literal["ppo", "fixed_cycle", "both"], | |
| demand_mode: Literal["preset", "custom"], | |
| total_vehicles: int | None, | |
| period: str | None, | |
| mix: dict[str, float] | None, | |
| n_runs: int = 1, | |
| deterministic: bool = True, | |
| seed: int | None = None, | |
| ) -> dict: | |
| run_dir = tempfile.mkdtemp(prefix="bantrel_job_") | |
| try: | |
| cfg_path = demand.prepare_run_config( | |
| run_dir=run_dir, net_path=net_path, mode=demand_mode, | |
| total_vehicles=total_vehicles, period=period, mix=mix, | |
| ) | |
| result = {"policy_requested": policy, "n_runs": n_runs} | |
| if policy in ("ppo", "both"): | |
| ppo_policy = PPOPolicy(checkpoint_path, deterministic=deterministic, seed=seed) | |
| ppo_metrics = _run_policy_episodes(ppo_policy, cfg_path, sumo_config_dir, n_runs) | |
| result["ppo"] = { | |
| "summary": _summarize(ppo_metrics), | |
| "timeseries": {k: v.tolist() for k, v in ppo_metrics.items()}, | |
| } | |
| if policy in ("fixed_cycle", "both"): | |
| fc_policy = FixedCyclePolicy(cycle_steps=3) | |
| fc_metrics = _run_policy_episodes(fc_policy, cfg_path, sumo_config_dir, n_runs) | |
| result["fixed_cycle"] = { | |
| "summary": _summarize(fc_metrics), | |
| "timeseries": {k: v.tolist() for k, v in fc_metrics.items()}, | |
| } | |
| if policy == "both": | |
| result["comparison"] = _verdict( | |
| result["ppo"]["summary"], result["fixed_cycle"]["summary"] | |
| ) | |
| return result | |
| finally: | |
| shutil.rmtree(run_dir, ignore_errors=True) |