KV caching in LLMs, clearly explained (with visuals):

Before understanding the internal details, look at the inference speed difference in the video:
- with KV caching → 9 seconds
- without KV caching → 42 seconds (~5x slower)
Let's dive in!
- Transformer produces hidden states for all tokens.
- Hidden states are projected to vocab space.
- Logits of the last token is used to generate the next token.
- Repeat for subsequent tokens.
Check this👇
None of the other hidden states are required.
Next, let's see how the last hidden state is computed within the transformer layer from the attention mechanism.
The last row of query-key-product involves:
- the last query vector.
- all key vectors.
Also, the last row of the final attention result involves:
- the last query vector.
- all key & value vectors.
Check this visual to understand better:
- query vector of the last token.
- all key & value vectors.
But, there's one more key insight here.
- The KV vectors used for ALL previous tokens do not change.
Thus, we just need to generate a KV vector for the token generated one step before.
Rest of the KV vectors can be retrieved from a cache to save compute and time.
To reiterate, instead of redundantly computing KV vectors of all context tokens, cache them.
To generate a token:
- Generate QKV vector for the token generated one step before.
- Get all other KV vectors from cache.
- Compute attention.
Check this👇
In fact, this is why ChatGPT takes some time to generate the first token than the subsequent tokens.
During that time, it is computing the KV cache of the prompt.
Llama3-70B has:
- total layers = 80
- hidden size = 8k
- max output size = 4k
Here:
- Every token takes up ~2.5 MB in KV cache.
- 4k tokens will take up 10.5 GB.
More users → more memory.
I'll cover KV optimization soon.
If you enjoyed this tutorial:
Find me → @_avichawla
Every day, I share tutorials and insights on DS, ML, LLMs, and RAGs.