From e577b50170156a879199f1b562e70db3d66c1c9e Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Sun, 6 Sep 2026 17:14:24 +0200 Subject: [PATCH 1/2] feat: add a JSD option to the distillation loss Adds kl_type=jsd to DistillationLossFn, interpolating between forward and reverse KL via a beta weight (0.5 by default gives the symmetric Jensen-Shannon divergence). Mirrors what forward/reverse already do for the top-k correction term. Signed-off-by: Kashif Rasul --- examples/configs/distillation_math.yaml | 3 +- nemo_rl/algorithms/loss/loss_functions.py | 52 ++++++++- tests/unit/algorithms/test_loss_functions.py | 105 +++++++++++++++++- .../reference_configs/distillation_math.yaml | 3 +- 4 files changed, 155 insertions(+), 8 deletions(-) diff --git a/examples/configs/distillation_math.yaml b/examples/configs/distillation_math.yaml index c288b50999c..300a5fd78d2 100644 --- a/examples/configs/distillation_math.yaml +++ b/examples/configs/distillation_math.yaml @@ -14,8 +14,9 @@ distillation: seed: 42 loss_fn: - kl_type: "mixed" # forward, reverse, mixed + kl_type: "mixed" # forward, reverse, mixed, jsd mixed_kl_weight: 0.5 # when kl_type is "mixed", this is the weight of the forward KL + jsd_beta: 0.5 # when kl_type is "jsd", interpolates forward KL (0) to reverse KL (1); 0.5 is the symmetric JSD zero_outside_topk: false # zero out the teacher logits outside the top k when calculate forward KL loss checkpointing: diff --git a/nemo_rl/algorithms/loss/loss_functions.py b/nemo_rl/algorithms/loss/loss_functions.py index 9e9cc872c0b..a44f26f251a 100755 --- a/nemo_rl/algorithms/loss/loss_functions.py +++ b/nemo_rl/algorithms/loss/loss_functions.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import math from typing import Any, NotRequired, Optional, TypedDict, TypeVar import torch @@ -1111,6 +1112,11 @@ class DistillationLossConfig(BaseModel, extra="allow"): kl_type: str = "mixed" mixed_kl_weight: float = 0.5 zero_outside_topk: bool = False + # Interpolation coefficient for kl_type="jsd" (generalized Jensen-Shannon + # divergence, see Eq. 1 of https://huggingface.co/papers/2306.13649): + # beta=0 reduces to forward KL, beta=1 to reverse KL, beta=0.5 (default) + # gives the symmetric JSD. Unused for other kl_type values. + jsd_beta: float = 0.5 class DistillationLossDataDict(TypedDict): @@ -1132,12 +1138,14 @@ def __init__(self, cfg: DistillationLossConfig): self.kl_type = cfg.kl_type self.mixed_kl_weight = cfg.mixed_kl_weight self.zero_outside_topk = cfg.zero_outside_topk + self.jsd_beta = cfg.jsd_beta self.log_infinitesimal = -100 - assert self.kl_type in ["forward", "reverse", "mixed"], "Invalid KL type" + assert self.kl_type in ["forward", "reverse", "mixed", "jsd"], "Invalid KL type" assert self.mixed_kl_weight >= 0 and self.mixed_kl_weight <= 1, ( "Invalid mixed KL weight" ) + assert self.jsd_beta >= 0 and self.jsd_beta <= 1, "Invalid JSD beta" def __call__( self, @@ -1152,8 +1160,14 @@ def __call__( student_probs = student_topk_logprobs.exp() # [B, S-1, k] teacher_probs = teacher_topk_logprobs.exp() # [B, S-1, k] + # jsd is forward KL at beta=0 and reverse KL at beta=1, so the topk + # correction below is skipped/unscaled the same way forward/reverse are. + is_forward_like = self.kl_type == "forward" or ( + self.kl_type == "jsd" and self.jsd_beta <= 0.0 + ) + loss_correction_term = torch.zeros_like(student_probs[..., 0]) # [B, S-1] - if self.zero_outside_topk and self.kl_type != "forward": + if self.zero_outside_topk and not is_forward_like: H_rest = H_all - (student_probs * student_topk_logprobs).sum(-1) P_rest = 1 - (student_probs.sum(-1)) # The entropy and prob of the rest of the tokens [B, S-1] @@ -1162,6 +1176,8 @@ def __call__( loss_correction_term = loss_correction_term * ( 1.0 - self.mixed_kl_weight ) + elif self.kl_type == "jsd" and self.jsd_beta < 1.0: + loss_correction_term = loss_correction_term * (1.0 - self.jsd_beta) if self.kl_type == "forward": per_token_kl = teacher_probs * ( @@ -1171,14 +1187,42 @@ def __call__( per_token_kl = student_probs * ( student_topk_logprobs - teacher_topk_logprobs ) - else: - # mixed KL + elif self.kl_type == "mixed": kl_forward = teacher_probs * (teacher_topk_logprobs - student_topk_logprobs) kl_reverse = student_probs * (student_topk_logprobs - teacher_topk_logprobs) per_token_kl = ( self.mixed_kl_weight * kl_forward + (1.0 - self.mixed_kl_weight) * kl_reverse ) + else: + # Generalized Jensen-Shannon divergence (Eq. 1 of + # https://huggingface.co/papers/2306.13649): beta * KL(teacher || M) + # + (1 - beta) * KL(student || M), M = beta * teacher + (1 - beta) * student. + # beta=0/1 degenerate to plain forward/reverse KL, since M collapses + # onto student/teacher and the mixture-weighted term above vanishes. + beta = self.jsd_beta + if beta <= 0.0: + per_token_kl = teacher_probs * ( + teacher_topk_logprobs - student_topk_logprobs + ) + elif beta >= 1.0: + per_token_kl = student_probs * ( + student_topk_logprobs - teacher_topk_logprobs + ) + else: + mixture_topk_logprobs = torch.logaddexp( + student_topk_logprobs + math.log1p(-beta), + teacher_topk_logprobs + math.log(beta), + ) + kl_teacher_to_mixture = teacher_probs * ( + teacher_topk_logprobs - mixture_topk_logprobs + ) + kl_student_to_mixture = student_probs * ( + student_topk_logprobs - mixture_topk_logprobs + ) + per_token_kl = ( + beta * kl_teacher_to_mixture + (1.0 - beta) * kl_student_to_mixture + ) per_token_kl = per_token_kl.sum(dim=-1) + loss_correction_term # [B, S-1] diff --git a/tests/unit/algorithms/test_loss_functions.py b/tests/unit/algorithms/test_loss_functions.py index 3b672ac1f42..39e774f5999 100644 --- a/tests/unit/algorithms/test_loss_functions.py +++ b/tests/unit/algorithms/test_loss_functions.py @@ -2191,7 +2191,7 @@ def setup_distillation_test_data(batch_size=2, seq_len=4, vocab_size=8, topk=64) return data, student_logits -@pytest.mark.parametrize("kl_type", ["forward", "reverse", "mixed"]) +@pytest.mark.parametrize("kl_type", ["forward", "reverse", "mixed", "jsd"]) @pytest.mark.parametrize("zero_outside_topk", [True, False]) def test_distillation_loss_different_settings(kl_type, zero_outside_topk): """Test different distillation loss settings.""" @@ -2201,6 +2201,7 @@ def test_distillation_loss_different_settings(kl_type, zero_outside_topk): DistillationLossConfig( kl_type=kl_type, mixed_kl_weight=0.3, + jsd_beta=0.3, zero_outside_topk=zero_outside_topk, ) ) @@ -2215,7 +2216,9 @@ def test_distillation_loss_different_settings(kl_type, zero_outside_topk): **loss_input, ) - # Verify loss + # Verify loss. jsd's exact value isn't pinned like the others above (no + # hand-derivable closed form for random top-k data) -- its correctness is + # covered by the boundary/symmetry/gradient tests below instead. if zero_outside_topk: if kl_type == "forward": assert torch.allclose(loss, torch.tensor(-0.9636520743370056)) @@ -2230,6 +2233,9 @@ def test_distillation_loss_different_settings(kl_type, zero_outside_topk): assert torch.allclose(loss, torch.tensor(0.5811167359352112)) elif kl_type == "mixed": assert torch.allclose(loss, torch.tensor(0.5802732110023499)) + if kl_type == "jsd": + assert not torch.isnan(loss) + assert not torch.isinf(loss) # Verify metrics dictionary assert isinstance(metrics, dict) @@ -2383,6 +2389,101 @@ def test_distillation_loss_edge_cases(): assert not torch.isinf(loss) +@pytest.mark.parametrize("zero_outside_topk", [True, False]) +def test_distillation_loss_jsd_matches_forward_reverse_at_boundaries( + zero_outside_topk, +): + """kl_type="jsd" must reduce to forward KL at beta=0 and reverse KL at beta=1.""" + data, student_logits = setup_distillation_test_data() + + def run(kl_type, jsd_beta=0.5): + loss_fn = DistillationLossFn( + DistillationLossConfig( + kl_type=kl_type, + jsd_beta=jsd_beta, + zero_outside_topk=zero_outside_topk, + ) + ) + loss_input, loss_data = prepare_loss_input(student_logits, data, loss_fn) + loss, _ = loss_fn( + data=loss_data, + global_valid_seqs=torch.sum(loss_data["sample_mask"]), + global_valid_toks=torch.sum( + loss_data["sample_mask"].unsqueeze(-1) * loss_data["token_mask"] + ), + **loss_input, + ) + return loss + + torch.testing.assert_close( + run("jsd", jsd_beta=0.0), run("forward"), atol=1e-5, rtol=1e-4 + ) + torch.testing.assert_close( + run("jsd", jsd_beta=1.0), run("reverse"), atol=1e-5, rtol=1e-4 + ) + + +def test_distillation_loss_jsd_symmetric_beta(): + """Generalized JSD at beta=0.5 is symmetric: swapping student <-> teacher + top-k distributions must give the same loss (forward/reverse KL are not). + """ + data, student_logits = setup_distillation_test_data() + + loss_fn = DistillationLossFn( + DistillationLossConfig(kl_type="jsd", jsd_beta=0.5, zero_outside_topk=False) + ) + loss_input, data = prepare_loss_input(student_logits, data, loss_fn) + loss, _ = loss_fn( + data=data, + global_valid_seqs=torch.sum(data["sample_mask"]), + global_valid_toks=torch.sum( + data["sample_mask"].unsqueeze(-1) * data["token_mask"] + ), + **loss_input, + ) + + swapped_input = dict(loss_input) + swapped_input["student_topk_logprobs"], swapped_input["teacher_topk_logprobs"] = ( + loss_input["teacher_topk_logprobs"], + loss_input["student_topk_logprobs"], + ) + swapped_loss, _ = loss_fn( + data=data, + global_valid_seqs=torch.sum(data["sample_mask"]), + global_valid_toks=torch.sum( + data["sample_mask"].unsqueeze(-1) * data["token_mask"] + ), + **swapped_input, + ) + + torch.testing.assert_close(loss, swapped_loss, atol=1e-5, rtol=1e-4) + + +def test_distillation_loss_jsd_gradient_flow(): + """Test gradient flow through the generalized JSD branch.""" + data, student_logits = setup_distillation_test_data() + student_logits.requires_grad_(True) + + loss_fn = DistillationLossFn( + DistillationLossConfig(kl_type="jsd", jsd_beta=0.5, zero_outside_topk=False) + ) + loss_input, data = prepare_loss_input(student_logits, data, loss_fn) + loss, _ = loss_fn( + data=data, + global_valid_seqs=torch.sum(data["sample_mask"]), + global_valid_toks=torch.sum( + data["sample_mask"].unsqueeze(-1) * data["token_mask"] + ), + **loss_input, + ) + loss.backward() + + assert student_logits.grad is not None + assert not torch.allclose( + student_logits.grad, torch.zeros_like(student_logits.grad) + ) + + def test_distillation_loss_fn_initialization(): """Test DistillationLossFn initialization.""" # Test with default values diff --git a/tests/unit/reference_configs/distillation_math.yaml b/tests/unit/reference_configs/distillation_math.yaml index 73455bb41ea..c2d0ad61baf 100644 --- a/tests/unit/reference_configs/distillation_math.yaml +++ b/tests/unit/reference_configs/distillation_math.yaml @@ -14,8 +14,9 @@ distillation: seed: 42 loss_fn: - kl_type: "mixed" # forward, reverse, mixed + kl_type: "mixed" # forward, reverse, mixed, jsd mixed_kl_weight: 0.5 # when kl_type is "mixed", this is the weight of the forward KL + jsd_beta: 0.5 # when kl_type is "jsd", interpolates forward KL (0) to reverse KL (1); 0.5 is the symmetric JSD zero_outside_topk: false # zero out the teacher logits outside the top k when calculate forward KL loss checkpointing: From e73cb92d276975c2447da40c96925c06c8668fb8 Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Tue, 22 Sep 2026 18:20:16 +0200 Subject: [PATCH 2/2] test: run the JSD tests on CPU and share their setup The three JSD tests inherited a CUDA-only helper, so they were skipped everywhere without a GPU and never actually ran. The JSD branch is pure tensor math, so setup_distillation_test_data now takes an optional device (default unchanged, still CUDA-or-skip) and these pass device=cpu. Also collapses the repeated prepare/invoke block in those three tests into one helper. The global_valid_* idiom appears 45 times in this file, so the other call sites are left alone rather than half-migrated. Signed-off-by: Kashif Rasul --- tests/unit/algorithms/test_loss_functions.py | 103 +++++++++---------- 1 file changed, 50 insertions(+), 53 deletions(-) diff --git a/tests/unit/algorithms/test_loss_functions.py b/tests/unit/algorithms/test_loss_functions.py index d27ae5477e6..528ee151f93 100644 --- a/tests/unit/algorithms/test_loss_functions.py +++ b/tests/unit/algorithms/test_loss_functions.py @@ -2153,16 +2153,25 @@ def test_clipped_pg_loss_gspo_importance_sampling_correction(): torch.testing.assert_close(actual_loss, expected_actor_loss, atol=1e-4, rtol=1e-3) -def setup_distillation_test_data(batch_size=2, seq_len=4, vocab_size=8, topk=64): - """Setup test data for distillation loss function tests.""" - if not torch.cuda.is_available(): - pytest.skip("No GPU available") +def setup_distillation_test_data( + batch_size=2, seq_len=4, vocab_size=8, topk=64, device=None +): + """Setup test data for distillation loss function tests. - device = "cuda" + Args: + device: Where to place the tensors. ``None`` keeps the historical + behaviour of requiring CUDA and skipping without it. Pass "cpu" + for branches that are pure tensor math and need no GPU. + """ + if device is None: + if not torch.cuda.is_available(): + pytest.skip("No GPU available") + device = "cuda" # Set seed for reproducibility torch.manual_seed(42) - torch.cuda.manual_seed_all(42) + if device == "cuda": + torch.cuda.manual_seed_all(42) # Create input data input_ids = torch.randint(0, vocab_size, (batch_size, seq_len), device=device) @@ -2191,6 +2200,22 @@ def setup_distillation_test_data(batch_size=2, seq_len=4, vocab_size=8, topk=64) return data, student_logits +def _run_distillation_loss(loss_fn, student_logits, data, loss_input_overrides=None): + """Prepare inputs and invoke a distillation loss, returning the scalar loss.""" + loss_input, loss_data = prepare_loss_input(student_logits, data, loss_fn) + if loss_input_overrides: + loss_input = {**loss_input, **loss_input_overrides} + loss, _ = loss_fn( + data=loss_data, + global_valid_seqs=torch.sum(loss_data["sample_mask"]), + global_valid_toks=torch.sum( + loss_data["sample_mask"].unsqueeze(-1) * loss_data["token_mask"] + ), + **loss_input, + ) + return loss, loss_input, loss_data + + @pytest.mark.parametrize("kl_type", ["forward", "reverse", "mixed", "jsd"]) @pytest.mark.parametrize("zero_outside_topk", [True, False]) def test_distillation_loss_different_settings(kl_type, zero_outside_topk): @@ -2394,7 +2419,7 @@ def test_distillation_loss_jsd_matches_forward_reverse_at_boundaries( zero_outside_topk, ): """kl_type="jsd" must reduce to forward KL at beta=0 and reverse KL at beta=1.""" - data, student_logits = setup_distillation_test_data() + data, student_logits = setup_distillation_test_data(device="cpu") def run(kl_type, jsd_beta=0.5): loss_fn = DistillationLossFn( @@ -2404,16 +2429,7 @@ def run(kl_type, jsd_beta=0.5): zero_outside_topk=zero_outside_topk, ) ) - loss_input, loss_data = prepare_loss_input(student_logits, data, loss_fn) - loss, _ = loss_fn( - data=loss_data, - global_valid_seqs=torch.sum(loss_data["sample_mask"]), - global_valid_toks=torch.sum( - loss_data["sample_mask"].unsqueeze(-1) * loss_data["token_mask"] - ), - **loss_input, - ) - return loss + return _run_distillation_loss(loss_fn, student_logits, data)[0] torch.testing.assert_close( run("jsd", jsd_beta=0.0), run("forward"), atol=1e-5, rtol=1e-4 @@ -2424,36 +2440,25 @@ def run(kl_type, jsd_beta=0.5): def test_distillation_loss_jsd_symmetric_beta(): - """Generalized JSD at beta=0.5 is symmetric: swapping student <-> teacher - top-k distributions must give the same loss (forward/reverse KL are not). - """ - data, student_logits = setup_distillation_test_data() + """Generalized JSD at beta=0.5 is symmetric. + Swapping the student and teacher top-k distributions must give the same + loss, which forward and reverse KL do not. + """ + data, student_logits = setup_distillation_test_data(device="cpu") loss_fn = DistillationLossFn( DistillationLossConfig(kl_type="jsd", jsd_beta=0.5, zero_outside_topk=False) ) - loss_input, data = prepare_loss_input(student_logits, data, loss_fn) - loss, _ = loss_fn( - data=data, - global_valid_seqs=torch.sum(data["sample_mask"]), - global_valid_toks=torch.sum( - data["sample_mask"].unsqueeze(-1) * data["token_mask"] - ), - **loss_input, - ) - swapped_input = dict(loss_input) - swapped_input["student_topk_logprobs"], swapped_input["teacher_topk_logprobs"] = ( - loss_input["teacher_topk_logprobs"], - loss_input["student_topk_logprobs"], - ) - swapped_loss, _ = loss_fn( - data=data, - global_valid_seqs=torch.sum(data["sample_mask"]), - global_valid_toks=torch.sum( - data["sample_mask"].unsqueeze(-1) * data["token_mask"] - ), - **swapped_input, + loss, loss_input, _ = _run_distillation_loss(loss_fn, student_logits, data) + swapped_loss, _, _ = _run_distillation_loss( + loss_fn, + student_logits, + data, + loss_input_overrides={ + "student_topk_logprobs": loss_input["teacher_topk_logprobs"], + "teacher_topk_logprobs": loss_input["student_topk_logprobs"], + }, ) torch.testing.assert_close(loss, swapped_loss, atol=1e-5, rtol=1e-4) @@ -2461,21 +2466,13 @@ def test_distillation_loss_jsd_symmetric_beta(): def test_distillation_loss_jsd_gradient_flow(): """Test gradient flow through the generalized JSD branch.""" - data, student_logits = setup_distillation_test_data() + data, student_logits = setup_distillation_test_data(device="cpu") student_logits.requires_grad_(True) - loss_fn = DistillationLossFn( DistillationLossConfig(kl_type="jsd", jsd_beta=0.5, zero_outside_topk=False) ) - loss_input, data = prepare_loss_input(student_logits, data, loss_fn) - loss, _ = loss_fn( - data=data, - global_valid_seqs=torch.sum(data["sample_mask"]), - global_valid_toks=torch.sum( - data["sample_mask"].unsqueeze(-1) * data["token_mask"] - ), - **loss_input, - ) + + loss, _, _ = _run_distillation_loss(loss_fn, student_logits, data) loss.backward() assert student_logits.grad is not None