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#