Distillers

.. py:function:: Distiller(distiller_type: str, **kwargs: ~typing.Any) -> ~silverspoon_kd.distillers.base_distiller.BaseDistiller

module:

silverspoon_kd.distillers

Factory that creates a distiller by type name.

A convenience wrapper that maps a string to the concrete distiller class. Useful for configuration-driven workflows where the distiller type comes from a config file or CLI argument.

Distiller types:

  • "blockwise" / "bkd" → [BlockwiseDistiller][silverspoon_kd.BlockwiseDistiller]

  • "holistic" / "hkd" → [HolisticDistiller][silverspoon_kd.HolisticDistiller]

  • "response_based" / "reskd" → [ResponseBasedDistiller][silverspoon_kd.ResponseBasedDistiller]

param distiller_type:

One of the type names listed above.

param **kwargs:

Arguments forwarded to the concrete distiller constructor. See the individual class docstrings for accepted parameters.

returns:

An instance of the selected distiller.

raises ValueError:

If distiller_type is not recognized.

Example::

distiller = Distiller(
    distiller_type="blockwise",
    teacher_model=teacher,
    alignments=alignments,
    args=TrainingArguments(output_dir="./out", backward_per_block=True),
    train_dataset=dataset,
)
distiller.train()

.. py:class:: BlockwiseDistiller(teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, alignments: list[~silverspoon_kd.alignments.alignment.Alignment], train_dataset=None, eval_dataset=None, data_collator=None, args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, auto_truncate: bool = False, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, student_models: dict[str, ~torch.nn.modules.module.Module] | None = None, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

canonical:

silverspoon_kd.distillers.blockwise_distiller.BlockwiseDistiller

Bases: :py:class:~silverspoon_kd.distillers.base_distiller.BaseDistiller

Block-Wise Knowledge Distillation (BKD), a type of Feature-Based Knowledge Distillation (FBKD).

A trainer that performs block-wise (layer-by-layer) knowledge distillation. The student is divided into blocks trained in isolation. Each student block receives input directly from the corresponding teacher block, treating each block as an independent regression problem.

Student blocks and their projectors are held inside a StudentBlocksContainer so that the standard Trainer save/load/zero_grad infrastructure works.

.. py:attribute:: BlockwiseDistiller.args

module:

silverspoon_kd.distillers

type:

~silverspoon_kd.training_arguments.TrainingArguments

.. py:method:: BlockwiseDistiller.init(teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, alignments: list[~silverspoon_kd.alignments.alignment.Alignment], train_dataset=None, eval_dataset=None, data_collator=None, args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, auto_truncate: bool = False, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, student_models: dict[str, ~torch.nn.modules.module.Module] | None = None, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

Initialize the BlockwiseDistiller.

param teacher_model:

The teacher model to distill from

param alignments:

List of Alignment instances

param train_dataset:

Training dataset

param eval_dataset:

Evaluation dataset

param data_collator:

Data collator for batching

param args:

TrainingArguments with distillation-specific parameters

param auto_truncate:

If True, automatically stop the teacher forward pass as soon as all aligned modules have been captured. For example, if a 24-layer teacher is aligned at layers 0-5, the remaining 18 layers are skipped — saving compute and memory. Defaults to False because exception-based truncation is incompatible with FSDP and torch.compile. Safe to enable for single-GPU or DDP setups without torch.compile.

param prepare_teacher_inputs:

Optional callable to transform teacher inputs

param student_models:

Optional dict mapping names to complete student models for checkpoint saving

param **kwargs:

Additional keyword arguments passed to Trainer

.. py:method:: BlockwiseDistiller.compute_distillation_loss(model, inputs, is_training)

module:

silverspoon_kd.distillers

Run teacher forward, iterate captured blocks, compute student losses.

.. py:method:: BlockwiseDistiller.create_optimizer(model=None)

module:

silverspoon_kd.distillers

Sync FSDP-wrapped block references before building per-alignment optimizers.

.. py:method:: BlockwiseDistiller.save_model(output_dir=None, _internal_call=False)

module:

silverspoon_kd.distillers

Reset FSDP states before save.

Per-block forward calls (bypassing root FSDP forward) leave inner FSDP units in BACKWARD_POST state and with stale _is_root flags. Reset both so state_dict() can run without assertion errors.

.. py:class:: HolisticDistiller(student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, alignments: list[~silverspoon_kd.alignments.alignment.Alignment], train_dataset=None, eval_dataset=None, data_collator=None, args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, auto_truncate: bool = False, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, prepare_student_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, overlap_alignment_loss: bool = False, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

canonical:

silverspoon_kd.distillers.holistic_distiller.HolisticDistiller

Bases: :py:class:~silverspoon_kd.distillers.base_distiller.BaseDistiller

Holistic Knowledge Distillation (HKD), a type of Feature-Based Knowledge Distillation (FBKD).

A trainer that performs end-to-end knowledge distillation. The entire student network is optimized jointly, allowing error gradients to propagate through the whole network.

Unlike [BlockwiseDistiller][silverspoon_kd.BlockwiseDistiller] which trains each block independently with immediate backpropagation, HolisticDistiller runs complete forward passes through both teacher and student models, accumulates alignment losses across all blocks, and performs a single backpropagation step.

This approach:

  1. Runs teacher model forward pass with capture engine to record activations

  2. Runs student model forward pass with capture engine to record activations

  3. Computes alignment losses between corresponding teacher-student layer pairs

  4. Sums all losses and performs a single backward pass

  5. Updates all student parameters together via a single optimizer covering the full student model plus any alignment projectors.

.. py:method:: HolisticDistiller.init(student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, alignments: list[~silverspoon_kd.alignments.alignment.Alignment], train_dataset=None, eval_dataset=None, data_collator=None, args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, auto_truncate: bool = False, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, prepare_student_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, overlap_alignment_loss: bool = False, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

Initialize the HolisticDistiller.

param student_model:

The student model to train

param teacher_model:

The teacher model to distill from

param alignments:

List of Alignment instances

param train_dataset:

Training dataset

param eval_dataset:

Evaluation dataset

param data_collator:

Data collator for batching

param args:

TrainingArguments with distillation-specific parameters

param auto_truncate:

If True, automatically stop each forward pass as soon as all aligned modules have been captured. For example, if a 24-layer teacher is aligned at layers 0-5, the remaining 18 layers are skipped — saving compute and memory. Defaults to False because exception-based truncation is incompatible with FSDP and torch.compile. Safe to enable for single-GPU or DDP setups without torch.compile.

param prepare_teacher_inputs:

Optional callable to transform inputs before passing to teacher

param prepare_student_inputs:

Optional callable to transform inputs before passing to student

param overlap_alignment_loss:

Opt-in optimisation, default False. When True and CUDA is available, each alignment’s loss kernel runs on a dedicated CUDA stream so it overlaps with the next student layer’s forward on the default stream. Gradients are equal to the False path within floating-point tolerance (the parity is verified in tests).

Off by default because the speed-up is small in practice and there are two real friction points to be aware of:

  • AccumulateGrad stream mismatch. PyTorch creates each parameter’s AccumulateGrad autograd node on whichever stream first produced its gradient. With overlap on, that’s the loss stream; on the next iteration backward hits it from the default stream and PyTorch warns about a stream mismatch. This adds an implicit synchronisation (eroding the perf gain) and breaks CUDA-graph capture — set this to False if you use torch.compile with CUDA graphs or anything else that requires a clean graph.

  • Floating-point reordering. Stream-level reordering of loss kernels can produce non-bitwise-identical losses vs the default-stream path. Tests use torch.allclose with a small tolerance for this reason.

Auto-disabled on CPU.

param **kwargs:

Keyword arguments passed to Trainer

.. py:method:: HolisticDistiller.create_optimizer(model=None)

module:

silverspoon_kd.distillers

Single optimizer over the full student + alignment projectors.

Delegates to HF Trainer’s default, then wires each alignment’s .optimizer at the main optimizer so that lazy auto-projectors created on the first forward pass are registered via Alignment._add_projector_params_to_optimizer.

.. py:method:: HolisticDistiller.compute_distillation_loss(model, inputs, is_training)

module:

silverspoon_kd.distillers

Run teacher and student forward passes, compute alignment losses.

Alignment losses are computed incrementally inside the student forward via :meth:_on_student_capture — each layer’s loss is accumulated as that layer’s hook fires, then the captured tensors are dropped. Peak alignment-activation memory is O(1) instead of O(N alignments). Gradients are mathematically identical to the previous all-at-once path (same weighted sum, same autograd graph).

When alpha > 0 and labels are present, the student model’s own loss (cross-entropy against labels) is mixed into the total: total = (1-alpha) * alignment_loss + alpha * hard_loss.

.. py:class:: ResponseBasedDistiller(student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, train_dataset=None, eval_dataset=None, data_collator=None, args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, soft_loss_fn: str | ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor] | None = None, soft_loss_fn_kwargs: dict[str, ~typing.Any] | None = None, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, prepare_student_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, use_liger_kernel: bool = False, output_head_layer: str | None = None, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

canonical:

silverspoon_kd.distillers.response_based_distiller.ResponseBasedDistiller

Bases: :py:class:~silverspoon_kd.distillers.base_distiller.BaseDistiller

Response-based knowledge distillation (ResKD).

A distiller for response-based (output-level) knowledge distillation. The student learns to match the teacher’s final output predictions (soft targets/logits), in contrast to feature-based approaches ([HolisticDistiller][silverspoon_kd.HolisticDistiller], [BlockwiseDistiller][silverspoon_kd.BlockwiseDistiller]) which match intermediate representations.

The soft loss function is user-provided (any callable matching the (student_logits, teacher_logits) → scalar signature used by all registered loss functions). The hard loss comes from the model’s own outputs.loss (the standard HuggingFace convention).

.. rubric:: Example

from silverspoon_kd.losses import kl_divergence_loss

args = TrainingArguments(alpha=0.5)
distiller = ResponseBasedDistiller(
    student_model=student,
    teacher_model=teacher,
    soft_loss_fn=kl_divergence_loss(temperature=4.0, chunk_size=1024),
    train_dataset=dataset,
    args=args,
)
distiller.train()

.. py:attribute:: ResponseBasedDistiller.args

module:

silverspoon_kd.distillers

type:

~silverspoon_kd.training_arguments.TrainingArguments

.. py:method:: ResponseBasedDistiller.init(student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, train_dataset=None, eval_dataset=None, data_collator=None, args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, soft_loss_fn: str | ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor] | None = None, soft_loss_fn_kwargs: dict[str, ~typing.Any] | None = None, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, prepare_student_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, use_liger_kernel: bool = False, output_head_layer: str | None = None, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

Initialize the ResponseBasedDistiller.

param student_model:

The student model to train. Should output logits.

param teacher_model:

The teacher model to distill from. Should output logits.

param train_dataset:

Training dataset

param eval_dataset:

Evaluation dataset

param data_collator:

Data collator for batching

param args:

TrainingArguments with training parameters. Contains alpha (soft/hard weighting), magnitude_aware_weighting, and device/dtype matching flags.

param soft_loss_fn:

Loss function for soft targets. Can be:

  • A callable (student_logits, teacher_logits) → scalar

  • A string name from the loss registry (e.g. “kl_div”, “jsd”)

  • None (defaults to kl_divergence_loss())

param soft_loss_fn_kwargs:

Optional keyword arguments forwarded to the loss factory when soft_loss_fn is a string name. Ignored when soft_loss_fn is already a callable.

param prepare_teacher_inputs:

Optional callable to transform inputs before passing to teacher.

param prepare_student_inputs:

Optional callable to transform inputs before passing to student.

param use_liger_kernel:

Whether to use Liger fused linear kernel for memory-efficient loss computation. When True, the soft loss is computed by a fused kernel that avoids materializing the full logit tensor. Requires the liger-kernel package and output_head_layer to be set. Default: False

param output_head_layer:

Attribute path to the output head linear layer (e.g. “lm_head” for language models, “classifier” for classification models). Required when use_liger_kernel=True.

param **kwargs:

Additional keyword arguments passed to Trainer

.. py:method:: ResponseBasedDistiller.compute_distillation_loss(model, inputs, is_training)

module:

silverspoon_kd.distillers

Run teacher/student forward passes and compute distillation loss.

.. py:method:: ResponseBasedDistiller.log(logs: dict[str, float], start_time: float | None = None) -> None

module:

silverspoon_kd.distillers

Override log to inject eval-specific component losses.

During evaluation, averaged component losses (eval_loss/soft, eval_loss/hard) are injected from the accumulator populated by compute_loss().

Training metrics are handled by BaseDistiller.log() via the accumulator.

param logs:

Dictionary of metrics to log

param start_time:

Optional start time for computing throughput metrics

.. py:method:: ResponseBasedDistiller.evaluate(eval_dataset: ~typing.Any | None = None, ignore_keys: list | None = None, metric_key_prefix: str = ‘eval’) -> dict[str, float]

module:

silverspoon_kd.distillers

Evaluate the student model, reporting component loss breakdown (loss/soft, loss/hard) alongside the aggregate eval_loss.

Optionally runs WeightWatcher analysis.

param eval_dataset:

Dataset to evaluate on (defaults to self.eval_dataset)

param ignore_keys:

Output keys to ignore

param metric_key_prefix:

Prefix for metric names

returns:

Dictionary of evaluation metrics

.. py:class:: BaseDistiller(teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, alignments: list[~silverspoon_kd.alignments.alignment.Alignment], args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, train_dataset=None, eval_dataset=None, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, prepare_student_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

canonical:

silverspoon_kd.distillers.base_distiller.BaseDistiller

Bases: :py:class:~transformers.trainer.Trainer

Abstract base class for knowledge distillation trainers.

This class extends HuggingFace’s Trainer to provide common functionality for all distillation approaches:

  • Metrics tracking (including per-layer metrics)

  • FLOP counting

  • Capture engine lifecycle management

  • Teacher input preparation

Subclasses must implement:

  • compute_distillation_loss(model, inputs, is_training): Define the distillation loss computation. Called by the template methods training_step() and compute_loss() provided by this base class.

Subclasses may optionally override:

  • _backward_loss(loss): Custom backward pass (default: accelerator.backward)

  • _after_training_step(loss): Hook called after each training step

  • _save_distiller_state(output_dir) / _load_distiller_state(checkpoint_dir): Custom checkpoint save/load logic (call super() when overriding)

  • _get_e2e_student_models(): Return full student models for end-to-end eval

  • _compute_e2e_eval_loss(eval_dataset, metric_key_prefix): Custom e2e eval

.. py:attribute:: BaseDistiller.args

module:

silverspoon_kd.distillers

type:

~silverspoon_kd.training_arguments.TrainingArguments

.. py:property:: BaseDistiller.student_model

module:

silverspoon_kd.distillers

Alias for self.model — the student model being trained.

.. py:method:: BaseDistiller.init(teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, alignments: list[~silverspoon_kd.alignments.alignment.Alignment], args: ~silverspoon_kd.training_arguments.TrainingArguments | None = None, train_dataset=None, eval_dataset=None, prepare_teacher_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, prepare_student_inputs: ~collections.abc.Callable[[dict[str, ~typing.Any]], dict[str, ~typing.Any]] | None = None, **kwargs: ~typing.Any)

module:

silverspoon_kd.distillers

Initialize the BaseDistiller.

param teacher_model:

The teacher model to distill from

param alignments:

List of Alignment instances (flat, one teacher-student pair each)

param args:

TrainingArguments with distillation-specific parameters

param train_dataset:

Training dataset

param eval_dataset:

Evaluation dataset

param prepare_teacher_inputs:

Optional callable to transform inputs before passing to teacher. Takes inputs dict and returns modified dict. If None, passes inputs as-is.

param **kwargs:

Additional keyword arguments passed to Trainer

.. py:attribute:: BaseDistiller.alignments

module:

silverspoon_kd.distillers

type:

list[~silverspoon_kd.alignments.alignment.Alignment]

.. py:attribute:: BaseDistiller.teacher_model

module:

silverspoon_kd.distillers

type:

~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module

.. py:attribute:: BaseDistiller.current_step_metrics

module:

silverspoon_kd.distillers

type:

dict[str, float | ~torch.Tensor]

.. py:attribute:: BaseDistiller.step_losses

module:

silverspoon_kd.distillers

type:

list[~torch.Tensor]

.. py:attribute:: BaseDistiller.eval_losses

module:

silverspoon_kd.distillers

type:

list[~torch.Tensor]

.. py:attribute:: BaseDistiller.flop_counter

module:

silverspoon_kd.distillers

type:

int

.. py:attribute:: BaseDistiller.flops_per_step

module:

silverspoon_kd.distillers

type:

int

.. py:method:: BaseDistiller.create_optimizer(model: ~torch.nn.modules.module.Module | None = None)

module:

silverspoon_kd.distillers

Override Trainer’s create_optimizer to use CompositeOptimizer.

Called by the Trainer after FSDP/DeepSpeed wrapping (if any), so parameter references are up-to-date at this point.

Under DeepSpeed, CompositeOptimizer is incompatible (DeepSpeed needs a single flat optimizer for ZeRO partitioning). Instead, we build a single standard optimizer with per-alignment param_groups tagged by _alignment_name, preserving per-alignment LR tracking and grad clipping.

param model:

The model to create the optimizer for. Passed by newer versions of the HF Trainer but unused here (we build optimizers from alignment parameters directly).

.. py:method:: BaseDistiller.create_scheduler(num_training_steps, optimizer=None)

module:

silverspoon_kd.distillers

Override Trainer’s create_scheduler to use CompositeScheduler.

.. py:method:: BaseDistiller.train(*args: ~typing.Any, **kwargs: ~typing.Any) -> ~typing.Any

module:

silverspoon_kd.distillers

Train the student model using knowledge distillation.

This method:

  1. Sets up teacher placement (if configured)

  2. Initializes FLOP counter

  3. Starts profiler (if enabled)

  4. Registers forward hooks and input capture wrappers

  5. Performs training

  6. Deregisters hooks after training

  7. Stops profiler and generates report (if enabled)

param *args:

Positional arguments passed to parent Trainer.train()

param **kwargs:

Keyword arguments passed to parent Trainer.train()

returns:

Training results from parent Trainer.train()

.. py:method:: BaseDistiller.training_step(model: ~torch.nn.modules.module.Module, inputs: dict[str, ~torch.Tensor | ~typing.Any], num_items_in_batch: ~torch.Tensor | None = None) -> ~torch.Tensor

module:

silverspoon_kd.distillers

Template training step: forward → loss → backward → clip → profiler → clear.

Subclasses implement compute_distillation_loss() for their specific logic.

.. py:method:: BaseDistiller.compute_loss(model: ~torch.nn.modules.module.Module, inputs: dict[str, ~torch.Tensor | ~typing.Any], return_outputs: bool = False, num_items_in_batch: ~torch.Tensor | None = None) -> ~torch.Tensor | tuple[~torch.Tensor, None]

module:

silverspoon_kd.distillers

Template compute_loss: forward → loss → store metrics → clear.

Subclasses implement compute_distillation_loss() for their specific logic.

.. py:method:: BaseDistiller.compute_distillation_loss(model: ~torch.nn.modules.module.Module, inputs: dict[str, ~typing.Any], is_training: bool) -> ~torch.Tensor

module:

silverspoon_kd.distillers

Compute the distillation loss. Subclasses must override this.

.. py:method:: BaseDistiller.prediction_step(model: ~torch.nn.modules.module.Module, inputs: dict[str, ~torch.Tensor | ~typing.Any], prediction_loss_only: bool, ignore_keys: list[str] | None = None) -> tuple[~torch.Tensor | None, ~torch.Tensor | None, ~torch.Tensor | None]

module:

silverspoon_kd.distillers

Override prediction_step to force compute_loss call even without labels.

The default prediction_step skips compute_loss when labels aren’t present, but distillation doesn’t use labels - it aligns hidden states.

When compute_metrics is set and labels are present in inputs, a separate student-only forward pass produces logits so that HF Trainer’s evaluation loop can call compute_metrics. This enables accuracy tracking and accuracy-based early stopping during distillation.

param model:

The student model

param inputs:

Dictionary of input tensors

param prediction_loss_only:

Whether only loss is needed

param ignore_keys:

Keys to ignore in outputs (unused)

returns:

Tuple of (loss, logits, labels). logits and labels are non-None only when compute_metrics is set and labels are available in the inputs.

.. py:method:: BaseDistiller.evaluate(eval_dataset: ~typing.Any | None = None, ignore_keys: list[str] | None = None, metric_key_prefix: str = ‘eval’) -> dict[str, float]

module:

silverspoon_kd.distillers

Evaluate the student model with per-layer metrics.

This method:

  1. Initializes per-layer metric collection

  2. Runs the standard evaluation loop (which calls compute_loss)

  3. Aggregates and logs per-layer metrics

  4. Returns combined metrics

param eval_dataset:

Dataset to evaluate on (defaults to self.eval_dataset)

param ignore_keys:

Output keys to ignore (unused in distillation)

param metric_key_prefix:

Prefix for metric names (e.g., “eval”, “test”)

returns:

Dictionary of evaluation metrics including per-layer losses

.. py:method:: BaseDistiller.log(logs: dict[str, float], start_time: float | None = None) -> None

module:

silverspoon_kd.distillers

Override log method to inject averaged training metrics and filter the global learning_rate metric.

The HuggingFace Trainer averages train/loss over logging_steps batches. This override applies the same averaging to all custom distiller metrics (per-layer losses, learning rates, gradient norms) that have been accumulated via _track_metric.

param logs:

Dictionary of metrics to log

param start_time:

Optional start time for computing throughput metrics

.. py:class:: SilentProgressCallback(max_str_len: int = 100)

module:

silverspoon_kd.distillers

canonical:

silverspoon_kd.distillers.base_distiller.SilentProgressCallback

Bases: :py:class:~transformers.trainer_callback.ProgressCallback

Progress callback that displays progress bars but doesn’t print per-step metrics.

This is useful for distillation where many per-module metrics would clutter the console output. Metrics are still logged to TensorBoard and other configured loggers.

.. py:method:: SilentProgressCallback.on_log(args, state, control, logs=None, **kwargs)

module:

silverspoon_kd.distillers

No-op override that suppresses per-step metric printing.