Utilities¶
SilverSpoon-KD provides several utility functions for model preparation, parameter management, and checkpoint loading. All utilities are available from the top-level package import.
reconfig_model¶
Creates a smaller student model by modifying the teacher’s configuration and optionally copying matching weights.
from silverspoon_kd import reconfig_model
student = reconfig_model(
model=teacher,
name_or_path="my-student",
diff={
"num_attention_heads": 8,
"num_key_value_heads": 4,
"intermediate_size": 1024,
},
copy_matching_weights=True,
freeze_copied_weights=True,
)
Parameters:
model— SourcePreTrainedModelto derive the student from.name_or_path— Thename_or_pathassigned to the new model’s config.diff— Dictionary of config attributes to change (e.g.,{"num_hidden_layers": 12}). Any attribute on the model’s config can be overridden.copy_matching_weights— IfTrue, copies weights from the source model for all parameters where shapes match. Useful when changing some dimensions (e.g.,num_attention_heads) but keeping others (e.g.,hidden_size), so embeddings, layer norms, and output heads are initialized from the teacher.freeze_copied_weights— IfTrue(andcopy_matching_weights=True), freezes all parameters that were copied. Only randomly-initialized parameters (those with shape mismatches) remain trainable.
prune_model¶
Performs structured pruning using importance-based weight selection. Unlike
reconfig_model (which creates a new model from scratch), prune_model retains
the most important weights from the original model based on L2 magnitude.
Requires the torch-pruning package (pip install torch-pruning).
from silverspoon_kd import prune_model
student = prune_model(
model=teacher,
name_or_path="my-pruned-student",
diff={
"num_attention_heads": 8,
"num_key_value_heads": 4,
"intermediate_size": 1024,
},
example_inputs={"input_ids": sample_input_ids},
output_transform=lambda x: x.logits,
freeze_copied_weights=True,
)
Parameters:
model— SourcePreTrainedModelto prune.name_or_path— Thename_or_pathassigned to the pruned model’s config.diff— Dictionary of config attributes defining the target dimensions.example_inputs— Required. Sample inputs forDependencyGraphtracing. Must be passable tomodel.forward().output_transform— Transform applied to model output for tracing (e.g.,lambda x: x.logits).ignored_layers— Modules to exclude from pruning. Auto-detected ifNone(embeddings, LM heads).num_heads— Dict mapping attention projection modules to their head count. Auto-detected ifNone.out_channel_groups— Dict mapping fused modules (e.g.,gate_up_proj) to their group count. Auto-detected ifNone.forward_fn— Custom forward function forDependencyGraphtracing.round_to— Round channel counts down to nearest multiple of this value, useful for GPU-aligned dimensions (e.g.,round_to=8).freeze_copied_weights— IfTrue, freezes parameters whose shapes were not changed by pruning. Structurally pruned parameters remain trainable.
Two-pass algorithm:
Attention head pruning — Q/K/V projections are pruned directly using known coupling rules (Q↔O, K↔V) to avoid dependency graph cycles caused by reshape operations in multi-head attention.
Width pruning — Non-attention modules (MLP layers, etc.) are pruned via Torch-Pruning’s
DependencyGraph, which handles architecture-agnostic coupling propagation.
freeze_parameters¶
Freezes model parameters whose names match any of the provided regex patterns.
from silverspoon_kd import freeze_parameters
# Freeze all embedding and layer norm parameters
freeze_parameters(model, [r"embed", r"layernorm", r"ln_"])
# Freeze specific layers and unfreeze everything else
freeze_parameters(
model,
[r"model\.layers\.[0-5]\."],
thaw_not_matched=True,
)
Parameters:
model— The model whose parameters to freeze.freeze_regex_patterns— List of regex strings. If a parameter name matches any pattern, it is set torequires_grad=False.thaw_not_matched— IfTrue, parameters not matching any pattern are explicitly set torequires_grad=True.
Distiller factory¶
The Distiller factory function returns the appropriate distiller instance
based on a string name:
from silverspoon_kd import Distiller
distiller = Distiller(
distiller_type="blockwise",
teacher_model=teacher,
alignments=alignments,
train_dataset=train_dataset,
args=args,
)
Valid distiller_type values:
Name |
Alias |
Class |
|---|---|---|
|
|
|
|
|
|
|
|
|
fuse_projectors_into_module¶
After training with projectors (due to teacher/student dimension mismatches), you can fuse the projectors into the module weights to eliminate runtime overhead during inference:
from silverspoon_kd import fuse_projectors_into_module
fused_module = fuse_projectors_into_module(
module=student_layer,
input_projector=input_proj,
output_projector=output_proj,
)
Supports nn.Linear with GenericLinearProjector and nn.Conv2d with
GenericConv2dProjector (1x1 convolutions only).
Checkpoint utilities¶
After training, load distilled weights back into the student model:
load_student_weights_from_checkpoint¶
Loads the trained student weights from a distillation checkpoint into a student
model. A BlockwiseDistiller checkpoint holds only the trained blocks, and each
one is loaded into its module of the student (modules that were not aligned are
left untouched); a HolisticDistiller or ResponseBasedDistiller checkpoint
holds the full student state dict. Projectors are not loaded:
from silverspoon_kd import load_student_weights_from_checkpoint
load_student_weights_from_checkpoint(
student_model=student,
checkpoint_dir="./output/checkpoint-1000",
student_model_name="Qwen/Qwen3-4B",
)
student.save_pretrained("./distilled-student")
Parameters:
student_model— The student to load weights into.checkpoint_dir— A checkpoint directory written by a distiller, e.g../output/checkpoint-1000, containingmodel.safetensors(orpytorch_model.bin).student_model_name— The student label recorded on the alignments during training.create_alignmentsuses the student’sname_or_path, so it must match exactly.strict— IfTrue(default), missing or unexpected keys raise aValueError; ifFalse, they are only reported in the return value.
Returns a dict with missing_keys and unexpected_keys.
load_student_with_projectors_from_checkpoint¶
Loads student blocks and their projector layers from a BlockwiseDistiller
checkpoint. Use this when the student was trained with projectors (dimension
mismatches between teacher and student). It takes the teacher and a student with
the trained architecture, works out which modules were aligned by inspecting the
checkpoint, and replaces each aligned teacher module in-place with
[input projector, student block, output projector]:
from transformers import AutoModelForCausalLM
from silverspoon_kd import (
load_student_with_projectors_from_checkpoint,
reconfig_model,
)
# The teacher used during training, and a student with the same architecture
# as the one that was trained (here: a narrower model derived from the teacher).
teacher = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-0.6B")
student = reconfig_model(
teacher,
name_or_path="my-student",
diff={"hidden_size": 512, "intermediate_size": 1536},
)
model_with_projectors = load_student_with_projectors_from_checkpoint(
teacher_model=teacher,
checkpoint_dir="./output/checkpoint-1000",
student_model_name="my-student",
student_model=student,
)
Parameters:
teacher_model— The teacher model. It is modified in-place and returned.checkpoint_dir— ABlockwiseDistillercheckpoint directory, containingmodel.safetensors(orpytorch_model.bin) and, when projectors were trained,projector_state.pt.student_model_name— The student label recorded on the alignments during training.create_alignmentsuses the student’sname_or_path, so it must match exactly.student_model— A student with the trained architecture and the same module structure as the teacher; its blocks provide the modules into which the checkpoint weights are loaded.device— Optional device for the loaded modules (defaults to the teacher’s device).
The returned model can be used directly for inference, or its projectors can be
fused into the student blocks with fuse_projectors_into_module.