"""Custom top-k for H100 (SM90). k==1 : argmax (1 HBM pass) k==8 : thread-local top-k + tree-merge (1 HBM pass) k>=16 : 4x8-bit radix select for kth ordered key + compact + bitonic multi-block partials for large n / small batch (cached workspace) Beats the cuTOPK / CUB baseline on L2-flushed roofline timing across the shape mix. """ from __future__ import annotations import os os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "9.0") import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline CUDA_SRC = r""" #include #include #include #include #include #include struct Pair { float val; int idx; }; __device__ __forceinline__ Pair make_pair(float v, int i) { Pair p; p.val = v; p.idx = i; return p; } __device__ __forceinline__ Pair neg_inf_pair() { return make_pair(-FLT_MAX, -1); } __device__ __forceinline__ bool better(const Pair& a, const Pair& b) { if (a.val != b.val) return a.val > b.val; if (a.idx < 0) return false; if (b.idx < 0) return true; return a.idx < b.idx; } __device__ __forceinline__ uint32_t float_to_ordered(float f) { uint32_t bits = __float_as_uint(f); uint32_t mask = (bits & 0x80000000u) ? 0xffffffffu : 0x80000000u; return bits ^ mask; } __device__ void bitonic_sort_desc(Pair* smem, int n, int tid, int nthreads) { for (int size = 2; size <= n; size <<= 1) { for (int stride = size >> 1; stride > 0; stride >>= 1) { for (int i = tid; i < n; i += nthreads) { int partner = i ^ stride; if (partner > i) { bool up = ((i & size) == 0); Pair a = smem[i], b = smem[partner]; bool swap = up ? better(b, a) : better(a, b); if (swap) { smem[i] = b; smem[partner] = a; } } } __syncthreads(); } } } // ======================== k==1 ======================== __global__ void argmax_kernel( const float* __restrict__ x, float* __restrict__ out_vals, int64_t* __restrict__ out_idxs, int n ) { const int row = blockIdx.x; const float* row_ptr = x + (size_t)row * n; float best_v = -FLT_MAX; int best_i = 0; for (int i = threadIdx.x; i < n; i += blockDim.x) { float v = row_ptr[i]; if (v > best_v || (v == best_v && i < best_i)) { best_v = v; best_i = i; } } unsigned mask = 0xffffffffu; #pragma unroll for (int off = 16; off > 0; off >>= 1) { float ov = __shfl_down_sync(mask, best_v, off); int oi = __shfl_down_sync(mask, best_i, off); if (ov > best_v || (ov == best_v && oi < best_i)) { best_v = ov; best_i = oi; } } __shared__ float s_val[32]; __shared__ int s_idx[32]; int lane = threadIdx.x & 31; int wid = threadIdx.x >> 5; if (lane == 0) { s_val[wid] = best_v; s_idx[wid] = best_i; } __syncthreads(); if (wid == 0) { int nwarps = (blockDim.x + 31) >> 5; best_v = (lane < nwarps) ? s_val[lane] : -FLT_MAX; best_i = (lane < nwarps) ? s_idx[lane] : 0; #pragma unroll for (int off = 16; off > 0; off >>= 1) { float ov = __shfl_down_sync(mask, best_v, off); int oi = __shfl_down_sync(mask, best_i, off); if (ov > best_v || (ov == best_v && oi < best_i)) { best_v = ov; best_i = oi; } } if (lane == 0) { out_vals[row] = best_v; out_idxs[row] = (int64_t)best_i; } } } // ======================== small-k ======================== template __device__ __forceinline__ void insert_topk(Pair (&buf)[K], float v, int idx) { if (v < buf[K-1].val) return; if (v == buf[K-1].val && buf[K-1].idx >= 0 && idx >= buf[K-1].idx) return; Pair item = make_pair(v, idx); int i = K - 1; #pragma unroll for (; i > 0; --i) { if (better(item, buf[i-1])) buf[i] = buf[i-1]; else break; } buf[i] = item; } template __device__ __forceinline__ void merge_topk(const Pair* a, const Pair* b, Pair* dst) { int ia = 0, ib = 0; #pragma unroll for (int j = 0; j < K; ++j) { if (ia < K && (ib >= K || better(a[ia], b[ib]))) dst[j] = a[ia++]; else dst[j] = b[ib++]; } } template __device__ void tree_merge_smem(Pair* smem, int tid, int nt) { for (int stride = 1; stride < nt; stride <<= 1) { __syncthreads(); if ((tid & ((stride << 1) - 1)) == 0) { int other = tid + stride; if (other < nt) { Pair tmp[K]; merge_topk(smem + tid * K, smem + other * K, tmp); #pragma unroll for (int j = 0; j < K; ++j) smem[tid * K + j] = tmp[j]; } } } __syncthreads(); } template __global__ void topk_smallk_kernel( const float* __restrict__ x, float* __restrict__ out_vals, int64_t* __restrict__ out_idxs, int n ) { extern __shared__ char raw[]; Pair* smem = reinterpret_cast(raw); const int row = blockIdx.x; const int tid = threadIdx.x; const int nt = blockDim.x; const float* row_ptr = x + (size_t)row * n; Pair local[K]; #pragma unroll for (int i = 0; i < K; ++i) local[i] = neg_inf_pair(); for (int i = tid; i < n; i += nt) insert_topk(local, row_ptr[i], i); #pragma unroll for (int j = 0; j < K; ++j) smem[tid * K + j] = local[j]; tree_merge_smem(smem, tid, nt); if (tid < K) { out_vals[row * K + tid] = smem[tid].val; out_idxs[row * K + tid] = (int64_t)smem[tid].idx; } } // ======================== radix select ======================== template __global__ void topk_radix_kernel( const float* __restrict__ x, float* __restrict__ out_vals, int64_t* __restrict__ out_idxs, int n ) { constexpr int RADIX = 256; // Need room for K gt + K eq + bitonic pad → 4*K is plenty (K<=64) constexpr int CAND_MAX = 256; extern __shared__ char raw[]; int* hist = reinterpret_cast(raw); Pair* cands = reinterpret_cast(hist + RADIX); __shared__ uint32_t s_prefix; __shared__ int s_needed; __shared__ uint32_t s_kth; const int row = blockIdx.x; const int tid = threadIdx.x; const int nt = blockDim.x; const float* row_ptr = x + (size_t)row * n; if (tid == 0) { s_prefix = 0; s_needed = K; } __syncthreads(); #pragma unroll for (int pass = 0; pass < 4; ++pass) { int shift = 24 - pass * 8; for (int i = tid; i < RADIX; i += nt) hist[i] = 0; __syncthreads(); uint32_t pmask = (pass == 0) ? 0u : (0xffffffffu << (32 - pass * 8)); uint32_t pfix = s_prefix; for (int i = tid; i < n; i += nt) { uint32_t key = float_to_ordered(row_ptr[i]); if (pass == 0 || (key & pmask) == pfix) atomicAdd(&hist[(key >> shift) & 255], 1); } __syncthreads(); if (tid == 0) { int remaining = s_needed; int chosen = 0; for (int d = RADIX - 1; d >= 0; --d) { if (hist[d] < remaining) remaining -= hist[d]; else { chosen = d; s_needed = remaining; break; } } s_prefix = (s_prefix & pmask) | ((uint32_t)chosen << shift); } __syncthreads(); } if (tid == 0) { s_kth = s_prefix; } __syncthreads(); uint32_t kth = s_kth; // Collect (correct under ties) in ONE pass: // count(key > kth) <= K-1 always. Store gt in cands[0..), eq in cands[K..). // Then compact: we need all gt + (K - n_gt) eq. __shared__ int s_n_gt; __shared__ int s_n_eq; if (tid == 0) { s_n_gt = 0; s_n_eq = 0; } __syncthreads(); // eq buffer starts at offset K (room for at most K-1 gt and K eq) for (int i = tid; i < n; i += nt) { float v = row_ptr[i]; uint32_t key = float_to_ordered(v); if (key > kth) { int pos = atomicAdd(&s_n_gt, 1); if (pos < K) cands[pos] = make_pair(v, i); } else if (key == kth) { int pos = atomicAdd(&s_n_eq, 1); if (pos < K) cands[K + pos] = make_pair(v, i); } } __syncthreads(); int n_gt = min(s_n_gt, K); int need_eq = K - n_gt; if (need_eq < 0) need_eq = 0; int n_eq_take = min(min(s_n_eq, need_eq), K); // Compact eq into cands[n_gt ..) if (tid < n_eq_take) { cands[n_gt + tid] = cands[K + tid]; } __syncthreads(); int nc = n_gt + n_eq_take; if (nc < K) { for (int i = nc + tid; i < K; i += nt) cands[i] = neg_inf_pair(); nc = K; } int n_pow2 = 1; while (n_pow2 < nc) n_pow2 <<= 1; if (n_pow2 < K) { n_pow2 = 1; while (n_pow2 < K) n_pow2 <<= 1; } for (int i = nc + tid; i < n_pow2; i += nt) cands[i] = neg_inf_pair(); __syncthreads(); bitonic_sort_desc(cands, n_pow2, tid, nt); __syncthreads(); if (tid < K) { out_vals[row * K + tid] = cands[tid].val; out_idxs[row * K + tid] = (int64_t)cands[tid].idx; } } template __global__ void topk_radix_partial_kernel( const float* __restrict__ x, float* __restrict__ partial_vals, int64_t* __restrict__ partial_idxs, int n, int blocks_per_row, int seg_size ) { constexpr int RADIX = 256; constexpr int CAND_MAX = 256; extern __shared__ char raw[]; int* hist = reinterpret_cast(raw); Pair* cands = reinterpret_cast(hist + RADIX); __shared__ uint32_t s_prefix; __shared__ int s_needed; __shared__ uint32_t s_kth; const int row = blockIdx.x / blocks_per_row; const int seg = blockIdx.x % blocks_per_row; const int tid = threadIdx.x; const int nt = blockDim.x; const int start = seg * seg_size; int end = min(start + seg_size, n); int seg_n = max(0, end - start); const float* base = x + (size_t)row * n + start; if (seg_n <= 0) { if (tid < K) { size_t off = ((size_t)row * blocks_per_row + seg) * K + tid; partial_vals[off] = -FLT_MAX; partial_idxs[off] = -1; } return; } int kk = min(K, seg_n); if (tid == 0) { s_prefix = 0; s_needed = kk; } __syncthreads(); #pragma unroll for (int pass = 0; pass < 4; ++pass) { int shift = 24 - pass * 8; for (int i = tid; i < RADIX; i += nt) hist[i] = 0; __syncthreads(); uint32_t pmask = (pass == 0) ? 0u : (0xffffffffu << (32 - pass * 8)); uint32_t pfix = s_prefix; for (int i = tid; i < seg_n; i += nt) { uint32_t key = float_to_ordered(base[i]); if (pass == 0 || (key & pmask) == pfix) atomicAdd(&hist[(key >> shift) & 255], 1); } __syncthreads(); if (tid == 0) { int remaining = s_needed; int chosen = 0; for (int d = RADIX - 1; d >= 0; --d) { if (hist[d] < remaining) remaining -= hist[d]; else { chosen = d; s_needed = remaining; break; } } s_prefix = (s_prefix & pmask) | ((uint32_t)chosen << shift); } __syncthreads(); } if (tid == 0) { s_kth = s_prefix; } __syncthreads(); uint32_t kth = s_kth; __shared__ int s_n_gt; __shared__ int s_n_eq; if (tid == 0) { s_n_gt = 0; s_n_eq = 0; } __syncthreads(); for (int i = tid; i < seg_n; i += nt) { float v = base[i]; uint32_t key = float_to_ordered(v); if (key > kth) { int pos = atomicAdd(&s_n_gt, 1); if (pos < K) cands[pos] = make_pair(v, start + i); } else if (key == kth) { int pos = atomicAdd(&s_n_eq, 1); if (pos < K) cands[K + pos] = make_pair(v, start + i); } } __syncthreads(); int n_gt = min(s_n_gt, K); int need_eq = kk - n_gt; if (need_eq < 0) need_eq = 0; int n_eq_take = min(min(s_n_eq, need_eq), K); if (tid < n_eq_take) cands[n_gt + tid] = cands[K + tid]; __syncthreads(); int nc = n_gt + n_eq_take; if (nc < kk) { for (int i = nc + tid; i < kk; i += nt) cands[i] = neg_inf_pair(); nc = kk; } int n_pow2 = 1; while (n_pow2 < nc) n_pow2 <<= 1; if (n_pow2 < K) { n_pow2 = 1; while (n_pow2 < K) n_pow2 <<= 1; } for (int i = nc + tid; i < n_pow2; i += nt) cands[i] = neg_inf_pair(); __syncthreads(); bitonic_sort_desc(cands, n_pow2, tid, nt); __syncthreads(); if (tid < K) { size_t off = ((size_t)row * blocks_per_row + seg) * K + tid; partial_vals[off] = cands[tid].val; partial_idxs[off] = (int64_t)cands[tid].idx; } } template __global__ void topk_merge_kernel( const float* __restrict__ partial_vals, const int64_t* __restrict__ partial_idxs, float* __restrict__ out_vals, int64_t* __restrict__ out_idxs, int blocks_per_row ) { extern __shared__ char raw[]; Pair* smem = reinterpret_cast(raw); const int row = blockIdx.x; const int n_cand = blocks_per_row * K; int n_pow2 = 1; while (n_pow2 < n_cand) n_pow2 <<= 1; for (int i = threadIdx.x; i < n_pow2; i += blockDim.x) { if (i < n_cand) { size_t off = (size_t)row * n_cand + i; smem[i] = make_pair(partial_vals[off], (int)partial_idxs[off]); } else smem[i] = neg_inf_pair(); } __syncthreads(); bitonic_sort_desc(smem, n_pow2, threadIdx.x, blockDim.x); __syncthreads(); if (threadIdx.x < K) { out_vals[row * K + threadIdx.x] = smem[threadIdx.x].val; out_idxs[row * K + threadIdx.x] = (int64_t)smem[threadIdx.x].idx; } } // ======================== workspace ======================== struct WS { float* v = nullptr; int64_t* i = nullptr; size_t cap = 0; }; static WS g_ws; static void ensure_ws(size_t n) { if (g_ws.cap >= n && g_ws.v) return; if (g_ws.v) { cudaFree(g_ws.v); cudaFree(g_ws.i); } cudaMalloc(&g_ws.v, n * sizeof(float)); cudaMalloc(&g_ws.i, n * sizeof(int64_t)); g_ws.cap = n; } // ======================== launchers ======================== void launch_argmax(torch::Tensor x, torch::Tensor vals, torch::Tensor idxs) { int n = (int)x.size(1); int t = (n >= 8192) ? 512 : 256; argmax_kernel<<>>( x.data_ptr(), vals.data_ptr(), idxs.data_ptr(), n); } template void launch_smallk(torch::Tensor x, torch::Tensor vals, torch::Tensor idxs) { int batch = (int)x.size(0); int n = (int)x.size(1); int threads = 256; if (n <= 4096) threads = 128; int smem = threads * K * (int)sizeof(Pair); cudaFuncSetAttribute(topk_smallk_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); topk_smallk_kernel<<>>( x.data_ptr(), vals.data_ptr(), idxs.data_ptr(), n); } template void launch_radix(torch::Tensor x, torch::Tensor vals, torch::Tensor idxs) { const int batch = (int)x.size(0); const int n = (int)x.size(1); auto stream = at::cuda::getCurrentCUDAStream(); bool multi = (batch == 1 && n >= 16384) || (batch <= 2 && n >= 65536); if (multi) { int threads = 256; int seg = 4096; int bpr = (n + seg - 1) / seg; int max_bpr = 1024 / K; if (bpr > max_bpr) { bpr = max_bpr; seg = (n + bpr - 1) / bpr; } int smem = 256 * (int)sizeof(int) + 256 * (int)sizeof(Pair); ensure_ws((size_t)batch * bpr * K); cudaFuncSetAttribute(topk_radix_partial_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); topk_radix_partial_kernel<<>>( x.data_ptr(), g_ws.v, g_ws.i, n, bpr, seg); int n_cand = bpr * K; int n_pow2 = 1; while (n_pow2 < n_cand) n_pow2 <<= 1; int mt = std::min(1024, n_pow2); topk_merge_kernel<<>>( g_ws.v, g_ws.i, vals.data_ptr(), idxs.data_ptr(), bpr); return; } int threads = 256; if (n >= 16384) threads = 512; int smem = 256 * (int)sizeof(int) + 256 * (int)sizeof(Pair); cudaFuncSetAttribute(topk_radix_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); topk_radix_kernel<<>>( x.data_ptr(), vals.data_ptr(), idxs.data_ptr(), n); } void topk_cuda(torch::Tensor x, torch::Tensor vals, torch::Tensor idxs, int64_t k) { TORCH_CHECK(x.is_cuda() && x.scalar_type() == torch::kFloat32, "x cuda fp32"); TORCH_CHECK(x.dim() == 2, "x 2D"); TORCH_CHECK(k >= 1 && k <= x.size(1), "bad k"); if (k == 1) { launch_argmax(x, vals, idxs); return; } if (k == 8) { launch_smallk<8>(x, vals, idxs); return; } if (k == 16) { launch_radix<16>(x, vals, idxs); return; } if (k == 32) { launch_radix<32>(x, vals, idxs); return; } if (k == 64) { launch_radix<64>(x, vals, idxs); return; } TORCH_CHECK(k <= 64, "k>64 unsupported"); auto v = torch::empty({x.size(0), 64}, vals.options()); auto i = torch::empty({x.size(0), 64}, idxs.options()); launch_radix<64>(x, v, i); vals.copy_(v.narrow(1, 0, k)); idxs.copy_(i.narrow(1, 0, k)); } """ CPP_SRC = r""" void topk_cuda(torch::Tensor x, torch::Tensor vals, torch::Tensor idxs, int64_t k); std::vector topk_forward(torch::Tensor x, int64_t k) { TORCH_CHECK(x.is_cuda(), "cuda"); auto vals = torch::empty({x.size(0), k}, x.options()); auto idxs = torch::empty({x.size(0), k}, torch::TensorOptions().dtype(torch::kInt64).device(x.device())); topk_cuda(x.contiguous(), vals, idxs, k); return {vals, idxs}; } """ _mod = None def _get_mod(): global _mod if _mod is None: _mod = load_inline( name="topk_h100_final4", cpp_sources=CPP_SRC, cuda_sources=CUDA_SRC, functions=["topk_forward"], extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"], verbose=False, ) return _mod class Model(nn.Module): def __init__(self, batch: int, n: int, k: int): super().__init__() self.batch, self.n, self.k = batch, n, k self.register_buffer("_dummy", torch.zeros(1)) def forward(self, x: torch.Tensor): return _get_mod().topk_forward(x, self.k) batch = 64 n = 8192 k = 8 def get_inputs(): return [torch.randn(batch, n, dtype=torch.float32)] def get_init_inputs(): return [batch, n, k]