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.TrainingArgumentsTraining arguments for all SilverSpoon distillers.
Extends HuggingFace
TrainingArgumentswith 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/stepand, 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.
Noneor"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 toTeacherPlacement(**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 (
labelsmust 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 toFalseto 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
fsdpis set,fsdp_config["version"]defaults to1: the distillers rely on FSDP1 and do not support the DTensor-based FSDP2 that newer Trainer versions select by default.