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
| """ | |
| gru_weather_policy.py | |
| ==================== | |
| Custom GRU feature extractor + zone-equivariant policy head for | |
| stable-baselines3 MaskablePPO. | |
| Architecture | |
| ------------ | |
| 1. GRUWeatherFeaturesExtractor: | |
| - Per-zone GRU over forecast_precip[zone, :] (14-day horizon) | |
| - Per-zone MLP over uncertainty + belief | |
| - Concatenate → zone-level feature vector | |
| - Stash zone scores and terminate logit for the policy head | |
| 2. ZoneEquivariantMaskablePolicy: | |
| - Overrides _get_action_dist_from_latent | |
| - Reads stashed zone scores + terminate logit from the extractor | |
| - Returns a Categorical distribution directly | |
| - This makes the policy permutation-equivariant across zones | |
| (inspecting zone 0 then zone 1 is the same as zone 1 then zone 0) | |
| WARNING | |
| ------- | |
| When using ZoneEquivariantMaskablePolicy, the ``net_arch`` pi layers are | |
| instantiated by SB3 inside the MLP extractor but are NEVER called at | |
| inference time because ``_get_action_dist_from_latent`` bypasses | |
| ``latent_pi`` entirely. The policy capacity is entirely in the extractor. | |
| The vf head still uses the ``net_arch`` vf layers normally. | |
| Dependencies | |
| ------------ | |
| pip install stable-baselines3 sb3-contrib torch | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| from typing import Any, Dict, List, Optional, Tuple, Type | |
| import gymnasium as gym | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| logger = logging.getLogger(__name__) | |
| # --------------------------------------------------------------------------- | |
| # SB3 availability | |
| # --------------------------------------------------------------------------- | |
| try: | |
| from stable_baselines3.common.torch_layers import BaseFeaturesExtractor | |
| from stable_baselines3.common.policies import MultiInputActorCriticPolicy | |
| _SB3_AVAILABLE = True | |
| except ImportError: | |
| _SB3_AVAILABLE = False | |
| BaseFeaturesExtractor = object # type: ignore[assignment,misc] | |
| MultiInputActorCriticPolicy = object # type: ignore[assignment,misc] | |
| try: | |
| from sb3_contrib.common.maskable.policies import ( | |
| MaskableMultiInputActorCriticPolicy, | |
| ) | |
| _MASKABLE_AVAILABLE = True | |
| except ImportError: | |
| _MASKABLE_AVAILABLE = False | |
| MaskableMultiInputActorCriticPolicy = object # type: ignore[assignment,misc] | |
| # --------------------------------------------------------------------------- | |
| # GRU feature extractor | |
| # --------------------------------------------------------------------------- | |
| class GRUWeatherFeaturesExtractor(BaseFeaturesExtractor): | |
| def __init__( | |
| self, | |
| observation_space: gym.spaces.Dict, | |
| features_dim: int = 128, | |
| hidden_size: int = 64, | |
| ): | |
| super().__init__(observation_space, features_dim=features_dim) | |
| self.hidden_size = hidden_size | |
| self._observation_space = observation_space | |
| precip_space = observation_space.spaces["forecast_precip"] | |
| self.n_zones = int(precip_space.shape[0]) | |
| self.horizon_days = int(precip_space.shape[1]) | |
| self.precip_gru = nn.GRU( | |
| input_size=1, | |
| hidden_size=hidden_size, | |
| num_layers=1, | |
| batch_first=True, | |
| ) | |
| self.static_mlp = nn.Sequential( | |
| nn.Linear(2, hidden_size), | |
| nn.Tanh(), | |
| ) | |
| self.zone_score = nn.Sequential( | |
| nn.Linear(hidden_size * 2, 64), | |
| nn.Tanh(), | |
| nn.Linear(64, 1), | |
| ) | |
| self.terminate_logit = nn.Linear(hidden_size * 2, 1) | |
| self.value_head = nn.Sequential( | |
| nn.Linear(self.n_zones * hidden_size * 2, 128), | |
| nn.Tanh(), | |
| nn.Linear(128, 1), | |
| ) | |
| def forward(self, observations: Dict[str, torch.Tensor]) -> torch.Tensor: | |
| precip = observations["forecast_precip"] | |
| batch_size = precip.shape[0] | |
| uncertainty = observations["forecast_uncertainty"] | |
| belief = observations["zone_belief"] | |
| precip_reshaped = precip.reshape(batch_size * self.n_zones, self.horizon_days, 1) | |
| _, gru_hidden = self.precip_gru(precip_reshaped) # [1, batch*n_zones, hidden_size] | |
| gru_features = gru_hidden.squeeze(0) # [batch*n_zones, hidden_size] | |
| static_input = torch.stack([uncertainty, belief], dim=-1) # [batch, n_zones, 2] | |
| static_input = static_input.reshape(batch_size * self.n_zones, 2) | |
| static_features = self.static_mlp(static_input) # [batch*n_zones, hidden_size] | |
| zone_features = torch.cat([gru_features, static_features], dim=-1) | |
| zone_scores = self.zone_score(zone_features).squeeze(-1) # [batch*n_zones] | |
| self._last_zone_scores = zone_scores.reshape(batch_size, self.n_zones) | |
| self._last_terminate_logit = self.terminate_logit(zone_features).squeeze(-1) # [batch*n_zones] | |
| self._last_terminate_logit = self._last_terminate_logit.reshape(batch_size, self.n_zones)[:, 0] | |
| global_features = zone_features.reshape(batch_size, self.n_zones, -1) | |
| global_features = global_features.reshape(batch_size, -1) | |
| return global_features | |
| def get_value(self, latent_vf: torch.Tensor) -> torch.Tensor: | |
| return self.value_head(latent_vf) | |
| # --------------------------------------------------------------------------- | |
| # Zone-equivariant policy head | |
| # --------------------------------------------------------------------------- | |
| class ZoneEquivariantMaskablePolicy(MaskableMultiInputActorCriticPolicy): | |
| def __init__( | |
| self, | |
| observation_space: gym.spaces.Dict, | |
| action_space: gym.spaces.Discrete, | |
| lr_schedule, | |
| net_arch: Optional[List[int]] = None, | |
| activation_fn: Type[nn.Module] = nn.Tanh, | |
| *args, | |
| **kwargs, | |
| ): | |
| super().__init__( | |
| observation_space, | |
| action_space, | |
| lr_schedule, | |
| net_arch=net_arch, | |
| activation_fn=activation_fn, | |
| *args, | |
| **kwargs, | |
| ) | |
| def _get_action_dist_from_latent(self, latent_pi: torch.Tensor) -> Any: | |
| features_extractor = self.features_extractor | |
| assert isinstance(features_extractor, GRUWeatherFeaturesExtractor) | |
| zone_scores = features_extractor._last_zone_scores # [batch, n_zones] | |
| terminate_logit = features_extractor._last_terminate_logit # [batch] | |
| logits = torch.cat([ | |
| zone_scores, | |
| terminate_logit.unsqueeze(-1), | |
| ], dim=-1) | |
| return self.action_dist.proba_distribution(action_logits=logits) | |
| # --------------------------------------------------------------------------- | |
| # Factory | |
| # --------------------------------------------------------------------------- | |
| def create_gru_weather_policy_kwargs( | |
| hidden_size: int = 64, | |
| features_dim: int = 128, | |
| ) -> Dict[str, Any]: | |
| if not _SB3_AVAILABLE: | |
| raise ImportError( | |
| "stable-baselines3 not installed. " | |
| "Run: pip install stable-baselines3" | |
| ) | |
| return { | |
| "features_extractor_class": GRUWeatherFeaturesExtractor, | |
| "features_extractor_kwargs": { | |
| "features_dim": features_dim, | |
| "hidden_size": hidden_size, | |
| }, | |
| "net_arch": dict(pi=[128, 64], vf=[128, 64]), | |
| } | |
| def get_equivariant_policy_class() -> Type[MaskableMultiInputActorCriticPolicy]: | |
| if not _MASKABLE_AVAILABLE: | |
| raise ImportError( | |
| "sb3-contrib not installed. " | |
| "Run: pip install sb3-contrib" | |
| ) | |
| return ZoneEquivariantMaskablePolicy | |
| # --------------------------------------------------------------------------- | |
| # Self-test | |
| # --------------------------------------------------------------------------- | |
| def _self_test() -> None: | |
| import numpy as np | |
| print("gru_weather_policy.py self-test") | |
| if not _SB3_AVAILABLE: | |
| print(" SKIP: stable-baselines3 not installed") | |
| return | |
| n_zones = 3 | |
| horizon_days = 14 | |
| obs_space = gym.spaces.Dict({ | |
| "forecast_precip": gym.spaces.Box( | |
| low=0, high=500, shape=(n_zones, horizon_days), dtype=np.float32 | |
| ), | |
| "forecast_uncertainty": gym.spaces.Box( | |
| low=0, high=1, shape=(n_zones,), dtype=np.float32 | |
| ), | |
| "zone_belief": gym.spaces.Box( | |
| low=0, high=1, shape=(n_zones,), dtype=np.float32 | |
| ), | |
| }) | |
| action_space = gym.spaces.Discrete(n_zones + 1) | |
| extractor = GRUWeatherFeaturesExtractor( | |
| observation_space=obs_space, | |
| features_dim=128, | |
| hidden_size=64, | |
| ) | |
| batch_size = 2 | |
| obs = { | |
| "forecast_precip": torch.randn(batch_size, n_zones, horizon_days), | |
| "forecast_uncertainty": torch.rand(batch_size, n_zones), | |
| "zone_belief": torch.rand(batch_size, n_zones), | |
| } | |
| features = extractor(obs) | |
| assert features.shape == (batch_size, n_zones * 64 * 2) | |
| assert hasattr(extractor, "_last_zone_scores") | |
| assert extractor._last_zone_scores.shape == (batch_size, n_zones) | |
| assert hasattr(extractor, "_last_terminate_logit") | |
| assert extractor._last_terminate_logit.shape == (batch_size,) | |
| print(" Feature extraction OK") | |
| if _MASKABLE_AVAILABLE: | |
| policy_class = get_equivariant_policy_class() | |
| assert policy_class is ZoneEquivariantMaskablePolicy | |
| print(" Policy class OK") | |
| kwargs = create_gru_weather_policy_kwargs(hidden_size=64) | |
| assert kwargs["features_extractor_class"] is GRUWeatherFeaturesExtractor | |
| assert kwargs["features_extractor_kwargs"]["hidden_size"] == 64 | |
| print(" Policy kwargs OK") | |
| print("All gru_weather_policy self-tests passed.") | |
| if __name__ == "__main__": | |
| _self_test() |