NLP Bascis
KV Cache
Overview
KV cache is an inference optimization used in autoregressive Transformer models.
During generation, the model produces one new token at a time. At decoding step (t), the model only needs the hidden state of the last token to predict the next token.
Without KV cache, however, the Transformer repeatedly processes the entire sequence from the beginning. This means that representations for previous tokens are recomputed at every decoding step.
KV cache avoids this redundant computation by storing the previously computed Key and Value tensors and reusing them in later decoding steps.
Main Idea of KV Cache
Suppose the model is generating the sequence:
Without KV cache, the input to the Transformer grows at every decoding step:
Accordingly, the Transformer also produces outputs for every token:
However, during autoregressive generation, only the last output is needed to predict the next token:
The other (t-1) outputs were already computed in previous decoding steps.
This means that without KV cache, the model repeatedly recomputes representations for old tokens even though only the newest token's output is actually needed.
The same redundancy occurs when computing the Key and Value tensors. At every step, (K) and (V) for all previous tokens are calculated again.
KV cache solves this by storing the previously calculated (K) and (V).
At step (t), instead of processing
the Transformer only processes the newly generated token:
It computes only
for that token.
The new (k_t) and (v_t) are then appended to the cached values:
Therefore, attention can use all previous Keys and Values without recomputing them.
The important distinction is:
while with KV cache:
The history is still available to attention through the cached
How Time Complexity Changes
The following analysis considers one decoding step at time (t), treating (N,D,d_k,d_v) as constants and focusing on how the cost grows with sequence length (t).
Without KV Cache
At time (t), the entire sequence (X_{1:t}) is processed again.
| Step | Shape / Operation | Time Complexity |
|---|---|---|
| (Q_{1:t}=X_{1:t}W_Q) | ((N,t,D)(D,d_k)) | (O(NtDd_k)\approx O(t)) |
| (K_{1:t}=X_{1:t}W_K) | ((N,t,D)(D,d_k)) | (O(NtDd_k)\approx O(t)) |
| (V_{1:t}=X_{1:t}W_V) | ((N,t,D)(D,d_v)) | (O(NtDd_v)\approx O(t)) |
| (Q_{1:t}K_{1:t}^{T}) | ((N,t,d_k)(N,d_k,t)) | (O(Nt^2d_k)\approx O(t^2)) |
| Softmax | ((N,t,t)) | (O(Nt^2)\approx O(t^2)) |
| (\operatorname{Softmax}(QK^T)V) | ((N,t,t)(N,t,d_v)) | (O(Nt^2d_v)\approx O(t^2)) |
| Output projection | ((N,t,d_v)W_O) | (O(Ntd_vD)\approx O(t)) |
The dominant attention operations are therefore:
for a single decoding step.
With KV Cache
With KV cache, only the newest token is projected into (q_t,k_t,v_t).
| Step | Shape / Operation | Time Complexity |
|---|---|---|
| (q_t=x_tW_Q) | ((N,1,D)(D,d_k)\rightarrow(N,1,d_k)) | (O(NDd_k)\approx O(1)) |
| (k_t=x_tW_K) | ((N,1,D)(D,d_k)\rightarrow(N,1,d_k)) | (O(NDd_k)\approx O(1)) |
| (v_t=x_tW_V) | ((N,1,D)(D,d_v)\rightarrow(N,1,d_v)) | (O(NDd_v)\approx O(1)) |
| (q_tK_{1:t}^{T}) | ((N,1,d_k)(N,d_k,t)\rightarrow(N,1,t)) | (O(Ntd_k)\approx O(t)) |
| Softmax | ((N,1,t)) | (O(Nt)\approx O(t)) |
| (\operatorname{Softmax}(q_tK^T)V_{1:t}) | ((N,1,t)(N,t,d_v)\rightarrow(N,1,d_v)) | (O(Ntd_v)\approx O(t)) |
| Output projection | ((N,1,d_v)W_O) | (O(Nd_vD)\approx O(1)) |
Therefore, the dominant attention computation at one decoding step becomes:
instead of
Time Complexity for Generating All (T) Tokens
Without KV cache, attention at decoding step (t) costs approximately:
Therefore:
The total attention computation for generating (T) tokens is:
Since
the total complexity is:
With KV Cache
With KV cache, attention at step (t) costs:
Therefore:
The total attention computation becomes:
Since
the total complexity is:
So, with respect to sequence length:
for autoregressive attention over the entire generation.
Simple Summary
Without KV Cache
At time (t), the model receives:
It calculates:
and
Attention is calculated using:
The MHA output has shape:
even though only its last position is needed for generating the next token.
With KV Cache
At time (t), the model receives only the newest token:
It calculates only:
and
The new (k_t) and (v_t) are appended to the previous KV cache:
Attention is then calculated using:
The MHA output has shape:
In short: