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
Update train_curriculum.py
Browse files- train_curriculum.py +11 -4
train_curriculum.py
CHANGED
|
@@ -47,11 +47,12 @@ except ImportError as _e:
|
|
| 47 |
BaseCallback = object
|
| 48 |
|
| 49 |
try:
|
| 50 |
-
from gru_weather_policy import create_gru_weather_policy_kwargs
|
| 51 |
_GRU_AVAILABLE = True
|
| 52 |
except ImportError:
|
| 53 |
_GRU_AVAILABLE = False
|
| 54 |
create_gru_weather_policy_kwargs = None
|
|
|
|
| 55 |
|
| 56 |
try:
|
| 57 |
from physics_dynamics import TemporalDynamicsModel, DynaRolloutBuffer, ZoneStateTensor
|
|
@@ -437,7 +438,8 @@ class DynaCallback(BaseCallback):
|
|
| 437 |
total_loss = 0.0
|
| 438 |
n_batches = 0
|
| 439 |
|
| 440 |
-
|
|
|
|
| 441 |
batch = pairs[i : i + batch_size]
|
| 442 |
curr_list = [p[0] for p in batch]
|
| 443 |
next_list = [p[1] for p in batch]
|
|
@@ -510,7 +512,11 @@ def _build_dyna_callback(
|
|
| 510 |
return None
|
| 511 |
|
| 512 |
try:
|
| 513 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 514 |
dynamics_model.eval()
|
| 515 |
dyna_buffer = DynaRolloutBuffer(
|
| 516 |
dynamics=dynamics_model,
|
|
@@ -636,9 +642,10 @@ def train_phase(
|
|
| 636 |
|
| 637 |
env = Monitor(make_weather_env(config))
|
| 638 |
|
|
|
|
| 639 |
if _GRU_AVAILABLE:
|
| 640 |
policy_kwargs = create_gru_weather_policy_kwargs(hidden_size=hidden_size)
|
| 641 |
-
policy =
|
| 642 |
logger.info("Using GRU policy (hidden_size=%d)", hidden_size)
|
| 643 |
else:
|
| 644 |
policy_kwargs = dict(net_arch=dict(pi=[128, 64], vf=[128, 64]))
|
|
|
|
| 47 |
BaseCallback = object
|
| 48 |
|
| 49 |
try:
|
| 50 |
+
from gru_weather_policy import create_gru_weather_policy_kwargs, get_equivariant_policy_class
|
| 51 |
_GRU_AVAILABLE = True
|
| 52 |
except ImportError:
|
| 53 |
_GRU_AVAILABLE = False
|
| 54 |
create_gru_weather_policy_kwargs = None
|
| 55 |
+
get_equivariant_policy_class = None
|
| 56 |
|
| 57 |
try:
|
| 58 |
from physics_dynamics import TemporalDynamicsModel, DynaRolloutBuffer, ZoneStateTensor
|
|
|
|
| 438 |
total_loss = 0.0
|
| 439 |
n_batches = 0
|
| 440 |
|
| 441 |
+
# FIX: iterate over all pairs, including the final partial batch
|
| 442 |
+
for i in range(0, len(pairs), batch_size):
|
| 443 |
batch = pairs[i : i + batch_size]
|
| 444 |
curr_list = [p[0] for p in batch]
|
| 445 |
next_list = [p[1] for p in batch]
|
|
|
|
| 512 |
return None
|
| 513 |
|
| 514 |
try:
|
| 515 |
+
# FIX: load onto the same device as training to avoid CPU/CUDA mismatch
|
| 516 |
+
import torch as _torch
|
| 517 |
+
dynamics_model = TemporalDynamicsModel.load(
|
| 518 |
+
str(model_path), device=_torch.device(device)
|
| 519 |
+
)
|
| 520 |
dynamics_model.eval()
|
| 521 |
dyna_buffer = DynaRolloutBuffer(
|
| 522 |
dynamics=dynamics_model,
|
|
|
|
| 642 |
|
| 643 |
env = Monitor(make_weather_env(config))
|
| 644 |
|
| 645 |
+
# FIX: use the zone-equivariant policy class instead of the generic string
|
| 646 |
if _GRU_AVAILABLE:
|
| 647 |
policy_kwargs = create_gru_weather_policy_kwargs(hidden_size=hidden_size)
|
| 648 |
+
policy = get_equivariant_policy_class()
|
| 649 |
logger.info("Using GRU policy (hidden_size=%d)", hidden_size)
|
| 650 |
else:
|
| 651 |
policy_kwargs = dict(net_arch=dict(pi=[128, 64], vf=[128, 64]))
|