From 7f0e8bb8c2b7b616fc2a719eab3fb63d88bcc399 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sat, 8 Aug 2026 23:50:59 +0800 Subject: [PATCH] fix: let flash backend handle 4D causal prefill mask - Treat 4D masks as causal (flash handles it natively), keep rejecting custom non-causal masks - Enables flash backend in benchmark --compare and real prefill path --- astrai/extension/attention_backend.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index c06cd09..2bc6510 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -140,7 +140,9 @@ def _backend_supports( return False if q.size(1) == 1 and kv_cache is not None: return True - return attn_mask is None + if attn_mask is None or is_causal: + return True + return attn_mask.dim() == 4 return True @@ -640,7 +642,7 @@ class FlashAttnBackend(AttentionBackend): k = repeat_kv(k, n_rep) v = repeat_kv(v, n_rep) - if attn_mask is not None and not is_causal: + if attn_mask is not None and not is_causal and attn_mask.dim() != 4: raise ValueError( "FlashAttnBackend does not support a custom attention mask; " "use a causal mask or select TorchNativeBackend." @@ -652,7 +654,10 @@ class FlashAttnBackend(AttentionBackend): "Install with `pip install flash-attn`." ) out = fa.flash_attn_func( - q.contiguous(), k.contiguous(), v.contiguous(), causal=is_causal + q.contiguous(), + k.contiguous(), + v.contiguous(), + causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4), ) return out.contiguous().flatten(2)