> This page is for version Nightly (default).
> For other versions, use one of these documentation indexes:
> - Nightly (default): https://docs.nvidia.com/nemo/automodel/nightly/llms.txt
> - Latest: https://docs.nvidia.com/nemo/automodel/latest/llms.txt
> - 0.5.0 · 26.06: https://docs.nvidia.com/nemo/automodel/v0.5/llms.txt
> - 0.4.0 · 26.04: https://docs.nvidia.com/nemo/automodel/v0.4/llms.txt

> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# 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:

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

| Name                                                                                                                                    | Description                                           |
| --------------------------------------------------------------------------------------------------------------------------------------- | ----------------------------------------------------- |
| [`KnowledgeDistillationRecipeForNextTokenPrediction`](#nemo_automodel-recipes-llm-kd-KnowledgeDistillationRecipeForNextTokenPrediction) | Fine-tune a student model via knowledge distillation. |

### Functions

| Name                                                                                                | Description                                                        |
| --------------------------------------------------------------------------------------------------- | ------------------------------------------------------------------ |
| [`_build_kd_loss_fn`](#nemo_automodel-recipes-llm-kd-_build_kd_loss_fn)                             | -                                                                  |
| [`_build_teacher_model`](#nemo_automodel-recipes-llm-kd-_build_teacher_model)                       | Build and initialize the teacher model for knowledge distillation. |
| [`_build_teacher_model_with_pp`](#nemo_automodel-recipes-llm-kd-_build_teacher_model_with_pp)       | Build a frozen teacher model with the supplied distributed setup.  |
| [`_verify_tokenizer_compatibility`](#nemo_automodel-recipes-llm-kd-_verify_tokenizer_compatibility) | -                                                                  |
| [`main`](#nemo_automodel-recipes-llm-kd-main)                                                       | Run the KD recipe from CLI or directly.                            |

### Data

[`logger`](#nemo_automodel-recipes-llm-kd-logger)

### API

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

**Bases:** [TrainFinetuneRecipeForNextTokenPrediction](/nemo-automodel/nemo_automodel/recipes/llm/train_ft#nemo_automodel-recipes-llm-train_ft-TrainFinetuneRecipeForNextTokenPrediction)

Fine-tune a student model via knowledge distillation.

```python
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.

```python
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.

```python
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._configure_teacher_pipeline() -> None
```

```python
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._create_distributed_setup() -> nemo_automodel.components.distributed.config.DistributedSetup
```

```python
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`** `bool` — default: True

Whether to run backward.

---

**Returns:**

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

```python
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`** `bool` — default: True

Whether the pipeline schedule runs backward.

---

```python
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

```python
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`.

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

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

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

Serve teacher forwards until the student mesh broadcasts stop.

```python
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 | None` — default: None

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

---

```python
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.

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

Run one pass over `self.val_dataloader`.

```python
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._setup_kd_state() -> None
```

```python
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction._should_setup_training_components() -> bool
```

```python
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

```python
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.

---

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

```python
nemo_automodel.recipes.llm.kd.KnowledgeDistillationRecipeForNextTokenPrediction.run_train_validation_loop()
```

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

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

Build student & teacher, dataloaders, optimizers, etc.

```python
nemo_automodel.recipes.llm.kd._build_kd_loss_fn(
    cfg_kd
)
```

```python
nemo_automodel.recipes.llm.kd._build_teacher_model(
    cfg_teacher,
    seed,
    has_packed_sequence,
    distributed_setup: nemo_automodel.components.distributed.config.DistributedSetup | None = None,
    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 | None` — default: None

Resolved distributed topology and policy object.

---

**`device`** — default: None

Device to place the teacher model on.

---

**Returns:**

The frozen teacher model ready for inference.

> **Note**
>
> The `offload_teacher_model` config option is not supported with this approach.
> Device placement is handled internally by NeMoAutoModelForCausalLM infrastructure.

```python
nemo_automodel.recipes.llm.kd._build_teacher_model_with_pp(
    cfg_teacher,
    seed: int,
    has_packed_sequence: bool,
    pipeline_config: nemo_automodel.components.distributed.pipelining.config.PipelineConfig,
    distributed_setup: nemo_automodel.components.distributed.config.DistributedSetup,
    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.

```python
nemo_automodel.recipes.llm.kd._verify_tokenizer_compatibility(
    student_cfg,
    teacher_cfg,
    trust_remote_code = True
)
```

```python
nemo_automodel.recipes.llm.kd.main(
    config_path = 'examples/llm_kd/llama3_2/l...
)
```

Run the KD recipe from CLI or directly.

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