mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Bound the stale-capture warning to once per capture and name its cause (#2492)
### What does this PR do?
Type of change: Bug fix
The activation-capture hooks latch `_intermediate_output` on every
forward, and only
`DistillationModel.compute_kd_loss()` ever clears it. Any training loop
that does not call
`compute_kd_loss()` therefore warns on every forward after the first,
once per hooked module:
```
plain transformers Trainer, 16 micro-batches
before: 30 warnings (Teacher 15 + Student 15)
after: 2 warnings (Teacher 1 + Student 1)
```
The count scales with forwards, not optimizer steps —
`gradient_accumulation_steps` of 1, 4 and
16 all produced 30 warnings for the same 16 forwards — so the
accumulator the report blamed is
not involved. A plain `transformers.Trainer` reproduces it because its
default `compute_loss`
reads the student's own CE from `outputs.loss` and never calls
`compute_kd_loss()`; the
warning is a symptom of that, but neither message said so. The student
message attributed the
re-forward to Activation Checkpointing only, and the teacher message
called the situation
"expected" while still raising `UserWarning`.
This PR:
- warns once per captured activation via
`_warn_once_about_stale_output`, re-armed wherever a
capture is cleared, so a consumed capture that is followed by another
unconsumed one is still
reported (Activation Checkpointing re-runs forwards and relies on that);
- states the cause and the consequence in both messages:
`compute_kd_loss()` did not run since
the previous forward, so no KD loss is applied;
- moves the three `_intermediate_output = None` reset sites onto one
`_clear_captured_output`
helper so the capture and its warning state cannot drift apart, and has
the layerwise teacher
hook share both helpers rather than keeping a second copy of the
warning.
### Usage
No API change. A loop that applies KD should call `compute_kd_loss()`
once per forward:
```python
class KDLossTrainer(Trainer):
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
outputs = model(**inputs)
loss = model.compute_kd_loss(student_loss=outputs.loss) # consumes the captures
return (loss, outputs) if return_outputs else loss
```
### Testing
`tests/unit/torch/distill/test_distill.py`:
- `test_duplicate_fwd_hook_call` now pins the bound under
`warnings.simplefilter("always")`
(3 forwards -> exactly 2 warnings) instead of relying on the
interpreter's per-message
deduplication, and asserts the message names `compute_kd_loss`;
- `test_stale_output_warning_rearms_after_consuming` covers the re-arm:
warn, consume, warn
again -> 4 warnings. This one passes before and after the change; it
guards against a future
"warn once ever" simplification silently muting the Activation
Checkpointing case.
`tests/unit/torch/distill/test_layerwise.py`:
- `test_layerwise_stale_output_warning_is_bounded` covers the layerwise
hooks (3 forwards ->
exactly 2 warnings).
The first and third tests fail on `main` (`assert 4 == 2`), verified in
a worktree of `main`
carrying these test files.
Measured with a `transformers.Trainer` over a 16-micro-batch run,
counting
`"already has an intermediate output stored"`:
| setup | before | after |
|---|---|---|
| plain `Trainer`, `gradient_accumulation_steps=1` | 30 | 2 |
| plain `Trainer`, `gradient_accumulation_steps=4` | 30 | 2 |
| plain `Trainer`, `gradient_accumulation_steps=16` | 30 | 2 |
| `Trainer` calling `compute_kd_loss()`, accum=4 | 0 | 0 |
```
$ python -m pytest tests/unit/torch/distill -q
34 passed
$ python -m pytest tests/unit/torch -q --ignore=tests/unit/torch/deploy
1 failed, 2673 passed, 18 skipped in 217.17s
```
The single failure is
`tests/unit/torch/quantization/plugins/test_huggingface.py::test_dbrx`,
which fails identically on `main` in this environment because
`transformers 5.17` is outside the
`transformers>=4.57,<5.15` range pinned in `pyproject.toml`. It is
unrelated to this change.
`pre-commit run --files <changed files>` passes every hook (ruff check,
ruff format, mypy,
bandit, insert-license, large files, line endings).
Note on severity, since it affects how the issue reads: Python already
deduplicates a warning
by (message, location), so with the default filters this shows up twice
rather than 30 times.
The repetition is user-visible under `-W always`,
`PYTHONWARNINGS=always`, pytest, or one
warning registry per DDP rank, and the count is what scales with epoch
length. The behavioural
fix that matters for all filters is the latch.
### Before your PR is "*Ready for review*"
Is this change backward compatible?: ✅
If you copied code from any other sources or added a new PIP dependency,
did you follow guidance in CONTRIBUTING.md: N/A
Did you write any new necessary tests?: ✅
Did you update Changelog?: N/A
Did you get Claude approval on this PR?: N/A
### Additional Information
Fixes #2487.
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
- **Bug Fixes**
- Improved handling of captured activations during distillation,
including clearing stale outputs and resetting warning behavior after
outputs are consumed.
- Stale-output warnings now appear once per affected activation and
clarify when knowledge-distillation loss was not applied.
- Warnings distinguish expected cases involving activation checkpointing
or teacher evaluation.
- **Tests**
- Expanded coverage for warning counts, warning reset behavior, and
stale-output handling in standard and layerwise distillation.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
Signed-off-by: Edwardssss <ed_129@qq.com>
This commit is contained in:
@@ -113,9 +113,9 @@ class DistillationModel(DynamicModule):
|
||||
def _register_hooks(self):
|
||||
"""Register hooks for intermediate tensors from teacher models and the student model."""
|
||||
for student_layer, teacher_layer in self._layers_to_loss:
|
||||
setattr(student_layer, "_intermediate_output", None)
|
||||
_clear_captured_output(student_layer)
|
||||
handle_s = student_layer.register_forward_hook(student_output_capture_fwd_hook)
|
||||
setattr(teacher_layer, "_intermediate_output", None)
|
||||
_clear_captured_output(teacher_layer)
|
||||
handle_t = teacher_layer.register_forward_hook(teacher_output_capture_fwd_hook)
|
||||
self._hook_handles.update([handle_s, handle_t])
|
||||
|
||||
@@ -188,8 +188,8 @@ class DistillationModel(DynamicModule):
|
||||
if self.training != mode:
|
||||
# When switching between train and eval, clear outputs
|
||||
for student_layer, teacher_layer in self._layers_to_loss:
|
||||
student_layer._intermediate_output = None
|
||||
teacher_layer._intermediate_output = None
|
||||
_clear_captured_output(student_layer)
|
||||
_clear_captured_output(teacher_layer)
|
||||
super().train(mode)
|
||||
|
||||
def state_dict(self, *args, **kwargs) -> dict[str, Any]:
|
||||
@@ -269,8 +269,8 @@ class DistillationModel(DynamicModule):
|
||||
for i, ((student_layer, teacher_layer), loss_fn) in enumerate(self._layers_to_loss.items()):
|
||||
out_s = student_layer._intermediate_output
|
||||
out_t = teacher_layer._intermediate_output
|
||||
student_layer._intermediate_output = None
|
||||
teacher_layer._intermediate_output = None
|
||||
_clear_captured_output(student_layer)
|
||||
_clear_captured_output(teacher_layer)
|
||||
|
||||
loss = loss_fn(out_s, out_t, **loss_fn_kwargs) # Student is pred, Teacher is target
|
||||
if loss_reduction_fn is not None:
|
||||
@@ -292,6 +292,24 @@ class DistillationModel(DynamicModule):
|
||||
return loss_total
|
||||
|
||||
|
||||
def _clear_captured_output(layer: nn.Module) -> None:
|
||||
"""Drop a captured activation and re-arm its stale-output warning."""
|
||||
layer._intermediate_output = None
|
||||
layer._intermediate_output_warned = False
|
||||
|
||||
|
||||
def _warn_once_about_stale_output(module: nn.Module, who: str, expectation: str) -> None:
|
||||
"""Report a forward that overwrites an unconsumed capture, once per captured activation."""
|
||||
if getattr(module, "_intermediate_output_warned", False):
|
||||
return
|
||||
module._intermediate_output_warned = True
|
||||
warnings.warn(
|
||||
f"{who}'s Module `{type(module).__name__}` already has an intermediate output stored:"
|
||||
" `DistillationModel.compute_kd_loss()` did not run since the previous forward, so no KD"
|
||||
f" loss is applied. {expectation}"
|
||||
)
|
||||
|
||||
|
||||
def student_output_capture_fwd_hook(module: nn.Module, input: Any, output: Any):
|
||||
"""A hook to capture layer output."""
|
||||
# NOTE: Defined externally to allow pickling during DDP initialization.
|
||||
@@ -299,9 +317,10 @@ def student_output_capture_fwd_hook(module: nn.Module, input: Any, output: Any):
|
||||
if getattr(module, "_only_teacher_fwd", False):
|
||||
return # Might be hooked on entire model fwd
|
||||
if module.training and module._intermediate_output is not None:
|
||||
warnings.warn(
|
||||
f"Student's Module `{type(module).__name__}` already has an intermediate output stored."
|
||||
" This is undesired behavior unless Activation Checkpointing is in use."
|
||||
_warn_once_about_stale_output(
|
||||
module,
|
||||
"Student",
|
||||
"Expected under Activation Checkpointing, which re-runs a forward.",
|
||||
)
|
||||
|
||||
module._intermediate_output = output
|
||||
@@ -313,9 +332,10 @@ def teacher_output_capture_fwd_hook(module: nn.Module, input: Any, output: Any):
|
||||
|
||||
if module._intermediate_output is not None:
|
||||
# NOTE: cannot tell if train or eval since teacher is always eval
|
||||
warnings.warn(
|
||||
f"Teacher's Module `{type(module).__name__}` already has an intermediate output stored."
|
||||
" This is expected when `DistillationModel.compute_kd_loss` is not called in eval mode."
|
||||
_warn_once_about_stale_output(
|
||||
module,
|
||||
"Teacher",
|
||||
"Expected in eval mode, where the loss is not required.",
|
||||
)
|
||||
|
||||
module._intermediate_output = output
|
||||
|
||||
@@ -20,7 +20,12 @@ from typing import Any
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
from .distillation_model import DistillationModel, student_output_capture_fwd_hook
|
||||
from .distillation_model import (
|
||||
DistillationModel,
|
||||
_clear_captured_output,
|
||||
_warn_once_about_stale_output,
|
||||
student_output_capture_fwd_hook,
|
||||
)
|
||||
|
||||
__all__ = ["LayerwiseDistillationModel"]
|
||||
|
||||
@@ -58,10 +63,10 @@ class LayerwiseDistillationModel(DistillationModel):
|
||||
for student_layer, teacher_layer in self._layers_to_loss:
|
||||
setattr(student_layer, "_teacher_layer", [teacher_layer])
|
||||
handle_s1 = student_layer.register_forward_pre_hook(student_input_bypass_fwd_hook)
|
||||
setattr(student_layer, "_intermediate_output", None)
|
||||
_clear_captured_output(student_layer)
|
||||
handle_s2 = student_layer.register_forward_hook(student_output_capture_fwd_hook)
|
||||
setattr(teacher_layer, "_intermediate_input", None)
|
||||
setattr(teacher_layer, "_intermediate_output", None)
|
||||
_clear_captured_output(teacher_layer)
|
||||
handle_t = teacher_layer.register_forward_hook(teacher_input_output_capture_fwd_hook)
|
||||
self._hook_handles.update([handle_s1, handle_s2, handle_t])
|
||||
|
||||
@@ -104,9 +109,10 @@ def teacher_input_output_capture_fwd_hook(module: nn.Module, input: Any, output:
|
||||
|
||||
if module._intermediate_output is not None:
|
||||
# NOTE: cannot tell if train or eval since teacher is always eval
|
||||
warnings.warn(
|
||||
f"Teacher's Module `{type(module).__name__}` already has an intermediate output stored."
|
||||
" This is expected when `DistillationModel.compute_kd_loss` is not called in eval mode."
|
||||
_warn_once_about_stale_output(
|
||||
module,
|
||||
"Teacher",
|
||||
"Expected in eval mode, where the loss is not required.",
|
||||
)
|
||||
|
||||
module._intermediate_input = input
|
||||
|
||||
@@ -269,10 +269,29 @@ def test_duplicate_fwd_hook_call(distillation_model):
|
||||
distillation_model.train()
|
||||
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter("always") # the latch, not the interpreter, must bound the count
|
||||
distillation_model(get_input_tensor())
|
||||
distillation_model(get_input_tensor())
|
||||
distillation_model(get_input_tensor())
|
||||
assert len(w) == 2 # one for student and one for teacher
|
||||
stale = [x for x in w if "already has an intermediate output stored" in str(x.message)]
|
||||
assert len(stale) == 2 # one for student and one for teacher
|
||||
assert all("compute_kd_loss" in str(x.message) for x in stale)
|
||||
|
||||
|
||||
def test_stale_output_warning_rearms_after_consuming(distillation_model):
|
||||
"""Activation Checkpointing re-runs forwards, so consuming a capture must re-arm the notice."""
|
||||
distillation_model.train()
|
||||
input_tensor = get_input_tensor()
|
||||
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter("always")
|
||||
distillation_model(input_tensor)
|
||||
distillation_model(input_tensor) # unconsumed -> reported once per module
|
||||
distillation_model.compute_kd_loss()
|
||||
distillation_model(input_tensor)
|
||||
distillation_model(input_tensor) # unconsumed again -> reported again
|
||||
stale = [x for x in w if "already has an intermediate output stored" in str(x.message)]
|
||||
assert len(stale) == 4
|
||||
|
||||
|
||||
def test_teacher_fwd_only(distillation_model):
|
||||
|
||||
@@ -232,3 +232,15 @@ def test_layerwise_gradient_flow():
|
||||
assert updated_any, (
|
||||
"No parameters were updated in 'features.2' or related layers during training"
|
||||
)
|
||||
|
||||
|
||||
def test_layerwise_stale_output_warning_is_bounded(layerwise_distillation_model):
|
||||
"""Both layerwise hooks must share the parent's latch instead of failing or repeating."""
|
||||
layerwise_distillation_model.train()
|
||||
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter("always")
|
||||
for _ in range(3):
|
||||
layerwise_distillation_model(get_input_tensor())
|
||||
stale = [x for x in w if "already has an intermediate output stored" in str(x.message)]
|
||||
assert len(stale) == 2 # teacher's input/output hook and the student's output hook
|
||||
|
||||
Reference in New Issue
Block a user