"""High-performance paged-attention decode for H100 (SM90). Triton GQA-shared kernel + CUDA-graph capture: - One CTA per (batch, kv_head, split) loads each K/V page once and attends all query heads that share that KV head (GQA reuse). - bf16 tensor-core MMA for QK / PV; fp32 online softmax. - Split-KV when B*Hkv underfills the GPU; reduce with online-softmax merge. - Model.forward captures a CUDA graph after the first call so subsequent replays (benchmark iters with fixed buffers) pay zero launch overhead. """ from __future__ import annotations import math import torch import torch.nn as nn import triton import triton.language as tl OP_TYPE = "attention" SUPPORTED_PRECISIONS = ["bf16"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] BATCH = 8 NUM_HEADS = 32 NUM_KV_HEADS = 8 HEAD_DIM = 128 SEQ_LEN = 1024 PAGE_SIZE = 16 @triton.jit def _gqa_kernel( Q_ptr, KV_ptr, BT_ptr, SL_ptr, Out_ptr, PartialO_ptr, PartialM_ptr, PartialL_ptr, stride_qb, stride_qh, stride_qd, stride_kv_block, stride_kv_tok, stride_kv_h, stride_kv_d, stride_bt_b, stride_bt_p, stride_ob, stride_oh, stride_od, stride_po_s, stride_po_b, stride_po_h, stride_po_d, stride_pm_s, stride_pm_b, stride_pm_h, scale, num_heads, num_kv_heads, PAGE_SIZE: tl.constexpr, HEAD_DIM: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, NUM_SPLITS: tl.constexpr, HAS_SPLITS: tl.constexpr, ): batch_id = tl.program_id(0) kv_head_id = tl.program_id(1) split_id = tl.program_id(2) group_size = num_heads // num_kv_heads q_head0 = kv_head_id * group_size seq_len = tl.load(SL_ptr + batch_id) tps = (seq_len + NUM_SPLITS - 1) // NUM_SPLITS start_tok = split_id * tps end_tok = tl.minimum(start_tok + tps, seq_len) offs_g = tl.arange(0, BLOCK_M) offs_d = tl.arange(0, HEAD_DIM) offs_n = tl.arange(0, BLOCK_N) mask_g = offs_g < group_size h_ids = q_head0 + offs_g if start_tok >= end_tok: if HAS_SPLITS: po = ( PartialO_ptr + split_id * stride_po_s + batch_id * stride_po_b + h_ids[:, None] * stride_po_h + offs_d[None, :] * stride_po_d ) tl.store(po, tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32), mask=mask_g[:, None]) pm = PartialM_ptr + split_id * stride_pm_s + batch_id * stride_pm_b + h_ids * stride_pm_h tl.store(pm, tl.full([BLOCK_M], float("-inf"), dtype=tl.float32), mask=mask_g) pl = PartialL_ptr + split_id * stride_pm_s + batch_id * stride_pm_b + h_ids * stride_pm_h tl.store(pl, tl.zeros([BLOCK_M], dtype=tl.float32), mask=mask_g) return q = tl.load( Q_ptr + batch_id * stride_qb + h_ids[:, None] * stride_qh + offs_d[None, :] * stride_qd, mask=mask_g[:, None], other=0.0, ) q_scaled = (q.to(tl.float32) * scale).to(tl.bfloat16) m_i = tl.full([BLOCK_M], float("-inf"), dtype=tl.float32) l_i = tl.zeros([BLOCK_M], dtype=tl.float32) acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32) n_tiles = (end_tok - start_tok + BLOCK_N - 1) // BLOCK_N for tile_id in tl.range(0, n_tiles): tok = start_tok + tile_id * BLOCK_N chunk = tl.minimum(BLOCK_N, end_tok - tok) abs_tok = tok + offs_n mask_n = offs_n < chunk page_idx = abs_tok // PAGE_SIZE page_off = abs_tok % PAGE_SIZE phys = tl.load( BT_ptr + batch_id * stride_bt_b + page_idx * stride_bt_p, mask=mask_n, other=0, ).to(tl.int64) tok_base = ( phys * stride_kv_block + page_off.to(tl.int64) * stride_kv_tok + kv_head_id * stride_kv_h ) k = tl.load( KV_ptr + tok_base[:, None] + offs_d[None, :] * stride_kv_d, mask=mask_n[:, None], other=0.0, ) v = tl.load( KV_ptr + tok_base[:, None] + (HEAD_DIM + offs_d[None, :]) * stride_kv_d, mask=mask_n[:, None], other=0.0, ) qk = tl.dot(q_scaled, tl.trans(k)) qk = tl.where(mask_n[None, :], qk, float("-inf")) m_ij = tl.max(qk, axis=1) m_new = tl.maximum(m_i, m_ij) alpha = tl.where(m_i == float("-inf"), 0.0, tl.exp(m_i - m_new)) p = tl.exp(qk - m_new[:, None]) p = tl.where(mask_n[None, :], p, 0.0) l_ij = tl.sum(p, axis=1) acc = acc * alpha[:, None] acc = tl.dot(p.to(tl.bfloat16), v, acc) l_i = l_i * alpha + l_ij m_i = m_new if HAS_SPLITS: po = ( PartialO_ptr + split_id * stride_po_s + batch_id * stride_po_b + h_ids[:, None] * stride_po_h + offs_d[None, :] * stride_po_d ) tl.store(po, acc, mask=mask_g[:, None]) pm = PartialM_ptr + split_id * stride_pm_s + batch_id * stride_pm_b + h_ids * stride_pm_h tl.store(pm, m_i, mask=mask_g) pl = PartialL_ptr + split_id * stride_pm_s + batch_id * stride_pm_b + h_ids * stride_pm_h tl.store(pl, l_i, mask=mask_g) else: l_safe = tl.where(l_i > 0.0, l_i, 1.0) tl.store( Out_ptr + batch_id * stride_ob + h_ids[:, None] * stride_oh + offs_d[None, :] * stride_od, (acc / l_safe[:, None]).to(tl.bfloat16), mask=mask_g[:, None], ) @triton.jit def _reduce_kernel( PartialO_ptr, PartialM_ptr, PartialL_ptr, Out_ptr, stride_po_s, stride_po_b, stride_po_h, stride_po_d, stride_pm_s, stride_pm_b, stride_pm_h, stride_ob, stride_oh, stride_od, NUM_SPLITS: tl.constexpr, HEAD_DIM: tl.constexpr, ): batch_id = tl.program_id(0) head_id = tl.program_id(1) offs_d = tl.arange(0, HEAD_DIM) m = tl.full([], float("-inf"), dtype=tl.float32) l = tl.full([], 0.0, dtype=tl.float32) acc = tl.zeros([HEAD_DIM], dtype=tl.float32) for s in tl.static_range(NUM_SPLITS): m_s = tl.load( PartialM_ptr + s * stride_pm_s + batch_id * stride_pm_b + head_id * stride_pm_h ) l_s = tl.load( PartialL_ptr + s * stride_pm_s + batch_id * stride_pm_b + head_id * stride_pm_h ) o_s = tl.load( PartialO_ptr + s * stride_po_s + batch_id * stride_po_b + head_id * stride_po_h + offs_d * stride_po_d ) m_new = tl.maximum(m, m_s) alpha = tl.where(m == float("-inf"), 0.0, tl.exp(m - m_new)) alpha_s = tl.where(m_s == float("-inf"), 0.0, tl.exp(m_s - m_new)) acc = acc * alpha + o_s * alpha_s l = l * alpha + l_s * alpha_s m = m_new tl.store( Out_ptr + batch_id * stride_ob + head_id * stride_oh + offs_d * stride_od, (acc / tl.where(l > 0.0, l, 1.0)).to(tl.bfloat16), ) class _Workspace: def __init__(self): self.key = None self.po = self.pm = self.pl = None def get(self, ns, B, H, D, device): key = (ns, B, H, D, str(device)) if key != self.key: self.po = torch.empty(ns, B, H, D, dtype=torch.float32, device=device) self.pm = torch.empty(ns, B, H, dtype=torch.float32, device=device) self.pl = torch.empty(ns, B, H, dtype=torch.float32, device=device) self.key = key return self.po, self.pm, self.pl _WS = _Workspace() def _pick_splits(batch: int, hkv: int, seq_len: int) -> int: """Split-KV heuristic for H100 (CUDA-graph friendly).""" base = max(1, batch * hkv) if base >= 192: return 1 target = 384 max_splits = max(1, seq_len // 64) desired = (target + base - 1) // base return min(max(1, desired), max_splits, 16) def paged_attention_decode( query: torch.Tensor, kv_cache: torch.Tensor, block_table: torch.Tensor, seq_lens: torch.Tensor, num_kv_heads: int, page_size: int, scale: float, out: torch.Tensor | None = None, max_seq: int | None = None, ) -> torch.Tensor: B, H, D = query.shape Hkv = num_kv_heads group_size = H // Hkv if out is None: out = torch.empty(B, H, D, dtype=query.dtype, device=query.device) if max_seq is None: max_seq = int(seq_lens.max().item()) if B > 0 else 0 num_splits = _pick_splits(B, Hkv, max_seq) BLOCK_M = max(16, triton.next_power_of_2(group_size)) n_ctas = B * Hkv * num_splits if D <= 64: BLOCK_N, num_warps, num_stages = 128, 4, 3 elif n_ctas >= 200 or max_seq >= 2048: BLOCK_N, num_warps, num_stages = 128, 4, 4 else: BLOCK_N, num_warps, num_stages = 64, 4, 4 kwargs = dict( stride_qb=query.stride(0), stride_qh=query.stride(1), stride_qd=query.stride(2), stride_kv_block=kv_cache.stride(0), stride_kv_tok=kv_cache.stride(1), stride_kv_h=kv_cache.stride(2), stride_kv_d=kv_cache.stride(3), stride_bt_b=block_table.stride(0), stride_bt_p=block_table.stride(1), stride_ob=out.stride(0), stride_oh=out.stride(1), stride_od=out.stride(2), scale=scale, num_heads=H, num_kv_heads=Hkv, PAGE_SIZE=page_size, HEAD_DIM=D, BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, NUM_SPLITS=num_splits, num_warps=num_warps, num_stages=num_stages, ) if num_splits == 1: _gqa_kernel[(B, Hkv, 1)]( query, kv_cache, block_table, seq_lens, out, None, None, None, stride_po_s=0, stride_po_b=0, stride_po_h=0, stride_po_d=0, stride_pm_s=0, stride_pm_b=0, stride_pm_h=0, HAS_SPLITS=False, **kwargs, ) else: po, pm, pl = _WS.get(num_splits, B, H, D, query.device) _gqa_kernel[(B, Hkv, num_splits)]( query, kv_cache, block_table, seq_lens, out, po, pm, pl, stride_po_s=po.stride(0), stride_po_b=po.stride(1), stride_po_h=po.stride(2), stride_po_d=po.stride(3), stride_pm_s=pm.stride(0), stride_pm_b=pm.stride(1), stride_pm_h=pm.stride(2), HAS_SPLITS=True, **kwargs, ) _reduce_kernel[(B, H)]( po, pm, pl, out, po.stride(0), po.stride(1), po.stride(2), po.stride(3), pm.stride(0), pm.stride(1), pm.stride(2), out.stride(0), out.stride(1), out.stride(2), NUM_SPLITS=num_splits, HEAD_DIM=D, num_warps=4, ) return out class Model(nn.Module): def __init__(self, batch, num_heads, num_kv_heads, head_dim, seq_len, page_size): super().__init__() assert num_heads % num_kv_heads == 0 self.batch = batch self.num_heads = num_heads self.num_kv_heads = num_kv_heads self.head_dim = head_dim self.seq_len = seq_len self.page_size = page_size self.group_size = num_heads // num_kv_heads self.scale = 1.0 / math.sqrt(head_dim) self.register_buffer("_dummy", torch.zeros(1, dtype=torch.bfloat16), persistent=False) self._graph = None self._graph_keys = None self._static_out = None def forward(self, query, kv_cache, block_table, seq_lens): keys = ( query.data_ptr(), kv_cache.data_ptr(), block_table.data_ptr(), seq_lens.data_ptr(), tuple(query.shape), tuple(kv_cache.shape), ) if ( self._graph is not None and self._graph_keys == keys and self._static_out is not None ): self._graph.replay() return self._static_out kw = dict( num_kv_heads=self.num_kv_heads, page_size=self.page_size, scale=self.scale, max_seq=self.seq_len, ) out = paged_attention_decode(query, kv_cache, block_table, seq_lens, **kw) if not query.is_cuda: return out try: if self._static_out is None or self._static_out.shape != out.shape: self._static_out = torch.empty_like(out) s = torch.cuda.Stream() s.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(s): paged_attention_decode( query, kv_cache, block_table, seq_lens, out=self._static_out, **kw, ) torch.cuda.current_stream().wait_stream(s) g = torch.cuda.CUDAGraph() with torch.cuda.graph(g): paged_attention_decode( query, kv_cache, block_table, seq_lens, out=self._static_out, **kw, ) self._graph = g self._graph_keys = keys g.replay() return self._static_out except Exception: self._graph = None self._graph_keys = None return out def get_inputs(): B, H, Hkv, D, L, P = BATCH, NUM_HEADS, NUM_KV_HEADS, HEAD_DIM, SEQ_LEN, PAGE_SIZE pages_per_seq = (L + P - 1) // P total_pages = max(B * pages_per_seq + 8, 64) query = torch.randn(B, H, D, dtype=torch.bfloat16) * 0.1 kv_cache = torch.randn(total_pages, P, Hkv, 2 * D, dtype=torch.bfloat16) * 0.1 perm = torch.randperm(total_pages)[: B * pages_per_seq].reshape(B, pages_per_seq).int() block_table = perm.contiguous() seq_lens = torch.full((B,), L, dtype=torch.int32) return [query, kv_cache, block_table, seq_lens] def get_init_inputs(): return [BATCH, NUM_HEADS, NUM_KV_HEADS, HEAD_DIM, SEQ_LEN, PAGE_SIZE]