diff --git a/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py b/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py index b3b6b1b..7680404 100644 --- a/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py +++ b/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py @@ -181,8 +181,8 @@ class DreamWaQPPO(PPO): """ if not hasattr(self, '_adaboot_reward_buf'): self._adaboot_reward_buf = [] - dones = storage.dones # (T, E) bool - rewards = storage.rewards # (T, E) 每步奖励 + dones = storage.dones.squeeze(-1) # (T, E, 1) → (T, E) + rewards = storage.rewards.squeeze(-1) # (T, E, 1) → (T, E) if dones is None or rewards is None: return T, E = dones.shape @@ -197,7 +197,8 @@ class DreamWaQPPO(PPO): if dones[t, e].item(): # episode 结束:累计从 start 到 t 的奖励 ep_sum = rewards[start:t + 1, e].sum().item() - ep_rewards.append(ep_sum) + if abs(ep_sum) < 1e9: # 过滤 NaN/Inf + ep_rewards.append(ep_sum) start = t + 1 if len(ep_rewards) < 8: @@ -213,12 +214,14 @@ class DreamWaQPPO(PPO): if len(self._adaboot_reward_buf) < 50: return - # 论文公式: CV = σ / μ + # 论文公式: CV = σ / μ(NaN 防护) buf = self._adaboot_reward_buf mean_r = sum(buf) / len(buf) var_r = sum((r - mean_r) ** 2 for r in buf) / len(buf) std_r = var_r ** 0.5 cv = std_r / (mean_r + 1e-6) + if not (0 <= cv < 1e6): # NaN 或 Inf → 保持当前 prob + return # CV → bootstrap 概率(CV 高→不稳定→多用 GT;映射系数 5.0 可调) self.actor._adaboot_prob = max(0.0, min(1.0, cv * 5.0)) diff --git a/motrix_rl/src/motrix_rl/tasks/go1.py b/motrix_rl/src/motrix_rl/tasks/go1.py index fd5be29..48d086e 100644 --- a/motrix_rl/src/motrix_rl/tasks/go1.py +++ b/motrix_rl/src/motrix_rl/tasks/go1.py @@ -105,14 +105,14 @@ class rslrl: class Go1DreamWaQWalkRslrlPpo(RslrlCfg): """Go1 DreamWaQ walk — CENet VAE + 不对称特权观测。""" - num_envs: int = 2048 + num_envs: int = 4096 # 论文原值 def __post_init__(self): runner = self.runner # Runner 设置(严格对齐上游 LeggedRobotCfgPPO + Go1RoughCfgPPO) runner.seed = 5 # 上游 seed=5 - runner.max_iterations = 5000 # 2048 envs 需要更多迭代补偿采样量 + runner.max_iterations = 1000 # 论文原值 runner.num_steps_per_env = 24 runner.experiment_name = "go1_dreamwaq_walk" runner.save_interval = 50