embedding_pruning#

Descriptor-owned residual hidden-width ranking, sorting, and slicing.

Classes

EmbeddingPruningSpec

Complete model-family contract for the language residual width.

PackedMinitronImportance

Exact Minitron hidden metric with original packed-sample boundaries.

TensorAxisRule

class EmbeddingPruningSpec#

Bases: object

Complete model-family contract for the language residual width.

__init__(hidden_size, legal_widths, alignment, tensor_rules, exempt_patterns=(), tie_groups=(), config_paths=(('hidden_size',),), residual_norm_patterns=(), permutation_group_size=1)#
Parameters:
  • hidden_size (int)

  • legal_widths (tuple[int, ...])

  • alignment (int)

  • tensor_rules (tuple[TensorAxisRule, ...])

  • exempt_patterns (tuple[str, ...])

  • tie_groups (tuple[tuple[str, ...], ...])

  • config_paths (tuple[tuple[str, ...], ...])

  • residual_norm_patterns (tuple[str, ...])

  • permutation_group_size (int)

Return type:

None

alignment: int#
audit_state_dict(state_dict)#
Parameters:

state_dict (Mapping[str, Tensor])

Return type:

dict[str, Any]

config_paths: tuple[tuple[str, ...], ...] = (('hidden_size',),)#
exempt_patterns: tuple[str, ...] = ()#
hidden_size: int#
legal_widths: tuple[int, ...]#
order_from_scores(scores)#

Return a legal residual permutation for the descriptor’s storage granularity.

Block-quantized inputs cannot move individual channels across quantization blocks without dequantizing and requantizing. Such descriptors rank whole blocks by summed channel importance and preserve channel order within a block.

Parameters:

scores (Tensor)

Return type:

Tensor

parameter_count(state_dict)#
Parameters:

state_dict (Mapping[str, Tensor])

Return type:

int

permutation_group_size: int = 1#
permute_state_dict(state_dict, order)#
Parameters:
  • state_dict (Mapping[str, Tensor])

  • order (Tensor)

Return type:

dict[str, Tensor]

residual_norm_patterns: tuple[str, ...] = ()#
rule_for(key)#

Return the unique descriptor rule for a canonical tensor key.

Parameters:

key (str)

Return type:

TensorAxisRule | None

slice_state_dict(state_dict, width, *, tp_size=1)#
Parameters:
  • state_dict (Mapping[str, Tensor])

  • width (int)

  • tp_size (int)

Return type:

dict[str, Tensor]

sliced_shape(key, shape, width)#

Return a tensor’s exact shape after residual-width slicing.

Parameters:
  • key (str)

  • shape (Sequence[int])

  • width (int)

Return type:

tuple[int, …]

tensor_rules: tuple[TensorAxisRule, ...]#
tie_groups: tuple[tuple[str, ...], ...] = ()#
update_config(config, width)#
Parameters:
  • config (Mapping[str, Any])

  • width (int)

Return type:

dict[str, Any]

update_config_object(config, width)#

Return a deep-copied config with every descriptor-owned width updated.

AutoModel/HF configuration trees mix PretrainedConfig objects and dictionaries. Physical width realization must update both the nested language config and any top-level mirror without round-tripping through to_dict(), which would lose the concrete config classes needed by model construction.

Parameters:
  • config (Any)

  • width (int)

Return type:

Any

validate_width(width, *, tp_size=1)#
Parameters:
  • width (int)

  • tp_size (int)

Return type:

int

class PackedMinitronImportance#

Bases: object

Exact Minitron hidden metric with original packed-sample boundaries.

__init__(hidden_size)#
Parameters:

hidden_size (int)

property sample_count: int#
scores()#
Return type:

Tensor

property site_names: tuple[str, ...]#
squared_sums()#
Return type:

dict[str, Tensor]

update(site_name, activations, *, cu_seqlens=None, token_mask=None)#
Parameters:
  • site_name (str)

  • activations (Tensor)

  • cu_seqlens (Tensor | None)

  • token_mask (Tensor | None)

Return type:

None

class TensorAxisRule#

Bases: object

__init__(pattern, axes, description, chunked_axes=(), grouped_axes=())#
Parameters:
  • pattern (str)

  • axes (tuple[int, ...])

  • description (str)

  • chunked_axes (tuple[tuple[int, int], ...])

  • grouped_axes (tuple[tuple[int, int], ...])

Return type:

None

axes: tuple[int, ...]#
chunked_axes: tuple[tuple[int, int], ...] = ()#
description: str#
grouped_axes: tuple[tuple[int, int], ...] = ()#
matches(key)#
Parameters:

key (str)

Return type:

bool

pattern: str#