Utilities¶
.. py:function:: freeze_parameters(model: ~torch.nn.modules.module.Module, freeze_regex_patterns: list[str], thaw_not_matched=False)
- module:
silverspoon_kd.utils
Freezes parameters whose names match any of the provided regex patterns. Parameters not matching any pattern will be set to trainable (requires_grad=True).
- param model:
The model whose parameters are to be frozen.
- type model:
torch.nn.Module
- param freeze_regex_patterns:
A list of regex strings. If a parameter’s name matches any of these patterns, it will be set to requires_grad=False.
- type freeze_regex_patterns:
List[str]
- param thaw_not_matched:
If True, parameters not matching any pattern will be unfrozen.
- type thaw_not_matched:
bool
.. py:function:: reconfig_model(model: ~transformers.modeling_utils.PreTrainedModel, name_or_path: str, diff: dict | None = None, copy_matching_weights: bool = False, freeze_copied_weights: bool = False) -> ~torch.nn.modules.module.Module
- module:
silverspoon_kd.utils
Reconfigures a model’s configuration based on a provided dictionary of differences. Uses the same model class as the input model to ensure compatibility.
- param model:
The source model to reconfigure from.
- param name_or_path:
The name_or_path to assign to the new model’s config.
- param diff:
Dictionary of config attributes to change (e.g., {“num_hidden_layers”: 12}).
- param copy_matching_weights:
If True, copies weights from the source model to the new model for all parameters where shapes match. This is useful when changing some dimensions (e.g., num_attention_heads, intermediate_size) but keeping others (e.g., hidden_size), allowing embeddings, layer norms, and output heads to be initialized from the source model.
- param freeze_copied_weights:
If True and copy_matching_weights is True, freezes (sets requires_grad=False) all parameters that were copied from the source model. This is ignored if copy_matching_weights is False.
- returns:
A new model with the modified configuration.
.. py:function:: prune_model(model: ~transformers.modeling_utils.PreTrainedModel, name_or_path: str, diff: dict | None = None, example_inputs: dict | None = None, output_transform: ~collections.abc.Callable | None = None, ignored_layers: list[~torch.nn.modules.module.Module] | None = None, num_heads: dict[~torch.nn.modules.module.Module, int] | None = None, out_channel_groups: dict[~torch.nn.modules.module.Module, int] | None = None, forward_fn: ~collections.abc.Callable | None = None, round_to: int | None = None, freeze_copied_weights: bool = False) -> ~torch.nn.modules.module.Module
- module:
silverspoon_kd.utils
Prune a model to a target config using importance-based weight selection.
Uses Torch-Pruning’s DependencyGraph to discover parameter couplings and propagate structural pruning correctly, while selecting which channels to keep based on L2 weight magnitude.
- param model:
Source pretrained model to prune.
- param name_or_path:
The name_or_path to assign to the pruned model’s config.
- param diff:
Dictionary of config attributes to change.
- param example_inputs:
Required inputs for DependencyGraph tracing.
- param output_transform:
Transform applied to model output for tracing (e.g., lambda x: x.logits).
- param ignored_layers:
Modules to exclude from pruning. Auto-detected if None.
- param num_heads:
Dict mapping attention projection modules to their head count. Auto-detected if None.
- param out_channel_groups:
Dict mapping fused modules to their group count. Auto-detected if None.
- param forward_fn:
Custom forward function for DependencyGraph tracing.
- param round_to:
Round channel counts down to nearest multiple of this value. Useful for GPU-aligned dimensions (e.g., round_to=4 or round_to=8).
- param freeze_copied_weights:
If True, freezes (sets requires_grad=False) all parameters whose shapes were not changed by pruning. Parameters that were structurally pruned remain trainable.
- returns:
A pruned copy of the model with target dimensions and retained weights.
.. py:function:: get_modules_by_names(model: ~torch.nn.modules.module.Module, regex_patterns: list[str]) -> list[~torch.nn.modules.module.Module]
- module:
silverspoon_kd.utils
Given a model and a list of regex patterns, return a list of modules whose names match any of the patterns.
- param model:
The model to search for modules.
- type model:
torch.nn.Module
- param regex_patterns:
A list of regex patterns to match against module names.
- type regex_patterns:
List[str]
- returns:
A list of modules whose names match any of the patterns.
- rtype:
List[torch.nn.Module]
.. py:function:: summarize_layer_names(keys: list[str]) -> list[str]
- module:
silverspoon_kd.utils
Recursively summarizes layer names into compact patterns. Example: [‘model.layers.0.attn.q_proj’, ‘model.layers.1.attn.q_proj’] -> [‘model.layers.[0-1].attn.q_proj’]
.. py:function:: partial_summarize_layer_names(keys: list[str]) -> list[str]
- module:
silverspoon_kd.utils
Groups a list of layer names into summarized patterns. Example: [‘layers.0.attn.q_proj’, ‘layers.1.attn.q_proj’] -> [‘layers.[0-1].attn.q_proj’]