File size: 5,737 Bytes
38ec721
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3ac1bd3
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
"""
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)