perf: use flash_attn_with_kvcache for contiguous cache decode

- Decode with contiguous cache uses flash_attn_with_kvcache instead of materializing full KV via gather + flash_attn_func
- _backend_supports allows FlashAttnBackend for decode (q_len==1) even with explicit mask
- Decode speedups vs TorchNative (B=1,4,8,16 mean): cuda 1.55x, flash 1.40x, torch_native 1.00x
- Read K/V directly from flat pool via cache_batch_idx + cache_seqlens, zero-copy view reshape
This commit is contained in:
2026-08-07 18:21:02 +08:00
parent 0e7fe57d96
commit 81788faef4
+53 -6
View File
@@ -170,7 +170,11 @@ def _backend_supports(
and q.size(-1) in (32, 64, 128, 256)
)
if isinstance(backend, FlashAttnBackend):
return flash_attn_available() and not (attn_mask is not None and not is_causal)
if not flash_attn_available():
return False
if q.size(1) == 1 and kv_cache is not None:
return True
return not (attn_mask is not None and not is_causal)
return True
@@ -554,15 +558,27 @@ class CudaBackend(AttentionBackend):
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
def _kv_cache_is_contiguous(kv_cache: "KVCache") -> bool:
return kv_cache.k_buffer.size(1) == kv_cache.req_to_token.size(
0
) * kv_cache.req_to_token.size(1)
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
class FlashAttnBackend(AttentionBackend):
"""FlashAttention (FA2/FA3) backend via the optional ``flash-attn`` package.
"""FlashAttention backend via the optional ``flash-attn`` package.
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``,
which reads K/V directly from the flat pool via cache_batch_idx +
cache_seqlens — no materialized KV gather.
Prefill / non-contiguous decode: falls back to KV gather +
``flash_attn_func``.
Uses the general ``flash_attn_func`` entry point for both prefill and
single-token decode, mirroring ``TorchNativeBackend``'s KV-cache gather.
This backend only does flash attention — inputs ``flash-attn`` cannot
express (missing package, custom attention mask, fp32, unsupported
head_dim) raise a clear error instead of silently falling back to torch.
express (missing package, custom attention mask on prefill, fp32,
unsupported head_dim) raise a clear error instead of silently falling
back to torch.
For a torch fallback, select ``TorchNativeBackend`` instead.
"""
@@ -602,6 +618,9 @@ class FlashAttnBackend(AttentionBackend):
is_causal: bool = False,
) -> Tensor:
if kv_cache is not None:
if q.size(1) == 1 and _kv_cache_is_contiguous(kv_cache):
return self._decode_with_kvcache(q, k, v, kv_cache, layer_id)
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
@@ -638,3 +657,31 @@ class FlashAttnBackend(AttentionBackend):
q.contiguous(), k.contiguous(), v.contiguous(), causal=is_causal
)
return out.contiguous().flatten(2)
def _decode_with_kvcache(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: "KVCache",
layer_id: int,
) -> Tensor:
max_batch = kv_cache.req_to_token.size(0)
max_seq = kv_cache.req_to_token.size(1)
n_kv = k.size(2)
k_cache = kv_cache.k_buffer[layer_id].view(max_batch, max_seq, n_kv, k.size(3))
v_cache = kv_cache.v_buffer[layer_id].view(max_batch, max_seq, n_kv, v.size(3))
fa = _get_flash_attn()
out = fa.flash_attn_with_kvcache(
q=q,
k_cache=k_cache,
v_cache=v_cache,
k=k,
v=v,
cache_seqlens=(kv_cache.seq_lens - 1).to(torch.int32),
cache_batch_idx=kv_cache.req_pool_indices.to(torch.int32),
causal=True,
)
return out.flatten(2)