Multi-Head Latent Attention

Instead of choosing how many heads share keys and values, compress them. Store one small latent vector per token and let every head expand its own keys and values from it. The cache shrinks like multi-query's, yet head diversity survives — the best of both worlds.

Large Language Models: From Transformers to Frontier Models

Every method so far saved cache memory the same way: by making heads share keys and values, and paying for it in lost diversity. Grouped-query attention shares less brutally than multi-query, but it is still a compromise on that one axis. This chapter is about a method that refuses the compromise. Multi-head latent attention (MLA) keeps a cache as small as multi-query's and lets every head keep its own distinct keys and values. It does this by changing the question entirely — from "how many heads share?" to "what, exactly, do we need to store?"

The reframing: store a summary, not the keys and values

Here is the shift in thinking. All along we have cached the keys and values themselves. But the keys and values are computed from the token — they are projections of the token's vector. What if, instead of storing the finished keys and values for every head, we store a single small summary of the token, and reconstruct each head's keys and values from that summary when needed?

That summary is a latent vector. For each token, MLA projects its vector down into a compact latent code — much smaller than the full set of per-head keys and values — and caches only that. Then, whenever attention needs the keys and values, it projects the latent code back up into each head's own keys and values, using a separate learned up-projection per head.

Two consequences follow immediately, and they are exactly the two things we could never get at once before:

  • The cache is tiny. We no longer store a key and value per head. We store one latent vector per token, whose size we choose. In a frontier-scale model this replaces a per-token footprint of (128 heads × 128 dims × 2, for K and V) with a single latent of a few hundred numbers — a reduction of roughly 50–60×, comparable to multi-query's saving.
  • Head diversity is preserved. Each head has its own up-projection, so each reconstructs a different set of keys and values from the shared latent. The heads are not forced to share keys and values at all — they share only the compact summary they each expand differently. This is the quality that multi-query and grouped-query gave up.
A token projected down into a small latent vector that is cached, then expanded by per-head up-projections into distinct keys and values for each head
MLA caches one small latent per token (down-projection), then each head expands it into its own keys and values (per-head up-projection). Small cache, yet diverse heads.

"But we added a step — how is that cheaper?"

A fair objection: we now project down to a latent and then up to keys and values — two matrix multiplies where before there was one. At face value that is more work and, if we cached the reconstructed keys and values, no memory saving at all. The trick that makes MLA actually pay off is an algebraic rearrangement, sometimes called absorption.

The key fact is that the up-projections are fixed after training — they are learned weights, constants at inference time. And attention is a chain of matrix multiplications: query times key, then times value. Because those multiplications can be re-associated, the fixed up-projection matrices can be folded into the neighbouring fixed matrices — the query projection on one side, the output projection on the other — once, ahead of time. After this folding, computing a query's attention scores no longer requires first reconstructing the keys at all: it can act directly on the cached latent. The up-projection never has to run at inference as a separate, cached step; it has been absorbed into weights that were going to be applied anyway.

The upshot: at inference, each new token is turned into its small latent (cached), and attention scores and value blends are computed straight from the latents using the pre-folded weights. You get the small latent cache and per-head-distinct keys and values and no extra cached computation. The intuition to keep is this:

Cache a compressed summary of each token; because the decompression is fixed weights, it can be absorbed into the surrounding maths, so the heads get their own keys and values for free.

Why 'latent'

A latent representation is a compact code that captures what matters about something in fewer numbers — an idea used across machine learning (autoencoders compress images this way, as the Deep Learning course's unsupervised pre-training lesson shows). MLA's insight was to bring that idea to the KV cache: the latent is a learned compression of everything a token needs to contribute to attention, and the model learns to compress and decompress it without losing what matters.

Best of both worlds, made concrete

Lay the four methods on the two axes we have been tracking — cache size and quality:

MethodCache sizeHead diversity / quality
Multi-headlargestbest
Multi-querysmallestworst
Grouped-querymediummedium
Latent (MLA)smallest-classnear multi-head

MLA is the first entry that is not on the straight line between the two extremes — it reaches down to a multi-query-sized cache while staying up near multi-head quality. That combination is a large part of why the models that introduced it could serve very long contexts cheaply without sacrificing accuracy. It is the clearest example in this whole course of an idea that does not trade off two goods but genuinely achieves both, by changing what is stored rather than how much is shared.

One loose thread

There is a catch we have glossed over, and it is the subject of the next chapter. The absorption trick relies on the keys being plain projections of the token that can be folded into fixed weights. But modern models inject position into the keys using rotary encodings (RoPE) from Module 5 — and a RoPE rotation depends on the token's position, so it is not a fixed weight that can be absorbed. Latent attention and rotary position, as described so far, do not combine cleanly. Resolving that collision is the final idea of this module.

Try it yourself
Transformer Lab: latent attention →

Compare a compressed latent cache with full multi-head attention at 128k tokens.

MediumMLA

How does MLA shrink the cache without forcing heads to share keys and values?

HardMLAabsorption

MLA adds a down-projection and an up-projection. Why doesn't that extra step cost more memory or compute at inference?

MediumMLA

Why is MLA described as 'off the line' between multi-head and multi-query attention?