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