Trained LM Heads for Layer-Pruning Examination
Per-layer language modeling heads used to examine intermediate layers of
7B/8B decoder models before pruning. Each checkpoint is an lm_head trained
to read out the hidden state at one specific layer, so you can measure how much
predictive signal that layer already carries and use the resulting curve to
decide where to cut.
These are diagnostic artifacts, not pruned model weights. Nothing here is
sparsified or pruned โ the tensors are dense and full-size. Do not use them as
drop-in replacements for a base model's lm_head.
Contents
Two sets, one per base model, covering layers 16โ31 (the back half of a 32-layer stack โ the region where early exit is plausible):
| Set | Base model | Tensor shape | dtype | Per-file | Total |
|---|---|---|---|---|---|
lm_head_checkpoints_L/ |
Llama 3 8B | (128256, 4096) |
float32 | 1.96 GiB | 31 GB |
lm_head_checkpoints_M/ |
Mistral 7B | (32768, 4096) |
float32 | 512 MiB | 8.1 GB |
lm_head_checkpoints_L/lm_head_prune16.pt ... lm_head_prune31.pt (16 files)
lm_head_checkpoints_M/lm_head_prune16.pt ... lm_head_prune31.pt (16 files)
The number in each filename is the layer index the head attaches to, not a training step and not a pruning ratio.
Vocabulary sizes are padded up from the tokenizers' native values (Llama 3: 128256; Mistral: 32000 โ 32768).
Format
Each .pt file is an OrderedDict with a single key:
{"weight": Tensor(vocab_size, 4096)}
No bias, no optimizer state, no metadata.
Usage
import torch
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained("meta-llama/Meta-Llama-3-8B")
LAYER = 24
sd = torch.load(f"lm_head_checkpoints_L/lm_head_prune{LAYER}.pt", map_location="cpu")
head = torch.nn.Linear(4096, sd["weight"].shape[0], bias=False)
head.load_state_dict(sd)
# Read out logits from layer LAYER's hidden state
out = model(input_ids, output_hidden_states=True)
logits = head(out.hidden_states[LAYER])
Sweep LAYER from 16 to 31 and compare perplexity or task accuracy against the
full model to find the depth at which quality degrades unacceptably.
Match the set to the base model โ the heads are not interchangeable between
L and M, and the vocab dimensions differ.