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_typeis 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.BaseDistillerBlock-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
StudentBlocksContainerso 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_POSTstate and with stale_is_rootflags. Reset both sostate_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.BaseDistillerHolistic 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:
Runs teacher model forward pass with capture engine to record activations
Runs student model forward pass with capture engine to record activations
Computes alignment losses between corresponding teacher-student layer pairs
Sums all losses and performs a single backward pass
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 theFalsepath 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
AccumulateGradautograd 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 toFalseif you usetorch.compilewith 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.allclosewith 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
.optimizerat the main optimizer so that lazy auto-projectors created on the first forward pass are registered viaAlignment._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 > 0and 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.BaseDistillerResponse-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) → scalarsignature used by all registered loss functions). The hard loss comes from the model’s ownoutputs.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) → scalarA 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_fnis a string name. Ignored whensoft_loss_fnis 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-kernelpackage andoutput_head_layerto 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.TrainerAbstract 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:
Sets up teacher placement (if configured)
Initializes FLOP counter
Starts profiler (if enabled)
Registers forward hooks and input capture wrappers
Performs training
Deregisters hooks after training
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_metricsis set and labels are present in inputs, a separate student-only forward pass produces logits so that HF Trainer’s evaluation loop can callcompute_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_metricsis 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:
Initializes per-layer metric collection
Runs the standard evaluation loop (which calls compute_loss)
Aggregates and logs per-layer metrics
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/lossoverlogging_stepsbatches. 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.ProgressCallbackProgress 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.