Cap DreamWaQ adaptive learning rate

This commit is contained in:
8x54zj-m
2026-07-22 19:58:19 +08:00
parent 944e91e28a
commit 509b3d635c
2 changed files with 9 additions and 1 deletions

View File

@@ -23,6 +23,7 @@ class DreamWaQPPO(PPO):
# VAE 损失权重(与上游 beta=1.0 一致) # VAE 损失权重(与上游 beta=1.0 一致)
vae_beta: float = 1.0 vae_beta: float = 1.0
max_learning_rate: float = 1e-3
@staticmethod @staticmethod
def construct_algorithm(obs: TensorDict, env, cfg: dict, device: str) -> "DreamWaQPPO": def construct_algorithm(obs: TensorDict, env, cfg: dict, device: str) -> "DreamWaQPPO":
@@ -133,7 +134,9 @@ class DreamWaQPPO(PPO):
if kl_mean > self.desired_kl * 2.0: if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5) self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0: elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5) self.learning_rate = min(
self.max_learning_rate, self.learning_rate * 1.5
)
for param_group in self.optimizer.param_groups: for param_group in self.optimizer.param_groups:
param_group['lr'] = self.learning_rate param_group['lr'] = self.learning_rate

View File

@@ -3,9 +3,14 @@ from types import SimpleNamespace
import numpy as np import numpy as np
import torch import torch
from motrix_rl.rslrl.torch.train.dreamwaq_ppo import DreamWaQPPO
from motrix_rl.rslrl.torch.wrap_vec_env import RslrlNpEnvWrap from motrix_rl.rslrl.torch.wrap_vec_env import RslrlNpEnvWrap
def test_dreamwaq_adaptive_learning_rate_is_capped_at_initial_rate():
assert DreamWaQPPO.max_learning_rate == 1e-3
def test_step_preserves_rewards_and_marks_only_truncations_as_timeouts(): def test_step_preserves_rewards_and_marks_only_truncations_as_timeouts():
state = SimpleNamespace( state = SimpleNamespace(
obs=np.zeros((2, 3), dtype=np.float32), obs=np.zeros((2, 3), dtype=np.float32),