[ENGINE] paged_attn.py: CCCL dispatch_reduce architecture port

Three changes from reading CCCL dispatch_reduce.cuh + kernel_reduce.cuh +
agent_reduce.cuh + grid_even_share.cuh + summary_statistics.cu:

1. V2 dispatch restored (was hardcoded use_v1=True)
   CCCL two-path: single-tile vs multi-tile (GridEvenShare).
   Threshold now uses BI-V100 SM count (16) for saturation calc.

2. _forward_decode_pytorch rewritten with CCCL patterns:
   agent_reduce ConsumeFullTile: reduced .contiguous() from 4 to 2.
   GridEvenShare RAKE tiling: adaptive _MAX_TILE_BLOCKS=1024.
   summary_statistics.cu compound reduce: online softmax {m,l,o}.

3. KV gather: permute(1,2,4,0,3) for K avoids intermediate alloc.

CCCL files read: dispatch_reduce.cuh, kernel_reduce.cuh,
agent_reduce.cuh, grid_even_share.cuh, summary_statistics.cu,
kernel_scan.cuh
This commit is contained in:
muh-engine
2026-08-05 09:22:46 +00:00
parent 821c59500d
commit 8c969ce7dc

View File

@@ -96,10 +96,30 @@ class PagedAttention:
) -> torch.Tensor:
"""Pure-PyTorch decode attention for long contexts (no hardware kernel).
paged_attention_v1 hangs on BI-V100 when max_seq_len > ~32K due to
shared memory limits. For decode, q_len=1 per sequence so no Q-tiling
is needed — the attention weight tensor is [H, 1, seq_len] which is
trivially small (~5 MB at 50K).
Architecture mirrors CCCL's three-layer reduce:
dispatch_reduce.cuh → kernel_reduce.cuh → agent_reduce.cuh
(work distribution) (kernel entry) (tile consumption)
CCCL agent_reduce.cuh has two key patterns we translate here:
1. ConsumeFullTile vectorized path: data loaded as VectorT in striped
access (no BlockLoad staging → no SMEM for data, only for BlockReduce
scratch). PyTorch equivalent: single reshape+view without .contiguous()
when possible; fall back to one .contiguous() per K/V gather.
2. ConsumeTiles with GridEvenShare STRIP_MINE: each CTA strides across
the input with stride = grid_size * tile_items. For decode (q_len=1),
we tile over KV blocks with adaptive tile_sz per the same
GridEvenShare formula: max_tiles = sm_count * subscription_factor.
3. summary_statistics.cu compound reduce: accumulator = {m, l, o}.
unary_op: score_tile → (max, sum_exp, weighted_V).
binary_op: online softmax merge with correction factor.
This is the Flash Attention online softmax — identical structure.
For decode, q_len=1 per sequence. The attention weight is [H, 1, seq_len]
which is small (~5 MB at 50K tokens). We tile over KV blocks to control
peak memory and apply online softmax (Flash Attention Algorithm 1) per tile.
Shapes
------
@@ -114,71 +134,146 @@ class PagedAttention:
block_size = value_cache.shape[3]
gqa_ratio = num_heads // num_kv_heads
orig_dtype = query.dtype
dev = query.device
output = torch.empty_like(query)
# ================================================================
# KV cache gather strategy — from CCCL agent_reduce.cuh
# CCCL GridEvenShare adaptive tile sizing for decode
#
# agent_reduce has two load paths:
# 1. Vectorized: aligned, contiguous, trivially relocatable, sizeof ≤ 8
# → loads VectorT (e.g. float4) in striped access
# 2. Scalar: fallback with CacheModifiedInputIterator
# dispatch_reduce.cuh:
# max_blocks = sm_occupancy × sm_count × subscription_factor
# tile_size = threads_per_block × items_per_thread
#
# PyTorch equivalent: .contiguous() ensures vectorized GPU memory access.
# The key optimization from agent_reduce is to minimize the number of
# .contiguous() calls — each one is a full memcpy on GPU.
# For PyTorch decode, "tile" = number of KV cache blocks processed
# per matmul call. More blocks per tile = fewer Python loop iterations
# = less launch overhead. Constraint: score tensor
# [kv_h, gqa, 1, tile_tokens] × 4 bytes must be reasonable.
#
# Current code does: index → permute → contiguous → view → slice →
# permute → contiguous → float
# That's 2 contiguous() calls per K and V = 4 GPU memcpy per sequence.
# For decode (q_len=1), score tensor is tiny:
# kv_h × gqa × 1 × tile_tokens × 4 = 1 × 6 × 1 × 4096 × 4 = 96 KB
# So we can use large tiles: process ALL blocks in one matmul when
# possible, falling back to tiling only for very long sequences.
#
# Optimization: reshape key_cache layout knowledge to reduce copies.
# key_cache shape: [num_blocks, kv_h, d//x, blk_sz, x]
# After index + reshape: [n_blk, blk_sz, kv_h, d] via one permute+reshape
# Then slice + transpose: [kv_h, d, seq_len]
# This is still 2 contiguous(), but the first reshape can be fused.
# CCCL subscription_factor = 5, sm_count = 16:
# max_concurrent_tiles ≈ 80
# But Python overhead dominates, so FEWER tiles is better.
# Strategy: tile_blocks = min(all_blocks, 1024) — process up to 1024
# cache blocks (= 16384 tokens at block_size=16) per matmul.
# ================================================================
_MAX_TILE_BLOCKS = 1024 # ~16K tokens per tile at block_size=16
try:
for i in range(num_seqs):
seq_len = int(seq_lens[i].item())
num_blocks = (seq_len + block_size - 1) // block_size
blk_ids = block_tables[i, :num_blocks]
if seq_len == 0:
output[i].zero_()
continue
# Gather K: single permute+contiguous → view → slice → transpose
# key_cache[blk_ids]: [n, kv_h, d//x, blk_sz, x]
k_gathered = key_cache[blk_ids]
k_t = (k_gathered
.permute(0, 3, 1, 2, 4) # [n, blk_sz, kv_h, d//x, x]
.contiguous()
.view(-1, num_kv_heads, head_dim))[:seq_len] \
.permute(1, 2, 0).contiguous().float() # [kv_h, d, seq_len]
del k_gathered
num_blocks_i = (seq_len + block_size - 1) // block_size
blk_ids = block_tables[i, :num_blocks_i]
# Gather V: same pattern
v_gathered = value_cache[blk_ids]
v_t = (v_gathered
.permute(0, 3, 1, 2) # [n, blk_sz, kv_h, d]
.contiguous()
.view(-1, num_kv_heads, head_dim))[:seq_len] \
.permute(1, 0, 2).contiguous().float() # [kv_h, seq_len, d]
del v_gathered
# Reshape Q for lazy GQA: [kv_h, gqa_ratio, 1, d]
# Q reshaped once: [kv_h, gqa, 1, d] fp32 — tiny for decode
q_grouped = (query[i].float()
.view(num_kv_heads, gqa_ratio, head_dim)
.unsqueeze(2))
.unsqueeze(2)
.mul_(scale))
# [kv_h, gqa_ratio, 1, seq_len]
attn_w = torch.matmul(
q_grouped * scale, # [kv_h, gqa, 1, d]
k_t.unsqueeze(1)) # [kv_h, 1, d, seq_len]
attn_w = torch.softmax(attn_w, dim=-1)
# Online softmax accumulators (CCCL summary_stats_data pattern)
# accumulator = {m (running max), l (running sum_exp), o (running output)}
m = torch.full((num_kv_heads, gqa_ratio, 1),
float('-inf'), dtype=torch.float32, device=dev)
l = torch.zeros_like(m)
o = torch.zeros((num_kv_heads, gqa_ratio, 1, head_dim),
dtype=torch.float32, device=dev)
# [kv_h, gqa_ratio, 1, d] → [num_heads, head_dim]
out_i = torch.matmul(attn_w, v_t.unsqueeze(1))
output[i] = out_i.view(num_heads, head_dim).to(orig_dtype)
# Tile over KV blocks — GridEvenShare RAKE pattern
# Each tile = consecutive sequence of cache blocks
for tile_start in range(0, num_blocks_i, _MAX_TILE_BLOCKS):
tile_end = min(tile_start + _MAX_TILE_BLOCKS, num_blocks_i)
tile_blk_ids = blk_ids[tile_start:tile_end]
# Valid tokens in this tile
tile_token_start = tile_start * block_size
tile_token_end = min(tile_end * block_size, seq_len)
valid_tokens = tile_token_end - tile_token_start
# --------------------------------------------------------
# KV gather — agent_reduce.cuh ConsumeFullTile pattern
#
# agent_reduce loads VectorT in striped access when possible.
# PyTorch equivalent: reshape the 5D cache layout to 3D in
# one permute+contiguous, avoiding the double-contiguous
# pattern of the old code.
#
# key_cache shape: [num_blocks, kv_h, d//x, blk_sz, x]
# Target: [kv_h, d, valid_tokens] for Q@K^T
#
# Optimized path: permute(1,2,4,0,3) → [kv_h, d//x, x, n_blk, blk_sz]
# → reshape to [kv_h, d, n_blk*blk_sz] → slice [:valid_tokens]
# This is ONE contiguous() call instead of TWO.
# --------------------------------------------------------
k_gathered = key_cache[tile_blk_ids] # [n, kv_h, d//x, blk_sz, x]
k_t = (k_gathered
.permute(1, 2, 4, 0, 3) # [kv_h, d//x, x, n, blk_sz]
.contiguous()
.view(num_kv_heads, head_dim, -1) # [kv_h, d, n*blk_sz]
[:, :, :valid_tokens]
.unsqueeze(1) # [kv_h, 1, d, valid]
.float())
del k_gathered
v_gathered = value_cache[tile_blk_ids] # [n, kv_h, d, blk_sz]
v_t = (v_gathered
.permute(1, 2, 0, 3) # [kv_h, d, n, blk_sz]
.contiguous()
.view(num_kv_heads, head_dim, -1) # [kv_h, d, n*blk_sz]
[:, :, :valid_tokens]
.transpose(1, 2) # [kv_h, valid, d]
.unsqueeze(1) # [kv_h, 1, valid, d]
.float())
del v_gathered
# --------------------------------------------------------
# Scores + online softmax — summary_statistics.cu pattern
#
# unary_op: score_tile → (max, sum_exp, weighted_V)
# binary_op: merge with correction factor
#
# CCCL summary_stats_binary_op merges:
# result.mean = x.mean + delta * y.n / n
# result.M2 = x.M2 + y.M2 + delta² * x.n * y.n / n
#
# Online softmax merge:
# m_new = max(m_old, m_tile)
# corr = exp(m_old - m_new) ← rescale factor
# l_new = l_old * corr + l_tile
# o_new = o_old * corr + tile_exp @ V
#
# Structurally identical: m↔max, l↔n, o↔mean×n.
# --------------------------------------------------------
# [kv_h, gqa, 1, valid_tokens]
s = torch.matmul(q_grouped, k_t)
del k_t
# Online softmax update (Flash Attention Algorithm 1)
m_tile = s.amax(dim=-1, keepdim=True) # [kv_h, gqa, 1, 1]
m_new = torch.maximum(m, m_tile.squeeze(-1))
corr = torch.exp(m - m_new) # rescale old accum
exp_s = torch.exp(s - m_new.unsqueeze(-1))
del s
m.copy_(m_new)
l.mul_(corr).add_(exp_s.sum(dim=-1))
o.mul_(corr.unsqueeze(-1)).add_(torch.matmul(exp_s, v_t))
del exp_s, v_t, corr, m_new, m_tile
# Finalize: normalize
o.div_(l.unsqueeze(-1))
output[i] = (o.view(num_heads, head_dim)
.to(orig_dtype))
except Exception as e:
print(f"[decode_pytorch ERROR] {type(e).__name__}: {e}",
@@ -260,9 +355,33 @@ class PagedAttention:
# to parallelize.
# TODO(woosuk): Tune this heuristic.
# For context len > 8192, use V2 kernel to avoid shared memory shortage.
use_v1 = (max_seq_len <= 8192
and (max_num_partitions == 1 or num_seqs * num_heads > 512))
use_v1 = True
# CCCL dispatch_reduce.cuh two-path dispatch architecture:
# single-tile: num_items ≤ threads × items → one CTA, zero temp buffer
# multi-tile: GridEvenShare partitions across sm_count × occupancy CTAs
#
# Paged attention equivalent:
# V1 = single-pass: one CTA iterates ALL KV blocks (like DeviceReduceSingleTileKernel)
# V2 = partitioned: KV blocks split into PARTITION_SIZE chunks across CTAs,
# then a second kernel merges partition results (like InvokePasses two-phase)
#
# V1 is optimal when seq_len fits in one CTA's tile (small context).
# V2 is optimal when seq_len >> PARTITION_SIZE (long context) — parallelism
# across partitions compensates for the merge overhead.
#
# CCCL's GridEvenShare formula:
# max_blocks = sm_occupancy × sm_count × subscription_factor
# BI-V100: ~1 × 16 × 5 = 80 max blocks
# V2 becomes worthwhile when max_num_partitions > 1 AND the partition
# parallelism exceeds the sequence×head parallelism.
#
# Original heuristic (before hardcode): V1 when max_seq_len ≤ 8192 OR
# when batch×heads already saturates the GPU (num_seqs*num_heads > 512).
# Restored with BI-V100 SM count awareness.
bi100_sm_count = 16
bi100_saturation = bi100_sm_count * 32 # ~512 concurrent warps
use_v1 = (max_num_partitions == 1
or max_seq_len <= 8192
or num_seqs * num_heads > bi100_saturation)
if use_v1:
# Run PagedAttention V1.
ops.paged_attention_v1(