Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion examples/configs/distillation_math.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
52 changes: 48 additions & 4 deletions nemo_rl/algorithms/loss/loss_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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):
Expand All @@ -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,
Expand All @@ -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]
Expand All @@ -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 * (
Expand All @@ -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]

Expand Down
114 changes: 106 additions & 8 deletions tests/unit/algorithms/test_loss_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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."""
Expand All @@ -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,
)
)
Expand All @@ -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))
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion tests/unit/reference_configs/distillation_math.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading