TRL documentation

Distillation Trainer

Hugging Face's logo
Join the Hugging Face community

and get access to the augmented documentation experience

to get started

Distillation Trainer

model badge

Overview

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

The abstract from the paper is the following:

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. It generates the student’s completions on-policy (optionally vLLM-powered) and matches the teacher over its full next-token distribution with a memory-efficient chunked Jensen-Shannon divergence loss, so the teacher’s dense distribution is never materialized in full.

This trainer was contributed by Carlos Miguel Patiño.

Quick start

This example demonstrates how to train a model using the distillation method. We distill a Qwen 2.5 0.5B Instruct model from a Qwen 2.5 1.5B Instruct teacher on the prompts from the UltraFeedback prompt dataset. You can view the data in the dataset here:

Below is the script to train the model.

# train_distillation.py
from datasets import load_dataset
from trl import DistillationTrainer

dataset = load_dataset("trl-lib/ultrafeedback-prompt", split="train")

trainer = DistillationTrainer(
    model="Qwen/Qwen2.5-0.5B-Instruct",
    teacher_model="Qwen/Qwen2.5-1.5B-Instruct",
    train_dataset=dataset,
)
trainer.train()

Execute the script using the following command:

accelerate launch train_distillation.py

Looking deeper into the distillation method

On-policy knowledge distillation trains a student to reproduce a teacher’s next-token distribution on completions the student generates itself, rather than on a fixed set of teacher outputs. Learning from its own generations lets the student correct its own mistakes, which generally outperforms off-policy distillation. This section breaks down how it works in practice, covering the two key steps: generating completions and computing the loss.

Generating completions

At each training step, the student generates a batch of completions for the sampled prompts.

Computing the loss

The loss is the generalized Jensen-Shannon divergence (JSD) between the student distribution pS p_S and the teacher distribution pT p_T over the generated completion tokens, interpolated by beta and defined as: Lβ=βDKL ⁣[pTpM]+(1β)DKL ⁣[pSpM],pM=(1β)pS+βpT, \mathcal{L}_\beta = \beta \, \mathbb{D}_{\mathrm{KL}}\!\left[ p_T \| p_M \right] + (1 - \beta) \, \mathbb{D}_{\mathrm{KL}}\!\left[ p_S \| p_M \right], \qquad p_M = (1 - \beta) \, p_S + \beta \, p_T,

where pM p_M is the β \beta -mixture of the two distributions. The endpoints reduce to the pure divergences: beta=0.0 gives the forward KL DKL ⁣[pTpS] \mathbb{D}_{\mathrm{KL}}\!\left[ p_T \| p_S \right] and beta=1.0 the reverse KL DKL ⁣[pSpT] \mathbb{D}_{\mathrm{KL}}\!\left[ p_S \| p_T \right] .

In practice, the projection to vocabulary logits and the divergence are computed in chunks, so peak activation memory does not scale with the full vocabulary × sequence-length logits tensor. See Reducing Memory Usage.

Expected dataset type

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

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

Logged metrics

While training and evaluating, we record the following metrics:

  • num_tokens: The total number of tokens processed so far, including both prompts and completions.
  • step_time: The average time (in seconds) taken per training step (including generation).
  • completions/mean_length: The average length of generated completions.
  • completions/min_length: The minimum length of generated completions.
  • completions/max_length: The maximum length of generated completions.
  • completions/mean_terminated_length: The average length of generated completions that terminate with EOS.
  • completions/min_terminated_length: The minimum length of generated completions that terminate with EOS.
  • completions/max_terminated_length: The maximum length of generated completions that terminate with EOS.
  • completions/clipped_ratio: The ratio of truncated (clipped) completions.
  • entropy: Average entropy of token predictions across generated completions (in nats). Not logged on the Liger fast path.

Customization

Speed up training with vLLM-powered generation

Generation is often the main bottleneck when training with on-policy methods. To accelerate generation, you can use vLLM, a high-throughput, low-latency inference engine for LLMs. To enable it, first install the package with

pip install trl[vllm]

We support two ways of using vLLM during training: colocate mode and server mode.

Option 1: Colocate mode

In this mode, vLLM runs inside the trainer process and shares GPU memory with the training model. This avoids launching a separate server and can improve GPU utilization, but may lead to memory contention on the training GPUs. This is the default mode.

from trl import DistillationConfig

training_args = DistillationConfig(
    ...,
    use_vllm=True,  # vllm_mode="colocate" by default
)

Option 2: Server mode

In this mode, vLLM runs in a separate process (and using separate GPUs) and communicates with the trainer via HTTP. This is ideal if you have dedicated GPUs for inference.

  1. Start the vLLM server:

    trl vllm-serve --model <model_name>
  2. Enable server mode in your training script:

    from trl import DistillationConfig
    
    training_args = DistillationConfig(
        ...,
        use_vllm=True,
        vllm_mode="server",
    )

Make sure that the server is using different GPUs than the trainer, otherwise you may run into NCCL errors. You can specify the GPUs to use with the CUDA_VISIBLE_DEVICES environment variable.

Depending on the model size and the overall GPU memory requirements for training, you may need to adjust the vllm_gpu_memory_utilization parameter in DistillationConfig to avoid underutilization or out-of-memory errors.

For more information, see Speeding up training with vLLM.

Train adapters with PEFT

We support tight integration with the 🤗 PEFT library, letting you train adapters and share them on the Hub rather than training the whole student.

from datasets import load_dataset
from trl import DistillationTrainer
from peft import LoraConfig

dataset = load_dataset("trl-lib/ultrafeedback-prompt", split="train")

trainer = DistillationTrainer(
    model="Qwen/Qwen2.5-0.5B-Instruct",
    teacher_model="Qwen/Qwen2.5-1.5B-Instruct",
    train_dataset=dataset,
    peft_config=LoraConfig(),
)
trainer.train()

The distillation loss reads lm_head.weight directly and runs the student backbone without going through PeftModel.forward(). Adapters on lm_head (via target_modules) and prompt-learning methods (PromptTuning, PrefixTuning, P-Tuning) are therefore rejected, since they would be silently ignored. To train the head, use modules_to_save=["lm_head"] instead.

Train with Liger Kernel

Liger Kernel is a collection of Triton kernels for LLM training that boosts multi-GPU throughput, cuts memory use, and works seamlessly with tools like FlashAttention, PyTorch FSDP, and DeepSpeed. For more information, see Liger Kernel Integration.

Set use_liger_kernel=True in the DistillationConfig to compute the JSD with the fused Liger kernel instead of the chunked path.

The fused Liger kernel cannot apply per-model logit_scale (e.g. Cohere) or final_logit_softcapping (e.g. Gemma), so it is rejected for models that set them — use the default chunked path for those.

Training Vision Language Models

DistillationTrainer supports distilling Vision-Language Models (VLMs) on multimodal datasets containing both text and images. Pass a VLM as both the student and the teacher, and provide a prompt-only dataset with either an image column (single image per sample) or an images column (list of images per sample). For more information on the expected dataset structure, see the Dataset Format — Vision datasets section.

Tested with:

  • Gemma 3 — e.g., google/gemma-3-4b-it
  • LLaVA-NeXT — e.g., llava-hf/llava-v1.6-mistral-7b-hf
  • Qwen2-VL — e.g., Qwen/Qwen2-VL-2B-Instruct
  • Qwen2.5-VL — e.g., Qwen/Qwen2.5-VL-3B-Instruct

Compatibility with all VLMs is not guaranteed. If you believe a model should be supported, feel free to open an issue on GitHub — or better yet, submit a pull request with the required changes.

Example script

Use examples/scripts/distillation.py to launch distillation training from the command line. The script supports full training and LoRA via the standard ModelConfig flags.

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

class trl.DistillationTrainer

< >

( model: str | PreTrainedModel | PeftModelteacher_model: str | transformers.modeling_utils.PreTrainedModel = Noneargs: trl.trainer.distillation_config.DistillationConfig | None = Nonetrain_dataset: datasets.arrow_dataset.Dataset | None = Noneeval_dataset: datasets.arrow_dataset.Dataset | dict[str, datasets.arrow_dataset.Dataset] | None = Noneprocessing_class: transformers.tokenization_utils_base.PreTrainedTokenizerBase | transformers.processing_utils.ProcessorMixin | None = Nonecallbacks: list[transformers.trainer_callback.TrainerCallback] | None = Noneoptimizers: tuple = (None, None)quantization_config: BitsAndBytesConfig | None = Nonepeft_config: typing.Optional[ForwardRef('PeftConfig')] = None )

Parameters

  • model (str or PreTrainedModel or PeftModel) — Model to be trained. Can be either:

    • A string, being the model id of a pretrained model hosted inside a model repo on huggingface.co, or a path to a directory containing model weights saved using save_pretrained, e.g., './my_model_directory/'. The model is loaded using <ModelArchitecture>.from_pretrained (where <ModelArchitecture> is derived from the model config) with the keyword arguments in args.model_init_kwargs. If dtype is not specified in args.model_init_kwargs, it defaults to float32. This differs from from_pretrained, where (since Transformers v5) the dtype is inferred from the model config.
    • A PreTrainedModel object. Only causal language models are supported.
    • A PeftModel object. Only causal language models are supported.
  • teacher_model (str or PreTrainedModel, optional) — Teacher model whose next-token distribution the student is trained to match. Can be a model id / path (loaded like model, using args.teacher_model_init_kwargs) or an instantiated PreTrainedModel. It must share the student’s vocabulary. May be omitted by subclasses that supply the teacher another way (e.g. a remote server).
  • args (DistillationConfig, optional) — Configuration for this trainer. If None, a default configuration is used.
  • train_dataset (Dataset or IterableDataset, optional) — Dataset to use for training. It must include a column "prompt". Any additional columns in the dataset is ignored. The format of the samples can be either:

    • Standard: Each sample contains plain text.
    • Conversational: Each sample contains structured messages (e.g., role and content).

    When train_dataset is an IterableDataset (e.g. a streaming dataset), max_steps must be set in the training arguments, since its length cannot be inferred and the total number of training steps is required to bound the training loop and configure the learning rate scheduler.

  • eval_dataset (Dataset, IterableDataset, DatasetDict, IterableDatasetDict or dict[str, Dataset | IterableDataset]) — Dataset to use for evaluation. It must meet the same requirements as train_dataset.
  • processing_class (PreTrainedTokenizerBase, ProcessorMixin, optional) — Processing class used to process the data. The padding side must be set to “left”. If None, the processing class is loaded from the model’s name with from_pretrained. A padding token, tokenizer.pad_token, must be set. If the processing class has not set a padding token, tokenizer.eos_token will be used as the default.
  • callbacks (list of TrainerCallback, optional) — List of callbacks to customize the training loop. Will add those to the list of default callbacks detailed in here.

    If you want to remove one of the default callbacks used, use the remove_callback method.

  • optimizers (tuple[torch.optim.Optimizer | None, torch.optim.lr_scheduler.LambdaLR | None], optional, defaults to (None, None)) — A tuple containing the optimizer and the scheduler to use. Will default to an instance of AdamW on your model and a scheduler given by get_linear_schedule_with_warmup controlled by args.
  • quantization_config (BitsAndBytesConfig, optional) — Quantization configuration used when loading the model from a model identifier. Combine with peft_config for QLoRA training. Ignored if the model is already instantiated.
  • peft_config (PeftConfig, optional) — PEFT configuration used to wrap the model. If None, the model is not wrapped.

Trainer for knowledge distillation. The student is trained on-policy — it generates the completions itself — to match the teacher’s next-token distribution under a generalized Jensen-Shannon divergence (interpolating forward KL, reverse KL, and JSD via beta), as introduced in On-Policy Distillation of Language Models.

Example:

>>> from trl import DistillationTrainer
>>> from datasets import load_dataset

>>> dataset = load_dataset("trl-lib/tldr", split="train")

>>> trainer = DistillationTrainer(
...     model="Qwen/Qwen2.5-0.5B-Instruct",
...     teacher_model="Qwen/Qwen2.5-1.5B-Instruct",
...     train_dataset=dataset,
... )
>>> trainer.train()

train

< >

( resume_from_checkpoint: str | bool | None = Nonetrial: optuna.Trial | dict[str, Any] | None = Noneignore_keys_for_eval: list[str] | None = None ) ~trainer_utils.TrainOutput

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

< >

( output_dir: str | None = None_internal_call: bool = False )

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

Will only save from the main process.

push_to_hub

< >

( commit_message: str | None = 'End of training'blocking: bool = Truetoken: str | None = Nonerevision: str | None = None**kwargs )

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.

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

DistillationConfig

class trl.DistillationConfig

< >

( output_dir: str | None = Noneper_device_train_batch_size: int = 8num_train_epochs: float = 3.0max_steps: int = -1learning_rate: float = 1e-06lr_scheduler_type: transformers.trainer_utils.SchedulerType | str = 'linear'lr_scheduler_kwargs: dict | str | None = Nonewarmup_steps: float = 0optim: transformers.training_args.OptimizerNames | str = 'adamw_torch_fused'optim_args: str | None = Noneweight_decay: float = 0.0adam_beta1: float = 0.9adam_beta2: float = 0.999adam_epsilon: float = 1e-08optim_target_modules: None | str | list[str] = Nonegradient_accumulation_steps: int = 1average_tokens_across_devices: bool = Truemax_grad_norm: float = 1.0label_smoothing_factor: float = 0.0bf16: bool | None = Nonefp16: bool = Falsebf16_full_eval: bool = Falsefp16_full_eval: bool = Falsetf32: bool | None = Nonegradient_checkpointing: bool = Truegradient_checkpointing_kwargs: dict[str, typing.Any] | str | None = Nonetorch_compile: bool = Falsetorch_compile_backend: str | None = Nonetorch_compile_mode: str | None = Noneuse_liger_kernel: bool = Falseliger_kernel_config: dict[str, bool] | None = Noneuse_cache: bool = Falseneftune_noise_alpha: float | None = Nonetorch_empty_cache_steps: int | None = Noneauto_find_batch_size: bool = Falselogging_strategy: transformers.trainer_utils.IntervalStrategy | str = 'steps'logging_steps: float = 10logging_first_step: bool = Falselog_on_each_node: bool = Truelogging_nan_inf_filter: bool = Trueinclude_num_input_tokens_seen: str | bool = 'no'log_level: str = 'passive'log_level_replica: str = 'warning'disable_tqdm: bool | None = Nonereport_to: None | str | list[str] = 'none'run_name: str | None = Noneproject: str = 'huggingface'trackio_space_id: str | None = Nonetrackio_bucket_id: str | None = Nonetrackio_static_space_id: typing.Union[str, NoneType, typing.Literal[False]] = Noneeval_strategy: transformers.trainer_utils.IntervalStrategy | str = 'no'eval_steps: float | None = Noneeval_delay: float = 0per_device_eval_batch_size: int = 8prediction_loss_only: bool = Falseeval_on_start: bool = Falseeval_do_concat_batches: bool = Trueeval_use_gather_object: bool = Falseeval_accumulation_steps: int | None = Noneinclude_for_metrics: list = <factory>batch_eval_metrics: bool = Falsesave_only_model: bool = Falsesave_strategy: transformers.trainer_utils.SaveStrategy | str = 'steps'save_steps: float = 500save_on_each_node: bool = Falsesave_total_limit: int | None = Noneenable_jit_checkpoint: bool = Falsepush_to_hub: bool = Falsehub_token: str | None = Nonehub_private_repo: bool | None = Nonehub_model_id: str | None = Nonehub_strategy: transformers.trainer_utils.HubStrategy | str = 'every_save'hub_always_push: bool = Falsehub_revision: str | None = Noneload_best_model_at_end: bool = Falsemetric_for_best_model: str | None = Nonegreater_is_better: bool | None = Noneignore_data_skip: bool = Falserestore_callback_states_from_checkpoint: bool = Falsefull_determinism: bool = Falseseed: int = 42data_seed: int | None = Noneuse_cpu: bool = Falseaccelerator_config: dict | str | None = Noneparallelism_config: accelerate.parallelism_config.ParallelismConfig | None = Nonedataloader_drop_last: bool = Falsedataloader_num_workers: int = 0dataloader_pin_memory: bool = Truedataloader_persistent_workers: bool = Falsedataloader_prefetch_factor: int | None = Nonedataloader_multiprocessing_context: str | None = Nonedataloader_in_order: bool = Trueremove_unused_columns: bool | None = Falselabel_names: list[str] | None = Nonetrain_sampling_strategy: str = 'random'length_column_name: str = 'length'ddp_find_unused_parameters: bool | None = Noneddp_bucket_cap_mb: int | None = Noneddp_broadcast_buffers: bool | None = Noneddp_static_graph: bool | None = Noneddp_backend: str | None = Noneddp_timeout: int = 1800fsdp: str | None = Nonefsdp_config: dict[str, typing.Any] | str | None = Nonedeepspeed: dict | str | None = Nonedebug: str | list[transformers.debug_utils.DebugOption] = ''skip_memory_metrics: bool = Truedo_train: bool = Falsedo_eval: bool = Falsedo_predict: bool = Falseresume_from_checkpoint: str | None = Nonelocal_rank: int = -1model_init_kwargs: dict[str, typing.Any] | str | None = Nonetrust_remote_code: bool = Falseteacher_model_name_or_path: str | None = Noneteacher_model_revision: str | None = Noneteacher_model_init_kwargs: dict[str, typing.Any] | str | None = Nonedisable_dropout: bool = Falsemax_completion_length: int | None = 512ds3_gather_for_generation: bool = Trueshuffle_dataset: bool | None = Truepad_to_multiple_of: int | None = Nonetemperature: float = 1.0top_p: float = 1.0top_k: int = 0min_p: float | None = Nonegeneration_kwargs: dict | None = Nonechat_template_kwargs: dict | None = Nonerepetition_penalty: float = 1.0cache_implementation: str | None = Noneuse_vllm: bool = Falsevllm_mode: str = 'colocate'vllm_model_impl: str = 'vllm'vllm_enable_sleep_mode: bool = Falsevllm_structured_outputs_regex: str | None = Nonevllm_server_base_url: str | None = Nonevllm_server_host: str = '0.0.0.0'vllm_server_port: int = 8000vllm_server_timeout: float = 240.0vllm_group_port: int = 51216vllm_gpu_memory_utilization: float = 0.3vllm_max_model_length: int | None = Nonevllm_tensor_parallel_size: int = 1beta: float = 1.0log_completions: bool = Falsenum_completions_to_print: int | None = Nonelog_unique_prompts: bool = False )

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 and 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 (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, 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, 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, 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 with parameters specific to knowledge distillation. All necessary fields are declared here.

Using HfArgumentParser we can turn this class into argparse arguments that can be specified on the command line.

These parameters have default values different from TrainingArguments:

  • logging_steps: Defaults to 10 instead of 500.
  • gradient_checkpointing: Defaults to True instead of False.
  • bf16: Defaults to True if fp16 is not set, instead of False.
  • learning_rate: Defaults to 1e-6 instead of 5e-5.
Update on GitHub