feat: GPU-native SWA + sparse decode attention kernels (CuTeDSL)
- native_swa_decode.py: BlackwellSWADecodeKernel
- CTA mapping: 1 CTA per (decode_token, q_head_group)
- Online softmax with KV tile streaming (16 tokens/tile)
- Pre-dequantized bf16 KV (fp8 dequant on host - MLIR cvt_fpext
requires 32-bit aligned vector, no scalar fp8->bf16 support)
- Cosine 0.9999+ vs PyTorch batched SDPA reference
- Fallback _fallback_batched_sdp when CuTeDSL unavailable
- native_sparse_decode.py: BlackwellSparseDecodeKernel
- Combined SWA + compressed KV in single attention pass
- Supports CSA (cr=4) and HCA (cr=128) layers
- Sink weight merge on host side
- Cosine 0.9999+ vs combined SDPA reference
- fp8_bf16.py: Documents MLIR limitation (cvt_fpext requires
vector<4xf8>, no scalar support). Pre-dequant is the workaround.
- vLLM wiring (attention.py):
- SWA-only layers: native_swa_decode_attention
- CSA/HCA layers: native_sparse_decode_attention with topk + attn_sink
- csa_attention.py updated to use native kernels
- Tests: test_decode_pipeline.py, test_sparse_decode.py both passing
This commit is contained in:
27
README.md
27
README.md
@@ -56,7 +56,25 @@ vLLM's internal kernels (FlashMLA, fp8_ds_mla, fused compressor, Triton indexer)
|
||||
|
||||
`CuTeDSLNvfp4Linear` — single-expert NVFP4 GEMM for shared experts and attention projections.
|
||||
|
||||
### ✅ Blackwell Attention (standalone, not yet in vLLM)
|
||||
### ✅ GPU-Native SWA Decode Attention (CuTeDSL)
|
||||
|
||||
`cutedsl/native_swa_decode.py` — `BlackwellSWADecodeKernel`:
|
||||
- CTA mapping: 1 CTA per (decode_token, q_head_group) — 8 groups × T tokens
|
||||
- Q loaded into registers, KV streamed in 16-token tiles through smem
|
||||
- Online softmax (max/exp/rescale/sum) across tiles
|
||||
- Pre-dequantized bf16 KV (fp8 dequant done on host, fused dequant is future work)
|
||||
- **Cosine 0.9999+** vs PyTorch batched SDPA reference
|
||||
|
||||
### ✅ GPU-Native Sparse + SWA Decode Attention (CuTeDSL)
|
||||
|
||||
`cutedsl/native_sparse_decode.py` — `BlackwellSparseDecodeKernel`:
|
||||
- Same CTA mapping as SWA kernel
|
||||
- Concatenated SWA + compressed KV in a single attention pass
|
||||
- Sink weight merge applied on host side
|
||||
- **Cosine 0.9999+** vs combined SDPA reference
|
||||
- Supports both CSA (cr=4) and HCA (cr=128) layers
|
||||
|
||||
### ✅ Blackwell Attention (standalone tests)
|
||||
|
||||
- `cutedsl/blackwell_attention.py` — KV cache write/read, full attention pipeline
|
||||
- `cutedsl/csa_attention.py` — CSA (cr=4) and HCA (cr=128) sparse attention
|
||||
@@ -180,8 +198,11 @@ The custom CUDA quantize kernel needs the **L2 activation global scale** (from t
|
||||
| What | Status | Notes |
|
||||
|------|--------|-------|
|
||||
| In-epilogue NVFP4 quantize (replace BF16 TMA with FP4 TMA) | 🔨 Future | Saves ~0.14ms/layer; requires register→GMEM mapping for FP4 output |
|
||||
| GPU-native KV cache + attention for vLLM | 🔨 Next | All standalone kernels work; need vLLM backend wiring |
|
||||
| vLLM model integration | 🔨 Next | Model definition, weight loading, attention backend |
|
||||
| GPU-native SWA decode attention | ✅ Done | CuTeDSL kernel, cosine 0.9999+ |
|
||||
| GPU-native sparse + SWA decode attention | ✅ Done | CuTeDSL kernel, cosine 0.9999+ |
|
||||
| vLLM Blackwell decode path | ✅ Done | _attention_impl_blackwell uses native SWA + sparse kernels |
|
||||
| Fuse fp8→bf16 dequant into CuTeDSL kernel | 🔨 Future | Currently pre-dequantized on host; need vectorized fp8 loads |
|
||||
| CSA/HCA sink weight merge in CuTeDSL | 🔨 Future | Applied on host for now; fuse into kernel for perf |
|
||||
|
||||
---
|
||||
|
||||
|
||||
Reference in New Issue
Block a user