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_function is 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 callable fn(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] = , 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

canonical:

silverspoon_kd.alignments.alignment.Alignment

Bases: :py:class:object

A 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 via module.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_function is a string name. Ignored when loss_function is 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:object

Extracts 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.BlockwiseDistiller saves only the trained student blocks (block_<name>.* entries). Each block is loaded into the matching module of student_model; modules that were not part of an alignment are left untouched.

  • class:

    ~silverspoon_kd.distillers.HolisticDistiller and

    class:

    ~silverspoon_kd.distillers.ResponseBasedDistiller save the full student state dict, which is loaded as a whole.

Projectors are not loaded. Use

func:

load_student_with_projectors_from_checkpoint when 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 (or pytorch_model.bin), e.g. ./runs/checkpoint-1000.

param student_model_name:

The student_model_name the 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 ValueError when 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_model under student_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.BlockwiseDistiller trains 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 its block_<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 (or pytorch_model.bin) and, when projectors were trained, projector_state.pt.

param student_model_name:

The student_model_name the 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_model under student_model_name, or a distilled module does not exist in teacher_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.Linear

A 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 a forward_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.Conv2d

A 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 a forward_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=0 and 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