mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-20 12:40:14 +00:00
[Doc][Attention] Fix MLA top-of-file comments (#37047)
Signed-off-by: wineandchord <[email protected]>
This commit is contained in:
@@ -14,7 +14,7 @@ MLA has two possible ways of computing, a data-movement friendly approach and a
|
||||
compute friendly approach. We generally want to use the compute friendly
|
||||
approach for "prefill" (i.e. the ratio Sq / Skv is relatively large, often near
|
||||
1) and the data-movement friendly approach for "decode" (i.e. the ratio
|
||||
Sq / Skv is small).
|
||||
Sq / Skv is small, often near 0).
|
||||
|
||||
NOTE what we deem small and large is currently determined by if it is labelled
|
||||
prefill or decode by the scheduler, but this is something we should probably
|
||||
@@ -28,7 +28,7 @@ Deepseek's MLA attention works the following way:
|
||||
* For decode (i.e. the memory friendly approach) the attention "simulates" a
|
||||
multi-head attention, while the compute is similar to multi-query attention.
|
||||
|
||||
Below is example of both paths assuming batchsize = 1
|
||||
Below is an example of both paths assuming batch size = 1
|
||||
|
||||
## More Extent Definitions:
|
||||
|
||||
@@ -77,13 +77,13 @@ v = (kv_c @ W_UV.view(Lkv, N * V)).view(Skv, N, V)
|
||||
|
||||
// MHA with QK headdim = P + R
|
||||
// V headdim = V
|
||||
// spda_o shape [Sq, N, V]
|
||||
spda_o = scaled_dot_product_attention(
|
||||
// sdpa_o shape [Sq, N, V]
|
||||
sdpa_o = scaled_dot_product_attention(
|
||||
torch.cat([q_nope, q_pe], dim=-1),
|
||||
torch.cat([k_nope, k_pe.unsqueeze(1).expand(-1, N, -1)], dim=-1),
|
||||
v
|
||||
)
|
||||
return spda_o @ W_O
|
||||
return sdpa_o @ W_O
|
||||
|
||||
NOTE: in the actual code,
|
||||
`kv_b_proj` is [W_UK; W_UV] concatenated per head
|
||||
@@ -105,16 +105,16 @@ k_pe = torch.cat([new_k_pe, cache_k_pe], dim=0)
|
||||
|
||||
// MQA with QK headdim = Lkv + R
|
||||
// V headdim = Lkv
|
||||
// spda_o shape [Sq, N, Lkv]
|
||||
// sdpa_o shape [Sq, N, Lkv]
|
||||
// NOTE: this is less compute-friendly since Lkv > P
|
||||
// but is more data-movement friendly since its MQA vs MHA
|
||||
spda_o = scaled_dot_product_attention(
|
||||
sdpa_o = scaled_dot_product_attention(
|
||||
torch.cat([ql_nope, q_pe], dim=-1),
|
||||
torch.cat([kv_c, k_pe], dim=-1),
|
||||
kv_c
|
||||
)
|
||||
|
||||
o = einsum("snl,lnv->snv", spda_o.reshape(-1, N, Lkv), W_UV)
|
||||
o = einsum("snl,lnv->snv", sdpa_o.reshape(-1, N, Lkv), W_UV)
|
||||
return o.view(-1, N * V) @ W_O
|
||||
|
||||
|
||||
@@ -153,7 +153,7 @@ curr_o, curr_lse = scaled_dot_product_attention(
|
||||
torch.cat([q_nope, q_pe], dim=-1),
|
||||
torch.cat([new_k_nope, new_k_pe.unsqueeze(1).expand(-1, N, -1)], dim=-1),
|
||||
new_v,
|
||||
casual=True,
|
||||
causal=True,
|
||||
return_softmax_lse=True
|
||||
)
|
||||
|
||||
@@ -173,7 +173,7 @@ for chunk_idx in range(cdiv(C, MCC)):
|
||||
cache_k_pe_chunk.unsqueeze(1).expand(-1, N, -1)],
|
||||
dim=-1),
|
||||
cache_v_chunk,
|
||||
casual=False,
|
||||
causal=False,
|
||||
return_softmax_lse=True
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user