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
| """ | |
| tests/test_weather_forecast_env.py | |
| ==================================== | |
| Integration tests for WeatherForecastEnv and NaNSafetyWrapper. | |
| These tests lock the contracts between the env and the zone_observation | |
| schema. | |
| Run with: pytest tests/ -v | |
| """ | |
| from __future__ import annotations | |
| import sys | |
| import os | |
| from datetime import datetime, timezone | |
| from unittest.mock import MagicMock, patch | |
| import numpy as np | |
| import pytest | |
| # --------------------------------------------------------------------------- | |
| # Path setup — allow running from repo root or tests/ directory | |
| # --------------------------------------------------------------------------- | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) | |
| from zone_observation import ( | |
| AlertLevel, | |
| EpisodeContext, | |
| ForecastConfig, | |
| RiskScore, | |
| ZoneObs, | |
| make_synthetic_episode_context, | |
| make_synthetic_forecast_result, | |
| make_synthetic_zone_obs, | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # crop_risk_scorer mock | |
| # | |
| # The env calls compute_risk_score(obs, forecast, config) and reads: | |
| # risk_score.alert_level (AlertLevel enum) | |
| # risk_score.alert_level.severity() (int) | |
| # risk_score.supply_shortfall_prob (float [0,1]) | |
| # risk_score.flood_risk (float [0,1]) | |
| # risk_score.drought_risk (float [0,1]) | |
| # | |
| # We inject a configurable factory so individual tests can control | |
| # what the scorer returns without re-importing the module each time. | |
| # --------------------------------------------------------------------------- | |
| def _make_mock_risk_score( | |
| alert_level: AlertLevel = AlertLevel.ADVISORY, | |
| supply_shortfall_prob: float = 0.4, | |
| flood_risk: float = 0.3, | |
| drought_risk: float = 0.1, | |
| ) -> RiskScore: | |
| return RiskScore( | |
| zone_id="mock_zone", | |
| scored_at=datetime.now(tz=timezone.utc), | |
| alert_level=alert_level, | |
| supply_shortfall_prob=supply_shortfall_prob, | |
| flood_risk=flood_risk, | |
| drought_risk=drought_risk, | |
| ) | |
| def mock_crop_risk_scorer(): | |
| mock_module = MagicMock() | |
| mock_module.compute_risk_score.return_value = _make_mock_risk_score() | |
| with patch.dict("sys.modules", {"crop_risk_scorer": mock_module}): | |
| yield mock_module | |
| def make_env(mock_crop_risk_scorer): | |
| def _factory(config: ForecastConfig = None, nan_wrapper: bool = True): | |
| from weather_forecast_env import make_weather_env | |
| return make_weather_env(config, use_nan_wrapper=nan_wrapper) | |
| return _factory | |
| # --------------------------------------------------------------------------- | |
| # Observation space contract | |
| # --------------------------------------------------------------------------- | |
| class TestObservationSpaceContract: | |
| def test_reset_obs_keys_are_complete(self, make_env): | |
| env = make_env() | |
| obs, _ = env.reset(seed=0) | |
| expected = { | |
| "zone_belief", "forecast_precip", "forecast_uncertainty", | |
| "action_mask", "prior_belief", "basin_context", | |
| } | |
| assert set(obs.keys()) == expected | |
| def test_reset_obs_shapes_match_space(self, make_env): | |
| cfg = ForecastConfig(n_zones=3, horizon_days=14) | |
| env = make_env(cfg) | |
| obs, _ = env.reset(seed=0) | |
| assert obs["zone_belief"].shape == (3,) | |
| assert obs["forecast_precip"].shape == (3, 14) | |
| assert obs["forecast_uncertainty"].shape == (3,) | |
| assert obs["action_mask"].shape == (4,) # n_zones + 1 | |
| assert obs["prior_belief"].shape == (1,) | |
| def test_step_obs_shapes_consistent_with_reset(self, make_env): | |
| cfg = ForecastConfig(n_zones=2, horizon_days=10) | |
| env = make_env(cfg) | |
| reset_obs, _ = env.reset(seed=1) | |
| step_obs, _, _, _, _ = env.step(0) | |
| for key in reset_obs: | |
| assert reset_obs[key].shape == step_obs[key].shape, ( | |
| f"Shape mismatch on '{key}' between reset and step" | |
| ) | |
| def test_obs_values_are_all_finite_after_reset(self, make_env): | |
| env = make_env() | |
| obs, _ = env.reset(seed=5) | |
| for key, arr in obs.items(): | |
| arr_f = arr.astype(float) | |
| assert np.all(np.isfinite(arr_f)), f"Non-finite value in obs['{key}'] after reset" | |
| def test_obs_values_are_all_finite_after_step(self, make_env): | |
| env = make_env() | |
| env.reset(seed=5) | |
| obs, _, _, _, _ = env.step(0) | |
| for key, arr in obs.items(): | |
| arr_f = arr.astype(float) | |
| assert np.all(np.isfinite(arr_f)), f"Non-finite value in obs['{key}'] after step" | |
| def test_zone_belief_in_unit_interval(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=4)) | |
| obs, _ = env.reset(seed=2) | |
| assert np.all(obs["zone_belief"] >= 0.0) | |
| assert np.all(obs["zone_belief"] <= 1.0) | |
| def test_forecast_precip_non_negative(self, make_env): | |
| env = make_env() | |
| obs, _ = env.reset(seed=3) | |
| assert np.all(obs["forecast_precip"] >= 0.0) | |
| def test_terminate_action_always_in_mask(self, make_env): | |
| cfg = ForecastConfig(n_zones=3) | |
| env = make_env(cfg) | |
| obs, _ = env.reset(seed=0) | |
| assert obs["action_mask"][env.terminate_action], ( | |
| "terminate_action must never be masked" | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Belief map dynamics | |
| # --------------------------------------------------------------------------- | |
| class TestBeliefMapDynamics: | |
| def test_belief_decreases_after_inspection(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1)) | |
| obs_before, _ = env.reset(seed=0) | |
| belief_before = obs_before["zone_belief"][0] | |
| obs_after, _, _, _, _ = env.step(0) | |
| belief_after = obs_after["zone_belief"][0] | |
| assert belief_after <= belief_before, ( | |
| f"Belief increased after inspection: {belief_before} → {belief_after}" | |
| ) | |
| def test_belief_never_goes_below_floor(self, make_env): | |
| cfg = ForecastConfig(n_zones=1, max_steps=100) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| for _ in range(50): | |
| obs, _, terminated, _, _ = env.step(0) | |
| assert obs["zone_belief"][0] >= cfg.belief_floor - 1e-6, ( | |
| f"Belief fell below floor: {obs['zone_belief'][0]} < {cfg.belief_floor}" | |
| ) | |
| if terminated: | |
| break | |
| def test_belief_initialised_from_composite_risk_on_reset(self, make_env): | |
| cfg = ForecastConfig(prior_belief=0.12, n_zones=2) | |
| env = make_env(cfg) | |
| obs, _ = env.reset(seed=0) | |
| assert obs["zone_belief"].shape == (2,) | |
| assert np.all(obs["zone_belief"] >= cfg.prior_belief * 0.5 - 1e-6) | |
| assert np.all(obs["zone_belief"] <= 1.0) | |
| def test_belief_reset_between_episodes(self, make_env): | |
| cfg = ForecastConfig(n_zones=1, max_steps=5) | |
| env = make_env(cfg) | |
| obs_seed1_first, _ = env.reset(seed=1) | |
| for _ in range(3): | |
| env.step(0) | |
| obs_seed1_again, _ = env.reset(seed=1) | |
| np.testing.assert_allclose( | |
| obs_seed1_first["zone_belief"], | |
| obs_seed1_again["zone_belief"], | |
| atol=1e-5, | |
| err_msg="Same seed must produce same initial belief (no leakage)", | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Reward contract | |
| # --------------------------------------------------------------------------- | |
| class TestRewardContract: | |
| def test_inspection_reward_is_finite(self, make_env): | |
| env = make_env() | |
| env.reset(seed=0) | |
| _, reward, _, _, _ = env.step(0) | |
| assert np.isfinite(reward) | |
| def test_termination_reward_is_finite(self, make_env): | |
| env = make_env() | |
| env.reset(seed=0) | |
| _, reward, _, _, _ = env.step(env.terminate_action) | |
| assert np.isfinite(reward) | |
| def test_reward_is_clipped_within_wrapper_bounds(self, make_env): | |
| env = make_env() | |
| env.reset(seed=0) | |
| for _ in range(3): | |
| _, reward, terminated, _, _ = env.step(0) | |
| assert -31_000.0 <= reward <= 15_500.0 | |
| if terminated: | |
| break | |
| def test_termination_reward_positive_when_high_belief( | |
| self, make_env, mock_crop_risk_scorer | |
| ): | |
| cfg = ForecastConfig(alert_value=100.0, false_alert_penalty=20.0, miss_penalty=50.0) | |
| mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score( | |
| supply_shortfall_prob=0.9, | |
| alert_level=AlertLevel.CRITICAL, | |
| ) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| env.unwrapped._belief_map[:] = 0.85 | |
| _, reward, terminated, _, info = env.step(env.terminate_action) | |
| assert terminated | |
| assert reward > 0, ( | |
| f"Expected positive termination reward for high belief, got {reward}" | |
| ) | |
| assert info["believed_p"] >= 0.8 | |
| def test_termination_reward_negative_when_low_risk( | |
| self, make_env, mock_crop_risk_scorer | |
| ): | |
| cfg = ForecastConfig(alert_value=100.0, false_alert_penalty=20.0) | |
| mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score( | |
| supply_shortfall_prob=0.02, | |
| alert_level=AlertLevel.NONE, | |
| ) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| _, reward, terminated, _, _ = env.step(env.terminate_action) | |
| assert terminated | |
| assert reward < 0, ( | |
| f"Expected negative termination reward for low-risk zone, got {reward}" | |
| ) | |
| def test_invalid_action_is_penalised(self, make_env): | |
| cfg = ForecastConfig(n_zones=1) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| env2 = make_env(ForecastConfig(n_zones=4)) | |
| ctx = make_synthetic_episode_context("z", seed=1) | |
| env2.reset(options={"context": ctx}) | |
| _, reward, _, _, info = env2.step(2) # padding zone | |
| assert reward < 0 or info.get("padding_action"), ( | |
| "Inspecting a padding zone should penalise or flag padding_action" | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Termination conditions | |
| # --------------------------------------------------------------------------- | |
| class TestTerminationConditions: | |
| def test_terminate_action_ends_episode(self, make_env): | |
| env = make_env() | |
| env.reset(seed=0) | |
| _, _, terminated, truncated, info = env.step(env.terminate_action) | |
| assert terminated | |
| assert not truncated | |
| assert info.get("early_termination") | |
| def test_max_steps_exhaustion_ends_episode(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=3)) | |
| env.reset(seed=0) | |
| for step in range(4): | |
| _, _, terminated, truncated, info = env.step(0) | |
| if terminated or truncated: | |
| assert step == 2, f"Expected termination at step 2, got {step}" | |
| assert info.get("budget_exhausted") | |
| break | |
| else: | |
| pytest.fail("Episode did not terminate after max_steps") | |
| def test_step_before_reset_raises(self, make_env): | |
| from weather_forecast_env import WeatherForecastEnv | |
| env = WeatherForecastEnv(ForecastConfig(n_zones=1)) | |
| with pytest.raises(RuntimeError, match="reset()"): | |
| env.step(0) | |
| def test_budget_saved_decreases_with_steps(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=10)) | |
| env.reset(seed=0) | |
| env.step(0) # 1 step | |
| _, _, _, _, info = env.step(env.terminate_action) | |
| assert info["budget_saved"] == 8 # max_steps - steps_taken(2) | |
| def test_miss_penalty_applied_when_ground_truth_present( | |
| self, make_env, mock_crop_risk_scorer | |
| ): | |
| from zone_observation import make_synthetic_episode_context, RiskScore, AlertLevel | |
| mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score( | |
| alert_level=AlertLevel.ADVISORY, | |
| supply_shortfall_prob=0.6, | |
| ) | |
| ctx = make_synthetic_episode_context("gt_zone", seed=7) | |
| gt = _make_mock_risk_score(supply_shortfall_prob=0.8, alert_level=AlertLevel.WARNING) | |
| object.__setattr__(gt, "zone_id", ctx.obs.zone_id) | |
| from zone_observation import EpisodeContext | |
| ctx = EpisodeContext( | |
| obs=ctx.obs, forecast=ctx.forecast, config=ctx.config, | |
| ground_truth=gt, zone_ids=ctx.zone_ids, data_source=ctx.data_source, | |
| ) | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=1, miss_penalty=40.0, alert_value=100.0)) | |
| env.reset(options={"context": ctx}) | |
| _, reward, terminated, _, info = env.step(env.terminate_action) | |
| assert terminated | |
| assert reward < 0.5, f"Expected miss penalty to pull reward down, got {reward}" | |
| # --------------------------------------------------------------------------- | |
| # EpisodeContext injection | |
| # --------------------------------------------------------------------------- | |
| class TestContextInjection: | |
| def test_injected_context_produces_valid_obs(self, make_env): | |
| ctx = make_synthetic_episode_context("inject_zone", seed=42) | |
| env = make_env(ForecastConfig(n_zones=1)) | |
| obs, _ = env.reset(options={"context": ctx}) | |
| assert obs["zone_belief"].shape == (1,) | |
| assert np.all(np.isfinite(obs["forecast_precip"])) | |
| def test_wrong_context_type_raises_type_error(self, make_env): | |
| env = make_env() | |
| env.reset(seed=0) | |
| with pytest.raises(TypeError): | |
| env.reset(options={"context": "not_an_episode_context"}) | |
| def test_two_different_contexts_produce_different_forecasts(self, make_env): | |
| cfg = ForecastConfig(n_zones=1) | |
| env = make_env(cfg) | |
| ctx_a = make_synthetic_episode_context("zone_A", flood=True, seed=1) | |
| ctx_b = make_synthetic_episode_context("zone_B", drought=True, seed=2) | |
| obs_a, _ = env.reset(options={"context": ctx_a}) | |
| obs_b, _ = env.reset(options={"context": ctx_b}) | |
| assert not np.array_equal(obs_a["forecast_precip"], obs_b["forecast_precip"]), ( | |
| "Different injected contexts must produce different forecast arrays" | |
| ) | |
| def test_synthetic_fallback_works_without_context(self, make_env): | |
| env = make_env() | |
| obs, info = env.reset(seed=7) | |
| assert obs is not None | |
| assert isinstance(info, dict) | |
| def test_context_forecast_horizon_matches_env_config(self, make_env): | |
| cfg = ForecastConfig(n_zones=1, horizon_days=14) | |
| obs_14d = make_synthetic_zone_obs("z", seed=1) | |
| fc_14d = make_synthetic_forecast_result("z", horizon_days=14, valid_time=obs_14d.valid_time) | |
| ctx = EpisodeContext(obs=obs_14d, forecast=fc_14d, config=cfg) | |
| env = make_env(cfg) | |
| obs, _ = env.reset(options={"context": ctx}) | |
| assert obs["forecast_precip"].shape == (1, 14) | |
| # --------------------------------------------------------------------------- | |
| # Soft reset — no cross-episode leakage | |
| # --------------------------------------------------------------------------- | |
| class TestSoftReset: | |
| def test_two_seeded_resets_differ(self, make_env): | |
| env = make_env() | |
| obs0, _ = env.reset(seed=0) | |
| obs1, _ = env.reset(seed=1) | |
| assert not np.array_equal( | |
| obs0["forecast_precip"], obs1["forecast_precip"] | |
| ), "Two resets with different seeds produced identical forecast_precip" | |
| def test_belief_array_cleared_between_episodes(self, make_env): | |
| cfg = ForecastConfig(n_zones=2, prior_belief=0.3, max_steps=20) | |
| env = make_env(cfg) | |
| obs_clean, _ = env.reset(seed=5) | |
| belief_clean = obs_clean["zone_belief"].copy() | |
| env.reset(seed=0) | |
| for _ in range(10): | |
| _, _, terminated, _, _ = env.step(0) | |
| if terminated: | |
| break | |
| obs_after, _ = env.reset(seed=5) | |
| np.testing.assert_allclose( | |
| obs_after["zone_belief"], belief_clean, atol=1e-5, | |
| err_msg="Belief leaked from previous episode into fresh reset", | |
| ) | |
| def test_forecast_array_cleared_between_episodes(self, make_env): | |
| cfg = ForecastConfig(n_zones=1, horizon_days=10) | |
| env = make_env(cfg) | |
| ctx_flood = make_synthetic_episode_context("z", flood=True, seed=1) | |
| ctx_drought = make_synthetic_episode_context("z", drought=True, seed=2) | |
| obs_flood, _ = env.reset(options={"context": ctx_flood}) | |
| obs_drought, _ = env.reset(options={"context": ctx_drought}) | |
| assert not np.array_equal( | |
| obs_flood["forecast_precip"], obs_drought["forecast_precip"] | |
| ), "forecast_precip not updated on second reset — stale data from first episode" | |
| def test_steps_taken_resets_to_zero(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=10)) | |
| env.reset(seed=0) | |
| for _ in range(3): | |
| env.step(0) | |
| env.reset(seed=1) | |
| done = False | |
| steps = 0 | |
| while not done and steps < 15: | |
| _, _, terminated, truncated, _ = env.step(0) | |
| done = terminated or truncated | |
| steps += 1 | |
| assert steps <= 11, "steps_taken not reset — episode ended prematurely" | |
| # --------------------------------------------------------------------------- | |
| # NaNSafetyWrapper | |
| # --------------------------------------------------------------------------- | |
| class TestNaNSafetyWrapper: | |
| def test_nan_in_belief_map_is_cleaned(self, make_env): | |
| from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv | |
| base = WeatherForecastEnv(ForecastConfig(n_zones=1)) | |
| wrapped = NaNSafetyWrapper(base) | |
| wrapped.reset(seed=0) | |
| base._belief_map[0] = float("nan") | |
| obs, _, _, _, _ = wrapped.step(wrapped.terminate_action) | |
| assert np.all(np.isfinite(obs["zone_belief"])), "NaN in belief_map not sanitised" | |
| def test_inf_in_forecast_is_cleaned(self, make_env): | |
| from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv | |
| base = WeatherForecastEnv(ForecastConfig(n_zones=1, horizon_days=5)) | |
| wrapped = NaNSafetyWrapper(base) | |
| wrapped.reset(seed=0) | |
| base._forecast_arr[0, 2] = float("inf") | |
| obs, _, _, _, _ = wrapped.step(0) | |
| assert np.all(np.isfinite(obs["forecast_precip"])), "Inf in forecast not sanitised" | |
| def test_non_finite_reward_replaced(self, mock_crop_risk_scorer): | |
| from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv | |
| base = WeatherForecastEnv(ForecastConfig(n_zones=1)) | |
| wrapped = NaNSafetyWrapper(base, reward_default=-999.0) | |
| wrapped.reset(seed=0) | |
| original_step = base.step | |
| def _bad_step(action): | |
| obs, _, term, trunc, info = original_step(action) | |
| return obs, float("inf"), term, trunc, info | |
| base.step = _bad_step | |
| _, reward, _, _, _ = wrapped.step(0) | |
| assert reward == -999.0 or np.isfinite(reward), ( | |
| "Non-finite reward from env was not replaced by NaNSafetyWrapper" | |
| ) | |
| def test_terminate_action_proxied_through_wrapper(self, mock_crop_risk_scorer): | |
| from weather_forecast_env import make_weather_env | |
| env = make_weather_env(use_nan_wrapper=True) | |
| env.reset(seed=0) | |
| assert hasattr(env, "terminate_action") | |
| assert isinstance(env.terminate_action, int) | |
| def test_nan_limit_raises_runtime_error(self, mock_crop_risk_scorer): | |
| from weather_forecast_env import NaNSafetyWrapper, WeatherForecastEnv | |
| base = WeatherForecastEnv(ForecastConfig(n_zones=1)) | |
| wrapped = NaNSafetyWrapper(base, nan_limit=0) | |
| wrapped.reset(seed=0) | |
| base._belief_map[0] = float("nan") | |
| with pytest.raises(RuntimeError, match="invalid values"): | |
| wrapped.step(0) | |
| # --------------------------------------------------------------------------- | |
| # End-to-end episode rollout | |
| # --------------------------------------------------------------------------- | |
| class TestEpisodeRollout: | |
| def test_single_episode_terminates(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=20)) | |
| env.reset(seed=0) | |
| done = False | |
| steps = 0 | |
| while not done: | |
| action = env.action_space.sample() | |
| _, _, terminated, truncated, _ = env.step(action) | |
| done = terminated or truncated | |
| steps += 1 | |
| assert steps <= 25, "Episode ran beyond max_steps + buffer" | |
| def test_10_episodes_all_terminate_cleanly(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=10)) | |
| for ep in range(10): | |
| env.reset(seed=ep) | |
| done = False | |
| steps = 0 | |
| while not done: | |
| _, _, terminated, truncated, info = env.step(0) | |
| done = terminated or truncated | |
| steps += 1 | |
| assert done, f"Episode {ep} never terminated" | |
| def test_cumulative_reward_in_info(self, make_env): | |
| env = make_env(ForecastConfig(n_zones=1, max_steps=5)) | |
| env.reset(seed=0) | |
| for _ in range(3): | |
| _, _, terminated, _, info = env.step(0) | |
| assert "cumulative_reward" in info | |
| assert np.isfinite(info["cumulative_reward"]) | |
| if terminated: | |
| break | |
| def test_obs_space_contains_all_returned_obs(self, make_env): | |
| from weather_forecast_env import WeatherForecastEnv | |
| cfg = ForecastConfig(n_zones=2, horizon_days=7, max_steps=5) | |
| base = WeatherForecastEnv(cfg) | |
| obs, _ = base.reset(seed=0) | |
| assert base.observation_space.contains(obs), ( | |
| "Reset obs not contained in observation_space" | |
| ) | |
| obs2, _, _, _, _ = base.step(0) | |
| assert base.observation_space.contains(obs2), ( | |
| "Step obs not contained in observation_space" | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Visited-zone masking | |
| # --------------------------------------------------------------------------- | |
| class TestVisitedZoneMasking: | |
| def test_visited_zone_is_masked_after_inspection(self, make_env): | |
| cfg = ForecastConfig(n_zones=3, max_steps=10) | |
| env = make_env(cfg) | |
| obs, _ = env.reset(seed=0) | |
| assert obs["action_mask"][0] | |
| obs, _, _, _, _ = env.step(0) | |
| assert not obs["action_mask"][0], "Zone 0 should be masked after inspection" | |
| def test_terminate_remains_valid_after_all_zones_visited(self, make_env): | |
| cfg = ForecastConfig(n_zones=2, max_steps=10) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| env.step(0) | |
| obs, _, terminated, _, _ = env.step(1) | |
| if not terminated: | |
| assert obs["action_mask"][env.terminate_action], ( | |
| "Terminate must remain valid even after all zones visited" | |
| ) | |
| def test_revisit_penalty_applied(self, make_env): | |
| cfg = ForecastConfig(n_zones=2, max_steps=10) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| _, r_first, _, _, _ = env.step(0) | |
| _, r_revisit, _, _, info = env.step(0) | |
| assert info.get("revisit_penalty") or r_revisit < 0, ( | |
| "Re-inspecting a visited zone must incur a penalty" | |
| ) | |
| def test_beliefs_differ_across_zones_on_reset(self, make_env): | |
| cfg = ForecastConfig(n_zones=4) | |
| env = make_env(cfg) | |
| found_diverse = False | |
| for seed in range(10): | |
| obs, _ = env.reset(seed=seed) | |
| beliefs = obs["zone_belief"] | |
| if not np.all(np.isclose(beliefs, beliefs[0], atol=1e-4)): | |
| found_diverse = True | |
| break | |
| assert found_diverse, ( | |
| "No seed in range(10) produced diverse zone beliefs — " | |
| "per-zone seeding not working" | |
| ) | |
| def test_beliefs_in_unit_interval_after_composite_risk_init(self, make_env): | |
| cfg = ForecastConfig(n_zones=4) | |
| env = make_env(cfg) | |
| for seed in range(5): | |
| obs, _ = env.reset(seed=seed) | |
| b = obs["zone_belief"] | |
| assert np.all(b >= 0.0) and np.all(b <= 1.0), ( | |
| f"Belief out of [0,1] at seed={seed}: {b}" | |
| ) | |
| def test_same_seed_produces_same_initial_beliefs(self, make_env): | |
| cfg = ForecastConfig(n_zones=3) | |
| env = make_env(cfg) | |
| obs_a, _ = env.reset(seed=7) | |
| obs_b, _ = env.reset(seed=7) | |
| np.testing.assert_allclose( | |
| obs_a["zone_belief"], obs_b["zone_belief"], atol=1e-5, | |
| err_msg="Same seed must produce same initial beliefs", | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Reward scale | |
| # --------------------------------------------------------------------------- | |
| class TestRewardScale: | |
| def test_inspection_and_termination_rewards_comparable( | |
| self, make_env, mock_crop_risk_scorer | |
| ): | |
| mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score( | |
| supply_shortfall_prob=0.5, | |
| alert_level=AlertLevel.ADVISORY, | |
| ) | |
| env = make_env(ForecastConfig(n_zones=2, max_steps=10)) | |
| env.reset(seed=0) | |
| _, step_reward, _, _, _ = env.step(0) | |
| env.reset(seed=0) | |
| _, term_reward, _, _, _ = env.step(env.terminate_action) | |
| ratio = abs(term_reward) / max(abs(step_reward), 1e-6) | |
| assert ratio < 100.0, ( | |
| f"Reward scale mismatch: step={step_reward:.3f}, " | |
| f"terminal={term_reward:.3f}, ratio={ratio:.1f}" | |
| ) | |
| def test_termination_reward_scaled_by_alert_value( | |
| self, make_env, mock_crop_risk_scorer | |
| ): | |
| mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score( | |
| supply_shortfall_prob=0.9, | |
| alert_level=AlertLevel.CRITICAL, | |
| ) | |
| cfg = ForecastConfig(alert_value=100.0, false_alert_penalty=20.0) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| _, reward, _, _, _ = env.step(env.terminate_action) | |
| assert abs(reward) < 20.0, ( | |
| f"Terminal reward {reward:.3f} appears unscaled (expected < 20)" | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Multi-zone properties | |
| # --------------------------------------------------------------------------- | |
| class TestMultiZoneProperties: | |
| def test_multi_zone_env_has_correct_action_space(self, make_env): | |
| cfg = ForecastConfig(n_zones=4) | |
| env = make_env(cfg) | |
| assert env.action_space.n == 5 # n_zones + terminate | |
| def test_n_zones_info_reflects_context(self, make_env): | |
| cfg = ForecastConfig(n_zones=3) | |
| env = make_env(cfg) | |
| _, info = env.reset(seed=0) | |
| assert info["n_zones"] == 3 | |
| def test_visited_info_in_step_output(self, make_env): | |
| cfg = ForecastConfig(n_zones=2, max_steps=5) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| _, _, _, _, info = env.step(0) | |
| assert "visited" in info | |
| assert info["visited"].shape == (2,) | |
| # --------------------------------------------------------------------------- | |
| # Per-zone forecast diversity, uncertainty decay | |
| # --------------------------------------------------------------------------- | |
| class TestPerZoneForecastDiversity: | |
| def test_forecast_rows_differ_across_zones(self, make_env): | |
| cfg = ForecastConfig(n_zones=4, horizon_days=14) | |
| env = make_env(cfg) | |
| obs, _ = env.reset(seed=0) | |
| precip = obs["forecast_precip"] | |
| all_same = all(np.allclose(precip[0], precip[i]) for i in range(1, 4)) | |
| assert not all_same, ( | |
| "All zones have identical forecast_precip — per-zone forecasts not wired" | |
| ) | |
| def test_forecast_diversity_across_seeds(self, make_env): | |
| cfg = ForecastConfig(n_zones=3, horizon_days=10) | |
| env = make_env(cfg) | |
| found_diverse = False | |
| for seed in range(5): | |
| obs, _ = env.reset(seed=seed) | |
| precip = obs["forecast_precip"] | |
| if not np.allclose(precip[0], precip[1], atol=0.1): | |
| found_diverse = True | |
| break | |
| assert found_diverse, "No seed produced diverse per-zone forecasts" | |
| def test_uncertainty_differs_across_zones(self, make_env): | |
| cfg = ForecastConfig(n_zones=4) | |
| env = make_env(cfg) | |
| found_diverse = False | |
| for seed in range(10): | |
| obs, _ = env.reset(seed=seed) | |
| unc = obs["forecast_uncertainty"] | |
| if not np.all(np.isclose(unc, unc[0], atol=1e-4)): | |
| found_diverse = True | |
| break | |
| assert found_diverse, "Forecast uncertainty is identical across all zones" | |
| class TestUncertaintyDecay: | |
| def test_uncertainty_decreases_after_inspection(self, make_env): | |
| cfg = ForecastConfig(n_zones=2, max_steps=10) | |
| env = make_env(cfg) | |
| obs_before, _ = env.reset(seed=0) | |
| unc_before = obs_before["forecast_uncertainty"][0] | |
| obs_after, _, _, _, _ = env.step(0) | |
| unc_after = obs_after["forecast_uncertainty"][0] | |
| assert unc_after < unc_before + 1e-6, ( | |
| f"Uncertainty did not decrease after inspection: " | |
| f"{unc_before:.4f} → {unc_after:.4f}" | |
| ) | |
| def test_uninspected_zone_uncertainty_unchanged(self, make_env): | |
| cfg = ForecastConfig(n_zones=3, max_steps=10) | |
| env = make_env(cfg) | |
| obs_before, _ = env.reset(seed=0) | |
| unc_zone1_before = obs_before["forecast_uncertainty"][1] | |
| obs_after, _, _, _, _ = env.step(0) | |
| unc_zone1_after = obs_after["forecast_uncertainty"][1] | |
| assert abs(unc_zone1_after - unc_zone1_before) < 1e-6, ( | |
| "Uninspected zone 1 uncertainty changed when zone 0 was inspected" | |
| ) | |
| def test_uncertainty_bounded_after_multiple_inspections(self, make_env): | |
| cfg = ForecastConfig(n_zones=1, max_steps=20) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| for _ in range(15): | |
| obs, _, terminated, _, _ = env.step(0) | |
| unc = obs["forecast_uncertainty"][0] | |
| assert 0.0 <= unc <= 1.0, f"Uncertainty out of [0,1]: {unc}" | |
| if terminated: | |
| break | |
| def test_uncertainty_reset_between_episodes(self, make_env): | |
| cfg = ForecastConfig(n_zones=1, max_steps=10) | |
| env = make_env(cfg) | |
| obs_init, _ = env.reset(seed=5) | |
| unc_init = obs_init["forecast_uncertainty"][0] | |
| for _ in range(5): | |
| _, _, terminated, _, _ = env.step(0) | |
| if terminated: | |
| break | |
| obs_fresh, _ = env.reset(seed=5) | |
| unc_fresh = obs_fresh["forecast_uncertainty"][0] | |
| assert abs(unc_fresh - unc_init) < 1e-5, ( | |
| f"Uncertainty not restored on reset: init={unc_init:.4f}, fresh={unc_fresh:.4f}" | |
| ) | |
| def test_lower_uncertainty_improves_termination_reward( | |
| self, make_env, mock_crop_risk_scorer | |
| ): | |
| mock_crop_risk_scorer.compute_risk_score.return_value = _make_mock_risk_score( | |
| supply_shortfall_prob=0.5, | |
| alert_level=AlertLevel.ADVISORY, | |
| ) | |
| cfg = ForecastConfig(n_zones=2, max_steps=10) | |
| env = make_env(cfg) | |
| env.reset(seed=0) | |
| _, reward_early, _, _, _ = env.step(env.terminate_action) | |
| env.reset(seed=0) | |
| env.step(0) | |
| env.step(1) | |
| _, reward_late, _, _, _ = env.step(env.terminate_action) | |
| assert reward_late >= reward_early - 1e-6, ( | |
| f"Terminating after inspection should reward at least as well: " | |
| f"early={reward_early:.4f}, late={reward_late:.4f}" | |
| ) |