(Specifically, there are two O(n^2) steps in an attention layer and KV caching makes the first one O(n) with caching - but the overall big-O is still O(n^2) because KV caching doesn't affect the second step.)
(Specifically, there are two O(n^2) steps in an attention layer and KV caching makes the first one O(n) with caching - but the overall big-O is still O(n^2) because KV caching doesn't affect the second step.)
softmax(QK^T/sqrt(d_k))V
So without KV-caching, if suppose you have N tokens, then you have Q: N x q_dim (leaving the batch size and n heads out for simplicity, they are constants for our purposes here)
K^T: k_dim x N
So QK^T: N x N <- this is the attention matrix
The scaling and softmax are irrelevant here, since they are elementwiseThen
V: N x v_dim
(QK^T)V: N x v_dim
This gives us the normal quadratic complexity of prefills (of course, in an actual implementation an attention mask is used to ensure that tokens cannot attend to prior tokens and you may have things like ALiBi).Note that during decoding, the dimensions change. Since it is an autoregressive model, we do not need to recompute the values and keys of prior tokens, only of the token that we are currently decoding. Of course, the token that we are decoding still attends to all prior tokens and itself.
Q: 1 x q_dim
K^T: k_dim x (N+1)
QK^T: 1 x (N+1) <- Note that attention in this step is not quadratic anymore, since
we only need to compute how the current token attends to prior
tokens, the representations of prior tokens are frozen. K comes
from the KV-cache.
V: (N + 1) x v_dim <- V comes from the KV-cache
(QK^T)V: 1 x v_dim <- Also not quadratic, the value is only computed for the current token.
So attention during a decoding step is O(N), so when decoding N tokens, it is O(N^2) overall.The point-wise feed-forward layer does not matter, in decoding it only needs to be computed for the representation of the token that we are currently generating. We don't need the representations of the preceding tokens for the next layer, since we have already cached their keys and values for each layer.
Disclaimer: I was one of the developers of a widely-used inference engine and implemented several of these optimizations.