DHDRL commited on
Commit
d96a958
·
verified ·
1 Parent(s): 899af8f

Update train_curriculum.py

Browse files
Files changed (1) hide show
  1. 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
- for i in range(0, len(pairs) - batch_size, batch_size):
 
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
- dynamics_model = TemporalDynamicsModel.load(str(model_path))
 
 
 
 
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 = "MultiInputPolicy"
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]))