### What does this PR do? Type of change: new example **Note:** This is **part 2 of 4** (builds on #1589): - **Part 1 (#1589):** Megatron-Bridge `quantize.py` + `export.py` support and tests. - **Part 2 (this PR):** extend `distill.py` for quantization-aware distillation (QAD) — load a quantized Megatron checkpoint as the student. - **Part 3:** https://github.com/NVIDIA/Model-Optimizer/pull/1601 - **Part 4:** repeat the NVFP4 + QAD experiments on a non-Nemotron model. Extends `examples/megatron_bridge/distill.py` to initialize the student from a **Megatron checkpoint** (a quantized checkpoint from `quantize.py`, or a pruned one) via `--student_megatron_path`, enabling **Quantization Aware Distillation (QAD)**: - `--student_hf_path` still builds the student architecture; `--student_megatron_path` supplies the (optionally quantized) weights. - For a quantized checkpoint, the ModelOpt quantize mode + base weights are restored onto the **plain student before the knowledge-distillation conversion** (`restore_sharded_modelopt_state` is a no-op once a model is already converted), so the distilled checkpoint stays exportable as a quantized model with `export.py`. **Upstream dependency / workaround:** `DistillationProvider.provide()` has no seam to transform the student before the KD conversion, so this patches `provide()` at the class level (via an `id()`-keyed registry, because the provider proxies instance-attribute assignment to its teacher once the teacher is set). A companion Megatron-Bridge PR adds a first-class `DistillationProvider.student_pre_conversion_hook`; from nemo:26.06 onwards the workaround should be removed and replaced with that hook (a removal note in `distill.py` documents exactly how). ### Usage ```bash # 1) PTQ -> quantized Megatron checkpoint (part 1) torchrun --nproc_per_node 2 quantize.py \ --hf_model_name_or_path Qwen/Qwen3-8B --quant_cfg fp8 --tp_size 2 \ --export_megatron_path /tmp/Qwen3-8B-FP8-megatron # 2) QAD: distill the quantized student from the unquantized teacher torchrun --nproc_per_node 8 distill.py \ --teacher_hf_path Qwen/Qwen3-8B \ --student_hf_path Qwen/Qwen3-8B \ --student_megatron_path /tmp/Qwen3-8B-FP8-megatron \ --data_paths 1.0 tokenized/data_text_document \ --train_iters 1000 --output_dir /output/qwen3_8b_qad # 3) export the distilled quantized checkpoint (part 1) torchrun --nproc_per_node 1 export.py \ --hf_model_name_or_path Qwen/Qwen3-8B \ --megatron_path /output/qwen3_8b_qad/checkpoints \ --export_unified_hf_path /tmp/qwen3_8b_qad_fp8_hf ``` ### Testing `tests/examples/megatron_bridge/test_qad.py` (validated on a 2-GPU NeMo `26.04` container): quantize a tiny Qwen3 at TP=2 → QAD distill from the quantized student → `export.py` to a unified HF checkpoint, asserting `hf_quant_config.json` is written (proves the quantize mode survived QAD). Includes a commented-out vLLM deployment check, validated locally (full flow passes; vLLM loads the export as `quantization=modelopt`). Existing normal/Puzzletron distillation tests still pass. ### Before your PR is "*Ready for review*" - Is this change backward compatible?: N/A (new example feature; default behavior unchanged when `--student_megatron_path` is not set) - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: N/A (no new dependencies) - Did you write any new necessary tests?: ✅ - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: ✅ - Did you get Claude approval on this PR?: ✅ ### Additional Information Depends on a companion Megatron-Bridge PR adding `DistillationProvider.student_pre_conversion_hook` (the upstream replacement for the class-level `provide()` workaround). The Nemotron-3 tutorial NVFP4 + QAD experiments ship in part 3. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Quantization Aware Distillation (QAD) workflow to recover accuracy of quantized Megatron students and distill from quantized checkpoints. * CLI option to initialize a distillation student from a Megatron checkpoint and a structure-only load path for bridging. * **Documentation** * Expanded runnable quantize → QAD → export guidance and best-practice tips. * **Tests** * End-to-end test validating quantize → QAD → export artifacts. * **Chores / UX** * Clearer rank-aware messages, improved tokenizer padding handling, and more consistent export behavior (fixed export dtype). <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com> Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
8.4 KiB
Knowledge Distillation
Knowledge Distillation is a machine learning technique where a compact "student" model learns to replicate the behavior of a larger, more complex "teacher" model to achieve comparable performance with improved efficiency.
Model Optimizer's Distillation is a set of wrappers and utilities to easily perform Knowledge Distillation among teacher and student models. Given a pretrained teacher model, Distillation has the potential to train a smaller student model faster and/or with higher accuracy than the student model could achieve on its own.
This section focuses on demonstrating how to apply Model Optimizer to perform knowledge distillation with ease.
| Section | Description | Link | Docs |
|---|---|---|---|
| Pre-Requisites | Required & optional packages to use this technique | [Link] | |
| Getting Started | Learn how to optimize your models using distillation to produce more intellegant smaller models | [Link] | [docs] |
| Support Matrix | View the support matrix to see compatibility and feature availability across different models | [Link] | |
| Distillation with Megatron-Bridge | Learn how to distill your models with Megatron-Bridge Framework | [Link] | [docs] |
| Distillation with Megatron-LM | Learn how to distill your models with Megatron-LM Framework | [Link] | |
| Distillation with Huggingface | Learn how to distill your models with Hugging Face | [Link] | [docs] |
| Resources | Extra links to relevant resources | [Link] |
Pre-Requisites
Docker
For Hugging Face models, please use the PyTorch docker image (e.g., nvcr.io/nvidia/pytorch:26.01-py3).
For Megatron-Bridge or Megatron-LM models, use the NeMo container (e.g., nvcr.io/nvidia/nemo:26.02) which has all the dependencies installed.
Visit our installation docs for more information.
Also follow the installation steps below to upgrade to the latest version of Model Optimizer and install example-specific dependencies.
Local Installation
For Hugging Face models, install Model Optimizer with hf dependencies using pip from PyPI and install the requirements for the example:
pip install -U nvidia-modelopt[hf]
pip install -r requirements.txt
Getting Started
Set up your base models
First obtain both a pretrained model to act as the teacher and a (usually smaller) model to serve as the student.
from transformers import AutoModelForCausalLM
# Define student & teacher
student_model = AutoModelForCausalLM.from_pretrained("student-model-id-or-path")
teacher_model = AutoModelForCausalLM.from_pretrained("teacher-model-id-or-path")
Set up the meta model
As Knowledge Distillation involves (at least) two models, ModelOpt simplifies the integration process by wrapping both student and teacher into one meta model.
Please see an example Distillation setup below. This example assumes the outputs of teacher_model and student_model are logits.
import modelopt.torch.distill as mtd
distillation_config = {
"teacher_model": teacher_model,
"criterion": mtd.LogitsDistillationLoss(), # callable receiving student and teacher outputs, in order
"loss_balancer": mtd.StaticLossBalancer(), # combines multiple losses; omit if only one distillation loss used
}
distillation_model = mtd.convert(student_model, mode=[("kd_loss", distillation_config)])
The teacher_model can be either a nn.Module, a callable which returns an nn.Module, or a tuple of (model_cls, args, kwargs). The criterion is the distillation loss used between student and teacher tensors. The loss_balancer determines how the original and distillation losses are combined (if needed).
See Distillation for more info.
Distill during training
To Distill from teacher to student, simply use the meta model in the usual training loop, while also using the meta model’s .compute_kd_loss() method to compute the distillation loss, in addition to the original user loss.
An example of Distillation training is given below:
# Setup the data loaders. As example:
train_loader = get_train_loader()
# Define user loss function. As example:
loss_fn = get_user_loss_fn()
for input, labels in train_dataloader:
distillation_model.zero_grad()
# Forward through the wrapped models
out = distillation_model(input)
# Same loss as originally present
loss = loss_fn(out, labels)
# Combine distillation and user losses
loss_total = distillation_model.compute_kd_loss(student_loss=loss)
loss_total.backward()
Note
DataParallel may break ModelOpt’s Distillation feature. Note that HuggingFace Trainer uses DataParallel by default.
Export trained model
The model can easily be reverted to its original class for further use (i.e deployment) without any ModelOpt modifications attached.
model = mtd.export(distillation_model)
Support Matrix
Current out of the box components
Loss criterion:
mtd.LogitsDistillationLoss()- Standard KL-Divergence on output logitsmtd.MGDLoss()- Masked Generative Distillation loss for 2D convolutional outputsmtd.MFTLoss()- KL-divergence loss with Minifinetuning threshold modification
Loss balancers:
mtd.StaticLossBalancer()- Combines original student loss and KD loss into a single weighted sum (without changing over time)
Supported Models
Note
The following are models that were confirmed to run with ModelOpt distillation, but it is absolutely not limited to these
| Model | type | confirmed compatible |
|---|---|---|
| Nemotron | mamba hybrid | ✅ |
| Llama 3 | llama | ✅ |
| Llama 4 | llama | ✅ |
| Gemma 2 | gemma | ✅ |
| Gemma 3 | gemma | ✅ |
| Phi 3 | phi | ✅ |
| Qwen 2 | qwen2 | ✅ |
| Qwen 3 | qwen3 | ✅ |
| Mamba | mamba | ✅ |
Knowledge Distillation (KD) in NVIDIA Megatron-Bridge Framework
Checkout the stand-alone distillation script in the examples/megatron_bridge/ for example scripts for KD with Megatron-Bridge which is generally more performant than the Hugging Face scripts.
Knowledge Distillation (KD) in NVIDIA Megatron-LM Framework
Checkout the Knowledge Distillation example in the Megatron-LM repository.
Knowledge Distillation (KD) for HuggingFace Models
In this e2e example we finetune Llama-3.2 models on the smol-smoltalk-Interaction-SFT dataset as a minimal example to demonstrate a simple way of integrating Model Optimizer's KD feature.
We replace normal supervised finetuning (SFT) of a Llama-3.2-1B base model by distilling information from Llama-3.2-3B-Instruct which has already been instruction-finetuned.
Note
We can fit the following in memory using FSDP enabled on 8x RTX 6000 (total ~400GB VRAM)
accelerate launch --config-file ./accelerate_config/fsdp2.yaml \
main.py \
--teacher_name_or_path 'meta-llama/Llama-3.2-3B-Instruct' \
--student_name_or_path 'meta-llama/Llama-3.2-1B' \
--output_dir ./llama3.2-distill \
--max_length 2048 \
--per_device_train_batch_size 4 \
--per_device_eval_batch_size 8 \
--max_steps 200 \
--logging_steps 5