nemo_automodel.recipes.llm.kd
nemo_automodel.recipes.llm.kd
Knowledge Distillation recipe for next-token prediction with NeMo AutoModel.
This recipe fine-tunes a student model using the logits of a frozen teacher model. It
extends FinetuneRecipeForNextTokenPrediction adding:
- teacher_model — an additional HF/NeMo model loaded in
evalmode - kd_loss_fn — KL-divergence between temperature-scaled distributions
- kd_ratio — linear mix between CE loss and KD loss
The training loop is copied from the parent class but the loss becomes: loss = (1-kd_ratio) * ce_loss + kd_ratio * kd_loss
Pipeline parallelism (PP) is supported. Teacher logits from every last-stage microbatch are captured via a lightweight closure and injected into the corresponding student pipeline microbatch.
The file exposes KnowledgeDistillationRecipeForNextTokenPrediction and a
main entry-point so it can be launched exactly the same way as other recipes:
python -m torch.distributed.run —nproc-per-node=8
nemo_automodel/recipes/llm/kd.py
-c examples/llm_kd/llama3_2/llama3_2_1b_kd.yaml
Module Contents
Classes
Functions
Data
API
Bases: TrainFinetuneRecipeForNextTokenPrediction
Fine-tune a student model via knowledge distillation.
Emit metadata for both KD consumers, including a teacher built later.
Adapt teacher stages and reject incompatible student/teacher mask layouts.
All separate-mesh ranks participate in the layout check before either side starts training. Metadata is always emitted by the KD dataloader, since the teacher is constructed after the student loader.
Run one non-PP student microbatch with KD.
Parameters:
Zero-based accumulation microbatch index.
Mapping containing input_ids and labels as tensors of
shape [batch, sequence]. Other tensor leaves may have
arbitrary rank and axis order.
Valid-label count across the optimizer step.
Number of accumulation microbatches in the step.
Whether to run backward.
Returns:
Tuple of scalar tensors containing detached mixed, KL, and CE loss.
PP path: run teacher eval to capture logits, then run student step/eval.
Teacher logits from the last PP stage are stored in self._current_teacher_logits
before the student schedule runs, so pp_kd_loss_fn can read them.
Transported teacher logits have shape [batch, sequence, vocab] and
are split into [microbatch, sequence, vocab] pipeline tensors.
Parameters:
Zero-based accumulation microbatch index.
Mapping containing input_ids and labels as tensors of
shape [batch, sequence]. Other tensor leaves may have
arbitrary rank and axis order.
Output list receiving one detached scalar tensor.
Valid-label count across the optimizer step.
Number of accumulation microbatches in the step.
Whether the pipeline schedule runs backward.
Request teacher logits for one student batch.
Parameters:
Mapping containing input_ids and labels as tensors of
shape [batch, sequence]. Other tensor leaves may have
arbitrary rank and axis order and are transported unchanged.
Returns: torch.Tensor
Replicated tensor of shape [batch, sequence, vocab] containing
Return a student pipeline loss_fn that combines CE and KD using teacher logits.
The wrapper reads self._current_teacher_logits which must be populated by
the teacher eval pass before each student step in _forward_backward_step_pp.
Run training loop; skip validation when PP is enabled (not yet supported).
Serve teacher forwards until the student mesh broadcasts stop.
Execute a single training step.
Parameters:
List of batches of training data.
Gradient clipping norm. Optional, if None will not clip gradients.
Execute a single training step when pipeline parallelism is enabled.
Run one pass over self.val_dataloader.
Run one teacher batch and materialize transport-ready logits.
Parameters:
Mapping containing input_ids and labels as tensors of
shape [batch, sequence]. Other tensor leaves may have
arbitrary rank and axis order. Tensors may initially reside on
CPU.
Returns: torch.Tensor | None
Detached tensor of shape [batch, sequence, vocab] on the teacher
Log metrics to wandb and other loggers.
Parameters:
MetricsSample object, containing: step: int, the current step. epoch: int, the current epoch. metrics: Dict[str, float], containing: “loss”: Training loss. “grad_norm”: Grad norm from the training step. “lr”: Learning rate. “mem”: Memory allocated. “tps”: Tokens per second. “tps_per_gpu”: Tokens per second per GPU. “num_label_tokens”: Number of label tokens.
Run the student loop or serve teacher forwards on a separate mesh.
Build student & teacher, dataloaders, optimizers, etc.
Build and initialize the teacher model for knowledge distillation.
Uses the same infrastructure as student model (NeMoAutoModelForCausalLM) but without PEFT, FP8, or QAT since the teacher should be frozen in full precision.
Parameters:
Configuration for teacher model instantiation.
Random seed for reproducibility.
Whether using packed sequences.
Resolved distributed topology and policy object.
Device to place the teacher model on.
Returns:
The frozen teacher model ready for inference.
The offload_teacher_model config option is not supported with this approach.
Device placement is handled internally by NeMoAutoModelForCausalLM infrastructure.
Build a frozen teacher model with the supplied distributed setup.
Teacher is built via build_model with pipeline_config so it becomes an AutoPipeline when PP is enabled. No PEFT/FP8/QAT. Teacher is frozen and set to eval mode.
Logit capture stores every last-stage microbatch in schedule order.
Parameters:
Configuration for teacher model instantiation.
Random seed for reproducibility.
Whether using packed sequences.
Pipeline configuration for the teacher.
Distributed setup for the teacher.
Whether to enable activation checkpointing.
Returns: Any
The frozen teacher AutoPipeline with a _teacher_logits_capture attribute.
Run the KD recipe from CLI or directly.