Early exit and layer skipping
Some words are decided long before the last layer, so a model can stop early and save real time, provided it was trained to be readable partway through and can repair the gaps it leaves behind.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Stop running layers as soon as the answer is settled, instead of always running all of them.
Checking whether the rice is done
You lift the lid and press a grain between your fingers. Sometimes it is ready after twelve minutes. Sometimes it needs twenty.
Nobody cooks rice by a stopwatch alone. You check, and you stop when it is done.
A language model does the opposite. Every word gets the same twelve, or forty, or eighty layers. Easy or hard, the cost is identical.
Where the waste is
Finishing "New" with "Delhi" is close to certain after a couple of layers. Choosing the next step in a piece of reasoning may need every layer there is.
Same cost for both. Running the full stack on the easy words is pure waste, and easy words are the majority.
The idea
Fit a small reader at several depths. After each layer, ask: is the answer already clear?
If it is, stop and emit. If it is not, keep going.
layer 1 -> layer 2 -> layer 3 -> ... -> layer 12
|
+-- clear enough? yes -> stop, answer now
no -> carry onThe part that makes this hard, and it is not the idea
Judging when the rice is done is a skill. You have to be able to read the pot.
A model that was trained to be read only at the end is not readable in the middle. The measurement in the developer section shows this plainly. Read an ordinary model's answer at layer six, and it is nowhere near the top of its list.
The fix is to train the model to be readable partway through, from the beginning. That has to be planned in, not added later.
The second problem, which is worse
Layers do not only produce an answer. They also leave notes that later words rely on.
Stop at layer six and layers seven onwards leave no notes. When the next word arrives and wants to look back, those notes are missing.
So you fill the gaps with a rough copy, or go back and compute them properly. The second choice spends what you saved. This bookkeeping, and not the idea, is why early exit is rare in production.
A trick that avoids both problems
Run the early layers to get a quick guess. Then run the full model once to check a whole batch of guesses at the same time.
Checking many guesses together is cheap. Wrong guesses are thrown away, and the answer is exactly what the full model would have said.
You get speed without any quality risk, because the full model still has the final word. Reported speed-ups on real tasks are roughly two times.
Remember this
- Easy words are settled early, and running every layer on them is waste.
- A model must be trained to be readable partway through. Ordinary models are not.
- Skipped layers leave gaps that later words need, and filling those gaps is the real difficulty.
What to learn next
- Mixture of depths — skipping layers instead of stopping at one.
- Knowledge distillation — training a small model to do the easy work instead.
- Latency and throughput — why per-token savings and per-batch savings differ.
Developer — Code and libraries.
Setup
pip install torch transformersWritten against PyTorch 2.5.1 and transformers 5.6.2, Python 3.10. Runs on CPU in seconds once GPT-2 has downloaded (about 500 MB).
Can you read an ordinary model halfway through?
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
torch.manual_seed(0)
tok = AutoTokenizer.from_pretrained("gpt2")
model = AutoModelForCausalLM.from_pretrained("gpt2").eval()
L = model.config.n_layer
print(f"model: gpt2 | {L} layers | torch: {torch.__version__}")
text = ("The train left the station at dawn. Fields of sugarcane ran past the window for "
"an hour. A vendor walked the aisle selling tea in paper cups. By noon the line "
"climbed into the hills and the air turned cool.") * 2
ids = tok(text, return_tensors="pt").input_ids
T = ids.shape[1]
print("tokens:", T)
with torch.no_grad():
hs = model(ids, output_hidden_states=True).hidden_states # L+1 tensors
# Read out a prediction from every layer using the model's own final head.
# This head was only ever trained on the LAST layer's output. That is the point.
def readout(h):
return model.lm_head(model.transformer.ln_f(h))[0].float()
final_logits = readout(hs[-1])
final = final_logits.argmax(-1)
print(f"\n{'layer':>6}{'top-1 = final':>15}{'rank of final token':>22}{'confidence':>12}")
for l in range(1, L + 1):
lg = readout(hs[l])
agree = (lg.argmax(-1) == final).float().mean()
rank = (lg > lg.gather(-1, final[:, None])).sum(-1).float() + 1 # 1 = already the top pick
conf = lg.softmax(-1).max(-1).values.mean()
print(f"{l:>6}{agree:>14.1%}{rank.median():>22.0f}{conf:>12.3f}")
print("\nthe rank column is the honest signal: the final answer climbs the list")
print("layer by layer. the confidence column is not trustworthy here, because")
print("the output head was never trained to read an intermediate layer.")model: gpt2 | 12 layers | torch: 2.5.1+cu121
tokens: 92
layer top-1 = final rank of final token confidence
1 5.4% 178 0.630
2 5.4% 115 0.570
3 10.9% 66 0.525
4 7.6% 58 0.492
5 9.8% 46 0.487
6 13.0% 55 0.492
7 9.8% 39 0.504
8 8.7% 45 0.670
9 8.7% 27 0.681
10 9.8% 31 0.716
11 15.2% 24 0.691
12 100.0% 1 0.094
the rank column is the honest signal: the final answer climbs the list
layer by layer. the confidence column is not trustworthy here, because
the output head was never trained to read an intermediate layer.Reading the output, including the parts that look wrong
The rank column falls from 178 to 24, then to 1. The final answer really does climb the candidate list as depth increases: 178th at layer 1, 66th at layer 3, 27th at layer 9, 24th at layer 11. Information accumulates gradually and monotonically enough to be convincing. That is the premise of early exit, confirmed.
The top-1 = final column stays around 10% until the last layer. Reading GPT-2 at layer 11 and taking the top prediction gives you the model's actual answer 15% of the time. Naive early exit on this model would be catastrophic.
Layer 11 reports 0.691 mean confidence and layer 12 reports 0.094. This is the trap. The intermediate layers look more confident than the final one and are far less correct. A confidence threshold applied to this readout would exit early, constantly, and be wrong.
The reason is that lm_head and ln_f were only ever trained on the final layer's representation. Applied to an intermediate one they produce a confidently wrong distribution. This is a well-known artefact of the "logit lens" technique, and it is exactly why CALM and LayerSkip train the model to be exitable rather than bolting exits onto a finished model.
The honest conclusion from this experiment is negative, and it is the useful one: you cannot early-exit a model that was not trained for it. Anyone who shows you a clean confidence curve from an off-the-shelf checkpoint has either trained exit heads or is measuring the artefact above.
What a trained early-exit model adds
Early exit loss. Compute the language-modelling loss at several depths during training, not only at the end, so intermediate representations are decodable by the shared head.
Layer dropout. Randomly drop layers during training, with higher rates for later layers, so the model learns to produce a usable answer without them.
A calibrated exit rule. CALM's contribution is largely here: which confidence measure, and how to connect a sequence-level quality constraint to per-token exit decisions with a statistical guarantee, rather than a threshold picked by eye.
The KV cache problem, which decides whether this ships
Exit at layer 6 of 32 and layers 7 to 32 never ran for that token, so they wrote no keys and values into the cache. The next token attends over the full stack and finds holes.
Three responses:
| Response | Cost |
|---|---|
| Copy layer 6's state up through layers 7 to 32 | approximate, cheap, degrades quality |
| Recompute the missing layers later, in a batch | correct, spends part of the saving |
| Never exit before the deepest exit any live token needs | correct, saves much less |
Add continuous batching and it gets worse: different sequences want to exit at different depths in the same step, so the batch is ragged and the GPU runs at the depth of the slowest member.
Self-speculative decoding, which sidesteps all of it
LayerSkip's approach. Use the first n layers as a fast draft model, generate several tokens, then run the full model once to verify them all in parallel.
Two properties make this the practical option. Verification is one forward pass over several positions, which is compute-bound and cheap. And the output is exactly what the full model would have produced, so there is no quality argument to have.
It also avoids a separate draft model entirely, since the draft is a prefix of the same network and shares its weights, compute and activations.
Common mistakes
Applying the final head to intermediate layers and trusting the confidence. Demonstrated above.
Ignoring the KV cache holes. The most common reason a working prototype fails to reproduce its speed-up on a real server.
Measuring per-token latency instead of throughput. With continuous batching, exiting early on one sequence does not help if the batch still runs deep for another.
Confusing early exit with depth pruning. Pruning removes layers permanently for all inputs. Early exit is per token, at run time. Different techniques with different failure modes.
Try it yourself
Change the text to something highly predictable, such as a repeated list of numbers, and re-measure the rank column. It should fall much faster with depth, because easy tokens really are decided early. That contrast, on the same model, is the strongest evidence for the idea that this script can produce.
What to learn next
- Mixture of depths — skipping layers instead of stopping at one.
- Knowledge distillation — training a small model to do the easy work instead.
- Latency and throughput — why per-token savings and per-batch savings differ.
Researcher — Mathematics and papers.
Formulation
Attach an exit head $h_l$ at layer $l$ and a confidence measure $c_l(x)$. Emit at the first layer where $c_l(x) \ge \tau_l$, otherwise continue. Compute per token becomes a random variable determined by the input.
Confidence measures in use: softmax margin between the top two candidates, the top-1 probability, cosine similarity between consecutive hidden states, and a trained lightweight classifier. CALM finds the trained classifier cheapest at equal quality and the softmax measures easiest to calibrate.
The two landmark papers
Schuster, Fisch, Gupta, Dehghani, Bahri, Tran, Tay and Metzler (2022), Confident Adaptive Language Modeling (arXiv:2207.07061, NeurIPS 2022 oral). The paper names the three problems precisely: what confidence measure to use, how to connect sequence-level constraints to per-token exit decisions, and how to handle hidden representations of early-exited tokens. Their framework provably maintains high performance under a chosen quality constraint, with reported speed-ups up to 3x on three generation tasks.
The calibration machinery is the part worth reading. Rather than tuning a threshold, they use distribution-free risk control to pick $\tau$ such that the full-model and early-exit outputs agree within a user-chosen tolerance with high probability. That converts an unbounded quality risk into a stated one.
Elhoushi et al. (2024), LayerSkip: Enabling Early Exit Inference and Self-Speculative Decoding (arXiv:2404.16710, ACL 2024). Training uses layer dropout with lower rates for earlier layers and higher for later ones, plus "an early exit loss where all transformer layers share the same exit". Inference exits early and then uses the remaining layers to "verify and correct" the prediction.
Reported speed-ups: 2.16x on CNN/DM summarisation, 1.82x on coding, 2.0x on TOPv2 semantic parsing. Because the draft and the verifier share weights and activations, memory overhead is far lower than with a separate draft model.
Sharing one exit head across all layers is the detail that makes the recipe cheap: no per-layer parameters, and every layer's representation is pushed toward a common decodable space.
The state problem, formally
Let $\ell_t$ be the exit layer for token $t$. Token $t' > t$ attends at layer $l$ over ${K_l^{(s)}, V_l^{(s)}}_{s \le t'}$, but $K_l^{(t)}$ exists only for $l \le \ell_t$.
Options and their properties:
- State propagation. Set $h_l^{(t)} = h_{\ell_t}^{(t)}$ for $l > \ell_t$ and compute keys and values from it. Cheap, and introduces an approximation whose error compounds along the sequence.
- Deferred recomputation. Fill the holes in a later batched pass. Exact, and consumes part of the saving.
- Monotonic exit constraint. Require $\ell_t$ non-decreasing over $t$, so holes never appear below a live prefix. Correct and substantially less adaptive.
CALM lists this as one of its three named challenges, which is a fair reflection of how much it dominates practical work.
Related but distinct
Depth pruning removes layers statically for all inputs, typically selecting by angular distance between a layer's input and output. Simpler, no state problem, and no adaptivity.
Mixture of depths skips individual layers while keeping the rest of the stack, so a skipped token can still be processed later. Early exit terminates the stack. Comparisons in the follow-up literature argue MoD preserves quality better for this reason — see mixture of depths.
Speculative decoding with a separate small draft model achieves the same goal with an entirely different mechanism. LayerSkip's self-speculation is the hybrid: the draft is a prefix of the verifier.
Recent directions
- GateSkip (arXiv:2510.13876) puts a small linear gate with a sigmoid on each attention and MLP branch of the residual stream, learning per-branch skipping rather than whole-layer exit.
- LayerRoute (arXiv:2606.01838) adapts a pretrained model post hoc via LoRA fine-tuning in minutes, routing at sequence level rather than token level, which removes both the causality and the ragged-batch problems at the cost of granularity.
The move toward sequence-level and branch-level decisions in both is telling. Token-level adaptive depth is the theoretically attractive formulation and the one that fights hardest with batched serving.
The standing assessment
The verification-based framing has won on practical grounds. Adaptive depth used as a draft gives a guaranteed-exact output and composes with continuous batching, while adaptive depth used as the answer requires calibration, state repair and ragged batching to all work at once.
That is why LayerSkip's self-speculative decoding is the form of this idea most likely to appear in a serving stack you use, and why plain confidence-threshold early exit remains largely a research setting.
What to learn next
- Mixture of depths — skipping layers instead of stopping at one.
- Knowledge distillation — training a small model to do the easy work instead.
- Latency and throughput — why per-token savings and per-batch savings differ.