Alignments¶
.. py:function:: create_alignments(teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, modules: str | list[str] | dict[str, str], loss_function: str | ~collections.abc.Callable | None = None, loss_function_kwargs: dict[str, ~typing.Any] | None = None, output_selector_index: int | None = None, auto_projector: bool = True, projector_init: str | ~collections.abc.Callable | None = None, max_grad_norm: float | int | None = None) -> list[~silverspoon_kd.alignments.alignment.Alignment]
- module:
silverspoon_kd.alignments
Create [Alignment][silverspoon_kd.Alignment] instances between teacher and student modules.
This is the primary API for creating alignments. Module names are matched using regex patterns against both models’ named modules.
- param teacher_model:
The teacher model
- param student_model:
The student model (singular)
- param modules:
Module matching specification:
str: Single regex matched against both teacher and student module names
List[str]: Multiple regexes, each matched against both models
Dict[str, str]: Teacher regex -> student replacement pattern (supports backreferences like \1, \2)
- param loss_function:
Optional loss function (string name or callable). Defaults to MSE.
- param loss_function_kwargs:
Optional keyword arguments forwarded to the loss factory when
loss_functionis a string name. Pass model attributes (weights, sub-modules) directly::loss_function="mahal_mse", loss_function_kwargs={ "weight_matrix": teacher.lm_head.weight, },- param output_selector_index:
Optional index for OutputSelector. When provided, sets both teacher_output_selector and student_output_selector to extract the element at this index from tuple outputs.
- param auto_projector:
If True (default), automatically infer and create projectors based on shape mismatches detected during the first forward pass.
- param projector_init:
Initialization for auto-created projectors.
None(default, kaiming),"normal"(std=0.02),"xavier", or a callablefn(module).- param max_grad_norm:
Maximum gradient norm for per-student clipping.
- returns:
List of [Alignment][silverspoon_kd.Alignment] instances
- raises ValueError:
If no modules match the patterns in either model
.. rubric:: Example
# Match same module names in both models alignments = create_alignments(teacher, student, r"model\.layers\.\d+") # Match specific layers alignments = create_alignments(teacher, student, ["model.layers.0", "model.layers.1"]) # Map teacher to student with different names alignments = create_alignments(teacher, student, { r"teacher\.layers\.(\d+)": r"student\.blocks\.\1" }) # Mahalanobis loss with teacher weight matrix alignments = create_alignments( teacher, student, r"model\.layers\.\d+", loss_function="mahal_mse", loss_function_kwargs={"weight_matrix": teacher.lm_head.weight}, )
.. py:class:: Alignment(teacher_block: ~torch.nn.modules.module.Module, student_block: ~torch.nn.modules.module.Module, teacher_model_name: str = ‘’, student_model_name: str = ‘’, teacher_module_name: str = ‘’, student_module_name: str = ‘’, teacher_output_selector: ~collections.abc.Callable[[~torch.Tensor], ~torch.Tensor] =
- module:
silverspoon_kd.alignments
- canonical:
silverspoon_kd.alignments.alignment.Alignment
Bases: :py:class:
objectA single teacher-module <-> student-module pair for knowledge distillation.
This class encapsulates the complete alignment between one teacher block and one student block, including output selectors, projectors, loss function, optimizer, scheduler, and device/dtype matching configuration.
.. py:method:: Alignment.init(teacher_block: ~torch.nn.modules.module.Module, student_block: ~torch.nn.modules.module.Module, teacher_model_name: str = ‘’, student_model_name: str = ‘’, teacher_module_name: str = ‘’, student_module_name: str = ‘’, teacher_output_selector: ~collections.abc.Callable[[~torch.Tensor], ~torch.Tensor] =
, student_output_selector: ~collections.abc.Callable[[~torch.Tensor], ~torch.Tensor] = <silverspoon_kd.alignments.output_selector.OutputSelector object>, input_projector: ~torch.nn.modules.module.Module | None = None, output_projector: ~torch.nn.modules.module.Module | None = None, auto_projector: bool = False, projector_init: str | ~collections.abc.Callable[[~torch.nn.modules.module.Module], None] | None = None, loss_function: str | ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor] | None = None, loss_function_kwargs: dict[str, ~typing.Any] | None = None, optimizer: ~torch.optim.optimizer.Optimizer | None = None, scheduler: ~torch.optim.lr_scheduler.LRScheduler | None = None, max_grad_norm: float | int | None = None, loss_weight: float = 1.0, auto_device_match: bool = False, auto_dtype_match: bool = False) - module:
silverspoon_kd.alignments
Initialize an Alignment instance.
- param teacher_block:
The teacher module to capture outputs from
- param student_block:
The student module to train
- param teacher_model_name:
Label for the teacher model (auto-set by [create_alignments][silverspoon_kd.create_alignments])
- param student_model_name:
Label for the student model (auto-set by [create_alignments][silverspoon_kd.create_alignments])
- param teacher_module_name:
Name of the teacher module (auto-set by [create_alignments][silverspoon_kd.create_alignments])
- param student_module_name:
Name of the student module (auto-set by [create_alignments][silverspoon_kd.create_alignments])
- param teacher_output_selector:
Function to select/extract teacher output before comparison
- param student_output_selector:
Function to select/extract student output before comparison
- param input_projector:
Optional module to project inputs before student forward pass
- param output_projector:
Optional module to project student output before loss computation
- param auto_projector:
If True, automatically infer and create projectors based on shape mismatches detected during the first forward pass
- param projector_init:
Initialization for auto-created projectors.
None(default): PyTorch default (kaiming_uniform)."normal":normal_(0, 0.02)(BERT/TextBrewer-style)."xavier": Xavier uniform.A callable
fn(module)applied viamodule.apply(fn).
- param loss_function:
Loss function for alignment loss (defaults to MSE). Can be a callable or a string name from the loss registry.
- param loss_function_kwargs:
Optional keyword arguments forwarded to the loss factory when
loss_functionis a string name. Ignored whenloss_functionis already a callable. Example:{"temperature": 3.0}or{"weight_matrix": tensor}.- param optimizer:
Optimizer for training the student block. If None, auto-created.
- param scheduler:
Learning rate scheduler. If None, auto-created.
- param max_grad_norm:
Maximum gradient norm for clipping
- param loss_weight:
Scalar weight for this alignment’s loss contribution when accumulating total loss. Only the ratios between weights matter, not their absolute values (e.g. [2, 3] gives the same gradient proportions as [0.4, 0.6]). Default: 1.0 (no scaling). Supported by HolisticDistiller; BlockwiseDistiller will warn if non-default weights are set.
- param auto_device_match:
Whether to auto-move inputs/teacher outputs to student device
- param auto_dtype_match:
Whether to auto-cast inputs/teacher outputs to student dtype
.. py:method:: Alignment.get_name() -> str
- module:
silverspoon_kd.alignments
Get the display name for this alignment.
Returns the student’s model_name.module_name, which is used as the metric key for per-layer loss reporting.
.. py:method:: Alignment.initialize_auto_projectors(input_args: tuple, input_kwargs: dict, teacher_output: ~torch.Tensor) -> None
- module:
silverspoon_kd.alignments
Initialize auto projectors on the first forward pass if auto_projector is enabled.
.. py:class:: OutputSelector(index: int = 0)
- module:
silverspoon_kd.alignments
- canonical:
silverspoon_kd.alignments.output_selector.OutputSelector
Bases: :py:class:
objectExtracts an element at a given index from tuple outputs.
Many PyTorch modules (e.g., attention layers) return tuples like
(hidden_states, attention_weights). This selector picks one element for comparison in an alignment. If the output is not a tuple, it is returned as-is... py:method:: OutputSelector.init(index: int = 0)
- module:
silverspoon_kd.alignments
Initialize the selector.
- param index:
The index of the element to extract from tuple outputs. Default:
0(the first element).
.. py:function:: load_student_weights_from_checkpoint(student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, checkpoint_dir: str | ~pathlib.Path, student_model_name: str, strict: bool = True) -> dict[str, list[str]]
- module:
silverspoon_kd.alignments
Load trained student weights from a distiller checkpoint into a student model.
Handles both checkpoint layouts the distillers write:
- class:
~silverspoon_kd.distillers.BlockwiseDistillersaves only the trained student blocks (block_<name>.*entries). Each block is loaded into the matching module ofstudent_model; modules that were not part of an alignment are left untouched.
- class:
~silverspoon_kd.distillers.HolisticDistillerand- class:
~silverspoon_kd.distillers.ResponseBasedDistillersave the full student state dict, which is loaded as a whole.
Projectors are not loaded. Use
- func:
load_student_with_projectors_from_checkpointwhen blockwise-trained blocks must run inside the teacher together with their projectors.- param student_model:
The student model to load weights into.
- param checkpoint_dir:
Directory containing
model.safetensors(orpytorch_model.bin), e.g../runs/checkpoint-1000.- param student_model_name:
The
student_model_namethe alignments were created with (e.g.'Qwen/Qwen3-0.6B'). Used to match blockwise entries to modules; ignored for end-to-end checkpoints.- param strict:
If True, raise
ValueErrorwhen any block (or the full state dict) has missing or unexpected keys. If False, return them instead.- returns:
Dictionary with
'missing_keys'and'unexpected_keys'lists, qualified with the module path for blockwise checkpoints.- raises FileNotFoundError:
If the checkpoint directory or model file is missing.
- raises ValueError:
If a blockwise entry matches no module of
student_modelunderstudent_model_name, or in strict mode when keys are missing or unexpected.
.. rubric:: Example
student_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-0.6B") result = load_student_weights_from_checkpoint( student_model, "./runs/checkpoint-1000", "Qwen/Qwen3-0.6B", ) student_model.save_pretrained("./trained_student")
.. py:function:: load_student_with_projectors_from_checkpoint(teacher_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, checkpoint_dir: str | ~pathlib.Path, student_model_name: str, student_model: ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module, device: ~torch.device | None = None) -> ~transformers.modeling_utils.PreTrainedModel | ~torch.nn.modules.module.Module
- module:
silverspoon_kd.alignments
Install blockwise-distilled student blocks, with their projectors, into the teacher.
- class:
~silverspoon_kd.distillers.BlockwiseDistillertrains student blocks in isolation, with input/output projectors bridging the dimension gap to the teacher. This function reads such a checkpoint, finds the teacher modules that were distilled from itsblock_<name>entries, and replaces each of them in place with[input_proj, student_block, output_proj](projectors are omitted when absent), so the returned model can be used directly for inference.- param teacher_model:
The teacher model; modified in place and returned.
- param checkpoint_dir:
Directory containing
model.safetensors(orpytorch_model.bin) and, when projectors were trained,projector_state.pt.- param student_model_name:
The
student_model_namethe alignments were created with (e.g.'google-bert/bert-base-uncased').- param student_model:
A student model with the trained architecture (e.g. smaller hidden dims); its modules are the templates for the blocks. Module names must match the teacher’s.
- param device:
Device for the loaded modules (defaults to the teacher’s).
- returns:
The teacher model with the student blocks and projectors installed.
- raises FileNotFoundError:
If the checkpoint directory or model file is missing.
- raises ValueError:
If the checkpoint is not a blockwise checkpoint, an entry matches no module of
student_modelunderstudent_model_name, or a distilled module does not exist inteacher_model.
.. rubric:: Example
from transformers import AutoModelForMaskedLM, AutoConfig # Load teacher teacher = AutoModelForMaskedLM.from_pretrained("bert-base-uncased") # Create student with smaller dimensions config = AutoConfig.from_pretrained("bert-base-uncased") config.hidden_size = 384 config.intermediate_size = 1536 config.num_attention_heads = 6 student = AutoModelForMaskedLM.from_config(config) # Load checkpoint with projectors result_model = load_student_with_projectors_from_checkpoint( teacher, "./runs/checkpoint-1000", "google-bert/bert-base-uncased", student ) output = result_model(inputs)
Projectors¶
.. py:class:: GenericLinearProjector(in_features: int, out_features: int, bias: bool = True, mode: str = ‘output’, apply_to_arg: int | None = None, apply_to_kwarg: str | None = None)
- module:
silverspoon_kd.alignments.projectors
- canonical:
silverspoon_kd.alignments.projectors.generic_linear_projector.GenericLinearProjector
Bases: :py:class:
~silverspoon_kd.alignments.projectors._mode_mixin._ProjectorModeMixin, :py:class:~torch.nn.modules.linear.LinearA linear projector for transformers and dense layers that supports both input and output modes.
Projects the last dimension of tensors (e.g., hidden states in transformers). Supports both input projection (modifying args/kwargs) and output projection (direct tensor projection).
- param in_features:
Input feature dimension
- param out_features:
Output feature dimension
- param bias:
Whether to include bias term (default: True)
- param mode:
‘input’ or ‘output’ (default: ‘output’)
- param apply_to_arg:
For input mode, which positional arg to project (default: 0)
- param apply_to_kwarg:
For input mode, which keyword arg to project (default: None)
Shape: - Input mode: Takes *args, **kwargs, returns (args, kwargs) - Output mode: Takes tensor, returns projected tensor - Tensor shape:
(*, in_features)->(*, out_features)where*is any dimensions.. rubric:: Example
# Output mode (default) projector = GenericLinearProjector(768, 512) hidden_states = torch.randn(32, 128, 768) projected = projector(hidden_states) projected.shape # torch.Size([32, 128, 512]) # Input mode projector = GenericLinearProjector( 768, 512, mode='input', apply_to_kwarg='hidden_states' ) args, kwargs = projector(hidden_states=torch.randn(32, 128, 768)) kwargs['hidden_states'].shape # torch.Size([32, 128, 512])
.. py:method:: GenericLinearProjector.init(in_features: int, out_features: int, bias: bool = True, mode: str = ‘output’, apply_to_arg: int | None = None, apply_to_kwarg: str | None = None)
- module:
silverspoon_kd.alignments.projectors
.. py:method:: GenericLinearProjector.forward(input: ~torch.Tensor | None = None, *args: ~typing.Any, **kwargs: ~typing.Any) -> ~typing.Any
- module:
silverspoon_kd.alignments.projectors
Project the input tensor, or the targeted arg/kwarg in input mode.
Output mode returns a :class:
Tensor; input mode returns the full(args, kwargs)tuple for aforward_pre_hook.
.. py:class:: GenericConv2dProjector(in_channels: int, out_channels: int, bias: bool = False, mode: str = ‘output’, apply_to_arg: int | None = None, apply_to_kwarg: str | None = None)
- module:
silverspoon_kd.alignments.projectors
- canonical:
silverspoon_kd.alignments.projectors.generic_conv2d_projector.GenericConv2dProjector
Bases: :py:class:
~silverspoon_kd.alignments.projectors._mode_mixin._ProjectorModeMixin, :py:class:~torch.nn.modules.conv.Conv2dA 1x1 convolution projector for CNNs that supports both input and output modes.
Projects the channel dimension of 4D tensors using 1x1 convolution. Preserves spatial dimensions while mapping from student channels to teacher channels.
- param in_channels:
Number of input channels (student)
- param out_channels:
Number of output channels (teacher)
- param bias:
Whether to include bias term (default: False, common for distillation)
- param mode:
‘input’ or ‘output’ (default: ‘output’)
- param apply_to_arg:
For input mode, which positional arg to project (default: 0)
- param apply_to_kwarg:
For input mode, which keyword arg to project (default: None)
Shape: - Input mode: Takes *args, **kwargs, returns (args, kwargs) - Output mode: Takes tensor, returns projected tensor - Tensor shape:
(batch, in_channels, H, W)->(batch, out_channels, H, W).. rubric:: Example
# Output mode (default) projector = GenericConv2dProjector(128, 512) features = torch.randn(8, 128, 14, 14) projected = projector(features) projected.shape # torch.Size([8, 512, 14, 14]) # Input mode projector = GenericConv2dProjector(128, 512, mode='input') args, kwargs = projector(torch.randn(8, 128, 14, 14)) args[0].shape # torch.Size([8, 512, 14, 14])
.. py:method:: GenericConv2dProjector.init(in_channels: int, out_channels: int, bias: bool = False, mode: str = ‘output’, apply_to_arg: int | None = None, apply_to_kwarg: str | None = None)
- module:
silverspoon_kd.alignments.projectors
.. py:method:: GenericConv2dProjector.forward(input: ~torch.Tensor | None = None, *args: ~typing.Any, **kwargs: ~typing.Any) -> ~typing.Any
- module:
silverspoon_kd.alignments.projectors
Project the input tensor, or the targeted arg/kwarg in input mode.
Output mode returns a :class:
Tensor; input mode returns the full(args, kwargs)tuple for aforward_pre_hook.
.. py:function:: fuse_projectors_into_module(module: ~torch.nn.modules.module.Module, input_projector: ~torch.nn.modules.module.Module | None = None, output_projector: ~torch.nn.modules.module.Module | None = None, inplace: bool = False) -> ~torch.nn.modules.module.Module
- module:
silverspoon_kd.alignments.projectors
Fuse input and output projectors into a module’s weights to eliminate runtime overhead.
The fusion follows the flow: prev_layer → output_projector → input_projector → [module]
This function absorbs both projectors into the module weights, creating a fused module that maintains student dimensions while performing the same transformation.
Mathematical formulation:
For Linear layers: W_fused = W_module @ W_input_proj @ W_output_proj
For Conv2d layers: Similar channel-wise fusion using einsum operations
- param module:
The module to fuse projectors into (nn.Linear or nn.Conv2d)
- param input_projector:
Optional input projector to fuse (applied after output_projector)
- param output_projector:
Optional output projector from preceding layer (applied first)
- param inplace:
Whether to modify the module in-place (default: False)
- returns:
A new fused module with projectors absorbed into weights
- raises ValueError:
If module or projector types are unsupported
- raises RuntimeError:
If projector configurations are incompatible with fusion
Supported combinations: - nn.Linear with GenericLinearProjector(s) - nn.Conv2d with GenericConv2dProjector(s) (1x1 conv only)
Limitations: - Projectors must be GenericLinearProjector or GenericConv2dProjector - Conv2d projectors must be 1x1 convolutions with stride=1, padding=0 - No non-linear activations between projectors - For a Conv2d module with zero padding, projector biases are fused exactly on interior pixels only: at the border, kernel taps that read padded zeros carried no projector bias before fusion but do after. Fusion is exact everywhere for
padding=0and for the non-zero padding modes (replicate,reflect,circular)... rubric:: Example
# Fuse linear projectors module = nn.Linear(256, 512) input_proj = GenericLinearProjector(128, 256, mode='input') output_proj = GenericLinearProjector(64, 128, mode='output') fused = fuse_projectors_into_module(module, input_proj, output_proj) # fused is now a Linear(64, 512) layer equivalent to output_proj -> input_proj -> module # Fuse CNN projectors module = nn.Conv2d(32, 64, kernel_size=3, padding=1) input_proj = GenericConv2dProjector(16, 32, mode='input') output_proj = GenericConv2dProjector(8, 16, mode='output') fused = fuse_projectors_into_module(module, input_proj, output_proj) # fused is Conv2d(8, 64, kernel_size=3, padding=1) with absorbed projectors