From aa17ea853e33d113933a4811cfc28608432d0f91 Mon Sep 17 00:00:00 2001 From: 8x54zj-m <8x54zj-m@motrixlab.local> Date: Tue, 30 Jun 2026 15:38:26 +0800 Subject: [PATCH] fix: VAE loss /num_mini_batches, KL use raw sum (align with upstream) --- .../src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) 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 bdb73db..6f9050d 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 @@ -192,17 +192,16 @@ class DreamWaQPPO(PPO): mse = nn.functional.mse_loss estimation_loss = mse(code_vel, vel_target) reconstruction_loss = mse(decode, obs_target) - # KL 散度:-0.5 * sum(1 + logvar - mean^2 - exp(logvar)) - # clamp logvar 防止 exp 溢出 + # KL 散度:上游用 raw sum(非 mean),然后除以 num_mini_batches logvar_latent = torch.clamp(logvar_latent, -20.0, 10.0) kl_loss = -0.5 * torch.sum( - 1 + logvar_latent - mean_latent.pow(2) - logvar_latent.exp(), dim=-1 - ).mean() + 1 + logvar_latent - mean_latent.pow(2) - logvar_latent.exp() + ) + # 与上游完全一致:除以 num_mini_batches autoenc_loss = ( estimation_loss + reconstruction_loss + self.vae_beta * kl_loss - ) - # 防止 NaN 传播 + ) / self.num_mini_batches if torch.isnan(autoenc_loss) or torch.isinf(autoenc_loss): return torch.tensor(0.0, device=obs_batch.device) return autoenc_loss