Spaces:
Runtime error
Runtime error
File size: 7,064 Bytes
e7a9f02 | 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 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 | #!/usr/bin/env python3
"""Train and validate the short-horizon density predictor.
The simulator is the data generator. Because it provides exact ground truth,
the model can be validated honestly: training and test use *disjoint seeds*,
and the report records the model's mean absolute error alongside the analytic
mass-balance baseline. If the model does not beat the baseline it is not used
at inference time.
Run: python scripts/train_predictor.py [--quick]
"""
from __future__ import annotations
import argparse
import datetime as dt
import json
import sys
from pathlib import Path
import numpy as np
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT / "backend"))
from flowtwin.config import MODEL_DIR, SETTINGS # noqa: E402
from flowtwin.prediction.features import ( # noqa: E402
FEATURE_NAMES,
analytic_projection,
build_feature_matrix,
)
from flowtwin.prediction.model import TrainingReport, fit_models # noqa: E402
from flowtwin.simulation.engine import Simulator # noqa: E402
from flowtwin.venue import compile_venue, load_scenario # noqa: E402
TRAIN_SEEDS = [42193, 1177, 90210, 5, 771]
TEST_SEEDS = [31337, 8080]
SCENARIOS = ["circuit_alpha_post_race", "barcelona_2022_egress", "circuit_alpha_arrival"]
#: Only sample every Nth step; consecutive steps are near-duplicates.
STEP_STRIDE = 4
#: Skip near-empty edges — they are trivially predictable and would dominate.
MIN_DENSITY = 0.05
def collect(scenario_id: str, seed: int, horizons: tuple[int, ...], max_steps: int
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Run one seeded scenario and return (features, targets, baseline)."""
scenario = load_scenario(scenario_id)
venue = compile_venue(scenario.venue_id)
sim = Simulator(venue, scenario, SETTINGS, seed=seed)
dt_s = sim.dt
horizon_steps = [int(round(h / dt_s)) for h in horizons]
max_h = max(horizon_steps)
feats: list[np.ndarray] = []
base: list[np.ndarray] = []
density_track: list[np.ndarray] = []
sample_at: list[int] = []
steps = min(max_steps, int(scenario.duration_s / dt_s))
for k in range(steps):
sim.step()
density_track.append(sim.state.edge_density.copy())
if k % STEP_STRIDE == 0:
feats.append(build_feature_matrix(sim))
base.append(analytic_projection(sim, horizons))
sample_at.append(k)
if sim.is_complete and k > max_h:
break
if not sample_at:
return (np.empty((0, len(FEATURE_NAMES)), np.float32),
np.empty((len(horizons), 0), np.float32),
np.empty((len(horizons), 0), np.float32))
n_track = len(density_track)
X_parts, Y_parts, B_parts = [], [], []
for j, k in enumerate(sample_at):
if k + max_h >= n_track:
break
now = density_track[k]
keep = now >= MIN_DENSITY
if not np.any(keep):
continue
X_parts.append(feats[j][keep])
Y_parts.append(np.stack([density_track[k + hs][keep] for hs in horizon_steps]))
B_parts.append(base[j][:, keep])
if not X_parts:
return (np.empty((0, len(FEATURE_NAMES)), np.float32),
np.empty((len(horizons), 0), np.float32),
np.empty((len(horizons), 0), np.float32))
return (np.concatenate(X_parts, axis=0),
np.concatenate(Y_parts, axis=1),
np.concatenate(B_parts, axis=1))
def gather(seeds: list[int], horizons: tuple[int, ...], max_steps: int, label: str):
Xs, Ys, Bs = [], [], []
for scenario_id in SCENARIOS:
for seed in seeds:
X, Y, B = collect(scenario_id, seed, horizons, max_steps)
if X.shape[0] == 0:
continue
Xs.append(X)
Ys.append(Y)
Bs.append(B)
print(f" [{label}] {scenario_id:<26} seed={seed:<8} rows={X.shape[0]:,}")
return (np.concatenate(Xs, axis=0),
np.concatenate(Ys, axis=1),
np.concatenate(Bs, axis=1))
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--quick", action="store_true",
help="fewer seeds and shorter runs, for a fast check")
args = ap.parse_args()
horizons = tuple(SETTINGS.prediction.horizons_s)
train_seeds = TRAIN_SEEDS[:2] if args.quick else TRAIN_SEEDS
test_seeds = TEST_SEEDS[:1] if args.quick else TEST_SEEDS
max_steps = 900 if args.quick else 2600
print(f"Horizons: {horizons}s train seeds {train_seeds} test seeds {test_seeds}")
print("Generating training data ...")
Xtr, Ytr, _ = gather(train_seeds, horizons, max_steps, "train")
print("Generating held-out data ...")
Xte, Yte, Bte = gather(test_seeds, horizons, max_steps, "test")
print(f"train rows {Xtr.shape[0]:,} test rows {Xte.shape[0]:,}")
predictor, model_name = fit_models(Xtr, Ytr, horizons, seed=7)
pred = predictor.predict(Xte)
mae_model, mae_base, rmse_model, r2_model, improvement = {}, {}, {}, {}, {}
for k, h in enumerate(horizons):
err_m = np.abs(pred[k] - Yte[k])
err_b = np.abs(Bte[k] - Yte[k])
mae_model[str(h)] = float(err_m.mean())
mae_base[str(h)] = float(err_b.mean())
rmse_model[str(h)] = float(np.sqrt(((pred[k] - Yte[k]) ** 2).mean()))
ss_res = float(((pred[k] - Yte[k]) ** 2).sum())
ss_tot = float(((Yte[k] - Yte[k].mean()) ** 2).sum())
r2_model[str(h)] = 1.0 - ss_res / max(ss_tot, 1e-9)
improvement[str(h)] = 100.0 * (mae_base[str(h)] - mae_model[str(h)]) / max(mae_base[str(h)], 1e-9)
print("\nhorizon MAE model MAE baseline improvement R²")
for h in horizons:
k = str(h)
print(f" +{h:>3}s {mae_model[k]:.4f} {mae_base[k]:.4f}"
f" {improvement[k]:+6.1f}% {r2_model[k]:.3f}")
if all(v <= 0 for v in improvement.values()):
print("\nModel did not beat the analytic baseline. Not saving; inference "
"will keep using the mass-balance projection.")
return
MODEL_DIR.mkdir(parents=True, exist_ok=True)
predictor.save(SETTINGS.prediction.model_path)
report = TrainingReport(
horizons_s=list(horizons),
n_train=int(Xtr.shape[0]),
n_test=int(Xte.shape[0]),
scenarios=SCENARIOS,
train_seeds=train_seeds,
test_seeds=test_seeds,
model_name=model_name,
mae_model=mae_model,
mae_baseline=mae_base,
rmse_model=rmse_model,
r2_model=r2_model,
improvement_pct=improvement,
feature_names=list(FEATURE_NAMES),
created_utc=dt.datetime.now(dt.timezone.utc).isoformat(timespec="seconds"),
)
SETTINGS.prediction.metrics_path.write_text(report.to_json(), encoding="utf-8")
print(f"\nSaved model -> {SETTINGS.prediction.model_path}")
print(f"Saved report -> {SETTINGS.prediction.metrics_path}")
if __name__ == "__main__":
main()
|