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
| """ | |
| weather_forecast_env.py | |
| ======================= | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import random | |
| from typing import Any, Dict, List, Optional, Tuple | |
| import numpy as np | |
| import gymnasium as gym | |
| from gymnasium import spaces | |
| import zone_observation as _zo | |
| assert _zo.SCHEMA_VERSION == 3, ( | |
| f"weather_forecast_env: zone_observation schema mismatch " | |
| f"(expected 3, got {_zo.SCHEMA_VERSION})" | |
| ) | |
| from zone_observation import ( | |
| BasinContext, | |
| DataSource, | |
| EpisodeContext, | |
| ForecastConfig, | |
| RiskScore, | |
| make_synthetic_forecast_result, | |
| make_synthetic_zone_obs, | |
| _stable_seed, | |
| ) | |
| _BASIN_CONTEXT_NEUTRAL = (0.0, 0.0, 0.0, 1013.25, 400.0, 2.0, -7.0, 0.0) | |
| _HELIO_REGIME_ORD = {"quiet": 0.0, "active": 1.0, "storm": 2.0} | |
| def basin_context_vector(bc: Optional[BasinContext]) -> np.ndarray: | |
| if bc is None: | |
| return np.array(_BASIN_CONTEXT_NEUTRAL, dtype=np.float32) | |
| goes = float(getattr(bc, "goes_xray_flux", 1e-7) or 1e-7) | |
| goes = max(goes, 1e-12) | |
| regime = str(getattr(bc, "helio_regime", "quiet") or "quiet").lower() | |
| return np.array( | |
| [ | |
| float(getattr(bc, "enso_oni", 0.0) or 0.0), | |
| float(getattr(bc, "iod_dmi", 0.0) or 0.0), | |
| float(getattr(bc, "itcz_latitude_deg", 0.0) or 0.0), | |
| float(getattr(bc, "mslp_regional_hpa", 1013.25) or 1013.25), | |
| float(getattr(bc, "solar_wind_speed_kms", 400.0) or 400.0), | |
| float(getattr(bc, "kp_index", 2.0) or 2.0), | |
| float(np.log10(goes)), | |
| float(_HELIO_REGIME_ORD.get(regime, 0.0)), | |
| ], | |
| dtype=np.float32, | |
| ) | |
| from crop_risk_scorer import compute_risk_score | |
| try: | |
| from product_alert_service import ( | |
| DEFAULT_PRODUCT_GATE, | |
| is_product_actionable as _product_gate_fn, | |
| ) | |
| except ImportError: | |
| DEFAULT_PRODUCT_GATE = None | |
| _product_gate_fn = None | |
| try: | |
| from real_episode_sampler import RealEpisodeIndex | |
| except ImportError: | |
| RealEpisodeIndex = None # training real-data sampling unavailable; | |
| # ForecastConfig.real_data_pkl_path is then | |
| # ignored and training stays fully synthetic, | |
| # same as before this module existed. | |
| logger = logging.getLogger(__name__) | |
| def _info_product_flags(risk_score: RiskScore) -> Dict[str, Any]: | |
| elevated = bool(risk_score.is_elevated()) | |
| if _product_gate_fn is not None and DEFAULT_PRODUCT_GATE is not None: | |
| product = bool(_product_gate_fn(risk_score, DEFAULT_PRODUCT_GATE)) | |
| else: | |
| product = bool(risk_score.is_product_actionable()) | |
| return { | |
| "elevated": elevated, | |
| "product_actionable": product, | |
| } | |
| class NaNSafetyWrapper(gym.Wrapper): | |
| def __init__( | |
| self, | |
| env: gym.Env, | |
| default_value: float = 0.0, | |
| reward_default: float = -100.0, | |
| nan_limit: int = 5000, | |
| verbose: bool = False, | |
| ) -> None: | |
| super().__init__(env) | |
| self.default_value = float(default_value) | |
| self.reward_default = float(reward_default) | |
| self.nan_count = 0 | |
| self.nan_limit = nan_limit | |
| self.verbose = verbose | |
| def _sanitize(self, value: Any, name: str = "") -> Any: | |
| if isinstance(value, dict): | |
| return {k: self._sanitize(v, f"{name}.{k}") for k, v in value.items()} | |
| if isinstance(value, np.ndarray): | |
| if not np.all(np.isfinite(value)): | |
| n = int((~np.isfinite(value)).sum()) | |
| self.nan_count += n | |
| if self.verbose: | |
| logger.warning("NaNSafetyWrapper: %d NaN/Inf in %s", n, name) | |
| return np.nan_to_num( | |
| value, | |
| nan=self.default_value, | |
| posinf=self.default_value, | |
| neginf=self.default_value, | |
| ) | |
| return value | |
| if isinstance(value, (float, np.floating)): | |
| if not np.isfinite(value): | |
| self.nan_count += 1 | |
| return self.default_value | |
| return float(value) | |
| if isinstance(value, bool): | |
| return value | |
| if isinstance(value, (int, np.integer)): | |
| return int(value) | |
| return value | |
| def step(self, action): | |
| obs, reward, terminated, truncated, info = self.env.step(action) | |
| obs = self._sanitize(obs, "obs") | |
| info = self._sanitize(info, "info") | |
| reward = reward if np.isfinite(reward) else self.reward_default | |
| reward = float(np.clip(reward, -31_000.0, 15_500.0)) | |
| if self.nan_count > self.nan_limit: | |
| raise RuntimeError( | |
| f"NaNSafetyWrapper: {self.nan_count} invalid values -- " | |
| f"check data pipeline." | |
| ) | |
| return obs, reward, terminated, truncated, info | |
| def reset(self, **kwargs): | |
| obs, info = self.env.reset(**kwargs) | |
| return self._sanitize(obs, "obs"), self._sanitize(info, "info") | |
| def terminate_action(self) -> int: | |
| return self.env.terminate_action | |
| def action_masks(self) -> np.ndarray: | |
| return self.env.action_masks() | |
| def _zone_event_flags(seed: int, clean_ratio: float) -> Dict[str, bool]: | |
| rng = random.Random(seed) | |
| if rng.random() < clean_ratio: | |
| return {"drought": False, "flood": False, "fungi": False} | |
| kind = rng.choice(["drought", "flood", "fungi"]) | |
| return { | |
| "drought": kind == "drought", | |
| "flood": kind == "flood", | |
| "fungi": kind == "fungi", | |
| } | |
| def _episode_event_plan( | |
| n_zones: int, | |
| base_seed: int, | |
| clean_ratio: float, | |
| spatial_corr: float = 0.85, | |
| ) -> List[Dict[str, bool]]: | |
| n = max(1, int(n_zones)) | |
| clean = {"drought": False, "flood": False, "fungi": False} | |
| rng = random.Random(base_seed + 7919) | |
| rho = float(max(0.0, min(1.0, spatial_corr))) | |
| p_event = 1.0 - float(max(0.0, min(1.0, clean_ratio))) | |
| if rng.random() >= p_event: | |
| return [dict(clean) for _ in range(n)] | |
| kind = rng.choice(["drought", "flood", "fungi"]) | |
| shared = { | |
| "drought": kind == "drought", | |
| "flood": kind == "flood", | |
| "fungi": kind == "fungi", | |
| } | |
| plan: List[Dict[str, bool]] = [] | |
| for i in range(n): | |
| if rng.random() < rho: | |
| plan.append(dict(shared)) | |
| else: | |
| plan.append(dict(clean)) | |
| if not any(any(f.values()) for f in plan): | |
| plan[rng.randrange(n)] = dict(shared) | |
| return plan | |
| def _make_multi_zone_context( | |
| cfg: ForecastConfig, | |
| seed: Optional[int], | |
| ) -> EpisodeContext: | |
| n = cfg.n_zones | |
| base = seed if seed is not None else 0 | |
| zone_ids = [f"zone_{i}" for i in range(n)] | |
| rho = float(getattr(cfg, "event_spatial_correlation", 0.85)) | |
| plan = _episode_event_plan(n, base, cfg.clean_episode_ratio, rho) | |
| ev0 = plan[0] | |
| obs_0 = make_synthetic_zone_obs(zone_ids[0], seed=base, **ev0) | |
| fc_0 = make_synthetic_forecast_result( | |
| zone_ids[0], valid_time=obs_0.valid_time, seed=base, | |
| drought=ev0["drought"], flood=ev0["flood"], | |
| ) | |
| return EpisodeContext( | |
| obs=obs_0, | |
| forecast=fc_0, | |
| config=cfg, | |
| zone_ids=zone_ids, | |
| data_source=DataSource.SYNTHETIC, | |
| ) | |
| def _make_per_zone_forecasts( | |
| zone_ids: List[str], | |
| valid_time: Any, | |
| horizon: int, | |
| base_seed: int, | |
| clean_ratio: float, | |
| spatial_corr: float = 0.85, | |
| event_plan: Optional[List[Dict[str, bool]]] = None, | |
| ) -> List[Any]: | |
| if event_plan is None: | |
| event_plan = _episode_event_plan( | |
| len(zone_ids), base_seed, clean_ratio, spatial_corr | |
| ) | |
| forecasts = [] | |
| for i, zid in enumerate(zone_ids): | |
| ev = event_plan[i] if i < len(event_plan) else _zone_event_flags( | |
| base_seed + i, clean_ratio | |
| ) | |
| fc = make_synthetic_forecast_result( | |
| zid, | |
| horizon_days=horizon, | |
| valid_time=valid_time, | |
| seed=base_seed + i * 31, | |
| drought=ev["drought"], | |
| flood=ev["flood"], | |
| ) | |
| forecasts.append(fc) | |
| return forecasts | |
| def _per_zone_beliefs( | |
| zone_ids: List[str], | |
| base_seed: int, | |
| prior: float, | |
| clean_ratio: float, | |
| spatial_corr: float = 0.85, | |
| event_plan: Optional[List[Dict[str, bool]]] = None, | |
| prior_weight: float = 0.70, | |
| ) -> np.ndarray: | |
| if event_plan is None: | |
| event_plan = _episode_event_plan( | |
| len(zone_ids), base_seed, clean_ratio, spatial_corr | |
| ) | |
| beliefs = [] | |
| for i, zid in enumerate(zone_ids): | |
| ev = event_plan[i] if i < len(event_plan) else _zone_event_flags( | |
| base_seed + i, clean_ratio | |
| ) | |
| obs = make_synthetic_zone_obs(zid, seed=base_seed + i * 17, **ev) | |
| signal = float(obs.composite_risk()) | |
| belief = prior_weight * prior + (1.0 - prior_weight) * signal | |
| beliefs.append(belief) | |
| return np.clip(np.array(beliefs, dtype=np.float32), 0.0, 1.0) | |
| def _resolved_zone_beliefs( | |
| zone_obs: List[Any], | |
| prior: float, | |
| prior_weight: float, | |
| ) -> np.ndarray: | |
| """Belief-map values for zone_obs that are already resolved (real data, | |
| whether externally injected via options={"context": ...} or sampled by | |
| RealEpisodeIndex during training) -- as opposed to _per_zone_beliefs() | |
| above, which generates synthetic obs itself from a seed. | |
| Single source of truth for the prior/signal blend across both | |
| real-data code paths in reset(), so they cannot drift apart the way | |
| the injected path and the synthetic path's hard-coded 0.70 literal | |
| once did. | |
| """ | |
| beliefs = np.zeros(len(zone_obs), dtype=np.float32) | |
| for i, zo in enumerate(zone_obs): | |
| signal = float(zo.composite_risk()) | |
| beliefs[i] = prior_weight * prior + (1.0 - prior_weight) * signal | |
| return np.clip(beliefs, 0.0, 1.0) | |
| def _make_per_zone_obs( | |
| zone_ids: List[str], | |
| base_seed: int, | |
| clean_ratio: float, | |
| spatial_corr: float = 0.85, | |
| event_plan: Optional[List[Dict[str, bool]]] = None, | |
| ) -> List[Any]: | |
| if event_plan is None: | |
| event_plan = _episode_event_plan( | |
| len(zone_ids), base_seed, clean_ratio, spatial_corr | |
| ) | |
| out = [] | |
| for i, zid in enumerate(zone_ids): | |
| ev = event_plan[i] if i < len(event_plan) else _zone_event_flags( | |
| base_seed + i, clean_ratio | |
| ) | |
| out.append(make_synthetic_zone_obs(zid, seed=base_seed + i * 17, **ev)) | |
| return out | |
| class WeatherForecastEnv(gym.Env): | |
| metadata = {"render_modes": ["human"]} | |
| def __init__(self, config: Optional[ForecastConfig] = None) -> None: | |
| super().__init__() | |
| self.config = config or ForecastConfig() | |
| self.max_zones = self.config.n_zones | |
| self.terminate_action = self.max_zones | |
| self.action_space = spaces.Discrete(self.max_zones + 1) | |
| self.observation_space = spaces.Dict({ | |
| "zone_belief": spaces.Box( | |
| 0.0, 1.0, (self.max_zones,), np.float32, | |
| ), | |
| "forecast_precip": spaces.Box( | |
| 0.0, 500.0, (self.max_zones, self.config.horizon_days), np.float32, | |
| ), | |
| "forecast_uncertainty": spaces.Box( | |
| 0.0, 1.0, (self.max_zones,), np.float32, | |
| ), | |
| "action_mask": spaces.Box( | |
| 0, 1, (self.max_zones + 1,), bool, | |
| ), | |
| "prior_belief": spaces.Box( | |
| 0.0, 1.0, (1,), np.float32, | |
| ), | |
| "basin_context": spaces.Box( | |
| low=np.array( | |
| [-5.0, -5.0, -30.0, 900.0, 200.0, 0.0, -9.0, 0.0], | |
| dtype=np.float32, | |
| ), | |
| high=np.array( | |
| [5.0, 5.0, 30.0, 1100.0, 900.0, 9.0, -3.0, 2.0], | |
| dtype=np.float32, | |
| ), | |
| dtype=np.float32, | |
| ), | |
| }) | |
| self._belief_map = np.full(self.max_zones, self.config.prior_belief, np.float32) | |
| self._action_mask = np.ones(self.max_zones + 1, dtype=bool) | |
| self._forecast_arr = np.zeros((self.max_zones, self.config.horizon_days), np.float32) | |
| self._uncertainty = np.zeros(self.max_zones, np.float32) | |
| self._uncertainty_init = np.zeros(self.max_zones, np.float32) | |
| self._visited = np.zeros(self.max_zones, dtype=bool) | |
| self._basin_context = np.array(_BASIN_CONTEXT_NEUTRAL, dtype=np.float32) | |
| self._context: Optional[EpisodeContext] = None | |
| self._zone_forecasts: List[Any] = [] | |
| self._zone_obs: List[Any] = [] | |
| self._zone_ids: List[str] = [] | |
| self._steps_taken: int = 0 | |
| self._cum_reward: float = 0.0 | |
| self._episode_count: int = 0 | |
| self._episode_seed: int = 0 | |
| self._context_source: str = "synthetic" | |
| self._real_index = None | |
| pkl_path = getattr(self.config, "real_data_pkl_path", None) | |
| if pkl_path: | |
| if RealEpisodeIndex is None: | |
| logger.warning( | |
| "WeatherForecastEnv: config.real_data_pkl_path=%r set, " | |
| "but real_episode_sampler.py is not importable -- " | |
| "training will remain fully synthetic.", | |
| pkl_path, | |
| ) | |
| else: | |
| try: | |
| self._real_index = RealEpisodeIndex(pkl_path) | |
| logger.info( | |
| "WeatherForecastEnv: real-data training sampling " | |
| "ENABLED from %s (%d eligible dates, " | |
| "real_data_ratio=%.2f, inject_noise=%s)", | |
| pkl_path, self._real_index.n_eligible_dates, | |
| self.config.real_data_ratio, self.config.inject_noise, | |
| ) | |
| except Exception as e: | |
| logger.error( | |
| "WeatherForecastEnv: failed to load " | |
| "real_data_pkl_path=%r (%s) -- training will " | |
| "remain fully synthetic.", pkl_path, e, | |
| ) | |
| self._real_index = None | |
| logger.info( | |
| "WeatherForecastEnv: n_zones=%d horizon=%dd max_steps=%d " | |
| "terminate_action=%d", | |
| self.max_zones, self.config.horizon_days, | |
| self.config.max_steps, self.terminate_action, | |
| ) | |
| def action_masks(self) -> np.ndarray: | |
| return self._action_mask.copy() | |
| def reset( | |
| self, | |
| *, | |
| seed: Optional[int] = None, | |
| options: Optional[Dict] = None, | |
| ) -> Tuple[Dict[str, np.ndarray], Dict[str, Any]]: | |
| super().reset(seed=seed) | |
| effective_seed = seed if seed is not None else ( | |
| (self.config.seed or 0) + self._episode_count | |
| ) | |
| self._episode_seed = effective_seed | |
| if options and "context" in options: | |
| ctx = options["context"] | |
| if not isinstance(ctx, EpisodeContext): | |
| raise TypeError( | |
| f"options['context'] must be EpisodeContext, " | |
| f"got {type(ctx).__name__}" | |
| ) | |
| self._context = ctx | |
| context_source = "injected" | |
| else: | |
| self._context = None | |
| context_source = "synthetic" | |
| if ( | |
| self._real_index is not None | |
| and float(self.config.real_data_ratio) > 0.0 | |
| ): | |
| # Dedicated RNG stream, offset from the streams | |
| # _maybe_shuffle_zone_order and other reset-time randomness | |
| # use, so enabling real-data sampling doesn't change the | |
| # synthetic path's own random draws when it IS taken. | |
| real_rng = random.Random(int(effective_seed) + 9001) | |
| if real_rng.random() < float(self.config.real_data_ratio): | |
| sampled = self._real_index.sample( | |
| real_rng, self.config.n_zones, self.config, | |
| inject_noise=bool(self.config.inject_noise), | |
| noise_scale=float(self.config.noise_scale), | |
| ) | |
| if sampled is not None: | |
| self._context = sampled | |
| context_source = "real_sampled" | |
| if self._context is None: | |
| self._context = _make_multi_zone_context(self.config, effective_seed) | |
| context_source = "synthetic" | |
| self._zone_ids = list(self._context.zone_ids) | |
| n_active = len(self._zone_ids) | |
| self._context_source = context_source # stamped onto every step's | |
| # info via _get_info(), not | |
| # just reset's -- SB3's | |
| # DummyVecEnv stores | |
| # auto-reset info separately | |
| # from what step() returns, | |
| # so anything only set in | |
| # reset()'s info dict is | |
| # invisible to a callback's | |
| # self.locals["infos"]. | |
| # Confirmed empirically. | |
| if n_active > self.max_zones: | |
| raise ValueError( | |
| f"EpisodeContext has {n_active} zones but env supports " | |
| f"max {self.max_zones}." | |
| ) | |
| # "injected" (explicit real-eval context) and "real_sampled" | |
| # (RealEpisodeIndex draw during training) both arrive as a fully | |
| # resolved EpisodeContext with real zone_obs/zone_forecasts already | |
| # attached -- they share the same validation and belief-computation | |
| # path below. Only "synthetic" generates obs/forecasts itself. | |
| use_resolved_path = context_source in ("injected", "real_sampled") | |
| self._forecast_arr.fill(0.0) | |
| self._uncertainty.fill(0.0) | |
| self._uncertainty_init.fill(0.0) | |
| self._visited.fill(False) | |
| self._steps_taken = 0 | |
| self._cum_reward = 0.0 | |
| self._episode_count += 1 | |
| self._basin_context[:] = basin_context_vector( | |
| self._context.basin_context if self._context is not None else None | |
| ) | |
| # Single source of truth for the prior/signal belief-blend weight -- | |
| # the injected, real-sampled, and synthetic reset paths below must | |
| # all read this same value, not their own copies of it. | |
| prior_w = float(getattr(self.config, "belief_prior_weight", 0.70)) | |
| if use_resolved_path: | |
| z_obs = self._context.resolved_zone_obs() | |
| z_fc = self._context.resolved_zone_forecasts() | |
| if len(z_obs) != n_active or len(z_fc) != n_active: | |
| raise ValueError( | |
| f"{context_source} EpisodeContext zone payload length " | |
| f"mismatch: n_active={n_active} resolved_obs={len(z_obs)} " | |
| f"resolved_fc={len(z_fc)} zone_ids={self._zone_ids}. " | |
| f"For multi-zone episodes supply zone_obs/zone_forecasts " | |
| f"parallel to zone_ids (do not pad one zone into N slots)." | |
| ) | |
| if n_active > 1 and not self._context.zone_obs: | |
| raise ValueError( | |
| f"{context_source} EpisodeContext refuses n_zones>1 " | |
| f"without explicit zone_obs lists (got n_active={n_active}, " | |
| f"zone_ids={self._zone_ids}). Padding the primary obs into " | |
| f"multiple slots is not multi-zone data." | |
| ) | |
| for i, zid in enumerate(self._zone_ids): | |
| if z_obs[i].zone_id != zid or z_fc[i].zone_id != zid: | |
| raise ValueError( | |
| f"{context_source} context zone mismatch at slot {i}: " | |
| f"zone_ids[{i}]={zid!r} obs={z_obs[i].zone_id!r} " | |
| f"fc={z_fc[i].zone_id!r}" | |
| ) | |
| self._zone_obs = list(z_obs) | |
| self._zone_forecasts = list(z_fc) | |
| prior = float(self.config.prior_belief) | |
| risks = _resolved_zone_beliefs(self._zone_obs, prior, prior_w) | |
| self._belief_map[:n_active] = risks[:n_active] | |
| self._belief_map[n_active:] = 0.0 | |
| else: | |
| rho = float(getattr(self.config, "event_spatial_correlation", 0.85)) | |
| plan = _episode_event_plan( | |
| n_active, | |
| effective_seed, | |
| self.config.clean_episode_ratio, | |
| rho, | |
| ) | |
| risks = _per_zone_beliefs( | |
| self._zone_ids, effective_seed, self.config.prior_belief, | |
| self.config.clean_episode_ratio, | |
| spatial_corr=rho, | |
| event_plan=plan, | |
| prior_weight=prior_w, | |
| ) | |
| self._belief_map[:n_active] = risks[:n_active] | |
| self._belief_map[n_active:] = 0.0 | |
| self._zone_forecasts = _make_per_zone_forecasts( | |
| self._zone_ids, | |
| self._context.obs.valid_time, | |
| self.config.horizon_days, | |
| effective_seed, | |
| self.config.clean_episode_ratio, | |
| spatial_corr=rho, | |
| event_plan=plan, | |
| ) | |
| self._zone_obs = _make_per_zone_obs( | |
| self._zone_ids, effective_seed, self.config.clean_episode_ratio, | |
| spatial_corr=rho, | |
| event_plan=plan, | |
| ) | |
| self._maybe_shuffle_zone_order(effective_seed, n_active) | |
| self._update_forecast_arrays() | |
| self._uncertainty_init[:] = self._uncertainty | |
| info = self._get_info() | |
| info["outcome_code"] = "RESET_OK" | |
| info["context_source"] = context_source | |
| return self._get_obs(), info | |
| def _maybe_shuffle_zone_order(self, effective_seed: int, n_active: int) -> None: | |
| if n_active <= 1: | |
| return | |
| if not bool(getattr(self.config, "shuffle_zone_order", True)): | |
| return | |
| if not self._zone_obs or not self._zone_forecasts: | |
| return | |
| if len(self._zone_obs) < n_active or len(self._zone_forecasts) < n_active: | |
| return | |
| order = list(range(n_active)) | |
| random.Random(int(effective_seed) + 4242).shuffle(order) | |
| self._zone_ids = [self._zone_ids[i] for i in order] | |
| self._zone_obs = [self._zone_obs[i] for i in order] | |
| self._zone_forecasts = [self._zone_forecasts[i] for i in order] | |
| reordered = np.zeros_like(self._belief_map) | |
| for new_i, old_i in enumerate(order): | |
| reordered[new_i] = self._belief_map[old_i] | |
| self._belief_map[:n_active] = reordered[:n_active] | |
| self._belief_map[n_active:] = 0.0 | |
| if self._context is not None: | |
| self._context = EpisodeContext( | |
| obs=self._zone_obs[0], | |
| forecast=self._zone_forecasts[0], | |
| config=self.config, | |
| zone_ids=list(self._zone_ids), | |
| basin_context=self._context.basin_context, | |
| data_source=getattr( | |
| self._context, "data_source", DataSource.SYNTHETIC | |
| ), | |
| ground_truth=getattr(self._context, "ground_truth", None), | |
| zone_obs=list(self._zone_obs), | |
| zone_forecasts=list(self._zone_forecasts), | |
| ) | |
| def _update_forecast_arrays(self) -> None: | |
| if self._context is None: | |
| return | |
| h = self.config.horizon_days | |
| n = min(len(self._zone_ids), self.max_zones) | |
| for zi in range(n): | |
| fr = ( | |
| self._zone_forecasts[zi] | |
| if zi < len(self._zone_forecasts) | |
| else self._context.forecast | |
| ) | |
| if fr.precip_mm and len(fr.precip_mm) != h: | |
| logger.warning( | |
| "_update_forecast_arrays: zone %d forecast length=%d != %d", | |
| zi, len(fr.precip_mm), h, | |
| ) | |
| if fr.precip_mm: | |
| vals = np.asarray(fr.precip_mm[:h], dtype=np.float32) | |
| self._forecast_arr[zi, :len(vals)] = vals | |
| if fr.precip_p90 and fr.precip_p10: | |
| p90 = np.asarray(fr.precip_p90[:h], dtype=np.float32) | |
| p10 = np.asarray(fr.precip_p10[:h], dtype=np.float32) | |
| spread = np.clip( | |
| (p90 - p10) / np.maximum(np.abs(p90), 1e-6), 0.0, 1.0 | |
| ) | |
| self._uncertainty[zi] = float(np.mean(spread)) | |
| def _compute_multi_zone_risk(self) -> RiskScore: | |
| if not self._zone_obs or not self._zone_ids: | |
| return compute_risk_score( | |
| self._context.obs, | |
| self._context.forecast, | |
| self._context.config, | |
| ) | |
| n_active = min(len(self._zone_ids), len(self._zone_obs), self.max_zones) | |
| best_score: Optional[RiskScore] = None | |
| for zi in range(n_active): | |
| zone_obs = self._zone_obs[zi] | |
| zone_fc = ( | |
| self._zone_forecasts[zi] | |
| if zi < len(self._zone_forecasts) | |
| else self._context.forecast | |
| ) | |
| score = compute_risk_score( | |
| zone_obs, | |
| zone_fc, | |
| self._context.config, | |
| ) | |
| if best_score is None or score.supply_shortfall_prob > best_score.supply_shortfall_prob: | |
| best_score = score | |
| return best_score if best_score is not None else compute_risk_score( | |
| self._context.obs, | |
| self._context.forecast, | |
| self._context.config, | |
| ) | |
| def step( | |
| self, action: int | |
| ) -> Tuple[Dict[str, np.ndarray], float, bool, bool, Dict[str, Any]]: | |
| if self._context is None: | |
| raise RuntimeError("WeatherForecastEnv: call reset() before step()") | |
| info: Dict[str, Any] = {} | |
| if action == self.terminate_action: | |
| risk_score = self._compute_multi_zone_risk() | |
| believed_p = float(np.max(self._belief_map[:len(self._zone_ids)])) | |
| reward = self._compute_termination_reward(risk_score, believed_p) | |
| self._steps_taken += 1 | |
| self._cum_reward += reward | |
| info["early_termination"] = True | |
| info["budget_saved"] = self.config.max_steps - self._steps_taken | |
| info["alert_level"] = risk_score.alert_level.value | |
| info.update(_info_product_flags(risk_score)) | |
| info["believed_p"] = believed_p | |
| info["supply_risk"] = risk_score.supply_shortfall_prob | |
| info["flood_risk"] = risk_score.flood_risk | |
| info["drought_risk"] = risk_score.drought_risk | |
| info.update(self._get_info()) | |
| info["outcome_code"] = "TERMINATED_EARLY" | |
| return self._get_obs(), float(reward), True, False, info | |
| if action < 0 or action >= self.max_zones: | |
| reward = -self.config.inspection_cost * 2.0 / self.config.alert_value | |
| info["invalid_action"] = True | |
| _branch = "invalid" | |
| elif action >= len(self._zone_ids): | |
| reward = -self.config.inspection_cost * 0.5 / self.config.alert_value | |
| info["padding_action"] = True | |
| _branch = "padding" | |
| elif self._visited[action]: | |
| reward = -self.config.inspection_cost * 3.0 / self.config.alert_value | |
| info["revisit_penalty"] = True | |
| _branch = "revisit" | |
| else: | |
| zone_id = self._zone_ids[action] | |
| reward = self._compute_zone_refinement_reward(action, zone_id) | |
| self._visited[action] = True | |
| info["zone_id"] = zone_id | |
| _branch = "inspect" | |
| self._steps_taken += 1 | |
| self._cum_reward += reward | |
| terminated = self._steps_taken >= self.config.max_steps | |
| if terminated: | |
| risk_score = self._compute_multi_zone_risk() | |
| believed_p = float(np.max(self._belief_map[:len(self._zone_ids)])) | |
| term_r = self._compute_termination_reward(risk_score, believed_p) | |
| reward += term_r | |
| self._cum_reward += term_r | |
| info["budget_exhausted"] = True | |
| info["believed_p"] = believed_p | |
| info["alert_level"] = risk_score.alert_level.value | |
| info.update(_info_product_flags(risk_score)) | |
| info["drought_risk"] = risk_score.drought_risk | |
| info["flood_risk"] = risk_score.flood_risk | |
| info["supply_risk"] = risk_score.supply_shortfall_prob | |
| if terminated: | |
| outcome_code = ( | |
| "INSPECT_AND_BUDGET_EXHAUSTED" if _branch == "inspect" | |
| else "TERMINATED_BUDGET" | |
| ) | |
| else: | |
| outcome_code = { | |
| "invalid": "PENALTY_INVALID", | |
| "padding": "PENALTY_PADDING", | |
| "revisit": "PENALTY_REVISIT", | |
| "inspect": "INSPECT_OK", | |
| }[_branch] | |
| info.update(self._get_info()) | |
| info["outcome_code"] = outcome_code | |
| return ( | |
| self._get_obs(), | |
| float(np.clip(reward, -31_000.0, 15_500.0)), | |
| terminated, False, info, | |
| ) | |
| def _belief_entropy(self, p: float) -> float: | |
| p = max(1e-9, min(1.0 - 1e-9, float(p))) | |
| return -(p * np.log(p) + (1.0 - p) * np.log(1.0 - p)) | |
| def _compute_zone_refinement_reward(self, idx: int, zone_id: str) -> float: | |
| current = float(self._belief_map[idx]) | |
| if idx < len(self._zone_obs): | |
| signal = float(self._zone_obs[idx].composite_risk()) | |
| else: | |
| signal = self.config.prior_belief | |
| if signal > current: | |
| updated = min(current + self.config.belief_increase_rate, signal) | |
| else: | |
| updated = max(current - self.config.belief_decrease_rate, signal) | |
| self._belief_map[idx] = float(np.clip( | |
| updated, self.config.belief_floor, 1.0, | |
| )) | |
| before_err = (current - signal) ** 2 | |
| after_err = (float(self._belief_map[idx]) - signal) ** 2 | |
| info_gain = max(0.0, before_err - after_err) | |
| decay = float(getattr(self.config, "uncertainty_decay", 0.70)) | |
| self._uncertainty[idx] = float( | |
| np.clip(self._uncertainty[idx] * decay, 0.0, 1.0) | |
| ) | |
| info_scale = float(getattr(self.config, "info_gain_scale", 5.0)) | |
| raw_reward = ( | |
| info_scale * info_gain | |
| - self.config.inspection_cost | |
| + self.config.zone_visit_bonus | |
| ) | |
| return float(raw_reward / self.config.alert_value) | |
| def _compute_termination_reward(self, risk_score: RiskScore, believed_p: float) -> float: | |
| cfg = self._context.config | |
| p_event = believed_p | |
| gain = p_event * (cfg.alert_value + cfg.miss_penalty) | |
| cost = (1.0 - p_event) * cfg.false_alert_penalty | |
| unc_scale = float(getattr(cfg, "uncertainty_penalty_scale", 5.0)) | |
| unc = unc_scale * float(np.mean(self._uncertainty)) | |
| n_active = len(self._zone_ids) | |
| n_visited = int(np.sum(self._visited[:n_active])) if n_active > 0 else 0 | |
| budget = max(1, min(n_active, int(self.config.max_steps))) | |
| unvisited_frac = (1.0 - n_visited / budget) if budget > 0 else 0.0 | |
| exploration_penalty = cfg.unvisited_zone_penalty * unvisited_frac | |
| base = (gain - cost - unc - exploration_penalty) / cfg.alert_value | |
| if self._context.ground_truth is not None: | |
| gt = self._context.ground_truth.supply_shortfall_prob | |
| if (gt > cfg.rational_termination_threshold | |
| and p_event < cfg.rational_termination_threshold): | |
| base -= cfg.miss_penalty / cfg.alert_value | |
| return float(base) | |
| def _get_obs(self) -> Dict[str, np.ndarray]: | |
| n_active = len(self._zone_ids) | |
| self._action_mask[:n_active] = ~self._visited[:n_active] | |
| self._action_mask[n_active:self.max_zones] = False | |
| self._action_mask[self.terminate_action] = True | |
| return { | |
| "zone_belief": self._belief_map.copy(), | |
| "forecast_precip": self._forecast_arr.copy(), | |
| "forecast_uncertainty": self._uncertainty.copy(), | |
| "action_mask": self._action_mask.copy(), | |
| "prior_belief": np.array([self.config.prior_belief], dtype=np.float32), | |
| "basin_context": self._basin_context.copy(), | |
| } | |
| def _get_info(self) -> Dict[str, Any]: | |
| n_active = len(self._zone_ids) | |
| active_belief = self._belief_map[:n_active] | |
| active_unc = self._uncertainty[:n_active] | |
| return { | |
| "episode_num": self._episode_count, | |
| "steps_taken": self._steps_taken, | |
| "cumulative_reward": float(self._cum_reward), | |
| "n_zones": n_active, | |
| "zone_ids": list(self._zone_ids), | |
| "visited": self._visited.copy(), | |
| "believed_p": float(np.max(active_belief)) if n_active else 0.0, | |
| "mean_belief": float(np.mean(active_belief)) if n_active else 0.0, | |
| "mean_uncertainty": float(np.mean(active_unc)) if n_active else 0.0, | |
| "n_visited": int(self._visited[:n_active].sum()) if n_active else 0, | |
| # Stamped on EVERY step's info via this method, not just reset's -- | |
| # see the comment in reset() where self._context_source is set for | |
| # why that distinction matters under SB3's VecEnv autoreset. | |
| "context_source": self._context_source, | |
| } | |
| def render(self, mode: Optional[str] = None) -> None: | |
| pass | |
| def close(self) -> None: | |
| pass | |
| def make_weather_env( | |
| config: Optional[ForecastConfig] = None, | |
| use_nan_wrapper: bool = True, | |
| ) -> gym.Env: | |
| env = WeatherForecastEnv(config) | |
| if use_nan_wrapper: | |
| env = NaNSafetyWrapper(env, verbose=False) | |
| return env | |
| def register_weather_environments() -> None: | |
| gym.register( | |
| id="WeatherForecast-v1", | |
| entry_point=lambda: make_weather_env(ForecastConfig()), | |
| max_episode_steps=None, | |
| ) | |
| logger.info("Registered WeatherForecast-v1 with Gymnasium") |