From c00b9755efaa57642efdd213fa458b4e21c3d092 Mon Sep 17 00:00:00 2001 From: Yizhuo Li <143763751+Dodojordi@users.noreply.github.com> Date: Sun, 9 Aug 2026 07:08:28 +0000 Subject: [PATCH] fix: skip KL backward graph when coefficient is zero --- slime/backends/megatron_utils/loss.py | 18 ++++++++++++++---- 1 file changed, 14 insertions(+), 4 deletions(-) diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index adfa5d67e6..1908c888d3 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -1053,18 +1053,28 @@ def policy_loss_function( if args.use_kl_loss: ref_log_probs = batch["ref_log_probs"] ref_log_probs = torch.cat(ref_log_probs, dim=0) + kl_loss_coef = float(args.kl_loss_coef) + with_kl_grad = kl_loss_coef != 0.0 + + # Keep reference KL metrics when the coefficient is zero, but detach + # their inputs so the metric-only path does not retain an actor + # backward graph or contaminate policy gradients with non-finite KL. + kl_log_probs = log_probs if with_kl_grad else log_probs.detach() + kl_old_log_probs = old_log_probs if with_kl_grad else old_log_probs.detach() + kl_ref_log_probs = ref_log_probs if with_kl_grad else ref_log_probs.detach() importance_ratio = None if args.use_unbiased_kl: - importance_ratio = torch.exp(log_probs - old_log_probs) + importance_ratio = torch.exp(kl_log_probs - kl_old_log_probs) kl = compute_approx_kl( - log_probs, - ref_log_probs, + kl_log_probs, + kl_ref_log_probs, kl_loss_type=args.kl_loss_type, importance_ratio=importance_ratio, ) kl_loss = sum_of_sample_mean(kl) - loss = loss + args.kl_loss_coef * kl_loss + if with_kl_grad: + loss = loss + kl_loss_coef * kl_loss # make sure the gradient could backprop correctly. if log_probs.numel() == 0: