"""Custom CUDA top-k for H100 (SM90). Top-k over the last dim of a 2D fp32 tensor, values+indices sorted descending. Two kernel strategies, plus a generic fallback: - small k (<= 8): per-thread register top-k (tile sort-8 + running merge with a whole-tile float reject), warp restore-merge tournament. - 16 <= k <= 64: per-warp distributed truncated bitonic top-W over register runs (strided layout, merge-truncate), block tree in smem. - k > 64: warp-cooperative sorted-list kernel (generic, robust fallback). All selection is done on packed 64-bit keys (ordered-float32 << 32 | idx), so values sort descending with a deterministic index tie-break. """ import os os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0" # container pins many archs; H100 only import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline OP_TYPE = "topk" SUPPORTED_PRECISIONS = ["fp32"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] _CUDA_SRC = r''' #include #include #include #include #include using u32 = unsigned int; using u64 = unsigned long long; #define DEV __device__ __forceinline__ // ---------------- key packing ---------------- // Map fp32 to a monotone u32 (order-preserving). -0.0 < +0.0, NaN sorts above // +inf (the usual "NaN largest" convention). DEV u32 f2ord(float f) { u32 b = __float_as_uint(f); return (b & 0x80000000u) ? ~b : (b | 0x80000000u); } DEV float ord2f(u32 o) { u32 b = (o & 0x80000000u) ? (o ^ 0x80000000u) : ~o; return __uint_as_float(b); } // sort key: high 32 bits = ordered value, low 32 = element index (within row). // Sorting keys descending == values descending (idx desc as tie-break). DEV u64 pack(float f, u32 idx) { return ((u64)f2ord(f) << 32) | (u64)idx; } // Pad key: below every real key (real keys are >= ((u64)0x007FFFFF << 32)). #define PADKEY 0ULL DEV void ce_desc(u64 &a, u64 &b) { u64 mx = a > b ? a : b; u64 mn = a > b ? b : a; a = mx; b = mn; } DEV void ce_asc (u64 &a, u64 &b) { u64 mx = a > b ? a : b; u64 mn = a > b ? b : a; a = mn; b = mx; } // ---------------- path A: per-thread insertion lists / tiles ---------------- template DEV void list_insert(u64 (&l)[KP], u64 key) { if (key <= l[KP - 1]) return; #pragma unroll for (int j = 0; j < KP; ++j) { bool gt = key > l[j]; u64 t = l[j]; l[j] = gt ? key : t; key = gt ? t : key; } } template DEV void clean_priv(u64 (&l)[KP]) { #pragma unroll for (int d = KP / 2; d >= 1; d >>= 1) #pragma unroll for (int j = 0; j < KP; ++j) if ((j & d) == 0 && (j + d) < KP) ce_desc(l[j], l[j + d]); } template <> DEV void clean_priv<1>(u64 (&)[1]) {} template DEV void sortnet_priv(u64 (&l)[N]) { #pragma unroll for (int s = 2; s <= N; s <<= 1) #pragma unroll for (int d = s >> 1; d >= 1; d >>= 1) #pragma unroll for (int i = 0; i < N; ++i) if ((i & d) == 0) { bool desc = (i & s) == 0; if (desc) ce_desc(l[i], l[i + d]); else ce_asc (l[i], l[i + d]); } } // Warp tournament: every lane holds its own sorted-desc list[KP]; afterwards // every lane holds the warp's top-KP (redundant). template DEV void warp_topk_priv(u64 (&l)[KP]) { #pragma unroll for (int s = 1; s < 32; s <<= 1) { #pragma unroll for (int j = 0; j < KP; ++j) { u64 p = __shfl_xor_sync(0xffffffffu, l[KP - 1 - j], s); l[j] = l[j] > p ? l[j] : p; } clean_priv(l); } } // ---------------- width-W distributed state (strided layout) ---------------- template DEV void dist_clean(u64 (&st)[RV], int lane, u32 mask) { #pragma unroll for (int d = W / 2; d >= 1; d >>= 1) { if (d >= 32) { const int t = d / 32; #pragma unroll for (int r = 0; r < RV; ++r) if ((r & t) == 0 && (r + t) < RV) ce_desc(st[r], st[r + t]); } else { #pragma unroll for (int r = 0; r < RV; ++r) { u64 pv = __shfl_xor_sync(mask, st[r], d); bool low = (lane & d) == 0; st[r] = low ? (st[r] > pv ? st[r] : pv) : (st[r] > pv ? pv : st[r]); } } } } template DEV void dist_merge_mem(u64 (&st)[RV], const u64* __restrict__ partner, int lane, u32 mask) { #pragma unroll for (int r = 0; r < RV; ++r) { if (lane + 32 * r < W) { int p = lane + 32 * r; u64 pv = partner[W - 1 - p]; st[r] = st[r] > pv ? st[r] : pv; } } dist_clean(st, lane, mask); } template DEV void dist_merge_reg(u64 (&st)[RV], const u64 (&pt)[RV], int lane, u32 mask) { #pragma unroll for (int r = 0; r < RV; ++r) { int p = lane + 32 * r; int q = W - 1 - p; u64 pv = __shfl_sync(mask, pt[q / 32], q % 32); st[r] = st[r] > pv ? st[r] : pv; } dist_clean(st, lane, mask); } // Sort R register runs (32 lane-strided keys each) jointly, descending. template DEV void warp_sort32_multi(u64 (&v)[R], int lane) { #pragma unroll for (int s = 2; s <= 32; s <<= 1) { #pragma unroll for (int d = s >> 1; d >= 1; d >>= 1) { u64 pv[R]; #pragma unroll for (int r = 0; r < R; ++r) pv[r] = __shfl_xor_sync(0xffffffffu, v[r], d); int ilow = lane & ~d; bool desc = (ilow & s) == 0; bool low = (lane & d) == 0; bool keepmax = (low == desc); #pragma unroll for (int r = 0; r < R; ++r) v[r] = keepmax ? (v[r] > pv[r] ? v[r] : pv[r]) : (v[r] > pv[r] ? pv[r] : v[r]); } } } // Truncated top-W of R register runs (N = 32R candidates, strided layout). template DEV void dist_top(u64 (&st)[W / 32], u64 (&v)[R], int lane) { constexpr int RV = W / 32; warp_sort32_multi(v, lane); if constexpr (RV == 1) { st[0] = v[0]; #pragma unroll for (int i = 1; i < R; ++i) { u64 pv = __shfl_sync(0xffffffffu, v[i], 31 - lane); st[0] = st[0] > pv ? st[0] : pv; dist_clean<1, 32>(st, lane, 0xffffffffu); } } else { st[0] = v[0]; st[1] = __shfl_sync(0xffffffffu, v[1], 31 - lane); dist_clean<2, 64>(st, lane, 0xffffffffu); #pragma unroll for (int i = 2; i < R; i += 2) { u64 w2[2]; w2[0] = v[i]; if (i + 1 < R) w2[1] = __shfl_sync(0xffffffffu, v[i + 1], 31 - lane); else w2[1] = PADKEY; dist_clean<2, 64>(w2, lane, 0xffffffffu); dist_merge_reg<2, 64>(st, w2, lane, 0xffffffffu); } } } // ---------------- shared block finalization ---------------- template DEV void block_finalize(u64 (&st)[RV], int row, int part, int parts, int k, u64* __restrict__ smem, u64* __restrict__ partials, u32* __restrict__ counters, float* __restrict__ out_v, int64_t* __restrict__ out_i) { const int tid = threadIdx.x; const int lane = tid & 31; const int warp = tid >> 5; const int NW = TPB / 32; const u32 mask = (W >= 32) ? 0xffffffffu : ((1u << W) - 1u); if (NW > 1) { #pragma unroll for (int r = 0; r < RV; ++r) if (lane + 32 * r < W) smem[warp * W + lane + 32 * r] = st[r]; __syncthreads(); #pragma unroll for (int off = NW / 2; off >= 1; off >>= 1) { if (warp < off) { const u64* partner = smem + (warp + off) * W; dist_merge_mem(st, partner, lane, mask); #pragma unroll for (int r = 0; r < RV; ++r) if (lane + 32 * r < W) smem[warp * W + lane + 32 * r] = st[r]; } __syncthreads(); } } if (parts > 1) { if (warp == 0) { u64* dst = partials + ((long)row * parts + part) * W; #pragma unroll for (int r = 0; r < RV; ++r) if (lane + 32 * r < W) dst[lane + 32 * r] = st[r]; } __syncthreads(); __shared__ int is_last; if (tid == 0) { __threadfence(); unsigned int old = atomicAdd(counters + row, 1u); is_last = (old == (unsigned int)(parts - 1)) ? 1 : 0; if (is_last) counters[row] = 0u; // self-reset for the next call } __syncthreads(); if (!is_last) return; // this block merges all `parts` partials: buffer partner keys first const u64* base = partials + (long)row * parts * W; u64 buf[8][RV]; int cnt = 0; for (int p = warp; p < parts && cnt < 8; p += NW, ++cnt) { #pragma unroll for (int r = 0; r < RV; ++r) { int q = lane + 32 * r; if (q < W) buf[cnt][r] = base[(long)p * W + (W - 1 - q)]; } } #pragma unroll for (int r = 0; r < RV; ++r) st[r] = PADKEY; for (int c = 0; c < cnt; ++c) { #pragma unroll for (int r = 0; r < RV; ++r) st[r] = st[r] > buf[c][r] ? st[r] : buf[c][r]; dist_clean(st, lane, mask); } for (int p = warp + 8 * NW; p < parts; p += NW) { u64 a[RV]; #pragma unroll for (int r = 0; r < RV; ++r) { int q = lane + 32 * r; if (q < W) a[r] = base[(long)p * W + (W - 1 - q)]; } #pragma unroll for (int r = 0; r < RV; ++r) st[r] = st[r] > a[r] ? st[r] : a[r]; dist_clean(st, lane, mask); } __syncthreads(); #pragma unroll for (int r = 0; r < RV; ++r) if (lane + 32 * r < W) smem[warp * W + lane + 32 * r] = st[r]; __syncthreads(); #pragma unroll for (int off = NW / 2; off >= 1; off >>= 1) { if (warp < off) { const u64* partner = smem + (warp + off) * W; dist_merge_mem(st, partner, lane, mask); #pragma unroll for (int r = 0; r < RV; ++r) if (lane + 32 * r < W) smem[warp * W + lane + 32 * r] = st[r]; } __syncthreads(); } } if (warp == 0) { float* ov = out_v + (long)row * k; int64_t* oi = out_i + (long)row * k; #pragma unroll for (int r = 0; r < RV; ++r) { int p = lane + 32 * r; if (p < W && p < k) { u64 key = st[r]; ov[p] = ord2f((u32)(key >> 32)); oi[p] = (int64_t)(u32)key; } } } } // ---------------- kernel A: small k ---------------- template __global__ void topk_kernel_a(const float* __restrict__ x, int n, int k, int parts, int elems_per_part, float* __restrict__ out_v, int64_t* __restrict__ out_i, u64* __restrict__ partials, u32* __restrict__ counters) { extern __shared__ u64 smem[]; const int row = blockIdx.x; const int part = blockIdx.y; const int tid = threadIdx.x; const int lane = tid & 31; const long rowbase = (long)row * n; int begin = part * elems_per_part; int end = begin + elems_per_part; if (end > n) end = n; u64 l[KP]; if constexpr (TILED) { u64 st[8]; #pragma unroll for (int j = 0; j < 8; ++j) st[j] = PADKEY; float thrF = -CUDART_INF_F; int idx = begin + tid; while (idx < end) { float vf[8]; #pragma unroll for (int c = 0; c < 8; ++c) { int id = idx + c * TPB; vf[c] = (id < end) ? __ldg(x + rowbase + id) : __int_as_float(0xff800000); } float fm = vf[0]; #pragma unroll for (int c = 1; c < 8; ++c) fm = fmaxf(fm, vf[c]); if (!(fm < thrF)) { u64 t[8]; #pragma unroll for (int c = 0; c < 8; ++c) { int id = idx + c * TPB; t[c] = (id < end) ? pack(vf[c], (u32)id) : PADKEY; } sortnet_priv<8>(t); #pragma unroll for (int j = 0; j < 8; ++j) st[j] = st[j] > t[7 - j] ? st[j] : t[7 - j]; clean_priv<8>(st); thrF = (st[KP - 1] == PADKEY) ? -CUDART_INF_F : ord2f((u32)(st[KP - 1] >> 32)); } idx += 8 * TPB; } #pragma unroll for (int j = 0; j < KP; ++j) l[j] = st[j]; } else { #pragma unroll for (int j = 0; j < KP; ++j) l[j] = PADKEY; int idx = begin + tid; while (idx < end) { float v[4]; int id0 = idx; #pragma unroll for (int c = 0; c < 4; ++c) { int id = id0 + c * TPB; v[c] = (id < end) ? __ldg(x + rowbase + id) : __int_as_float(0xff800000); } #pragma unroll for (int c = 0; c < 4; ++c) { int id = id0 + c * TPB; if (id < end) list_insert(l, pack(v[c], (u32)id)); } idx += 4 * TPB; } } warp_topk_priv(l); constexpr int W = KP; u64 st[1]; st[0] = l[0]; #pragma unroll for (int j = 1; j < KP; ++j) st[0] = (lane == j) ? l[j] : st[0]; block_finalize<1, W, TPB>(st, row, part, parts, k, smem, partials, counters, out_v, out_i); } // ---------------- kernel B: distributed truncated bitonic top-W ---------------- template __global__ void topk_kernel_b(const float* __restrict__ x, int n, int k, int parts, int elems_per_part, float* __restrict__ out_v, int64_t* __restrict__ out_i, u64* __restrict__ partials, u32* __restrict__ counters) { extern __shared__ u64 smem[]; const int row = blockIdx.x; const int part = blockIdx.y; const int tid = threadIdx.x; const int lane = tid & 31; const long rowbase = (long)row * n; int begin = part * elems_per_part; int end = begin + elems_per_part; if (end > n) end = n; constexpr int RV = (W + 31) / 32; u64 v[R]; #pragma unroll for (int j = 0; j < R; ++j) { int idx = begin + tid + j * TPB; u64 key = PADKEY; if (idx < end) key = pack(__ldg(x + rowbase + idx), (u32)idx); v[j] = key; } static_assert(32 * R >= W, "sort must cover state width"); u64 st[RV]; dist_top(st, v, lane); block_finalize(st, row, part, parts, k, smem, partials, counters, out_v, out_i); } // ---------------- kernel F: generic warp list kernel (k > 64) ---------------- // Each warp keeps a sorted-desc list of k keys in smem; warp-cooperative // insertion for any element above the list minimum. One block per row. // Robustness path (rare/hypothetical shapes); simplicity over speed. __global__ void topk_kernel_f(const float* __restrict__ x, int n, int k, float* __restrict__ out_v, int64_t* __restrict__ out_i) { extern __shared__ u64 smem[]; const int row = blockIdx.x; const int tid = threadIdx.x; const int lane = tid & 31; const int warp = tid >> 5; const int NW = blockDim.x >> 5; const long rowbase = (long)row * n; u64* list = smem + (long)warp * (k + 32); for (int i = lane; i < k; i += 32) list[i] = PADKEY; __syncwarp(); // uniform trip count: lane idx = base + lane; guards per element for (int base = warp * 32; base < n; base += blockDim.x) { int idx = base + lane; u64 key = (idx < n) ? pack(__ldg(x + rowbase + idx), (u32)idx) : PADKEY; unsigned ballot = __ballot_sync(0xffffffffu, key > list[k - 1]); while (ballot) { int src = __ffs(ballot) - 1; ballot &= ballot - 1; u64 cand = __shfl_sync(0xffffffffu, key, src); int lo = 0, hi = k - 1; while (lo < hi) { int mid = (lo + hi) >> 1; if (list[mid] < cand) hi = mid; else lo = mid + 1; } for (int j = k - 1; j > lo; --j) list[j] = list[j - 1]; list[lo] = cand; __syncwarp(); } } __syncthreads(); // merge the NW warp lists into warp 0's list (single-lane linear merge) if (warp == 0 && lane == 0) { u64* a = list; for (int w = 1; w < NW; ++w) { u64* b = smem + (long)w * (k + 32); u64* tmp = smem + (long)NW * (k + 32); // scratch area int i = 0, j = 0; for (int t = 0; t < k; ++t) { if (j >= k || (i < k && a[i] >= b[j])) tmp[t] = a[i++]; else tmp[t] = b[j++]; } for (int t = 0; t < k; ++t) a[t] = tmp[t]; } } __syncthreads(); if (warp == 0) { float* ov = out_v + (long)row * k; int64_t* oi = out_i + (long)row * k; for (int j = lane; j < k; j += 32) { ov[j] = ord2f((u32)(list[j] >> 32)); oi[j] = (int64_t)(u32)list[j]; } } } // ---------------- host ---------------- static inline int next_pow2(int x) { int p = 1; while (p < x) p <<= 1; return p; } struct Workspace { torch::Tensor buf; u64* partials = nullptr; u32* counters = nullptr; long partial_capacity = 0; long counter_capacity = 0; }; static Workspace g_ws; static void ensure_workspace(const torch::TensorOptions& opt, long partial_keys, long rows) { if (partial_keys <= g_ws.partial_capacity && rows <= g_ws.counter_capacity) return; long pk = std::max(partial_keys, g_ws.partial_capacity * 2); long cr = std::max(rows, g_ws.counter_capacity * 2); if (pk < 1 << 16) pk = 1 << 16; if (cr < 1024) cr = 1024; long bytes = pk * 8 + cr * 4; g_ws.buf = at::zeros({bytes}, opt.dtype(at::kByte)); g_ws.partials = reinterpret_cast(g_ws.buf.data_ptr()); g_ws.counters = reinterpret_cast(reinterpret_cast(g_ws.buf.data_ptr()) + pk * 8); g_ws.partial_capacity = pk; g_ws.counter_capacity = cr; } std::tuple run(torch::Tensor x, long k_) { TORCH_CHECK(x.is_cuda(), "x must be a CUDA tensor"); TORCH_CHECK(x.scalar_type() == at::kFloat, "x must be fp32"); TORCH_CHECK(x.dim() == 2, "x must be 2D"); auto xc = x.is_contiguous() ? x : x.contiguous(); const int rows = (int)xc.size(0); const int n = (int)xc.size(1); const int k = (int)k_; TORCH_CHECK(k >= 1 && k <= n, "need 1 <= k <= n"); auto stream = at::cuda::getCurrentCUDAStream(); auto vals = at::detail::empty_cuda({rows, k}, at::kFloat, x.device(), ::std::nullopt); auto idxs = at::detail::empty_cuda({rows, k}, at::kLong, x.device(), ::std::nullopt); float* ov = vals.data_ptr(); int64_t* oi = idxs.data_ptr(); const float* xp = xc.data_ptr(); if (k <= 8) { const int KP = next_pow2(k); const int T = 256; int B = next_pow2(std::max(1, std::min((64 + rows - 1) / rows, 1 << 30))); int maxB = (n + T - 1) / T; if (maxB < 1) maxB = 1; if (B > maxB) { int p = 1; while (p * 2 <= maxB) p *= 2; B = p; } if (B < 1) B = 1; int elems = (n + B - 1) / B; dim3 grid(rows, B); size_t smem = (T / 32) * KP * 8; u64* partials = nullptr; u32* counters = nullptr; if (B > 1) { ensure_workspace(x.options(), (long)rows * B * KP, rows); partials = g_ws.partials; counters = g_ws.counters; } #define LAUNCH_A(KK, TT) topk_kernel_a 1)><<>>(xp, n, k, B, elems, ov, oi, partials, counters) switch (KP) { case 1: LAUNCH_A(1, 256); break; case 2: LAUNCH_A(2, 256); break; case 4: LAUNCH_A(4, 256); break; case 8: LAUNCH_A(8, 256); break; } #undef LAUNCH_A } else if (k <= 64) { const int W = k <= 32 ? 32 : 64; const int T = 256; int denom = (W == 32 ? (T * 7 / 2) : (T * 8)); if (denom < 1) denom = 1; int B0 = (n + denom / 2) / denom; if (B0 < 1) B0 = 1; int pb = 1; while (pb * 2 <= B0) pb <<= 1; if (pb * 3 < B0 * 2) pb <<= 1; int B = pb; if (B > 64) B = 64; if (B < 1) B = 1; int elems = (n + B - 1) / B; int R = (elems + T - 1) / T; if (R < 4) R = 4; while (R > 32 && B < 64) { B *= 2; elems = (n + B - 1) / B; R = (elems + T - 1) / T; } if (R > 32) R = 32; R = next_pow2(R); dim3 grid(rows, B); size_t smem = (T / 32) * W * 8; u64* partials = nullptr; u32* counters = nullptr; if (B > 1) { ensure_workspace(x.options(), (long)rows * B * W, rows); partials = g_ws.partials; counters = g_ws.counters; } #define LAUNCH_B(WW, RR) topk_kernel_b<<>>(xp, n, k, B, elems, ov, oi, partials, counters) switch (W * 100 + R) { case 3204: LAUNCH_B(32, 4); break; case 3208: LAUNCH_B(32, 8); break; case 3216: LAUNCH_B(32, 16); break; case 3232: LAUNCH_B(32, 32); break; case 6404: LAUNCH_B(64, 4); break; case 6408: LAUNCH_B(64, 8); break; case 6416: LAUNCH_B(64, 16); break; case 6432: LAUNCH_B(64, 32); break; default: TORCH_CHECK(false, "bad kernel B config"); } #undef LAUNCH_B } else { TORCH_CHECK(k <= 2048, "k too large"); const int T = 256; size_t smem = (size_t)(T / 32 + 1) * ((int)k + 32) * 8; cudaFuncSetAttribute(topk_kernel_f, cudaFuncAttributeMaxDynamicSharedMemorySize, 220000); topk_kernel_f<<>>(xp, n, k, ov, oi); } return {torch::Tensor(vals), torch::Tensor(idxs)}; } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("run", &run); } ''' _mod = load_inline( name="topk_h100", cpp_sources="", cuda_sources=_CUDA_SRC, verbose=False, extra_cuda_cflags=["-O3"], ) class Model(nn.Module): """Top-k over the last dim of a 2D fp32 tensor. Input: x (batch, n) fp32. Output: values (batch, k) fp32 sorted descending, indices (batch, k) int64. """ 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)) self._run = _mod.run def forward(self, x: torch.Tensor): return self._run(x, self.k) # bypass nn.Module.__call__ hook machinery (hot path) __call__ = forward # Module-level shims rebuilt by check.py / benchmark.py per shape. batch = 64 n = 8192 k = 8 def get_inputs(): x = torch.randn(batch, n, dtype=torch.float32) return [x] def get_init_inputs(): return [batch, n, k]