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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support