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’]