diff --git a/src/coolrl_lost_cities/games/classic/ismcts/config.py b/src/coolrl_lost_cities/games/classic/ismcts/config.py index 8e3aeae..e08588d 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/config.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/config.py @@ -62,6 +62,11 @@ class TrainingConfig(StrictModel): interleave_max_batch: int = 64 num_workers: int = 1 worker_device: str = "cpu" + # Multiplier on the value-head MSE loss (already normalized by value_scale**2). + # Default 1.0 keeps current behavior; raising it (e.g. 50-100) makes the value + # head learn faster relative to policy loss. Useful when value_prediction_error + # is large but loss/value is tiny because of the normalization. + value_loss_weight: float = 1.0 @field_validator( "games_per_iter", diff --git a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py index 2f263b7..d9cf695 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py @@ -261,7 +261,8 @@ class IsMctsTrainer: policy_loss = -(pi * log_probs).sum(dim=-1).mean() v_scale = float(self.network.value_scale) value_loss = nn.functional.mse_loss(value_pred / v_scale, value_target / v_scale) - loss = policy_loss + value_loss + value_weight = float(self.config.training.value_loss_weight) + loss = policy_loss + value_weight * value_loss self.optimizer.zero_grad(set_to_none=True) loss.backward() if self.config.optimization.grad_clip > 0: