Reinforcement Learning
stable-baselines3
deep-reinforcement-learning
agricultural-ai
weather-modelling
curriculum-learning
edge-ai
Instructions to use DHDRL/monsoon-rl with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- stable-baselines3
How to use DHDRL/monsoon-rl with stable-baselines3:
from huggingface_sb3 import load_from_hub checkpoint = load_from_hub( repo_id="DHDRL/monsoon-rl", filename="{MODEL FILENAME}.zip", ) - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python3 | |
| """ | |
| evaluate_checkpoint.py | |
| ======================= | |
| Synthetic alert-quality eval for a trained checkpoint. | |
| Critical design points | |
| ---------------------- | |
| 1. Product-gate flags come from compute_risk_score on episode ZoneObs at | |
| terminate — identical for checkpoint vs zero_inspect on the same seed. | |
| Identical product P/R is *expected*; it is not evidence of agent skill. | |
| 2. Policy-sensitive signal is believed_p (max zone belief at terminate) | |
| and episode length. Inspections update belief; zero_inspect keeps the | |
| reset-time blend only. | |
| 3. Ground truth must use the same regional event model as training | |
| (_episode_event_plan), not independent _zone_event_flags. | |
| 4. Eval ForecastConfig must match train (clean_episode_ratio + | |
| event_spatial_correlation) or base rates and EV tables are meaningless. | |
| Usage | |
| ----- | |
| python evaluate_checkpoint.py \\ | |
| --checkpoint run_nz3_c090/final_model.zip \\ | |
| --n-zones 3 --max-steps 250 --n-episodes 200 \\ | |
| --clean-episode-ratio 0.90 --event-spatial-correlation 0.85 \\ | |
| --also-zero-inspect | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| from collections import Counter | |
| from typing import Any, Dict, List, Optional | |
| import numpy as np | |
| import zone_observation as _zo | |
| assert _zo.SCHEMA_VERSION == 3, ( | |
| f"evaluate_checkpoint: zone_observation schema mismatch " | |
| f"(expected 3, got {_zo.SCHEMA_VERSION})" | |
| ) | |
| from zone_observation import AlertLevel, ForecastConfig | |
| from weather_forecast_env import ( | |
| make_weather_env, | |
| _episode_event_plan, | |
| ) | |
| def _episode_is_risky( | |
| effective_seed: int, | |
| n_zones: int, | |
| clean_ratio: float, | |
| spatial_corr: float, | |
| ) -> bool: | |
| """Match training: regional event plan, not independent per-zone draws.""" | |
| plan = _episode_event_plan( | |
| n_zones, effective_seed, clean_ratio, spatial_corr | |
| ) | |
| return any(any(f.values()) for f in plan) | |
| def _confusion(tp: int, fp: int, fn: int, tn: int) -> Dict[str, float]: | |
| precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0 | |
| recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0 | |
| f1 = ( | |
| 2 * precision * recall / (precision + recall) | |
| if (precision + recall) > 0 | |
| else 0.0 | |
| ) | |
| return dict( | |
| tp=tp, fp=fp, fn=fn, tn=tn, | |
| precision=precision, recall=recall, f1=f1, | |
| n_positive=tp + fn, | |
| n_negative=tn + fp, | |
| n=tp + fp + fn + tn, | |
| base_rate=(tp + fn) / max(tp + fp + fn + tn, 1), | |
| alert_rate=(tp + fp) / max(tp + fp + fn + tn, 1), | |
| ) | |
| def _day_level_ev( | |
| alert: bool, | |
| risky: bool, | |
| alert_value: float, | |
| false_alert_penalty: float, | |
| miss_penalty: float, | |
| ) -> float: | |
| if alert and risky: | |
| return alert_value | |
| if alert and not risky: | |
| return -false_alert_penalty | |
| if (not alert) and risky: | |
| return -miss_penalty | |
| return 0.0 | |
| def evaluate( | |
| config: ForecastConfig, | |
| n_episodes: int, | |
| eval_seed_base: int = 500_000, | |
| model=None, | |
| policy: str = "checkpoint", | |
| ) -> Dict[str, Any]: | |
| if policy == "checkpoint" and model is None: | |
| raise ValueError("policy=checkpoint requires a loaded model") | |
| env = make_weather_env(config, use_nan_wrapper=True) | |
| base_env = env | |
| if hasattr(env, "env"): | |
| base_env = env.env | |
| terminate_action = int(getattr(base_env, "terminate_action", config.n_zones)) | |
| lengths: List[int] = [] | |
| a_tp = a_fp = a_fn = a_tn = 0 | |
| p_tp = p_fp = p_fn = p_tn = 0 | |
| b_tp = b_fp = b_fn = b_tn = 0 | |
| alert_levels = Counter() | |
| product_flags = Counter() | |
| drought_risks: List[float] = [] | |
| flood_risks: List[float] = [] | |
| max_risks: List[float] = [] | |
| believed_list: List[float] = [] | |
| believed_risky: List[float] = [] | |
| believed_clean: List[float] = [] | |
| adv_ev = prod_ev = belief_ev = 0.0 | |
| always_ev = never_ev = oracle_ev = 0.0 | |
| av = float(config.alert_value) | |
| fap = float(config.false_alert_penalty) | |
| mp = float(config.miss_penalty) | |
| rational = float(config.rational_termination_threshold) | |
| belief_bar = max(rational, float(config.prior_belief) + 0.05, 0.20) | |
| rho = float(getattr(config, "event_spatial_correlation", 0.85)) | |
| n_risky = 0 | |
| for i in range(n_episodes): | |
| eval_seed = eval_seed_base + i | |
| obs, info = env.reset(seed=eval_seed) | |
| risky = _episode_is_risky( | |
| eval_seed, | |
| config.n_zones, | |
| config.clean_episode_ratio, | |
| rho, | |
| ) | |
| if risky: | |
| n_risky += 1 | |
| done = False | |
| ep_len = 0 | |
| final_severity = 0 | |
| product = False | |
| drought = flood = 0.0 | |
| alert_level = "none" | |
| believed_p = float(config.prior_belief) | |
| while not done: | |
| if policy == "zero_inspect": | |
| action = terminate_action | |
| else: | |
| action_masks = env.action_masks() | |
| action, _ = model.predict( | |
| obs, action_masks=action_masks, deterministic=True | |
| ) | |
| action = int(action) | |
| obs, reward, terminated, truncated, info = env.step(action) | |
| ep_len += 1 | |
| done = terminated or truncated | |
| if "alert_level" in info: | |
| alert_level = str(info["alert_level"]) | |
| final_severity = AlertLevel(alert_level).severity() | |
| if "product_actionable" in info: | |
| product = bool(info["product_actionable"]) | |
| if "drought_risk" in info: | |
| drought = float(info["drought_risk"]) | |
| if "flood_risk" in info: | |
| flood = float(info["flood_risk"]) | |
| if "believed_p" in info: | |
| believed_p = float(info["believed_p"]) | |
| elif "zone_belief" in obs: | |
| n_act = config.n_zones | |
| believed_p = float(np.max(obs["zone_belief"][:n_act])) | |
| lengths.append(ep_len) | |
| alert_levels[alert_level] += 1 | |
| product_flags[str(product)] += 1 | |
| drought_risks.append(drought) | |
| flood_risks.append(flood) | |
| max_risks.append(max(drought, flood)) | |
| believed_list.append(believed_p) | |
| if risky: | |
| believed_risky.append(believed_p) | |
| else: | |
| believed_clean.append(believed_p) | |
| alerted_adv = final_severity >= AlertLevel.ADVISORY.severity() | |
| if risky and alerted_adv: | |
| a_tp += 1 | |
| elif risky and not alerted_adv: | |
| a_fn += 1 | |
| elif (not risky) and alerted_adv: | |
| a_fp += 1 | |
| else: | |
| a_tn += 1 | |
| alerted_prod = product | |
| if risky and alerted_prod: | |
| p_tp += 1 | |
| elif risky and not alerted_prod: | |
| p_fn += 1 | |
| elif (not risky) and alerted_prod: | |
| p_fp += 1 | |
| else: | |
| p_tn += 1 | |
| alerted_belief = believed_p >= belief_bar | |
| if risky and alerted_belief: | |
| b_tp += 1 | |
| elif risky and not alerted_belief: | |
| b_fn += 1 | |
| elif (not risky) and alerted_belief: | |
| b_fp += 1 | |
| else: | |
| b_tn += 1 | |
| adv_ev += _day_level_ev(alerted_adv, risky, av, fap, mp) | |
| prod_ev += _day_level_ev(alerted_prod, risky, av, fap, mp) | |
| belief_ev += _day_level_ev(alerted_belief, risky, av, fap, mp) | |
| always_ev += _day_level_ev(True, risky, av, fap, mp) | |
| never_ev += _day_level_ev(False, risky, av, fap, mp) | |
| oracle_ev += _day_level_ev(risky, risky, av, fap, mp) | |
| p_event = 1.0 - float(config.clean_episode_ratio) | |
| analytic_corr = p_event | |
| analytic_iid = 1.0 - (config.clean_episode_ratio ** config.n_zones) | |
| def _mean(xs: List[float]) -> float: | |
| return float(np.mean(xs)) if xs else 0.0 | |
| return { | |
| "policy": policy, | |
| "mean_len": float(np.mean(lengths)) if lengths else 0.0, | |
| "n_episodes": n_episodes, | |
| "n_risky": n_risky, | |
| "empirical_base_rate": n_risky / max(n_episodes, 1), | |
| "analytic_base_rate_iid": analytic_iid, | |
| "analytic_base_rate_corr": analytic_corr, | |
| "clean_episode_ratio": config.clean_episode_ratio, | |
| "event_spatial_correlation": rho, | |
| "n_zones": config.n_zones, | |
| "belief_bar": belief_bar, | |
| "advisory": _confusion(a_tp, a_fp, a_fn, a_tn), | |
| "product": _confusion(p_tp, p_fp, p_fn, p_tn), | |
| "belief": _confusion(b_tp, b_fp, b_fn, b_tn), | |
| "alert_level_counts": dict(alert_levels), | |
| "product_flag_counts": dict(product_flags), | |
| "drought_risk_mean": _mean(drought_risks), | |
| "flood_risk_mean": _mean(flood_risks), | |
| "max_hazard_mean": _mean(max_risks), | |
| "believed_p_mean": _mean(believed_list), | |
| "believed_p_risky_mean": _mean(believed_risky), | |
| "believed_p_clean_mean": _mean(believed_clean), | |
| "believed_p_sep": _mean(believed_risky) - _mean(believed_clean), | |
| "economics": { | |
| "alert_value": av, | |
| "false_alert_penalty": fap, | |
| "miss_penalty": mp, | |
| "rational_threshold": rational, | |
| "ev_advisory_policy": adv_ev, | |
| "ev_product_policy": prod_ev, | |
| "ev_belief_policy": belief_ev, | |
| "ev_always_alert": always_ev, | |
| "ev_never_alert": never_ev, | |
| "ev_oracle": oracle_ev, | |
| "ev_always_per_ep": always_ev / max(n_episodes, 1), | |
| "ev_never_per_ep": never_ev / max(n_episodes, 1), | |
| "ev_product_per_ep": prod_ev / max(n_episodes, 1), | |
| "ev_belief_per_ep": belief_ev / max(n_episodes, 1), | |
| "ev_advisory_per_ep": adv_ev / max(n_episodes, 1), | |
| "oracle_minus_always": oracle_ev - always_ev, | |
| "oracle_minus_always_frac": ( | |
| (oracle_ev - always_ev) / max(abs(always_ev), 1e-9) | |
| ), | |
| }, | |
| } | |
| def _print_block(name: str, c: Dict[str, float]) -> None: | |
| print(f"--- {name} ---") | |
| print( | |
| f" P={c['precision']:.3f} R={c['recall']:.3f} F1={c['f1']:.3f} " | |
| f"alert_rate={c['alert_rate']:.3f}" | |
| ) | |
| print( | |
| f" tp={int(c['tp'])} fp={int(c['fp'])} " | |
| f"fn={int(c['fn'])} tn={int(c['tn'])}" | |
| ) | |
| def _print_result(label: str, m: Dict[str, Any], args: argparse.Namespace) -> None: | |
| print(f"\n===== {label} =====") | |
| print(f"policy={m['policy']}") | |
| if getattr(args, "checkpoint", None) and m["policy"] == "checkpoint": | |
| print(f"checkpoint: {args.checkpoint}") | |
| print( | |
| f"n_zones={m['n_zones']} max_steps={args.max_steps} " | |
| f"n_episodes={m['n_episodes']}" | |
| ) | |
| print( | |
| f"mean_ep_len = {m['mean_len']:.2f} " | |
| f"(structural ceiling ~{m['n_zones'] + 1})" | |
| ) | |
| print( | |
| f"risky base rate empirical={m['empirical_base_rate']:.3f} " | |
| f"analytic_corr≈{m['analytic_base_rate_corr']:.3f} " | |
| f"analytic_iid={m['analytic_base_rate_iid']:.3f} " | |
| f"(clean={m['clean_episode_ratio']:.3f} rho={m['event_spatial_correlation']:.3f})" | |
| ) | |
| print( | |
| f"hazard means drought={m['drought_risk_mean']:.3f} " | |
| f"flood={m['flood_risk_mean']:.3f} max={m['max_hazard_mean']:.3f}" | |
| ) | |
| print( | |
| f"believed_p mean={m['believed_p_mean']:.3f} " | |
| f"risky={m['believed_p_risky_mean']:.3f} " | |
| f"clean={m['believed_p_clean_mean']:.3f} " | |
| f"sep={m['believed_p_sep']:+.3f} " | |
| f"bar={m['belief_bar']:.3f}" | |
| ) | |
| print(f"alert_level counts: {m['alert_level_counts']}") | |
| print(f"product_actionable counts: {m['product_flag_counts']}") | |
| print() | |
| _print_block("ADVISORY+ (legacy; often saturates)", m["advisory"]) | |
| print() | |
| _print_block( | |
| "PRODUCT GATE (scorer on episode obs — policy-insensitive by design)", | |
| m["product"], | |
| ) | |
| print() | |
| _print_block( | |
| f"BELIEF GATE (believed_p ≥ {m['belief_bar']:.2f} — policy-sensitive)", | |
| m["belief"], | |
| ) | |
| print() | |
| e = m["economics"] | |
| print("--- day-level EV ---") | |
| print( | |
| f" economics: alert={e['alert_value']} false={e['false_alert_penalty']} " | |
| f"miss={e['miss_penalty']}" | |
| ) | |
| print( | |
| f" always_alert total_ev={e['ev_always_alert']:+.1f} " | |
| f"per_ep={e['ev_always_per_ep']:+.3f}" | |
| ) | |
| print( | |
| f" never_alert total_ev={e['ev_never_alert']:+.1f} " | |
| f"per_ep={e['ev_never_per_ep']:+.3f}" | |
| ) | |
| print( | |
| f" product_gate total_ev={e['ev_product_policy']:+.1f} " | |
| f"per_ep={e['ev_product_per_ep']:+.3f}" | |
| ) | |
| print( | |
| f" belief_gate total_ev={e['ev_belief_policy']:+.1f} " | |
| f"per_ep={e['ev_belief_per_ep']:+.3f}" | |
| ) | |
| print(f" oracle total_ev={e['ev_oracle']:+.1f}") | |
| print( | |
| f" oracle−always = {e['oracle_minus_always']:+.1f} " | |
| f"({100 * e['oracle_minus_always_frac']:.1f}% of |always|)" | |
| ) | |
| def _print_comparison(trained: Dict[str, Any], zero: Dict[str, Any]) -> None: | |
| print("\n===== POLICY SENSITIVITY =====") | |
| tp, zp = trained["product"], zero["product"] | |
| tb, zb = trained["belief"], zero["belief"] | |
| print("product gate (expect IDENTICAL — scorer on episode obs):") | |
| print( | |
| f" trained P={tp['precision']:.3f} R={tp['recall']:.3f} " | |
| f"F1={tp['f1']:.3f} mean_len={trained['mean_len']:.2f}" | |
| ) | |
| print( | |
| f" zero_inspect P={zp['precision']:.3f} R={zp['recall']:.3f} " | |
| f"F1={zp['f1']:.3f} mean_len={zero['mean_len']:.2f}" | |
| ) | |
| same_prod = ( | |
| int(tp["tp"]) == int(zp["tp"]) | |
| and int(tp["fp"]) == int(zp["fp"]) | |
| and int(tp["fn"]) == int(zp["fn"]) | |
| and int(tp["tn"]) == int(zp["tn"]) | |
| ) | |
| print(f" product identical: {same_prod} (expected True)") | |
| print() | |
| print("belief gate (should DIFFER if inspections change beliefs):") | |
| print( | |
| f" trained P={tb['precision']:.3f} R={tb['recall']:.3f} " | |
| f"F1={tb['f1']:.3f} sep={trained['believed_p_sep']:+.3f}" | |
| ) | |
| print( | |
| f" zero_inspect P={zb['precision']:.3f} R={zb['recall']:.3f} " | |
| f"F1={zb['f1']:.3f} sep={zero['believed_p_sep']:+.3f}" | |
| ) | |
| same_bel = ( | |
| int(tb["tp"]) == int(zb["tp"]) | |
| and int(tb["fp"]) == int(zb["fp"]) | |
| and int(tb["fn"]) == int(zb["fn"]) | |
| and int(tb["tn"]) == int(zb["tn"]) | |
| ) | |
| print(f" belief identical: {same_bel}") | |
| print() | |
| if same_prod and not same_bel: | |
| print( | |
| " RESULT: product policy-insensitive (expected); belief gate differs →\n" | |
| " inspections change terminal belief. Use belief metrics + ep_len as\n" | |
| " the synthetic skill signal." | |
| ) | |
| elif same_prod and same_bel: | |
| print( | |
| " RESULT: both product and belief match zero_inspect.\n" | |
| " Either inspections do not move belief enough, or reset-time\n" | |
| " composite blend already encodes the event (common). Check\n" | |
| " believed_p sep and ep_len; real L1 eval remains the decisive test." | |
| ) | |
| else: | |
| print(" RESULT: product differs (unexpected — check env product path).") | |
| e = trained["economics"] | |
| print() | |
| print("--- calibration pressure ---") | |
| print( | |
| f" oracle−always = {e['oracle_minus_always']:+.1f} " | |
| f"({100 * e['oracle_minus_always_frac']:.1f}% of |always EV|)" | |
| ) | |
| if abs(e["oracle_minus_always_frac"]) < 0.20: | |
| print( | |
| " GAP < 20%: always-alert near oracle under this base rate.\n" | |
| " Prefer higher clean_episode_ratio for selective policies." | |
| ) | |
| def main() -> None: | |
| p = argparse.ArgumentParser( | |
| description="Alert-quality eval (product + belief gate + matched train config)" | |
| ) | |
| p.add_argument("--checkpoint", default=None) | |
| p.add_argument( | |
| "--policy", | |
| choices=("checkpoint", "zero_inspect"), | |
| default="checkpoint", | |
| ) | |
| p.add_argument("--also-zero-inspect", action="store_true") | |
| p.add_argument("--n-zones", type=int, required=True) | |
| p.add_argument("--max-steps", type=int, required=True) | |
| p.add_argument("--n-episodes", type=int, default=200) | |
| p.add_argument("--eval-seed-base", type=int, default=500_000) | |
| p.add_argument("--device", default="auto") | |
| p.add_argument( | |
| "--clean-episode-ratio", | |
| type=float, | |
| default=0.90, | |
| help="Must match train (default 0.90 for correlated runs)", | |
| ) | |
| p.add_argument( | |
| "--event-spatial-correlation", | |
| type=float, | |
| default=0.85, | |
| help="Must match train", | |
| ) | |
| args = p.parse_args() | |
| if args.policy == "checkpoint" and not args.checkpoint: | |
| p.error("--checkpoint is required when --policy checkpoint") | |
| config = ForecastConfig( | |
| n_zones=args.n_zones, | |
| max_steps=args.max_steps, | |
| soft_reset=True, | |
| clean_episode_ratio=args.clean_episode_ratio, | |
| event_spatial_correlation=args.event_spatial_correlation, | |
| ) | |
| model = None | |
| if args.policy == "checkpoint" or args.also_zero_inspect: | |
| if args.checkpoint: | |
| from sb3_contrib import MaskablePPO | |
| model = MaskablePPO.load(args.checkpoint, device=args.device) | |
| trained_m = None | |
| if args.policy == "checkpoint": | |
| trained_m = evaluate( | |
| config, | |
| args.n_episodes, | |
| args.eval_seed_base, | |
| model=model, | |
| policy="checkpoint", | |
| ) | |
| _print_result("TRAINED CHECKPOINT", trained_m, args) | |
| if args.policy == "zero_inspect" or args.also_zero_inspect: | |
| zero_m = evaluate( | |
| config, | |
| args.n_episodes, | |
| args.eval_seed_base, | |
| model=None, | |
| policy="zero_inspect", | |
| ) | |
| _print_result("ZERO-INSPECT CONTROL", zero_m, args) | |
| if trained_m is not None: | |
| _print_comparison(trained_m, zero_m) | |
| elif trained_m is not None: | |
| e = trained_m["economics"] | |
| print() | |
| print("--- interpretation hints ---") | |
| print( | |
| " Product gate is policy-insensitive on synthetic data by design.\n" | |
| " Re-run with --also-zero-inspect to compare belief gate + ep_len." | |
| ) | |
| print( | |
| f" oracle−always = {e['oracle_minus_always']:+.1f} " | |
| f"({100 * e['oracle_minus_always_frac']:.1f}% of |always|)" | |
| ) | |
| if __name__ == "__main__": | |
| main() | |