Skip to content

vllm.v1.attention.ops.turboquant_soa.triton_turboquant_decode_v2

Optimized Triton TurboQuant decode attention (v2).

FLUTE-paper optimizations applied
  1. 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.
  2. 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.
  3. exp2 instead of exp: scores pre-scaled by log2(e) so the hardware- native exp2 instruction replaces the more expensive exp.
  4. Wider index extraction: for 4-bit MSE, two adjacent 4-bit indices share a byte. One byte load yields both, eliminating redundant loads.
  5. Centroids pre-warmed in L1 at kernel start.
  6. 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:

_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
def _get_pair_lut(centroids: torch.Tensor) -> torch.Tensor:
    """Return a fresh pair-LUT for ``centroids`` on each call.

    The LUT is tiny (N*N*2 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.
    """
    return build_pair_lut(centroids)

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.

Source code in vllm/v1/attention/ops/turboquant_soa/triton_turboquant_decode_v2.py
def build_pair_lut(centroids: torch.Tensor) -> torch.Tensor:
    """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 [N*N, 2] for kernel indexing: pair_lut[i*N + j].

    For 4-bit MSE (N=16): 16*16*2*4 = 2048 bytes — fits in L1/smem.
    For 3-bit MSE (N=8):   8*8*2*4  = 512 bytes.
    """
    N = centroids.shape[0]
    # pair_lut[i,j,0] = centroids[i], pair_lut[i,j,1] = centroids[j]
    c = centroids.float()
    lut = torch.empty(N, N, 2, dtype=torch.float32, device=centroids.device)
    lut[:, :, 0] = c[:, None]
    lut[:, :, 1] = c[None, :]
    return lut.reshape(N * N, 2).contiguous()