Training Arguments

.. py:class:: TrainingArguments(*args, count_flops: bool = False, use_weightwatcher: bool = False, report_to: str | list[str] | None = None, teacher_placement: ~typing.Any = None, overlap_teacher_forward: bool | None = None, enable_profiling: bool = False, profiling_output_dir: str | None = None, profiling_wait: int = 20, profiling_warmup: int = 3, profiling_active: int = 3, profiling_repeat: int = 1, profiling_with_stack: bool = False, e2e_eval_loss: str | None = None, alpha: float = 0.0, auto_device_match: bool = False, auto_dtype_match: bool = False, deepcopy_captured_args_and_kwargs: bool = False, backward_per_block: bool = False, magnitude_aware_weighting: bool = False, clip_projectors: bool = True, **kwargs)

module:

silverspoon_kd

canonical:

silverspoon_kd.training_arguments.TrainingArguments

Bases: :py:class:~transformers.training_args.TrainingArguments

Training arguments for all SilverSpoon distillers.

Extends HuggingFace TrainingArguments with distillation-specific parameters. Every parameter has a safe default, so the same class works for any distiller type ([BlockwiseDistiller][silverspoon_kd.BlockwiseDistiller], [HolisticDistiller][silverspoon_kd.HolisticDistiller], [ResponseBasedDistiller][silverspoon_kd.ResponseBasedDistiller]). Parameters unused by a particular distiller are simply ignored.

Example::

# Works for any distiller — no need to pick a subclass.
args = TrainingArguments(
    output_dir="./output",
    num_train_epochs=3,
    alpha=0.5,               # ResponseBased: soft/hard weighting
    backward_per_block=True,  # Blockwise: per-block backward
)

.. py:method:: TrainingArguments.init(*args, count_flops: bool = False, use_weightwatcher: bool = False, report_to: str | list[str] | None = None, teacher_placement: ~typing.Any = None, overlap_teacher_forward: bool | None = None, enable_profiling: bool = False, profiling_output_dir: str | None = None, profiling_wait: int = 20, profiling_warmup: int = 3, profiling_active: int = 3, profiling_repeat: int = 1, profiling_with_stack: bool = False, e2e_eval_loss: str | None = None, alpha: float = 0.0, auto_device_match: bool = False, auto_dtype_match: bool = False, deepcopy_captured_args_and_kwargs: bool = False, backward_per_block: bool = False, magnitude_aware_weighting: bool = False, clip_projectors: bool = True, **kwargs)

module:

silverspoon_kd

Initialize TrainingArguments with distillation parameters.

param count_flops:

Measure the FLOPs of the first training step and log them as flops/step and, cumulatively, flops/total.

param use_weightwatcher:

Whether to use WeightWatcher for model analysis.

param report_to:

Where to report training metrics. Defaults to [] (no reporting).

param teacher_placement:

How the teacher model is distributed.

  • None or "replicated" (default): full copy on every rank.

  • "sharded": FSDP full_shard across all ranks.

  • TeacherPlacement(teacher_only_devices=[...], strategy=...): dedicated teacher GPUs with PP/TP/sharded strategy.

  • dict: auto-converted to TeacherPlacement(**dict).

param overlap_teacher_forward:

Whether to overlap teacher and student forward passes on separate CUDA streams. Auto-detected when None (default).

param enable_profiling:

Enable CPU/GPU/memory profiling via torch.profiler.

param profiling_output_dir:

Directory for trace files. Defaults to {output_dir}/profiling.

param profiling_wait:

Steps to skip before profiling begins.

param profiling_warmup:

Warmup steps (not traced).

param profiling_active:

Steps to actively trace.

param profiling_repeat:

Repeat count for wait/warmup/active cycle.

param profiling_with_stack:

Record Python call stacks in traces.

param e2e_eval_loss:

End-to-end evaluation loss mode.

  • None: disabled (default).

  • "forward": use the student model’s own forward loss.

param alpha:

ResponseBased / Holistic. Weight of the student’s own task loss (labels must be present) relative to the distillation loss. alpha=0: distillation loss only. alpha=1: task loss only. total = (1-alpha) * distillation + alpha * task. Default: 0.0.

param auto_device_match:

ResponseBased only. Automatically move teacher logits to the student device before loss computation.

param auto_dtype_match:

ResponseBased only. Automatically cast teacher logits to student dtype before loss computation.

param deepcopy_captured_args_and_kwargs:

Feature-based only. Deepcopy captured inputs in the capture engine. Needed when model hooks modify tensors in-place. Default: False.

param backward_per_block:

Blockwise only. Run backward after each block instead of summing all losses first. Reduces peak memory from O(N blocks) to O(1 block). Default: False.

param magnitude_aware_weighting:

Holistic / ResponseBased. Normalize each loss component by its detached magnitude before applying weights, ensuring gradient contributions are proportional to the specified ratios regardless of raw magnitudes. Default: False.

param clip_projectors:

Feature-based (Holistic + Blockwise). Whether to include projector parameters in gradient clipping. When True (default), projector gradients contribute to the global norm used for clipping. Set to False to clip only the student model parameters, excluding projector params from the norm calculation. This can increase effective student gradient magnitudes when projector gradients are large. Default: True.

param **kwargs:

Additional keyword arguments for HF TrainingArguments. When fsdp is set, fsdp_config["version"] defaults to 1: the distillers rely on FSDP1 and do not support the DTensor-based FSDP2 that newer Trainer versions select by default.