vllm.v1.attention.ops.turboquant_soa.triton_turboquant_decode_v2
¶
Optimized Triton TurboQuant decode attention (v2).
FLUTE-paper optimizations applied
- Grouped Q heads: grid over (B, Hk, splits) instead of (B, Hq, splits). Each program loads BLOCK_M Q heads sharing a KV head into a 2D tile, enabling tl.dot on tensor cores (MFMA/WMMA) for both Q·K and P·V.
- Vectorized pair LUT: precompute
pair_table[i][j] = (T[i], T[j])offline. At runtime, extract adjacent index pairs and fetch two dequantized centroids with a single gather, halving LUT lookups. - exp2 instead of exp: scores pre-scaled by log2(e) so the hardware- native exp2 instruction replaces the more expensive exp.
- Wider index extraction: for 4-bit MSE, two adjacent 4-bit indices share a byte. One byte load yields both, eliminating redundant loads.
- Centroids pre-warmed in L1 at kernel start.
- BLOCK_KV = TILE_SIZE raised to 16-32 (from 4), reducing loop iterations and softmax rescaling overhead.
Stage 2 is reused unchanged from triton_decode_attention.py.
Functions:
-
build_pair_lut–Build vectorized pair lookup table.
_get_pair_lut(centroids)
¶
Return a fresh pair-LUT for centroids on each call.
The LUT is tiny (NN2 fp32, e.g. 2KB for 4-bit MSE) so the build cost is negligible compared to attention. We avoid caching by data_ptr() because CUDA allocator memory reuse across different centroid tensors can silently return a stale LUT (subtle correctness bug). If this ever shows up on a profile, cache by a hash-of-values fingerprint instead.
Source code in vllm/v1/attention/ops/turboquant_soa/triton_turboquant_decode_v2.py
build_pair_lut(centroids)
¶
Build vectorized pair lookup table.
For N centroids, returns a [N, N, 2] float32 tensor where pair_lut[i, j] = (centroids[i], centroids[j]). Flattened to [NN, 2] for kernel indexing: pair_lut[iN + j].
For 4-bit MSE (N=16): 161624 = 2048 bytes — fits in L1/smem. For 3-bit MSE (N=8): 8824 = 512 bytes.