ple_pruning#

Global per-layer-embedding (PLE) channel ranking and tensor surgery.

Classes

PLEPruningSpec

Own every tensor coupled to one global PLE channel dimension.

class PLEPruningSpec#

Bases: object

Own every tensor coupled to one global PLE channel dimension.

Gemma-style PLE uses one shared RMSNorm across all layer chunks. Therefore every layer must use the same permutation and retained width. Per-layer contribution scores are summed before sorting to preserve this invariant.

__init__(language_prefix, layer_template, num_layers, width, layer_gate_name='per_layer_input_gate', layer_projection_name='per_layer_projection', model_embedding_name='embed_tokens_per_layer', model_projection_name='per_layer_model_projection', model_norm_name='per_layer_projection_norm')#
Parameters:
  • language_prefix (str)

  • layer_template (str)

  • num_layers (int)

  • width (int)

  • layer_gate_name (str)

  • layer_projection_name (str)

  • model_embedding_name (str)

  • model_projection_name (str)

  • model_norm_name (str)

Return type:

None

language_prefix: str#
layer_gate_name: str = 'per_layer_input_gate'#
layer_prefix(layer_idx)#
Parameters:

layer_idx (int)

Return type:

str

layer_projection_name: str = 'per_layer_projection'#
layer_score_key(layer_idx)#
Parameters:

layer_idx (int)

Return type:

str

layer_template: str#
model_embedding_name: str = 'embed_tokens_per_layer'#
model_norm_name: str = 'per_layer_projection_norm'#
model_projection_name: str = 'per_layer_model_projection'#
num_layers: int#
order_from_score_logs(score_logs)#
Parameters:

score_logs (dict[str, dict])

Return type:

Tensor

permute_state_dict(state_dict, order)#
Parameters:
  • state_dict (dict[str, Tensor])

  • order (Tensor)

Return type:

tuple[dict[str, Tensor], set[str]]

slice_state_dict(state_dict, target_width)#
Parameters:
  • state_dict (dict[str, Tensor])

  • target_width (int)

Return type:

dict[str, Tensor]

width: int#