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
1 change: 1 addition & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -113,6 +113,7 @@ Changelog

**Bug Fixes**

- Nested ``only_student_forward()`` and ``only_teacher_forward()`` contexts now preserve the enclosing execution mode, including when an inner context is disabled or raises an exception.
- Fix Megatron unified HF export of MoE models with grouped-GEMM experts when only the experts are quantized (e.g. ``nvfp4_experts_only-*`` recipes): ``hf_quant_config.json`` and the ``quantization_config`` in ``config.json`` were not written, so the quantized experts were served as unquantized weights. Re-export such checkpoints.
- Fix DFlash conversion on NoPE targets whose config leaves ``rope_theta`` unset.
- Fix offline DFlash training failing to reconstruct the target logits when the captured hidden states are stored in a different dtype than the target's weights.
Expand Down
7 changes: 5 additions & 2 deletions modelopt/torch/distill/distillation_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,22 +166,25 @@ def hide_loss_modules(self, enable=True):
@contextmanager
def only_teacher_forward(self, enable=True):
"""Context manager to temporarily disable forward passes on the student model."""
# restore the previous mode so nested contexts keep working
only_teacher_fwd = self._only_teacher_fwd
if enable:
self._only_teacher_fwd = True
try:
yield
finally:
self._only_teacher_fwd = False
self._only_teacher_fwd = only_teacher_fwd

@contextmanager
def only_student_forward(self, enable=True):
"""Context manager to temporarily run forward passes only on the student model."""
only_student_fwd = self._only_student_fwd
if enable:
self._only_student_fwd = True
try:
yield
finally:
self._only_student_fwd = False
self._only_student_fwd = only_student_fwd

def train(self, mode: bool = True):
"""Override to prevent warnings of stored intermediate outputs in future forwards."""
Expand Down
69 changes: 69 additions & 0 deletions tests/unit/torch/distill/test_distill.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@

import inspect
import warnings
from contextlib import nullcontext

import pytest
import torch
Expand All @@ -26,6 +27,17 @@
import modelopt.torch.opt as mto


class LinearModel(nn.Module):
def __init__(self, scale):
super().__init__()
self.linear = nn.Linear(2, 2, bias=False) #Wx
with torch.no_grad():
self.linear.weight.copy_(scale * torch.eye(2))

def forward(self, inputs):
return self.linear(inputs)


def get_input_tensor():
"""Dummy input tensor."""
return torch.rand(2, 3, 112, 112)
Expand Down Expand Up @@ -310,6 +322,63 @@ def test_student_fwd_only(distillation_model):
assert distillation_model._intermediate_output is not None


@pytest.mark.parametrize("context_name", ["only_student_forward", "only_teacher_forward"])
@pytest.mark.parametrize("inner_enabled", [False, True])
@pytest.mark.parametrize("raise_in_inner", [False, True])
def test_nested_forward_contexts(context_name, inner_enabled, raise_in_inner):
student, teacher = LinearModel(1), LinearModel(2)
model = mtd.convert(
student, mode=[("kd_loss", {"teacher_model": teacher, "criterion": nn.MSELoss()})]
)
context = getattr(model, context_name)
inputs = torch.ones(1, 2)
calls = {"student": 0, "teacher": 0}

def count_student(module, inputs):
calls["student"] += 1

def count_teacher(module, inputs):
calls["teacher"] += 1

model.linear.register_forward_pre_hook(count_student)
teacher.linear.register_forward_pre_hook(count_teacher)
expected = inputs if context_name == "only_student_forward" else 2 * inputs

with context():
torch.testing.assert_close(model(inputs), expected)
with (
pytest.raises(RuntimeError, match="inner failure") if raise_in_inner else nullcontext(),
context(enable=inner_enabled),
):
torch.testing.assert_close(model(inputs), expected)
if raise_in_inner:
raise RuntimeError("inner failure")
torch.testing.assert_close(model(inputs), expected)

assert calls["teacher"] == (0 if context_name == "only_student_forward" else 3)
assert calls["student"] == (3 if context_name == "only_student_forward" else 0)
torch.testing.assert_close(model(inputs), inputs)
assert calls["teacher"] == (1 if context_name == "only_student_forward" else 4)
assert calls["student"] == (4 if context_name == "only_student_forward" else 1)

with pytest.raises(RuntimeError, match="outer failure"), context():
raise RuntimeError("outer failure")
with context(enable=False):
torch.testing.assert_close(model(inputs), inputs)
assert calls["teacher"] == (2 if context_name == "only_student_forward" else 5)
assert calls["student"] == (5 if context_name == "only_student_forward" else 2)

other_context = (
model.only_teacher_forward
if context_name == "only_student_forward"
else model.only_student_forward
)
with context():
with other_context():
torch.testing.assert_close(model(inputs), inputs)
torch.testing.assert_close(model(inputs), expected)


def test_train_eval_mode_switch(distillation_model):
distillation_model.train()
assert distillation_model._intermediate_output is None
Expand Down