Skip to content

feat: add a JSD option to the distillation loss - #4026

Open
kashif wants to merge 4 commits into
NVIDIA-NeMo:mainfrom
kashif:feat/distillation-jsd-loss
Open

kashif wants to merge 4 commits into
NVIDIA-NeMo:mainfrom
kashif:feat/distillation-jsd-loss

Conversation

@kashif

@kashif kashif commented Sep 6, 2026

Copy link
Copy Markdown

What does this PR do ?

Adds a Jensen-Shannon divergence option to the distillation loss, alongside the existing forward/reverse/mixed KL choices.

Issues

The distillation loss currently only supports forward KL, reverse KL, or a fixed linear blend of the two. This adds kl_type="jsd" with a jsd_beta weight: beta=0 is forward KL, beta=1 is reverse KL, and 0.5 (default) gives the symmetric, bounded Jensen-Shannon divergence - a genuinely different loss shape than blending two KL terms, since both sides are measured against a shared mixture distribution instead of each other directly.

Usage

loss_fn:
    kl_type: "jsd"
    jsd_beta: 0.5

Before your PR is "Ready for review"

Pre checks:

  • Make sure you read and followed Contributor guidelines
  • Did you write any new necessary tests?
  • Did you run the unit tests and functional tests locally? Visit our Testing Guide for how to run tests
  • Did you add or update any necessary documentation? Visit our Document Development Guide for how to write, build and test the docs.

Additional Information

  • No GPU available in my environment to run the full test suite (the existing distillation tests are GPU-gated); I validated the math directly against the CPU loss function instead. Would appreciate a CI run / second pair of eyes on the new tests.

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 <kashif.rasul@gmail.com>
@kashif
kashif requested review from a team as code owners September 6, 2026 15:14
@copy-pr-bot

copy-pr-bot Bot commented Sep 6, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@svcnvidia-nemo-ci svcnvidia-nemo-ci added the waiting-on-maintainers Waiting on maintainers to respond label Sep 8, 2026
…-loss

# Conflicts:
#	nemo_rl/algorithms/loss/loss_functions.py
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 <kashif.rasul@gmail.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

community-request waiting-on-maintainers Waiting on maintainers to respond

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants