monsoon-rl / calib_clean_ratio_ev.py
DHDRL's picture
Update calib_clean_ratio_ev.py
0184a3f verified
Raw
History Blame
5.4 kB
#!/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())