#!/usr/bin/env python3 """calib_clean_ratio_ev.py — toy EV + optional rollouts (max belief, not mean).""" from __future__ import annotations import argparse import sys from pathlib import Path def toy_ev(q, alert=100.0, false=20.0, miss=200.0): always = alert * q - false * (1.0 - q) never = -miss * q oracle = alert * q gap = oracle - always rel = gap / abs(always) if always != 0 else float("inf") return dict(q=q, always=always, never=never, oracle=oracle, gap=gap, relative_gap=rel) def analytic_q_independent(clean_ratio, n_zones): """Legacy i.i.d. per-zone: P(any dirty) = 1 - c^n.""" return 1.0 - (clean_ratio ** n_zones) def analytic_q_correlated(clean_ratio, n_zones, rho=0.85): """Regional event model: P(any dirty) = 1 - [(1-p) + p*(1-rho)^n] with guarantee ≥1 dirty when event fires → ≈ p when rho high. Approximate with p = 1-c (full correlation limit).""" p = 1.0 - clean_ratio # Exact without guarantee: 1 - ((1-p) + p*(1-rho)**n) q_soft = 1.0 - ((1.0 - p) + p * ((1.0 - rho) ** n_zones)) # With ≥1-dirty guarantee when event fires, q = p q_hard = p return q_soft, q_hard def main(argv=None): p = argparse.ArgumentParser() p.add_argument("--n-zones", type=int, default=3) p.add_argument("--ratios", default="0.70,0.85,0.90,0.95") p.add_argument("--rho", type=float, default=0.85, help="event_spatial_correlation") p.add_argument("--alert", type=float, default=100.0) p.add_argument("--false", type=float, default=20.0) p.add_argument("--miss", type=float, default=200.0) p.add_argument("--rollouts", type=int, default=0) p.add_argument("--seed", type=int, default=0) args = p.parse_args(argv) ratios = [float(x.strip()) for x in args.ratios.split(",") if x.strip()] n = args.n_zones rho = args.rho print("=== Toy EV vs clean_episode_ratio ===") print(f"n_zones={n} rho={rho} economics: alert={args.alert} false={args.false} miss={args.miss}") print(f"{'clean':>7s} {'q_iid':>8s} {'q_corr':>8s} {'always@corr':>12s} {'oracle':>10s} {'gap':>10s} {'rel_gap':>9s}") for c in ratios: q_iid = analytic_q_independent(c, n) q_soft, q_hard = analytic_q_correlated(c, n, rho) # Use hard (guarantee) as primary when rho is high q = q_hard if rho >= 0.5 else q_soft ev = toy_ev(q, args.alert, args.false, args.miss) print( f"{c:7.3f} {q_iid:8.4f} {q:8.4f} {ev['always']:12.2f} {ev['oracle']:10.2f} " f"{ev['gap']:10.2f} {ev['relative_gap']:8.1%}" ) print() print("q_corr ≈ 1-clean under regional events with ≥1-zone guarantee (rho high).") print("q_iid is the old independent formula — do not use for calibration when rho>0.") print("Prefer clean in {0.90, 0.95}; 0.97 is stretch after A/B look stable.") if args.rollouts <= 0: return 0 root = Path(__file__).resolve().parent sys.path.insert(0, str(root)) try: import numpy as np from zone_observation import ForecastConfig from weather_forecast_env import make_weather_env, _episode_event_plan except ImportError as e: print(f"rollouts skipped: {e}", file=sys.stderr) return 0 print() print("=== Synthetic rollouts (terminate immediately; belief = max not mean) ===") for c in ratios: cfg = ForecastConfig( n_zones=n, max_steps=n + 2, clean_episode_ratio=c, event_spatial_correlation=rho, alert_value=args.alert, false_alert_penalty=args.false, miss_penalty=args.miss, seed=args.seed, ) env = make_weather_env(cfg, use_nan_wrapper=True) base = env.unwrapped if hasattr(env, "unwrapped") else env risky = multi = 0 term_rewards, gains, costs = [], [], [] for i in range(args.rollouts): obs, info = env.reset(seed=args.seed + i) n_active = len(getattr(base, "_zone_ids", []) or list(range(n))) zone_obs = getattr(base, "_zone_obs", []) dirty = 0 for zo in zone_obs: try: if float(zo.composite_risk()) > float(cfg.prior_belief) + 0.05: dirty += 1 except Exception: pass if dirty > 0: risky += 1 if dirty >= 2: multi += 1 term = base.terminate_action _, rew, _, _, _ = env.step(int(term)) term_rewards.append(float(rew)) # Match env termination: max over active zone beliefs believed = float(np.max(obs["zone_belief"][:n_active])) if "zone_belief" in obs else float(cfg.prior_belief) gains.append(believed * (cfg.alert_value + cfg.miss_penalty) / cfg.alert_value) costs.append((1.0 - believed) * cfg.false_alert_penalty / cfg.alert_value) q_hat = risky / args.rollouts m_hat = multi / args.rollouts print( f"clean={c:.3f} empirical_any_risky={q_hat:.3f} empirical_multi_dirty={m_hat:.3f} " f"analytic_q≈{1-c:.3f} mean_term_reward={sum(term_rewards)/len(term_rewards):+.3f} " f"mean_gain/alert={sum(gains)/len(gains):+.3f} mean_cost/alert={sum(costs)/len(costs):+.3f}" ) return 0 if __name__ == "__main__": raise SystemExit(main())