ple_pruning
Global per-layer-embedding (PLE) channel ranking and tensor surgery.
Classes
Own every tensor coupled to one global PLE channel dimension. |
- class PLEPruningSpec
Bases:
objectOwn 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