diff --git a/examples/configs/distillation_math.yaml b/examples/configs/distillation_math.yaml index b263dd57448..ba0718cdee2 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 d62335a0616..6e7b049fdec 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 import warnings from typing import TYPE_CHECKING, Any, NotRequired, Optional, TypedDict, TypeVar @@ -1828,6 +1829,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): @@ -1849,12 +1855,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, @@ -1869,8 +1877,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] @@ -1879,6 +1893,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 * ( @@ -1888,14 +1904,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 9820f2c7c79..e0cc83a37ed 100644 --- a/tests/unit/algorithms/test_loss_functions.py +++ b/tests/unit/algorithms/test_loss_functions.py @@ -2152,16 +2152,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) @@ -2190,7 +2199,23 @@ 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"]) +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): """Test different distillation loss settings.""" @@ -2200,6 +2225,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, ) ) @@ -2214,7 +2240,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)) @@ -2229,6 +2257,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) @@ -2382,6 +2413,73 @@ 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(device="cpu") + + 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, + ) + ) + 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 + ) + 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 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, 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) + + +def test_distillation_loss_jsd_gradient_flow(): + """Test gradient flow through the generalized JSD branch.""" + 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, _, _ = _run_distillation_loss(loss_fn, student_logits, data) + 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 0d47ffa4d25..f80bbf916d7 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: