# Distillation Trainer

## Overview

The Distillation Trainer implements on-policy knowledge distillation as described in [On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes](https://huggingface.co/papers/2306.13649) by Rishabh Agarwal, Nino Vieillard, Yongchao Zhou, Piotr Stanczyk, Sabela Ramos, Matthieu Geist, and Olivier Bachem.

> Knowledge distillation (KD) is widely used for compressing a teacher model to reduce its inference cost and memory footprint, by training a smaller student model. However, current KD methods for auto-regressive sequence models suffer from distribution mismatch between output sequences seen during training and those generated by the student during inference. To address this issue, we introduce Generalized Knowledge Distillation (GKD). Instead of solely relying on a fixed set of output sequences, GKD trains the student on its self-generated output sequences by leveraging feedback from the teacher on such sequences. Unlike supervised KD approaches, GKD also offers the flexibility to employ alternative loss functions between the student and teacher, which can be useful when the student lacks the expressivity to mimic the teacher's distribution.

The `DistillationTrainer` trains a smaller student model to match a teacher's next-token distribution on the student's own on-policy generations, extending the ideas from the `GKDTrainer`. A generation buffer decouples the training microbatch size from the generation batch size, letting vLLM batch many prompts in a single call across gradient accumulation steps.

> [!NOTE]
> The Distillation Trainer is currently part of the `trl.experimental` namespace. APIs may change without notice while the feature is iterated on.

## Quick start

```python
from datasets import load_dataset
from trl.experimental.distillation import DistillationConfig, DistillationTrainer

# 1. Load dataset and format as a prompt-only column
dataset = load_dataset("openai/gsm8k", "main", split="train")
dataset = dataset.map(
    lambda x: {"prompt": [{"role": "user", "content": x["question"]}]},
    remove_columns=dataset.column_names,
)

# 2. Configure distillation
config = DistillationConfig(
    output_dir="results/distill-qwen-gsm8k",
    num_train_epochs=1,
    bf16=True,
    save_strategy="no",
    # Distillation
    beta=1.0,                       # reverse KL
    # Teacher
    teacher_model_init_kwargs={"dtype": "bfloat16"},
)

# 3. Train
trainer = DistillationTrainer(
    model="Qwen/Qwen2.5-1.5B-Instruct",
    teacher_model="Qwen/Qwen2.5-7B-Instruct",
    args=config,
    train_dataset=dataset,
)
trainer.train()
trainer.save_model()
```

## Usage tips

The [experimental.distillation.DistillationTrainer](/docs/trl/main/en/distillation_trainer#trl.experimental.distillation.DistillationTrainer) trains the student fully on-policy: the student generates its own completions and learns to match the teacher's next-token distribution on them. The key parameter is set via [experimental.distillation.DistillationConfig](/docs/trl/main/en/distillation_trainer#trl.experimental.distillation.DistillationConfig):

* `beta`: controls the interpolation in the Generalized Jensen-Shannon Divergence. When `beta=0.0` the loss approximates forward KL divergence, while `beta=1.0` approximates reverse KL divergence. Values in between interpolate.

### On-policy generation

Fully on-policy training generally outperforms off-policy distillation because the student learns from its own mistakes rather than imitating trajectories it may never produce. The generation buffer keeps this efficient: prompts across gradient accumulation steps are batched into a single vLLM call.

### Expected dataset type

The dataset should be formatted as a [conversational](dataset_formats#conversational) [prompt-only](dataset_formats#prompt-only) dataset. The student generates its own completions on-policy, so only the prompt is needed:

```python
{"prompt": [{"role": "user", "content": "What color is the sky?"}]}
```

## Example script

Use [`examples/scripts/distillation.py`](https://github.com/huggingface/trl/blob/main/examples/scripts/distillation.py) to launch distillation training from the command line. The script supports full training and LoRA via the standard `ModelConfig` flags.

```bash
# Full training:
python examples/scripts/distillation.py \
    --model_name_or_path Qwen/Qwen2.5-0.5B-Instruct \
    --teacher_model_name_or_path Qwen/Qwen2.5-1.5B-Instruct \
    --dataset_name trl-lib/ultrafeedback-prompt \
    --learning_rate 2e-5 \
    --per_device_train_batch_size 4 \
    --gradient_accumulation_steps 8 \
    --output_dir distilled-model \
    --num_train_epochs 1
```

```bash
# LoRA:
python examples/scripts/distillation.py \
    --model_name_or_path Qwen/Qwen2.5-0.5B-Instruct \
    --teacher_model_name_or_path Qwen/Qwen2.5-1.5B-Instruct \
    --dataset_name trl-lib/ultrafeedback-prompt \
    --learning_rate 2e-4 \
    --per_device_train_batch_size 4 \
    --gradient_accumulation_steps 8 \
    --output_dir distilled-model \
    --num_train_epochs 1 \
    --use_peft \
    --lora_r 64 \
    --lora_alpha 16
```

## DistillationTrainer[[trl.experimental.distillation.DistillationTrainer]]

#### trl.experimental.distillation.DistillationTrainer[[trl.experimental.distillation.DistillationTrainer]]

```python
trl.experimental.distillation.DistillationTrainer(model: typing.Union[transformers.modeling_utils.PreTrainedModel, torch.nn.Module, str], teacher_model: typing.Union[transformers.modeling_utils.PreTrainedModel, torch.nn.Module, str] = None, args: trl.experimental.distillation.distillation_config.DistillationConfig | None = None, train_dataset: datasets.arrow_dataset.Dataset | None = None, eval_dataset: datasets.arrow_dataset.Dataset | dict[str, datasets.arrow_dataset.Dataset] | None = None, processing_class: transformers.tokenization_utils_base.PreTrainedTokenizerBase | transformers.image_processing_utils.BaseImageProcessor | transformers.feature_extraction_utils.FeatureExtractionMixin | transformers.processing_utils.ProcessorMixin | None = None, callbacks: list[transformers.trainer_callback.TrainerCallback] | None = None, optimizers: tuple = (None, None), quantization_config: BitsAndBytesConfig | None = None, peft_config: typing.Optional[ForwardRef('PeftConfig')] = None)
```

[Source](https://github.com/huggingface/trl/blob/main/trl/experimental/distillation/distillation_trainer.py#L282)

Trainer for knowledge distillation from a teacher model to a student model.

Supports:
- Generalized JSD loss (forward KL, reverse KL, or interpolated JSD via `beta`)
- On-policy distillation: the student generates completions, the teacher scores them
- Local teacher model
- Student on-policy generation via vLLM or model.generate()
- Liger kernel for memory-efficient fused JSD loss

#### train[[trl.experimental.distillation.DistillationTrainer.train]]

```python
train(resume_from_checkpoint: str | bool | None = None, trial: optuna.Trial | dict[str, Any] | None = None, ignore_keys_for_eval: list[str] | None = None)
```

[Source](https://github.com/huggingface/trl/blob/main/transformers/trainer.py#L1335)

**Parameters:**

resume_from_checkpoint (`str` or `bool`, *optional*) : If a `str`, local path to a saved checkpoint as saved by a previous instance of `Trainer`. If a `bool` and equals `True`, load the last checkpoint in *args.output_dir* as saved by a previous instance of `Trainer`. If present, training will resume from the model/optimizer/scheduler states loaded here.

trial (`optuna.Trial` or `dict[str, Any]`, *optional*) : The trial run or the hyperparameter dictionary for hyperparameter search.

ignore_keys_for_eval (`list[str]`, *optional*) : A list of keys in the output of your model (if it is a dictionary) that should be ignored when gathering predictions for evaluation during the training.

**Returns:** `~trainer_utils.TrainOutput`

Object containing the global step count, training loss, and metrics.

Main training entry point.

#### save_model[[trl.experimental.distillation.DistillationTrainer.save_model]]

```python
save_model(output_dir: str | None = None, _internal_call: bool = False)
```

[Source](https://github.com/huggingface/trl/blob/main/transformers/trainer.py#L3780)

Will save the model, so you can reload it using `from_pretrained()`.

Will only save from the main process.

#### push_to_hub[[trl.experimental.distillation.DistillationTrainer.push_to_hub]]

```python
push_to_hub(commit_message: str | None = 'End of training', blocking: bool = True, token: str | None = None, revision: str | None = None, **kwargs)
```

[Source](https://github.com/huggingface/trl/blob/main/transformers/trainer.py#L4027)

**Parameters:**

commit_message (`str`, *optional*, defaults to `"End of training"`) : Message to commit while pushing.

blocking (`bool`, *optional*, defaults to `True`) : Whether the function should return only when the `git push` has finished.

token (`str`, *optional*, defaults to `None`) : Token with write permission to overwrite Trainer's original args.

revision (`str`, *optional*) : The git revision to commit from. Defaults to the head of the "main" branch.

kwargs (`dict[str, Any]`, *optional*) : Additional keyword arguments passed along to `~Trainer.create_model_card`.

**Returns:**

The URL of the repository where the model was pushed if `blocking=False`, or a `Future` object tracking the
progress of the commit if `blocking=True`.

Upload `self.model` and `self.processing_class` to the 🤗 model hub on the repo `self.args.hub_model_id`.

## DistillationConfig[[trl.experimental.distillation.DistillationConfig]]

#### trl.experimental.distillation.DistillationConfig[[trl.experimental.distillation.DistillationConfig]]

```python
trl.experimental.distillation.DistillationConfig(output_dir: str | None = None, per_device_train_batch_size: int = 8, num_train_epochs: float = 3.0, max_steps: int = -1, learning_rate: float = 1e-06, lr_scheduler_type: transformers.trainer_utils.SchedulerType | str = 'linear', lr_scheduler_kwargs: dict | str | None = None, warmup_steps: float = 0, optim: transformers.training_args.OptimizerNames | str = 'adamw_torch_fused', optim_args: str | None = None, weight_decay: float = 0.0, adam_beta1: float = 0.9, adam_beta2: float = 0.999, adam_epsilon: float = 1e-08, optim_target_modules: None | str | list[str] = None, gradient_accumulation_steps: int = 1, average_tokens_across_devices: bool = True, max_grad_norm: float = 1.0, label_smoothing_factor: float = 0.0, bf16: bool | None = None, fp16: bool = False, bf16_full_eval: bool = False, fp16_full_eval: bool = False, tf32: bool | None = None, gradient_checkpointing: bool = True, gradient_checkpointing_kwargs: dict[str, typing.Any] | str | None = None, torch_compile: bool = False, torch_compile_backend: str | None = None, torch_compile_mode: str | None = None, use_liger_kernel: bool = False, liger_kernel_config: dict[str, bool] | None = None, use_cache: bool = False, neftune_noise_alpha: float | None = None, torch_empty_cache_steps: int | None = None, auto_find_batch_size: bool = False, logging_strategy: transformers.trainer_utils.IntervalStrategy | str = 'steps', logging_steps: float = 10, logging_first_step: bool = False, log_on_each_node: bool = True, logging_nan_inf_filter: bool = True, include_num_input_tokens_seen: str | bool = 'no', log_level: str = 'passive', log_level_replica: str = 'warning', disable_tqdm: bool | None = None, report_to: None | str | list[str] = 'none', run_name: str | None = None, project: str = 'huggingface', trackio_space_id: str | None = None, trackio_bucket_id: str | None = None, trackio_static_space_id: typing.Union[str, NoneType, typing.Literal[False]] = None, eval_strategy: transformers.trainer_utils.IntervalStrategy | str = 'no', eval_steps: float | None = None, eval_delay: float = 0, per_device_eval_batch_size: int = 8, prediction_loss_only: bool = False, eval_on_start: bool = False, eval_do_concat_batches: bool = True, eval_use_gather_object: bool = False, eval_accumulation_steps: int | None = None, include_for_metrics: list = <factory>, batch_eval_metrics: bool = False, save_only_model: bool = False, save_strategy: transformers.trainer_utils.SaveStrategy | str = 'steps', save_steps: float = 500, save_on_each_node: bool = False, save_total_limit: int | None = None, enable_jit_checkpoint: bool = False, push_to_hub: bool = False, hub_token: str | None = None, hub_private_repo: bool | None = None, hub_model_id: str | None = None, hub_strategy: transformers.trainer_utils.HubStrategy | str = 'every_save', hub_always_push: bool = False, hub_revision: str | None = None, load_best_model_at_end: bool = False, metric_for_best_model: str | None = None, greater_is_better: bool | None = None, ignore_data_skip: bool = False, restore_callback_states_from_checkpoint: bool = False, full_determinism: bool = False, seed: int = 42, data_seed: int | None = None, use_cpu: bool = False, accelerator_config: dict | str | None = None, parallelism_config: accelerate.parallelism_config.ParallelismConfig | None = None, dataloader_drop_last: bool = False, dataloader_num_workers: int = 0, dataloader_pin_memory: bool = True, dataloader_persistent_workers: bool = False, dataloader_prefetch_factor: int | None = None, remove_unused_columns: bool = False, label_names: list[str] | None = None, train_sampling_strategy: str = 'random', length_column_name: str = 'length', ddp_find_unused_parameters: bool | None = None, ddp_bucket_cap_mb: int | None = None, ddp_broadcast_buffers: bool | None = None, ddp_static_graph: bool | None = None, ddp_backend: str | None = None, ddp_timeout: int = 1800, fsdp: str | None = None, fsdp_config: dict[str, typing.Any] | str | None = None, deepspeed: dict | str | None = None, debug: str | list[transformers.debug_utils.DebugOption] = '', skip_memory_metrics: bool = True, do_train: bool = False, do_eval: bool = False, do_predict: bool = False, resume_from_checkpoint: str | None = None, warmup_ratio: float | None = None, logging_dir: str | None = None, local_rank: int = -1, model_init_kwargs: dict[str, typing.Any] | str | None = None, trust_remote_code: bool = False, teacher_model_name_or_path: str | None = None, teacher_model_revision: str | None = None, teacher_model_init_kwargs: dict[str, typing.Any] | str | None = None, disable_dropout: bool = False, max_completion_length: int | None = 512, ds3_gather_for_generation: bool = True, shuffle_dataset: bool | None = True, pad_to_multiple_of: int | None = None, temperature: float = 1.0, top_p: float = 1.0, top_k: int = 0, min_p: float | None = None, generation_kwargs: dict | None = None, chat_template_kwargs: dict | None = None, repetition_penalty: float = 1.0, cache_implementation: str | None = None, use_vllm: bool = False, vllm_mode: str = 'colocate', vllm_model_impl: str = 'vllm', vllm_enable_sleep_mode: bool = False, vllm_structured_outputs_regex: str | None = None, vllm_server_base_url: str | None = None, vllm_server_host: str = '0.0.0.0', vllm_server_port: int = 8000, vllm_server_timeout: float = 240.0, vllm_group_port: int = 51216, vllm_gpu_memory_utilization: float = 0.3, vllm_max_model_length: int | None = None, vllm_tensor_parallel_size: int = 1, beta: float = 1.0, log_completions: bool = False, num_completions_to_print: int | None = None, log_unique_prompts: bool = False)
```

[Source](https://github.com/huggingface/trl/blob/main/trl/experimental/distillation/distillation_config.py#L22)

**Parameters that control the model and the teacher model:**

model_init_kwargs (`str` or `dict[str, Any]`, *optional*) : Keyword arguments for `AutoModelForCausalLM.from_pretrained`, used when the `model` argument of the trainer is provided as a string.

trust_remote_code (`bool`, *optional*, defaults to `False`) : Whether to allow loading models and tokenizers that ship custom Python code from the Hub. Forwarded to [from_pretrained](https://huggingface.co/docs/transformers/main/en/model_doc/auto#transformers.AutoModelForCausalLM.from_pretrained) and [from_pretrained](https://huggingface.co/docs/transformers/main/en/model_doc/auto#transformers.AutoTokenizer.from_pretrained), for both the student and teacher.

teacher_model_name_or_path (`str`, *optional*) : Model name or path for the teacher model. Used when the teacher is loaded locally.

teacher_model_revision (`str`, *optional*) : Model revision of the teacher model (e.g., branch name, tag, or commit hash).

teacher_model_init_kwargs (`str` or `dict[str, Any]`, *optional*) : Keyword arguments passed to `AutoModelForCausalLM.from_pretrained` when instantiating the teacher model from a string.

disable_dropout (`bool`, *optional*, defaults to `False`) : Whether to disable dropout in the student model during training.

**Parameters that control the data preprocessing:**

remove_unused_columns (`bool`, *optional*, defaults to `False`) : Whether to only keep the column `"prompt"` in the dataset. The trainer consumes the raw prompt column and generates completions on-policy, so it defaults to `False`.

max_completion_length (`int` or `None`, *optional*, defaults to `512`) : Maximum number of tokens to generate per completion during on-policy generation.

ds3_gather_for_generation (`bool`, *optional*, defaults to `True`) : This setting applies to DeepSpeed ZeRO-3. If enabled, the policy model weights are gathered for generation, improving generation speed. However, disabling this option allows training models that exceed the VRAM capacity of a single GPU, albeit at the cost of slower generation. Disabling this option is not compatible with vLLM generation.

shuffle_dataset (`bool`, *optional*, defaults to `True`) : Whether to shuffle the training dataset.

pad_to_multiple_of (`int`, *optional*) : If set, the prompts ids and completions ids will be padded to a multiple of this value.

**Parameters that control generation:**

temperature (`float`, *optional*, defaults to `1.0`) : Temperature for sampling during generation and for computing the distillation loss. Higher values produce softer probability distributions.

top_p (`float`, *optional*, defaults to `1.0`) : Top-p (nucleus) sampling parameter for on-policy generation.

top_k (`int`, *optional*, defaults to `0`) : Top-k sampling parameter for on-policy generation. `0` disables top-k filtering.

min_p (`float`, *optional*) : Minimum token probability, which will be scaled by the probability of the most likely token. It must be a value between `0.0` and `1.0`. Typical values are in the `0.01-0.2` range.

generation_kwargs (`dict[str, Any]`, *optional*) : Additional keyword arguments to pass to [GenerationConfig](https://huggingface.co/docs/transformers/main/en/main_classes/text_generation#transformers.GenerationConfig) (if using transformers) or `SamplingParams` (if using vLLM) when sampling completions. This can be used to further customize the generation behavior, such as setting `suppress_tokens`, `num_beams`, etc. If it contains keys that conflict with the other generation parameters (like `min_p`, `top_p`, etc.), they will override them.

chat_template_kwargs (`dict[str, Any]`, *optional*) : Additional keyword arguments to pass to the `apply_chat_template` function when generating completions.

repetition_penalty (`float`, *optional*, defaults to `1.0`) : Float that penalizes new tokens based on whether they appear in the prompt and the generated text so far. Values > `1.0` encourage the model to use new tokens, while values < `1.0` encourage the model to repeat tokens.

cache_implementation (`str`, *optional*) : Implementation of the cache method for faster generation when `use_vllm` is set to `False`.

**Parameters that control generation acceleration powered by vLLM:**

use_vllm (`bool`, *optional*, defaults to `False`) : Whether to use vLLM for generating on-policy completions from the student model.

vllm_mode (`str`, *optional*, defaults to `"colocate"`) : Mode for student vLLM integration. Either `"server"` or `"colocate"`.

vllm_model_impl (`str`, *optional*, defaults to `"vllm"`) : Model implementation backend for vLLM. Use `"vllm"` or `"transformers"`.

vllm_enable_sleep_mode (`bool`, *optional*, defaults to `False`) : Enable vLLM sleep mode to offload student weights during the optimizer step.

vllm_structured_outputs_regex (`str` or `None`, *optional*) : Regex pattern for vLLM structured outputs.

**Parameters that control the vLLM server (only used when `vllm_mode` is `"server"`):**

vllm_server_base_url (`str` or `None`, *optional*) : Base URL for the student vLLM server. If provided, `vllm_server_host` and `vllm_server_port` are ignored.

vllm_server_host (`str`, *optional*, defaults to `"0.0.0.0"`) : Host of the student vLLM server.

vllm_server_port (`int`, *optional*, defaults to `8000`) : Port of the student vLLM server.

vllm_server_timeout (`float`, *optional*, defaults to `240.0`) : Timeout for connecting to the student vLLM server.

vllm_group_port (`int`, *optional*, defaults to `51216`) : Port for the vLLM weight-update group (NCCL communicator).

**Parameters that control colocated vLLM execution (only used when `vllm_mode` is `"colocate"`):**

vllm_gpu_memory_utilization (`float`, *optional*, defaults to `0.3`) : GPU memory utilization for the colocated student vLLM engine.

vllm_max_model_length (`int` or `None`, *optional*) : Maximum model sequence length for the colocated vLLM engine.

vllm_tensor_parallel_size (`int`, *optional*, defaults to `1`) : Tensor parallel size for the colocated student vLLM engine.

**Parameters that control the training:**

beta (`float`, *optional*, defaults to `1.0`) : Interpolation coefficient for the Generalized Jensen-Shannon Divergence loss. When `0.0`, the loss is the forward KL divergence. When `1.0`, the loss is the reverse KL divergence. When `0.5`, it is the standard JSD. Unlike GRPO's `beta` (a KL-penalty coefficient against a reference model), here it selects the divergence itself; there is no reference-model KL penalty.

**Parameters that control the logging:**

log_completions (`bool`, *optional*, defaults to `False`) : Whether to log a sample of (prompt, completion) pairs every `logging_steps` steps. If `rich` is installed, it prints the sample. If `wandb` and/or `trackio` logging is enabled, it logs it to `wandb` and/or `trackio`.

num_completions_to_print (`int`, *optional*) : Number of completions to print with `rich`. If `None`, all completions are logged.

log_unique_prompts (`bool`, *optional*, defaults to `False`) : Whether to log unique prompts. If `True`, only unique prompts are logged. If `False`, all prompts are logged.

Configuration class for the `DistillationTrainer`.

Extends [TrainingArguments](https://huggingface.co/docs/transformers/main/en/main_classes/trainer#transformers.TrainingArguments) with parameters specific to knowledge distillation. All necessary
fields are declared here.

Using [HfArgumentParser](https://huggingface.co/docs/transformers/main/en/internal/trainer_utils#transformers.HfArgumentParser) we can turn this class into
[argparse](https://docs.python.org/3/library/argparse#module-argparse) arguments that can be specified on the
command line.

> [!NOTE]
> Some [TrainingArguments](https://huggingface.co/docs/transformers/main/en/main_classes/trainer#transformers.TrainingArguments) defaults are overridden: `learning_rate` defaults to `1e-6` (instead of
> `5e-5`) and `remove_unused_columns` defaults to `False` (instead of `True`).

