"""Paged attention decode for SM120 (RTX PRO 6000 Blackwell), single-query decode. Two-level custom implementation (no library attention dispatch): 1. Primary: a hand-written CUDA C++ kernel (built via load_inline). - Grid (batch * num_kv_heads * splits). Each CTA processes one KV head of one batch element over a token partition, computing the whole GQA group of G query heads together so every KV byte is loaded exactly once. - SIMT math: each warp strides over tokens; per token it loads the K row (8B/lane), dot-reduces q.k with shuffles, runs online softmax in fp32, and accumulates p*V fragments in registers; a smem combine merges warp partials. L2 cache-hint loads (createpolicy + ld.global.L2::cache_hint) pick evict-first streaming or evict-last reuse per shape. - The cross-split reduction is fused into the same launch with a last-CTA-done semaphore (threadfence + atomic counter), so the whole forward is one kernel launch. A CUDA graph replay removes launch gaps. 2. Fallback: a Triton split-KV kernel (used only if the extension build fails). KV cache layout: (num_blocks, page_size, num_kv_heads, head_dim * 2), last dim packs [K | V]. block_table[b] lists the physical pages of batch element b. """ import math import torch import torch.nn as nn OP_TYPE = "attention" SUPPORTED_PRECISIONS = ["bf16"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] # --- Shape knobs (overridden by check.py / benchmark.py from shapes.py) ---- BATCH = 8 NUM_HEADS = 32 NUM_KV_HEADS = 8 HEAD_DIM = 128 SEQ_LEN = 1024 PAGE_SIZE = 16 _NUM_SMS = 188 # RTX PRO 6000 Blackwell # --------------------------------------------------------------------------- # CUDA C++ kernel source # --------------------------------------------------------------------------- _CUDA_SRC = r""" #include #include #include #include #define DEV_INLINE __device__ __forceinline__ #define MAX_SPLITS 64 DEV_INLINE unsigned long long make_policy(int which, float frac) { unsigned long long pol = 0; if (which == 2) { asm volatile("createpolicy.fractional.L2::evict_last.b64 %0, %1;" : "=l"(pol) : "f"(frac)); } else if (which == 1) { asm volatile("createpolicy.fractional.L2::evict_first.b64 %0, %1;" : "=l"(pol) : "f"(frac)); } return pol; } template DEV_INLINE uint2 ldg8(const void* p, unsigned long long pol) { uint2 r; if constexpr (POLICY == 0) { r = __ldg((const uint2*)p); } else { asm volatile("ld.global.L2::cache_hint.v2.b32 {%0,%1}, [%2], %3;" : "=r"(r.x), "=r"(r.y) : "l"(p), "l"(pol)); } return r; } template DEV_INLINE unsigned ldg4(const void* p, unsigned long long pol) { unsigned r; if constexpr (POLICY == 0) { r = __ldg((const unsigned*)p); } else { asm volatile("ld.global.L2::cache_hint.b32 %0, [%1], %2;" : "=r"(r) : "l"(p), "l"(pol)); } return r; } template DEV_INLINE uint4 ldg16(const void* p, unsigned long long pol) { uint4 r; if constexpr (POLICY == 0) { r = __ldg((const uint4*)p); } else { asm volatile("ld.global.L2::cache_hint.v4.b32 {%0,%1,%2,%3}, [%4], %5;" : "=r"(r.x), "=r"(r.y), "=r"(r.z), "=r"(r.w) : "l"(p), "l"(pol)); } return r; } DEV_INLINE void unpack_bf16x2(unsigned u, float& a, float& b) { __nv_bfloat162 h = *reinterpret_cast<__nv_bfloat162*>(&u); float2 f = __bfloat1622float2(h); a = f.x; b = f.y; } // One CTA handles (batch b, kv-head kvh, split). The G query heads of the GQA // group are computed together so each KV byte is loaded once. Warp w strides // over the tokens of the split; a smem combine merges warp partials; if // SPLITS > 1 the last CTA per (b,kvh) merges the split partials (semaphore). template __global__ void __launch_bounds__(WARPS * 32) paged_decode_kernel( const __nv_bfloat16* __restrict__ q, // (B, H, D) const __nv_bfloat16* __restrict__ kv, // (pages, P, Hkv, 2D) [K|V] const int* __restrict__ block_table, // (B, maxbt) const int* __restrict__ seq_lens, // (B,) float* __restrict__ partial, // (B, Hkv, S, G, D) float* __restrict__ ml, // (B, Hkv, S, G, 2): m, l int* __restrict__ counters, // (B, Hkv) __nv_bfloat16* __restrict__ out, // (B, H, D) const int H, const int Hkv, const int S, const int maxbt, const float scale, const long stride_page, const float frac) { constexpr int E = D / 32; // elems per lane (4 for D=128, 2 for D=64) constexpr int T = WARPS * 32; const int tid = threadIdx.x; const int warp = tid >> 5, lane = tid & 31; const int pid = blockIdx.x; const int kvh = pid % Hkv; const int tmp = pid / Hkv; const int split = tmp % S; const int b = tmp / S; const int L = __ldg(seq_lens + b); const int tps = (L + S - 1) / S; const int start = split * tps; const int end = min(start + tps, L); unsigned long long pol = make_policy(POLICY, frac); // Load the GQA group's queries into smem as fp32. __shared__ float qf_s[G][D]; { const __nv_bfloat16* qg = q + (long)b * H * D + (long)kvh * G * D; for (int i = tid; i < G * D / 8; i += T) { uint4 r = ldg16(qg + i * 8, pol); const __nv_bfloat16* h = reinterpret_cast(&r); #pragma unroll for (int j = 0; j < 8; ++j) qf_s[(i * 8 + j) / D][(i * 8 + j) % D] = __bfloat162float(h[j]); } } __syncthreads(); // qf_s written by a subset of threads float m[G], l[G], acc[G][E]; #pragma unroll for (int g = 0; g < G; ++g) { m[g] = -INFINITY; l[g] = 0.f; #pragma unroll for (int i = 0; i < E; ++i) acc[g][i] = 0.f; } const long kvs_stride = (long)Hkv * 2 * D; const int kvh_off = kvh * 2 * D; const int* bt_row = block_table + (long)b * maxbt; // Preload the page indices this split touches into smem: removes the // dependent global load from the per-token address chain. constexpr int MAX_SPAN = 544; __shared__ int pages_s[MAX_SPAN]; const int p0 = start / P; const int p1 = (end - 1) / P; const int npages = (end > start) ? (p1 - p0 + 1) : 0; if (npages <= MAX_SPAN) { for (int i = tid; i < npages; i += T) pages_s[i] = __ldg(bt_row + p0 + i); } __syncthreads(); for (int n = start + warp; n < end; n += WARPS) { const int slot = n % P; int page; if (npages <= MAX_SPAN) page = pages_s[n / P - p0]; else page = __ldg(bt_row + n / P); const __nv_bfloat16* base = kv + (long)page * stride_page + (long)slot * kvs_stride + kvh_off + lane * E; float kf[E], vf[E]; if constexpr (E == 4) { uint2 kr = ldg8(base, pol); uint2 vr = ldg8(base + D, pol); unpack_bf16x2(kr.x, kf[0], kf[1]); unpack_bf16x2(kr.y, kf[2], kf[3]); unpack_bf16x2(vr.x, vf[0], vf[1]); unpack_bf16x2(vr.y, vf[2], vf[3]); } else { unsigned kr = ldg4(base, pol); unsigned vr = ldg4(base + D, pol); unpack_bf16x2(kr, kf[0], kf[1]); unpack_bf16x2(vr, vf[0], vf[1]); } float sg[G]; #pragma unroll for (int g = 0; g < G; ++g) { float s = 0.f; #pragma unroll for (int i = 0; i < E; ++i) s += qf_s[g][lane * E + i] * kf[i]; sg[g] = s; } #pragma unroll for (int off = 16; off > 0; off >>= 1) { #pragma unroll for (int g = 0; g < G; ++g) sg[g] += __shfl_xor_sync(0xffffffffu, sg[g], off); } #pragma unroll for (int g = 0; g < G; ++g) { const float sc = sg[g] * scale; const float m_new = fmaxf(m[g], sc); const float alpha = __expf(m[g] - m_new); const float p = __expf(sc - m_new); l[g] = l[g] * alpha + p; #pragma unroll for (int i = 0; i < E; ++i) acc[g][i] = acc[g][i] * alpha + p * vf[i]; m[g] = m_new; } } // ---- cross-warp combine ---- __shared__ float wacc[WARPS][G][D]; __shared__ float wml[WARPS][G][2]; #pragma unroll for (int g = 0; g < G; ++g) { #pragma unroll for (int i = 0; i < E; ++i) wacc[warp][g][lane * E + i] = acc[g][i]; } if (lane == 0) { #pragma unroll for (int g = 0; g < G; ++g) { wml[warp][g][0] = m[g]; wml[warp][g][1] = l[g]; } } __syncthreads(); __shared__ float wgt_s[WARPS][G]; __shared__ float Mglob[G]; __shared__ float denom[G]; if (tid < G) { const int g = tid; float M = -INFINITY; #pragma unroll for (int w = 0; w < WARPS; ++w) M = fmaxf(M, wml[w][g][0]); float Lsum = 0.f; #pragma unroll for (int w = 0; w < WARPS; ++w) { const float ww = __expf(wml[w][g][0] - M); wgt_s[w][g] = ww; Lsum += wml[w][g][1] * ww; } Mglob[g] = M; denom[g] = Lsum; } __syncthreads(); const long grp = ((long)b * Hkv + kvh) * S + split; if constexpr (DIRECT) { // Single split: write final bf16 output. for (int idx = tid; idx < G * D; idx += T) { const int g = idx / D, d = idx % D; float val = 0.f; #pragma unroll for (int w = 0; w < WARPS; ++w) val += wacc[w][g][d] * wgt_s[w][g]; const float Lsum = denom[g]; const float o = (Lsum > 0.f) ? val / Lsum : 0.f; out[(long)b * H * D + (long)(kvh * G + g) * D + d] = __float2bfloat16(o); } return; } // Store unnormalized combined partial + (M, l) for the cross-split reduce. for (int idx = tid; idx < G * D; idx += T) { const int g = idx / D, d = idx % D; float val = 0.f; #pragma unroll for (int w = 0; w < WARPS; ++w) val += wacc[w][g][d] * wgt_s[w][g]; partial[grp * G * D + idx] = val; } if (tid < G) { ml[grp * G * 2 + tid * 2 + 0] = Mglob[tid]; ml[grp * G * 2 + tid * 2 + 1] = denom[tid]; } __syncthreads(); __threadfence(); __shared__ int is_last_s; if (tid == 0) { const unsigned old = atomicAdd((unsigned*)&counters[b * Hkv + kvh], 1u); is_last_s = (old == (unsigned)S - 1); } __syncthreads(); if (!is_last_s) return; __threadfence(); // ---- last CTA: merge split partials ---- __shared__ float sml[MAX_SPLITS][G][2]; const long grp0 = ((long)b * Hkv + kvh) * S; for (int i = tid; i < S * G * 2; i += T) { ((float*)sml)[i] = ml[grp0 * G * 2 + i]; } __syncthreads(); for (int idx = tid; idx < G * D; idx += T) { const int g = idx / D, d = idx % D; float M = -INFINITY; for (int s = 0; s < S; ++s) M = fmaxf(M, sml[s][g][0]); float Lsum = 0.f, val = 0.f; for (int s = 0; s < S; ++s) { const float w = __expf(sml[s][g][0] - M); Lsum += sml[s][g][1] * w; val += partial[(grp0 + s) * G * D + idx] * w; } const float o = (Lsum > 0.f) ? val / Lsum : 0.f; out[(long)b * H * D + (long)(kvh * G + g) * D + d] = __float2bfloat16(o); } if (tid == 0) counters[b * Hkv + kvh] = 0; // reset for the next launch } // ---------------- host dispatch ---------------- template static void dispatch_warps_policy( int warps, int policy, bool direct, dim3 grid, cudaStream_t stream, const __nv_bfloat16* q, const __nv_bfloat16* kv, const int* bt, const int* sl, float* partial, float* ml, int* counters, __nv_bfloat16* out, int H, int Hkv, int S, int maxbt, float scale, long stride_page, float frac) { constexpr int Pg = 16; #define LAUNCH(W, POL, DIR) \ paged_decode_kernel<<>>( \ q, kv, bt, sl, partial, ml, counters, out, H, Hkv, S, maxbt, scale, stride_page, frac) #define CASE_POL(W, DIR) \ if (policy == 2) { LAUNCH(W, 2, DIR); } \ else if (policy == 1) { LAUNCH(W, 1, DIR); } \ else { LAUNCH(W, 0, DIR); } if (direct) { if (warps == 16) { CASE_POL(16, true); } else if (warps == 8) { CASE_POL(8, true); } else { CASE_POL(4, true); } } else { if (warps == 16) { CASE_POL(16, false); } else if (warps == 8) { CASE_POL(8, false); } else { CASE_POL(4, false); } } #undef CASE_POL #undef LAUNCH } void paged_decode_forward( torch::Tensor q, torch::Tensor kv, torch::Tensor block_table, torch::Tensor seq_lens, torch::Tensor partial, torch::Tensor ml, torch::Tensor counters, torch::Tensor out, int64_t S, int64_t warps, int64_t policy, double frac) { const int B = q.size(0); const int H = q.size(1); const int D = q.size(2); const int Hkv = kv.size(2); const int Pg = kv.size(1); const int G = H / Hkv; const int maxbt = block_table.size(1); const float scale = 1.0f / sqrtf((float)D); const long stride_page = (long)Pg * Hkv * 2 * D; dim3 grid(B * Hkv * (int)S); bool direct = (S == 1); cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream(); auto q_p = reinterpret_cast(q.data_ptr()); auto kv_p = reinterpret_cast(kv.data_ptr()); auto bt_p = block_table.data_ptr(); auto sl_p = seq_lens.data_ptr(); auto partial_p = partial.data_ptr(); auto ml_p = ml.data_ptr(); auto counters_p = counters.data_ptr(); auto out_p = reinterpret_cast<__nv_bfloat16*>(out.data_ptr()); TORCH_CHECK(Pg == 16, "page size 16 only"); if (D == 128) { if (G == 8) dispatch_warps_policy<128, 8>((int)warps, (int)policy, direct, grid, stream, q_p, kv_p, bt_p, sl_p, partial_p, ml_p, counters_p, out_p, H, Hkv, (int)S, maxbt, scale, stride_page, (float)frac); else dispatch_warps_policy<128, 4>((int)warps, (int)policy, direct, grid, stream, q_p, kv_p, bt_p, sl_p, partial_p, ml_p, counters_p, out_p, H, Hkv, (int)S, maxbt, scale, stride_page, (float)frac); } else { if (G == 8) dispatch_warps_policy<64, 8>((int)warps, (int)policy, direct, grid, stream, q_p, kv_p, bt_p, sl_p, partial_p, ml_p, counters_p, out_p, H, Hkv, (int)S, maxbt, scale, stride_page, (float)frac); else dispatch_warps_policy<64, 4>((int)warps, (int)policy, direct, grid, stream, q_p, kv_p, bt_p, sl_p, partial_p, ml_p, counters_p, out_p, H, Hkv, (int)S, maxbt, scale, stride_page, (float)frac); } } """ _CPP_DECL = ("void paged_decode_forward(torch::Tensor q, torch::Tensor kv, " "torch::Tensor block_table, torch::Tensor seq_lens, torch::Tensor partial, " "torch::Tensor ml, torch::Tensor counters, torch::Tensor out, " "int64_t S, int64_t warps, int64_t policy, double frac);") _cuda_ext = None _cuda_build_failed = False def _cleanup_stale_build_lock(name: str, max_age_s: float = 120.0) -> None: """Remove a stale torch FileBaton lock left behind by a killed build. torch's cpp_extension waits forever on a baton 'lock' file; if the process that created it died (or its build dir vanished mid-build), every later import deadlocks. If the lock is older than max_age_s and no compiler is running for it, drop it. """ import os, time try: from torch.utils.cpp_extension import _get_build_directory lock = os.path.join(_get_build_directory(name, verbose=False), "lock") if os.path.exists(lock) and time.time() - os.path.getmtime(lock) > max_age_s: try: os.unlink(lock) except OSError: pass except Exception: pass _SO_BASENAME = "paged_decode_sm120_v2" def _local_so_dir(): import os return os.path.dirname(os.path.abspath(__file__)) def _get_cuda_ext(): global _cuda_ext, _cuda_build_failed if _cuda_ext is None and not _cuda_build_failed: # Fast path: a previously built .so kept next to this file (immune to # external cache wipes of the torch_extensions directory). try: import os from torch.utils.cpp_extension import _import_module_from_library d = _local_so_dir() if os.path.exists(os.path.join(d, _SO_BASENAME + ".so")): _cuda_ext = _import_module_from_library(_SO_BASENAME, d, True) return _cuda_ext except Exception: _cuda_ext = None try: import os, shutil from torch.utils.cpp_extension import load_inline, _get_build_directory _cleanup_stale_build_lock(_SO_BASENAME) _cuda_ext = load_inline( name=_SO_BASENAME, cpp_sources=_CPP_DECL, cuda_sources=_CUDA_SRC, functions=["paged_decode_forward"], extra_cuda_cflags=["-O3", "--use_fast_math", "--threads=8", "-gencode=arch=compute_120a,code=sm_120a"], verbose=False, ) try: bdir = _get_build_directory(_SO_BASENAME, verbose=False) so = os.path.join(bdir, _SO_BASENAME + ".so") if os.path.exists(so): shutil.copy(so, os.path.join(_local_so_dir(), _SO_BASENAME + ".so")) except Exception: pass except Exception: _cuda_build_failed = True return _cuda_ext # --------------------------------------------------------------------------- # Triton fallback kernels # --------------------------------------------------------------------------- import triton # noqa: E402 import triton.language as tl # noqa: E402 @triton.jit def _paged_decode_main( q_ptr, kv_ptr, bt_ptr, sl_ptr, part_ptr, lse_ptr, H: tl.constexpr, HKV: tl.constexpr, G: tl.constexpr, D: tl.constexpr, P: tl.constexpr, MAXBT, SCALE, BLOCK_N: tl.constexpr, SPLITS: tl.constexpr, EV: tl.constexpr, ): pid = tl.program_id(0) nsplits: tl.constexpr = SPLITS kvh = pid % HKV tmp = pid // HKV split = tmp % nsplits b = tmp // nsplits L = tl.load(sl_ptr + b) tokens_per_split = (L + nsplits - 1) // nsplits start = split * tokens_per_split end = tl.minimum(start + tokens_per_split, L) offs_m = tl.arange(0, 16) offs_d = tl.arange(0, D) m_valid = offs_m < G q = tl.load( q_ptr + b * H * D + (kvh * G + offs_m)[:, None] * D + offs_d[None, :], mask=m_valid[:, None], other=0.0, ) m_i = tl.full((16,), float("-inf"), dtype=tl.float32) l_i = tl.zeros((16,), dtype=tl.float32) acc = tl.zeros((16, D), dtype=tl.float32) bt_row = bt_ptr + b * MAXBT kv_base = kv_ptr + kvh * 2 * D for n0 in range(start, end, BLOCK_N): offs_n = n0 + tl.arange(0, BLOCK_N) n_valid = offs_n < end pages = tl.load(bt_row + offs_n // P, mask=n_valid, other=0) slot = offs_n % P kv_off = (pages[:, None] * (P * HKV * 2 * D) + slot[:, None] * (HKV * 2 * D)) if EV == 1: k = tl.load(kv_base + kv_off + offs_d[None, :], mask=n_valid[:, None], other=0.0, eviction_policy="evict_first") v = tl.load(kv_base + D + kv_off + offs_d[None, :], mask=n_valid[:, None], other=0.0, eviction_policy="evict_first") elif EV == 2: k = tl.load(kv_base + kv_off + offs_d[None, :], mask=n_valid[:, None], other=0.0, eviction_policy="evict_last") v = tl.load(kv_base + D + kv_off + offs_d[None, :], mask=n_valid[:, None], other=0.0, eviction_policy="evict_last") else: k = tl.load(kv_base + kv_off + offs_d[None, :], mask=n_valid[:, None], other=0.0) v = tl.load(kv_base + D + kv_off + offs_d[None, :], mask=n_valid[:, None], other=0.0) s = tl.dot(q, tl.trans(k)) * SCALE s = tl.where(n_valid[None, :], s, float("-inf")) m_new = tl.maximum(m_i, tl.max(s, 1)) alpha = tl.exp(m_i - m_new) p = tl.exp(s - m_new[:, None]) l_i = l_i * alpha + tl.sum(p, 1) acc = acc * alpha[:, None] + tl.dot(p.to(k.dtype), v) m_i = m_new l_safe = tl.where(l_i == 0.0, 1.0, l_i) out = acc / l_safe[:, None] part_base = part_ptr + ((b * HKV + kvh) * nsplits + split) * 16 * D tl.store(part_base + offs_m[:, None] * D + offs_d[None, :], out, mask=m_valid[:, None]) lse = m_i + tl.log(l_safe) lse = tl.where(l_i == 0.0, float("-inf"), lse) tl.store(lse_ptr + ((b * HKV + kvh) * nsplits + split) * 16 + offs_m, lse, mask=m_valid) @triton.jit def _paged_decode_reduce( part_ptr, lse_ptr, o_ptr, H: tl.constexpr, HKV: tl.constexpr, G: tl.constexpr, D: tl.constexpr, SPLITS: tl.constexpr, BLOCK_S: tl.constexpr, ): pid = tl.program_id(0) m = pid % 16 tmp = pid // 16 kvh = tmp % HKV b = tmp // HKV if m >= G: return offs_d = tl.arange(0, D) offs_s = tl.arange(0, BLOCK_S) s_valid = offs_s < SPLITS grp = (b * HKV + kvh) * SPLITS lse = tl.load(lse_ptr + grp * 16 + offs_s * 16 + m, mask=s_valid, other=float("-inf")) m_g = tl.max(lse, 0) w = tl.exp(lse - m_g) w = tl.where(s_valid, w, 0.0) w_sum = tl.sum(w, 0) part = tl.load( part_ptr + grp * 16 * D + offs_s[:, None] * (16 * D) + m * D + offs_d[None, :], mask=s_valid[:, None], other=0.0, ) acc = tl.sum(part * w[:, None], 0) / tl.maximum(w_sum, 1e-30) tl.store(o_ptr + b * H * D + (kvh * G + m) * D + offs_d, acc.to(o_ptr.dtype.element_ty)) # --------------------------------------------------------------------------- # Module # --------------------------------------------------------------------------- # Hand-tuned per-shape configs: (splits, warps, policy). policy: 0 default, # 1 evict_first (large streaming caches), 2 evict_last (small caches kept # resident in L2 across decode steps). _CONFIG_TABLE = { # (B, H, Hkv, D, L, P): (splits, warps, policy) (8, 32, 8, 128, 1024, 16): (11, 8, 2), (32, 32, 8, 128, 2048, 16): (4, 4, 1), (4, 64, 8, 128, 4096, 16): (16, 16, 2), (16, 32, 8, 128, 1535, 16): (16, 16, 2), (8, 16, 4, 64, 2000, 16): (16, 8, 2), } def _default_config(batch, num_heads, num_kv_heads, head_dim, seq_len, page_size): """Pick (splits, warps, policy) for a shape.""" hit = _CONFIG_TABLE.get((batch, num_heads, num_kv_heads, head_dim, seq_len, page_size)) if hit is not None: return hit kv_bytes = batch * seq_len * num_kv_heads * head_dim * 2 * 2 base = batch * num_kv_heads # enough CTAs to fill the machine a couple times over splits = max(1, min((seq_len + 63) // 64, (2 * _NUM_SMS) // base)) if splits > 1 and base * splits < _NUM_SMS: splits = max(1, min((seq_len + 63) // 64, _NUM_SMS // base)) warps = 8 # evict_last keeps hot KV resident in L2 across decode steps (small caches); # evict_first streams large caches without thrashing. policy = 2 if kv_bytes <= 24 * 1024 * 1024 else 1 return splits, warps, policy class Model(nn.Module): """Single-query paged attention decode. Forward inputs (all on device): query: (batch, num_heads, head_dim) bf16 kv_cache: (num_blocks, page_size, num_kv_heads, head_dim * 2) bf16 block_table: (batch, max_blocks) int32 seq_lens: (batch,) int32 Output: attn_out: (batch, num_heads, head_dim) bf16 """ def __init__( self, batch: int, num_heads: int, num_kv_heads: int, head_dim: int, seq_len: int, page_size: int, ): super().__init__() assert num_heads % num_kv_heads == 0, "num_heads must be a multiple of num_kv_heads (GQA)" 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.frac = 100 self.splits, self.warps, self.policy = _default_config( batch, num_heads, num_kv_heads, head_dim, seq_len, page_size) self._ext = _get_cuda_ext() self._ws = None self._graphs: dict = {} # Triton fallback config self.block_n = 64 self.block_s = triton.next_power_of_2(self.splits) self.warps_main = 4 self.stages_main = 2 self.ev_policy = self.policy def _workspace(self, device): if self._ws is None: B, Hkv, S, G, D = (self.batch, self.num_kv_heads, self.splits, self.group_size, self.head_dim) self._ws = ( torch.zeros(B, Hkv, S, G, D, dtype=torch.float32, device=device), torch.zeros(B, Hkv, S, G, 2, dtype=torch.float32, device=device), torch.zeros(B, Hkv, dtype=torch.int32, device=device), ) return self._ws def _triton_ws(self, device): if not hasattr(self, "_triton_ws_bufs"): B, Hkv, S, D = (self.batch, self.num_kv_heads, self.splits, self.head_dim) self._triton_ws_bufs = ( torch.empty(B, Hkv, S, 16, D, dtype=torch.float32, device=device), torch.empty(B, Hkv, S, 16, dtype=torch.float32, device=device), ) return self._triton_ws_bufs def _forward_impl(self, query, kv_cache, block_table, seq_lens): B, H, D = query.shape out = torch.empty(B, H, D, dtype=query.dtype, device=query.device) if self._ext is not None: partial, ml, counters = self._workspace(query.device) self._ext.paged_decode_forward( query, kv_cache, block_table, seq_lens, partial, ml, counters, out, self.splits, self.warps, self.policy, self.frac / 100.0) return out # Triton fallback Hkv, G, P = self.num_kv_heads, self.group_size, self.page_size S = self.splits part, lse = self._triton_ws(query.device) _paged_decode_main[(B * Hkv * S,)]( query, kv_cache, block_table, seq_lens, part, lse, H=H, HKV=Hkv, G=G, D=D, P=P, MAXBT=block_table.stride(0), SCALE=self.scale, BLOCK_N=self.block_n, SPLITS=S, EV=self.ev_policy, num_warps=self.warps_main, num_stages=self.stages_main, ) _paged_decode_reduce[(B * Hkv * 16,)]( part, lse, out, H=H, HKV=Hkv, G=G, D=D, SPLITS=S, BLOCK_S=self.block_s, num_warps=2, ) return out def forward(self, query, kv_cache, block_table, seq_lens): key = (query.data_ptr(), kv_cache.data_ptr(), block_table.data_ptr(), seq_lens.data_ptr()) ent = self._graphs.get(key) if ent is None: # First sighting of this pointer set: run eager once to compile / # warm everything, then capture a graph for subsequent replays. self._forward_impl(query, kv_cache, block_table, seq_lens) try: g = torch.cuda.CUDAGraph() with torch.cuda.graph(g): static_out = self._forward_impl( query, kv_cache, block_table, seq_lens) ent = (g, static_out, (query, kv_cache, block_table, seq_lens)) self._graphs[key] = ent g.replay() return static_out except Exception: return self._forward_impl(query, kv_cache, block_table, seq_lens) g, static_out, _refs = ent g.replay() return static_out def get_inputs(): """Build random paged inputs for the current module-level shape knobs.""" B = BATCH H = NUM_HEADS Hkv = NUM_KV_HEADS D = HEAD_DIM L = SEQ_LEN P = 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]