"""SM100 paged-attention decode kernel. The CUDA kernel uses one warp per query head and one thread block per (batch, KV head, sequence partition). The query heads in a GQA group share each packed K/V page through shared memory, so every cache element is fetched from HBM only once per partition. A second kernel performs the stable merge of the independently normalized partitions. """ from __future__ import annotations import os from pathlib import Path # The run harness puts compiler-lock wrappers ahead of the actual toolkit. # Point cpp_extension at the installed CUDA toolkit before it is imported. if Path("/usr/local/cuda-12.8/bin/nvcc").exists(): os.environ["CUDA_HOME"] = "/usr/local/cuda-12.8" try: import ninja os.environ["PATH"] = ninja.BIN_DIR + os.pathsep + os.environ.get("PATH", "") except ImportError: pass import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline 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 _CPP_SRC = r""" #include void paged_attention_cuda( torch::Tensor query, torch::Tensor kv_cache, torch::Tensor block_table, torch::Tensor seq_lens, torch::Tensor partial_out, torch::Tensor partial_max, torch::Tensor partial_sum, torch::Tensor output, int64_t group_size, int64_t head_dim, int64_t page_size, int64_t partition_size); """ _CUDA_SRC = r""" #include #include #include #include #include #include namespace { constexpr int PAGE = 16; __device__ __forceinline__ void copy_async_16(void* dst, const void* src) { const unsigned int dst_shared = __cvta_generic_to_shared(dst); asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" [REDACTED: IP] "r"(dst_shared), "l"(src)); } __device__ __forceinline__ void copy_async_commit() { asm volatile("cp.async.commit_group;\n" [REDACTED: IP]); } __device__ __forceinline__ void copy_async_wait() { asm volatile("cp.async.wait_group 0;\n" [REDACTED: IP]); } template __global__ void partition_kernel( const __nv_bfloat16* __restrict__ query, const __nv_bfloat16* __restrict__ kv_cache, const int* __restrict__ block_table, const int* __restrict__ seq_lens, float* __restrict__ partial_out, float* __restrict__ partial_max, float* __restrict__ partial_sum, int batch, int num_heads, int num_kv_heads, int max_pages, int num_partitions) { const int linear = blockIdx.x; const int part = linear % num_partitions; const int t0 = linear / num_partitions; const int kv_head = t0 % num_kv_heads; const int b = t0 / num_kv_heads; const int warp = threadIdx.x >> 5; const int lane = threadIdx.x & 31; const int q_head = kv_head * G + warp; const int start = part * PARTITION; const int length = seq_lens[b]; const int end = min(start + PARTITION, length); // D is at most 128 in the shape set, hence four scalars per lane. constexpr int COLS = D / 32; float q[COLS]; float acc[COLS]; #pragma unroll for (int i = 0; i < COLS / 2; ++i) { const long long q_offset = ((long long)b * num_heads + q_head) * D + lane * COLS + 2 * i; const float2 q_pair = __bfloat1622float2( *reinterpret_cast(query + q_offset)); q[2 * i] = q_pair.x; q[2 * i + 1] = q_pair.y; acc[2 * i] = 0.0f; acc[2 * i + 1] = 0.0f; } float m = -INFINITY; float l = 0.0f; extern __shared__ __align__(16) unsigned char smem_raw[]; __nv_bfloat16* smem = reinterpret_cast<__nv_bfloat16*>(smem_raw); if (start < length) { const int first_page = start / PAGE; constexpr int [REDACTED credential assignment] * D) / 8; // int4 = 8 bf16 constexpr int PAGE_VECS = PAGE * VECS_PER_TOKEN; // Prime one half of the double buffer. const int first_physical_page = block_table[b * max_pages + first_page]; for (int vi = threadIdx.x; vi < PAGE_VECS; vi += blockDim.x) { const int [REDACTED credential assignment] / VECS_PER_TOKEN; const int [REDACTED credential assignment] - token * VECS_PER_TOKEN; const long long scalar_offset = ((((long long)first_physical_page * PAGE + token) * num_kv_heads + kv_head) * (2 * D)) + vector_in_token * 8; copy_async_16(reinterpret_cast(smem) + vi, kv_cache + scalar_offset); } copy_async_commit(); copy_async_wait(); __syncthreads(); // Partition and page boundaries are aligned (64/128 and 16 respectively). #pragma unroll for (int page_in_part = 0; page_in_part < PARTITION / PAGE; ++page_in_part) { const int token_base = start + page_in_part * PAGE; if (token_base >= end) break; __nv_bfloat16* page_smem = smem + (page_in_part & 1) * PAGE * 2 * D; // Prefetch the next gathered page while the warps consume this one. const int next_token_base = token_base + PAGE; if (next_token_base < end) { const int next_physical_page = block_table[b * max_pages + first_page + page_in_part + 1]; int4* next_smem = reinterpret_cast( smem + ((page_in_part + 1) & 1) * PAGE * 2 * D); for (int vi = threadIdx.x; vi < PAGE_VECS; vi += blockDim.x) { const int [REDACTED credential assignment] / VECS_PER_TOKEN; const int [REDACTED credential assignment] - token * VECS_PER_TOKEN; const long long scalar_offset = ((((long long)next_physical_page * PAGE + token) * num_kv_heads + kv_head) * (2 * D)) + vector_in_token * 8; copy_async_16(next_smem + vi, kv_cache + scalar_offset); } copy_async_commit(); } const int valid = min(PAGE, end - token_base); for (int [REDACTED credential assignment]; token < valid; ++token) { float dot = 0.0f; #pragma unroll for (int i = 0; i < COLS / 2; ++i) { const int d = lane * COLS + 2 * i; const float2 kval = __bfloat1622float2( *reinterpret_cast( page_smem + token * (2 * D) + d)); dot = fmaf(q[2 * i], kval.x, dot); dot = fmaf(q[2 * i + 1], kval.y, dot); } #pragma unroll for (int offset = 16; offset > 0; offset >>= 1) dot += __shfl_down_sync(0xffffffffu, dot, offset); const float score = __shfl_sync(0xffffffffu, dot, 0) * (1.0f / sqrtf((float)D)); const float new_m = fmaxf(m, score); const float alpha = __expf(m - new_m); const float beta = __expf(score - new_m); l = l * alpha + beta; #pragma unroll for (int i = 0; i < COLS / 2; ++i) { const int d = lane * COLS + 2 * i; const float2 vval = __bfloat1622float2( *reinterpret_cast( page_smem + token * (2 * D) + D + d)); acc[2 * i] = acc[2 * i] * alpha + beta * vval.x; acc[2 * i + 1] = acc[2 * i + 1] * alpha + beta * vval.y; } m = new_m; } copy_async_wait(); __syncthreads(); } } const long long stat_idx = ((long long)b * num_heads + q_head) * num_partitions + part; if (lane == 0) { partial_max[stat_idx] = m; partial_sum[stat_idx] = l; } const long long out_base = stat_idx * D; #pragma unroll for (int i = 0; i < COLS / 2; ++i) *reinterpret_cast(partial_out + out_base + lane * COLS + 2 * i) = make_float2(acc[2 * i], acc[2 * i + 1]); } template __global__ void reduce_kernel( const float* __restrict__ partial_out, const float* __restrict__ partial_max, const float* __restrict__ partial_sum, __nv_bfloat16* __restrict__ output, int num_heads, int num_partitions) { const int bh = blockIdx.x; const int lane = threadIdx.x; constexpr int COLS = D / 32; const long long stat_base = (long long)bh * num_partitions; float global_m = -INFINITY; if (lane == 0) { for (int p = 0; p < num_partitions; ++p) global_m = fmaxf(global_m, partial_max[stat_base + p]); } global_m = __shfl_sync(0xffffffffu, global_m, 0); float denom = 0.0f; float result[COLS]; #pragma unroll for (int i = 0; i < COLS; ++i) result[i] = 0.0f; for (int p = 0; p < num_partitions; ++p) { float factor = 0.0f; if (lane == 0) { factor = __expf(partial_max[stat_base + p] - global_m); denom += factor * partial_sum[stat_base + p]; } factor = __shfl_sync(0xffffffffu, factor, 0); const long long in_base = (stat_base + p) * D; #pragma unroll for (int i = 0; i < COLS; ++i) result[i] += factor * partial_out[in_base + lane * COLS + i]; } float inv_denom = 0.0f; if (lane == 0) inv_denom = 1.0f / denom; inv_denom = __shfl_sync(0xffffffffu, inv_denom, 0); const long long output_base = (long long)bh * D; #pragma unroll for (int i = 0; i < COLS / 2; ++i) *reinterpret_cast<__nv_bfloat162*>( output + output_base + lane * COLS + 2 * i) = __floats2bfloat162_rn(result[2 * i] * inv_denom, result[2 * i + 1] * inv_denom); } template void launch_typed( const torch::Tensor& query, const torch::Tensor& kv_cache, const torch::Tensor& block_table, const torch::Tensor& seq_lens, torch::Tensor& partial_out, torch::Tensor& partial_max, torch::Tensor& partial_sum, torch::Tensor& output) { const int batch = query.size(0); const int num_heads = query.size(1); const int num_kv_heads = num_heads / G; const int max_pages = block_table.size(1); const int num_partitions = partial_max.size(2); cudaStream_t stream = at::cuda::getCurrentCUDAStream(); const int blocks = batch * num_kv_heads * num_partitions; partition_kernel<<>>( reinterpret_cast(query.data_ptr()), reinterpret_cast(kv_cache.data_ptr()), block_table.data_ptr(), seq_lens.data_ptr(), partial_out.data_ptr(), partial_max.data_ptr(), partial_sum.data_ptr(), batch, num_heads, num_kv_heads, max_pages, num_partitions); reduce_kernel<<>>( partial_out.data_ptr(), partial_max.data_ptr(), partial_sum.data_ptr(), reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), num_heads, num_partitions); } } // namespace void paged_attention_cuda( torch::Tensor query, torch::Tensor kv_cache, torch::Tensor block_table, torch::Tensor seq_lens, torch::Tensor partial_out, torch::Tensor partial_max, torch::Tensor partial_sum, torch::Tensor output, int64_t group_size, int64_t head_dim, int64_t page_size, int64_t partition_size) { TORCH_CHECK(query.is_cuda() && kv_cache.is_cuda(), "CUDA tensors required"); TORCH_CHECK(query.scalar_type() == torch::kBFloat16, "query must be bf16"); TORCH_CHECK(kv_cache.scalar_type() == torch::kBFloat16, "cache must be bf16"); TORCH_CHECK(page_size == PAGE, "this kernel requires page_size=16"); if (head_dim == 128 && group_size == 4) { if (partition_size == 64) launch_typed<128, 4, 64>(query, kv_cache, block_table, seq_lens, partial_out, partial_max, partial_sum, output); else launch_typed<128, 4, 128>(query, kv_cache, block_table, seq_lens, partial_out, partial_max, partial_sum, output); } else if (head_dim == 128 && group_size == 8) { launch_typed<128, 8, 128>(query, kv_cache, block_table, seq_lens, partial_out, partial_max, partial_sum, output); } else if (head_dim == 64 && group_size == 4) { launch_typed<64, 4, 64>(query, kv_cache, block_table, seq_lens, partial_out, partial_max, partial_sum, output); } else { TORCH_CHECK(false, "unsupported (head_dim, group_size)"); } } """ def _load_extension(): # Keep JIT artifacts local to the submitted problem workspace. build_dir = Path(__file__).resolve().parent / ".paged_attention_build" build_dir.mkdir(exist_ok=True) os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "10.0") return load_inline( name="paged_attention_sm100_v1", cpp_sources=_CPP_SRC, cuda_sources=_CUDA_SRC, functions=["paged_attention_cuda"], extra_cflags=["-O3"], extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"], with_cuda=True, build_directory=str(build_dir), verbose=False, ) _EXT = _load_extension() class Model(nn.Module): 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 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 # The smaller workloads benefit from more independent partitions; # larger batches amortize better with 128-token partitions. self.partition_size = ( 64 if head_dim == 64 or (batch == 8 and self.group_size == 4) else 128 ) self.num_partitions = ( seq_len + self.partition_size - 1 ) // self.partition_size # Non-persistent workspaces preserve the reference's empty state_dict. self.register_buffer( "_partial_out", torch.empty(batch, num_heads, self.num_partitions, head_dim, dtype=torch.float32), persistent=False, ) self.register_buffer( "_partial_max", torch.empty(batch, num_heads, self.num_partitions, dtype=torch.float32), persistent=False, ) self.register_buffer( "_partial_sum", torch.empty(batch, num_heads, self.num_partitions, dtype=torch.float32), persistent=False, ) self.register_buffer( "_output", torch.empty(batch, num_heads, head_dim, dtype=torch.bfloat16), persistent=False, ) def forward( self, query: torch.Tensor, kv_cache: torch.Tensor, block_table: torch.Tensor, seq_lens: torch.Tensor, ) -> torch.Tensor: _EXT.paged_attention_cuda( query, kv_cache, block_table, seq_lens, self._partial_out, self._partial_max, self._partial_sum, self._output, self.group_size, self.head_dim, self.page_size, self.partition_size, ) return self._output def get_inputs(): batch, num_heads, num_kv_heads = BATCH, NUM_HEADS, NUM_KV_HEADS head_dim, seq_len, page_size = HEAD_DIM, SEQ_LEN, PAGE_SIZE pages_per_seq = (seq_len + page_size - 1) // page_size total_pages = max(batch * pages_per_seq + 8, 64) query = torch.randn(batch, num_heads, head_dim, dtype=torch.bfloat16) * 0.1 kv_cache = torch.randn( total_pages, page_size, num_kv_heads, 2 * head_dim, dtype=torch.bfloat16, ) * 0.1 block_table = torch.randperm(total_pages)[:batch * pages_per_seq] block_table = block_table.reshape(batch, pages_per_seq).int().contiguous() seq_lens = torch.full((batch,), seq_len, 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]