feat: unify attention backend with multi-dim mask support

- Add attention() functional entry delegating to active backend
- GQA/MLA forward calls attention() instead of inline cache/SDPA
- CUDA kernels support 2D/3D/4D mask via mask_h_stride field
- CudaBackend.fwd_decode builds 2D padding mask for mixed seq_lens
- KVCache.max_len precomputed in bind_tasks to avoid GPU sync
- batch==1 decode short-circuits mask=None
- Split tests into conftest, test_backend, test_backend_equivalence, test_kernel_mask
- 440 tests pass, L20 decode 1.44-1.60x speedup vs torch native
This commit is contained in:
2026-07-30 20:38:34 +08:00
parent 97114b95a4
commit 3067a8e1a6
19 changed files with 438 additions and 81 deletions
+2
View File
@@ -103,6 +103,7 @@ inline void set_default_strides(P& p) {
p.kv_stride_l = p.head_dim;
p.kv_stride_d = 1;
p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}
@@ -114,6 +115,7 @@ inline void set_default_paged_strides(P& p) {
p.q_stride_l = p.head_dim;
p.q_stride_d = 1;
p.mask_b_stride = p.kv_len;
p.mask_h_stride = 0;
p.mask_q_stride = 0;
}