nemo_automodel.recipes.llm.kd

View as Markdown

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:

  1. teacher_model — an additional HF/NeMo model loaded in eval mode
  2. kd_loss_fn — KL-divergence between temperature-scaled distributions
  3. 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

NameDescription
KnowledgeDistillationRecipeForNextTokenPredictionFine-tune a student model via knowledge distillation.

Functions

NameDescription
_build_kd_loss_fn-
_build_teacher_modelBuild and initialize the teacher model for knowledge distillation.
_build_teacher_model_with_ppBuild a frozen teacher model with the supplied distributed setup.
_verify_tokenizer_compatibility-
mainRun the KD recipe from CLI or directly.

Data

logger

API

class nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction()

Bases: TrainFinetuneRecipeForNextTokenPrediction

Fine-tune a student model via knowledge distillation.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._configure_packing() -> nemo_automodel.components.models.common.packing.PackingCapabilities

Emit metadata for both KD consumers, including a teacher built later.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._configure_teacher_packing() -> None

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.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._configure_teacher_pipeline() -> None
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._create_distributed_setup() -> nemo_automodel.components.distributed.config.DistributedSetup
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._forward_backward_step(
idx,
batch,
num_label_tokens,
num_batches,
is_train: bool = True
)

Run one non-PP student microbatch with KD.

Parameters:

idx

Zero-based accumulation microbatch index.

batch

Mapping containing input_ids and labels as tensors of shape [batch, sequence]. Other tensor leaves may have arbitrary rank and axis order.

num_label_tokens

Valid-label count across the optimizer step.

num_batches

Number of accumulation microbatches in the step.

is_train
boolDefaults to True

Whether to run backward.

Returns:

Tuple of scalar tensors containing detached mixed, KL, and CE loss.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._forward_backward_step_pp(
idx,
batch,
loss_buffer,
num_label_tokens,
num_batches,
is_train: bool = True
)

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:

idx

Zero-based accumulation microbatch index.

batch

Mapping containing input_ids and labels as tensors of shape [batch, sequence]. Other tensor leaves may have arbitrary rank and axis order.

loss_buffer

Output list receiving one detached scalar tensor.

num_label_tokens

Valid-label count across the optimizer step.

num_batches

Number of accumulation microbatches in the step.

is_train
boolDefaults to True

Whether the pipeline schedule runs backward.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._get_separate_teacher_logits(
batch: dict[str, typing.Any]
) -> torch.Tensor

Request teacher logits for one student batch.

Parameters:

batch
dict[str, Any]

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

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._make_pp_kd_loss_wrapper()

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.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._run_student_train_validation_loop()

Run training loop; skip validation when PP is enabled (not yet supported).

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._run_teacher_worker() -> None

Serve teacher forwards until the student mesh broadcasts stop.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._run_train_optim_step(
batches,
max_grad_norm: float | None = None
)

Execute a single training step.

Parameters:

batches

List of batches of training data.

max_grad_norm
float | NoneDefaults to None

Gradient clipping norm. Optional, if None will not clip gradients.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._run_train_optim_step_pp(
batches,
max_grad_norm: float | None = None
)

Execute a single training step when pipeline parallelism is enabled.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._run_validation_epoch(
val_dataloader
)

Run one pass over self.val_dataloader.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._setup_kd_state() -> None
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._should_setup_training_components() -> bool
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._teacher_forward_separate(
batch: dict[str, typing.Any]
) -> torch.Tensor | None

Run one teacher batch and materialize transport-ready logits.

Parameters:

batch
dict[str, Any]

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

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction.log_train_metrics(
log_data
) -> float

Log metrics to wandb and other loggers.

Parameters:

log_data

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.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction.log_val_metrics(
val_name,
log_data,
metric_logger = None
)
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction.run_train_validation_loop()

Run the student loop or serve teacher forwards on a separate mesh.

nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction.setup()

Build student & teacher, dataloaders, optimizers, etc.

nemo_automodel.recipes.llm.kd._build_kd_loss_fn(
cfg_kd
)
nemo_automodel.recipes.llm.kd._build_teacher_model(
cfg_teacher,
seed,
has_packed_sequence,
device = None
)

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:

cfg_teacher

Configuration for teacher model instantiation.

seed

Random seed for reproducibility.

has_packed_sequence

Whether using packed sequences.

distributed_setup
DistributedSetup | NoneDefaults to None

Resolved distributed topology and policy object.

device
Defaults to None

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.

nemo_automodel.recipes.llm.kd._build_teacher_model_with_pp(
cfg_teacher,
seed: int,
has_packed_sequence: bool,
activation_checkpointing: bool
) -> typing.Any

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:

cfg_teacher

Configuration for teacher model instantiation.

seed
int

Random seed for reproducibility.

has_packed_sequence
bool

Whether using packed sequences.

pipeline_config
PipelineConfig

Pipeline configuration for the teacher.

distributed_setup
DistributedSetup

Distributed setup for the teacher.

activation_checkpointing
bool

Whether to enable activation checkpointing.

Returns: Any

The frozen teacher AutoPipeline with a _teacher_logits_capture attribute.

nemo_automodel.recipes.llm.kd._verify_tokenizer_compatibility(
student_cfg,
teacher_cfg,
trust_remote_code = True
)
nemo_automodel.recipes.llm.kd.main(
config_path = 'examples/llm_kd/llama3_2/l...
)

Run the KD recipe from CLI or directly.

nemo_automodel.recipes.llm.kd.logger = logging.getLogger(__name__)