From 81788faef4ff8cd9490af934e3648be855efa178 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 7 Aug 2026 18:19:17 +0800 Subject: [PATCH] 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 --- astrai/extension/attention_backend.py | 59 ++++++++++++++++++++++++--- 1 file changed, 53 insertions(+), 6 deletions(-) diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index c053e69..a672303 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -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)