KernelBench hard · RTX PRO 6000

TopK Bitonic Grok 4.6

3.68%geomean peak fraction across shapes

manually audited: clean

Custom CUDA top-k: k=1 block argmax, k=8/16 insertion+merge, k=32 warp bitonic, k=64 / long rows CUB DeviceTopK plus warp sort. No graph cache. 0.0368 sits in the known launch-overhead band for this problem.

harnessgrokagent session1h 12mtotal wall1h 15mcheck2sbenchmark1soutput tokensregimememory

Per-shape vs governing ceilingeach shape graded against whichever binds — fp32 compute or HBM bandwidth

1×131072×640.021 ms1.4%0.02 TB/s · 1% of 1.8 TB/s HBM · also 0 TFLOPS (0% of compute)
64×8192×80.017 ms7.1%0.13 TB/s · 7% of 1.8 TB/s HBM · also 0 TFLOPS (0% of compute)
32×16384×320.033 ms3.5%0.06 TB/s · 4% of 1.8 TB/s HBM · also 0 TFLOPS (0% of compute)
16×12000×160.026 ms1.6%0.03 TB/s · 2% of 1.8 TB/s HBM · also 0 TFLOPS (0% of compute)
128×4096×10.010 ms12.0%0.22 TB/s · 12% of 1.8 TB/s HBM · also 0 TFLOPS (0% of compute)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(1.4% · 7.1% · 3.5% · 1.6% · 12.0%) = 3.7%

Kernel source (redacted)
"""Custom CUDA top-k (largest, sorted) for RTX PRO 6000 / SM120.

k=1: vectorized block argmax
k=8,16: per-thread insertion + pairwise list merge
k=32: warp-shuffle bitonic select over row chunks
k=64 / long rows: CUB DeviceTopK (AIR) + tiny warp sort
"""
from __future__ import annotations

import ctypes
import hashlib
import os
import subprocess
import sys
from pathlib import Path

import torch
import torch.nn as nn

# Module-level shims (same contract as reference.py).
batch = 64
n = 8192
k = 8

_CUDA_SRC = r"""
#include <cuda_runtime.h>
#include <cfloat>
#include <cstdint>

#include <cub/device/device_topk.cuh>
#include <cuda/__execution/determinism.h>
#include <cuda/__execution/output_ordering.h>
#include <cuda/__execution/require.h>
#include <cuda/__stream/get_stream.h>
#include <cuda/std/__execution/env.h>
#include <thrust/iterator/counting_iterator.h>

// ---------------------------------------------------------------------------
// Warp bitonic
// ---------------------------------------------------------------------------
__device__ __forceinline__ void warp_sort32_desc(float& key, int& val) {
    const int lane = threadIdx.x & 31;
    #pragma unroll
    for (int size = 2; size <= 32; size <<= 1) {
        #pragma unroll
        for (int stride = size >> 1; stride > 0; stride >>= 1) {
            const float ok = __shfl_xor_sync(0xffffffff, key, stride);
            const int ov = __shfl_xor_sync(0xffffffff, val, stride);
            const bool dir_desc = (lane & size) == 0;
            const bool is_high = (lane & stride) != 0;
            const bool want = dir_desc ? (is_high ? key > ok : key < ok)
                                       : (is_high ? key < ok : key > ok);
            if (want) { key = ok; val = ov; }
        }
    }
}

__device__ __forceinline__ void warp_merge32_desc(float& key, int& val) {
    const int lane = threadIdx.x & 31;
    #pragma unroll
    for (int stride = 16; stride > 0; stride >>= 1) {
        const float ok = __shfl_xor_sync(0xffffffff, key, stride);
        const int ov = __shfl_xor_sync(0xffffffff, val, stride);
        const bool is_high = (lane & stride) != 0;
        const bool want = is_high ? key > ok : key < ok;
        if (want) { key = ok; val = ov; }
    }
}

__device__ __forceinline__ void warp_sort64_desc(float& k0, float& k1, int& v0, int& v1) {
    float k[2] = {k0, k1};
    int v[2] = {v0, v1};
    const int lane = threadIdx.x & 31;
    #pragma unroll
    for (int size = 2; size <= 64; size <<= 1) {
        #pragma unroll
        for (int stride = size >> 1; stride > 0; stride >>= 1) {
            const float sk0 = k[0], sk1 = k[1];
            const int sv0 = v[0], sv1 = v[1];
            #pragma unroll
            for (int r = 0; r < 2; ++r) {
                const int i = (r << 5) + lane;
                const int partner = i ^ stride;
                const int pr = partner >> 5;
                const int plane = partner & 31;
                const float src_k = (pr == 0) ? sk0 : sk1;
                const int src_v = (pr == 0) ? sv0 : sv1;
                const float ok = __shfl_sync(0xffffffff, src_k, plane);
                const int ov = __shfl_sync(0xffffffff, src_v, plane);
                const float cur = (r == 0) ? sk0 : sk1;
                const bool dir_desc = (i & size) == 0;
                const bool is_high = (i & stride) != 0;
                const bool want = dir_desc ? (is_high ? cur > ok : cur < ok)
                                           : (is_high ? cur < ok : cur > ok);
                if (want) { k[r] = ok; v[r] = ov; }
            }
        }
    }
    k0 = k[0]; k1 = k[1]; v0 = v[0]; v1 = v[1];
}

__device__ __forceinline__ void wsel32_add(
        float val, bool keep, int idx,
        float& wk, int& wv, float& thresh, int keep_lane) {
    if (__any_sync(0xffffffff, keep)) {
        float tk = keep ? val : -FLT_MAX;
        int tv = keep ? idx : -1;
        warp_sort32_desc(tk, tv);
        const float ok = __shfl_xor_sync(0xffffffff, tk, 31);
        const int ov = __shfl_xor_sync(0xffffffff, tv, 31);
        if (ok > wk) { wk = ok; wv = ov; }
        warp_merge32_desc(wk, wv);
        thresh = __shfl_sync(0xffffffff, wk, keep_lane);
    }
}

__device__ __forceinline__ void wsel64_add(
        float val, bool keep, int idx,
        float& wk0, float& wk1, int& wv0, int& wv1, float& thresh) {
    if (__any_sync(0xffffffff, keep)) {
        float tk = keep ? val : -FLT_MAX;
        int tv = keep ? idx : -1;
        warp_sort32_desc(tk, tv);
        const float ok = __shfl_xor_sync(0xffffffff, tk, 31);
        const int ov = __shfl_xor_sync(0xffffffff, tv, 31);
        if (ok > wk1) { wk1 = ok; wv1 = ov; }
        warp_sort64_desc(wk0, wk1, wv0, wv1);
        thresh = __shfl_sync(0xffffffff, wk1, 31);
    }
}

template <int BLOCK>
__device__ __forceinline__ void smem_bitonic_desc(float* keys, int* vals, int n_pow2) {
    const int tid = threadIdx.x;
    for (int size = 2; size <= n_pow2; size <<= 1) {
        for (int stride = size >> 1; stride > 0; stride >>= 1) {
            __syncthreads();
            for (int i = tid; i < n_pow2; i += BLOCK) {
                const int j = i ^ stride;
                if (j > i) {
                    const bool dir_desc = (i & size) == 0;
                    const float ki = keys[i], kj = keys[j];
                    if (dir_desc ? (ki < kj) : (ki > kj)) {
                        keys[i] = kj; keys[j] = ki;
                        const int vi = vals[i], vj = vals[j];
                        vals[i] = vj; vals[j] = vi;
                    }
                }
            }
        }
    }
    __syncthreads();
}

// ---------------------------------------------------------------------------
// Insertion top-k for KK in {8,16}
// ---------------------------------------------------------------------------
template <int KK>
__device__ __forceinline__ void insert_desc(float* lk, int* li, float val, int idx) {
    if (val <= lk[KK - 1]) return;
    lk[KK - 1] = val;
    li[KK - 1] = idx;
    #pragma unroll
    for (int t = KK - 2; t >= 0; --t) {
        if (lk[t] < lk[t + 1]) {
            const float tv = lk[t]; lk[t] = lk[t + 1]; lk[t + 1] = tv;
            const int ti = li[t]; li[t] = li[t + 1]; li[t + 1] = ti;
        }
    }
}

template <int KK>
__device__ __forceinline__ void merge_two_desc(float* a, int* ai, const float* b, const int* bi) {
    float tv[KK];
    int ti[KK];
    int ia = 0, ib = 0;
    #pragma unroll
    for (int t = 0; t < KK; ++t) {
        const bool take_a = (ia < KK) && (ib >= KK || a[ia] >= b[ib]);
        if (take_a) { tv[t] = a[ia]; ti[t] = ai[ia]; ++ia; }
        else        { tv[t] = b[ib]; ti[t] = bi[ib]; ++ib; }
    }
    #pragma unroll
    for (int t = 0; t < KK; ++t) { a[t] = tv[t]; ai[t] = ti[t]; }
}

template <int BLOCK, int KK, bool WRITE_I64>
__global__ void topk_insert_kernel(
        const float* __restrict__ x,
        float* __restrict__ out_v,
        long long* __restrict__ out_i64,
        int* __restrict__ out_i32,
        int nn, int chunks) {
    constexpr int N_WARPS = BLOCK / 32;
    __shared__ float cand_v[BLOCK * KK];
    __shared__ int   cand_i[BLOCK * KK];
    __shared__ float warp_v[N_WARPS * KK];
    __shared__ int   warp_i[N_WARPS * KK];

    const int chunk = (int)blockIdx.x;
    const int row = (int)blockIdx.y;
    const int chunk_sz = (nn + chunks - 1) / chunks;
    const int col0 = chunk * chunk_sz;
    int col1 = col0 + chunk_sz;
    if (col1 > nn) col1 = nn;
    if (col0 >= nn) return;

    const float* rowp = x + (size_t)row * nn;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int tid = threadIdx.x;

    float lk[KK];
    int li[KK];
    #pragma unroll
    for (int t = 0; t < KK; ++t) { lk[t] = -FLT_MAX; li[t] = -1; }

    int a0 = (col0 + 3) & ~3;
    int a1 = col1 & ~3;
    if (a0 > col1) a0 = col1;
    if (a1 < a0) a1 = a0;

    for (int idx = col0 + tid; idx < a0; idx += BLOCK)
        insert_desc<KK>(lk, li, __ldg(rowp + idx), idx);
    {
        const int nvec = (a1 - a0) >> 2;
        const float4* vptr = reinterpret_cast<const float4*>(rowp + a0);
        for (int vec = tid; vec < nvec; vec += BLOCK) {
            const float4 v = __ldg(vptr + vec);
            const int b = a0 + (vec << 2);
            insert_desc<KK>(lk, li, v.x, b);
            insert_desc<KK>(lk, li, v.y, b + 1);
            insert_desc<KK>(lk, li, v.z, b + 2);
            insert_desc<KK>(lk, li, v.w, b + 3);
        }
    }
    for (int idx = a1 + tid; idx < col1; idx += BLOCK)
        insert_desc<KK>(lk, li, __ldg(rowp + idx), idx);

    #pragma unroll
    for (int t = 0; t < KK; ++t) {
        cand_v[tid * KK + t] = lk[t];
        cand_i[tid * KK + t] = li[t];
    }
    __syncwarp();
    for (int off = 1; off < 32; off <<= 1) {
        if ((lane & ((off << 1) - 1)) == 0) {
            merge_two_desc<KK>(cand_v + tid * KK, cand_i + tid * KK,
                               cand_v + (tid + off) * KK, cand_i + (tid + off) * KK);
        }
        __syncwarp();
    }
    if (lane == 0) {
        #pragma unroll
        for (int t = 0; t < KK; ++t) {
            warp_v[warp * KK + t] = cand_v[tid * KK + t];
            warp_i[warp * KK + t] = cand_i[tid * KK + t];
        }
    }
    __syncthreads();
    for (int off = 1; off < N_WARPS; off <<= 1) {
        if ((warp & ((off << 1) - 1)) == 0 && lane == 0) {
            merge_two_desc<KK>(warp_v + warp * KK, warp_i + warp * KK,
                               warp_v + (warp + off) * KK, warp_i + (warp + off) * KK);
        }
        __syncthreads();
    }
    if (tid < KK) {
        const int out_base = WRITE_I64 ? (row * KK + tid)
                                       : ((row * chunks + chunk) * KK + tid);
        out_v[out_base] = warp_v[tid];
        if constexpr (WRITE_I64) out_i64[out_base] = (long long)warp_i[tid];
        else out_i32[out_base] = warp_i[tid];
    }
}

// ---------------------------------------------------------------------------
// Warp-select for KK in {32,64}
// ---------------------------------------------------------------------------
template <int BLOCK, int KK, bool WRITE_I64>
__global__ void topk_wsel_kernel(
        const float* __restrict__ x,
        float* __restrict__ out_v,
        long long* __restrict__ out_i64,
        int* __restrict__ out_i32,
        int nn, int chunks) {
    constexpr int N_WARPS = BLOCK / 32;
    constexpr bool WIDE = (KK > 32);
    constexpr int Q = WIDE ? 64 : 32;
    constexpr int KEEP_LANE = (KK <= 32) ? (KK - 1) : 31;

    __shared__ float sm_v[N_WARPS * Q];
    __shared__ int   sm_i[N_WARPS * Q];

    const int chunk = (int)blockIdx.x;
    const int row = (int)blockIdx.y;
    const int chunk_sz = (nn + chunks - 1) / chunks;
    const int col0 = chunk * chunk_sz;
    int col1 = col0 + chunk_sz;
    if (col1 > nn) col1 = nn;
    if (col0 >= nn) return;

    const float* rowp = x + (size_t)row * nn;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;

    float thresh = -FLT_MAX;
    float wk0 = -FLT_MAX, wk1 = -FLT_MAX;
    int wv0 = -1, wv1 = -1;

    int a0 = (col0 + 3) & ~3;
    int a1 = col1 & ~3;
    if (a0 > col1) a0 = col1;
    if (a1 < a0) a1 = a0;

    auto add = [&](float val, bool valid, int idx) {
        const bool keep = valid && (val > thresh);
        if constexpr (WIDE) wsel64_add(val, keep, idx, wk0, wk1, wv0, wv1, thresh);
        else                wsel32_add(val, keep, idx, wk0, wv0, thresh, KEEP_LANE);
    };

    {
        const int idx = col0 + threadIdx.x;
        const bool valid = idx < a0;
        add(valid ? __ldg(rowp + idx) : -FLT_MAX, valid, idx);
    }
    {
        const int nvec = (a1 - a0) >> 2;
        const float4* vptr = reinterpret_cast<const float4*>(rowp + a0);
        for (int base = 0; base < nvec; base += BLOCK) {
            const int vec = base + threadIdx.x;
            const bool valid = vec < nvec;
            float4 v; int g = 0;
            if (valid) { v = __ldg(vptr + vec); g = a0 + (vec << 2); }
            else v.x = v.y = v.z = v.w = -FLT_MAX;
            add(v.x, valid, g);
            add(v.y, valid, g + 1);
            add(v.z, valid, g + 2);
            add(v.w, valid, g + 3);
        }
    }
    {
        const int idx = a1 + threadIdx.x;
        const bool valid = idx < col1;
        add(valid ? __ldg(rowp + idx) : -FLT_MAX, valid, idx);
    }

    sm_v[warp * Q + lane] = wk0;
    sm_i[warp * Q + lane] = wv0;
    if constexpr (WIDE) {
        sm_v[warp * Q + 32 + lane] = wk1;
        sm_i[warp * Q + 32 + lane] = wv1;
    }
    __syncthreads();
    smem_bitonic_desc<BLOCK>(sm_v, sm_i, N_WARPS * Q);

    if (threadIdx.x < KK) {
        const int out_base = WRITE_I64 ? (row * KK + threadIdx.x)
                                       : ((row * chunks + chunk) * KK + threadIdx.x);
        out_v[out_base] = sm_v[threadIdx.x];
        if constexpr (WRITE_I64) out_i64[out_base] = (long long)sm_i[threadIdx.x];
        else out_i32[out_base] = sm_i[threadIdx.x];
    }
}

// Merge of partial (value, index) pairs via insertion (small n_cand).
template <int BLOCK, int KK>
__global__ void topk_merge_insert_kernel(
        const float* __restrict__ in_v,
        const int* __restrict__ in_i,
        float* __restrict__ out_v,
        long long* __restrict__ out_i,
        int n_cand) {
    constexpr int N_WARPS = BLOCK / 32;
    __shared__ float cand_v[BLOCK * KK];
    __shared__ int   cand_i[BLOCK * KK];
    __shared__ float warp_v[N_WARPS * KK];
    __shared__ int   warp_i[N_WARPS * KK];

    const int row = (int)blockIdx.x;
    const float* rv = in_v + (size_t)row * n_cand;
    const int* ri = in_i + (size_t)row * n_cand;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int tid = threadIdx.x;

    float lk[KK];
    int li[KK];
    #pragma unroll
    for (int t = 0; t < KK; ++t) { lk[t] = -FLT_MAX; li[t] = -1; }
    for (int i = tid; i < n_cand; i += BLOCK)
        insert_desc<KK>(lk, li, rv[i], ri[i]);

    #pragma unroll
    for (int t = 0; t < KK; ++t) {
        cand_v[tid * KK + t] = lk[t];
        cand_i[tid * KK + t] = li[t];
    }
    __syncwarp();
    for (int off = 1; off < 32; off <<= 1) {
        if ((lane & ((off << 1) - 1)) == 0) {
            merge_two_desc<KK>(cand_v + tid * KK, cand_i + tid * KK,
                               cand_v + (tid + off) * KK, cand_i + (tid + off) * KK);
        }
        __syncwarp();
    }
    if (lane == 0) {
        #pragma unroll
        for (int t = 0; t < KK; ++t) {
            warp_v[warp * KK + t] = cand_v[tid * KK + t];
            warp_i[warp * KK + t] = cand_i[tid * KK + t];
        }
    }
    __syncthreads();
    for (int off = 1; off < N_WARPS; off <<= 1) {
        if ((warp & ((off << 1) - 1)) == 0 && lane == 0) {
            merge_two_desc<KK>(warp_v + warp * KK, warp_i + warp * KK,
                               warp_v + (warp + off) * KK, warp_i + (warp + off) * KK);
        }
        __syncthreads();
    }
    if (tid < KK) {
        out_v[row * KK + tid] = warp_v[tid];
        out_i[row * KK + tid] = (long long)warp_i[tid];
    }
}

template <int BLOCK>
__global__ void argmax_kernel(
        const float* __restrict__ x,
        float* __restrict__ out_v,
        long long* __restrict__ out_i,
        int nn) {
    const int row = (int)blockIdx.x;
    const float* rowp = x + (size_t)row * nn;
    float best = -FLT_MAX;
    int besti = 0;
    const int nvec = nn >> 2;
    const float4* vptr = reinterpret_cast<const float4*>(rowp);
    for (int i = threadIdx.x; i < nvec; i += BLOCK) {
        const float4 v = __ldg(vptr + i);
        const int b = i << 2;
        if (v.x > best) { best = v.x; besti = b; }
        if (v.y > best) { best = v.y; besti = b + 1; }
        if (v.z > best) { best = v.z; besti = b + 2; }
        if (v.w > best) { best = v.w; besti = b + 3; }
    }
    for (int j = (nvec << 2) + threadIdx.x; j < nn; j += BLOCK) {
        const float v = __ldg(rowp + j);
        if (v > best) { best = v; besti = j; }
    }
    #pragma unroll
    for (int off = 16; off > 0; off >>= 1) {
        const float ov = __shfl_xor_sync(0xffffffff, best, off);
        const int oi = __shfl_xor_sync(0xffffffff, besti, off);
        if (ov > best) { best = ov; besti = oi; }
    }
    __shared__ float sbest[BLOCK / 32];
    __shared__ int sidx[BLOCK / 32];
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    if (lane == 0) { sbest[warp] = best; sidx[warp] = besti; }
    __syncthreads();
    if (warp == 0) {
        const int nwarps = BLOCK / 32;
        best = (lane < nwarps) ? sbest[lane] : -FLT_MAX;
        besti = (lane < nwarps) ? sidx[lane] : 0;
        #pragma unroll
        for (int off = 16; off > 0; off >>= 1) {
            const float ov = __shfl_xor_sync(0xffffffff, best, off);
            const int oi = __shfl_xor_sync(0xffffffff, besti, off);
            if (ov > best) { best = ov; besti = oi; }
        }
        if (lane == 0) {
            out_v[row] = best;
            out_i[row] = (long long)besti;
        }
    }
}

// ---------------------------------------------------------------------------
// k=64 over a short chunk: each thread holds <=16 values, pairwise merge
// growing from 16 -> 32 -> 64, then across warps.
// ---------------------------------------------------------------------------
template <bool WRITE_I64>
__global__ void topk_k64_kernel(
        const float* __restrict__ x,
        float* __restrict__ out_v,
        long long* __restrict__ out_i64,
        int* __restrict__ out_i32,
        int nn, int chunks) {
    constexpr int BLOCK = 256;
    constexpr int KK = 64;
    constexpr int N_WARPS = 8;
    constexpr int LOCAL = 16;

    // Per-warp working area: 32 slots of up to 64 (value,index).
    // 8 * 32 * 64 * 8 bytes = 128KB -- too big.
    // Instead: 32 slots of 16 initially in cand[BLOCK*16], then
    // compact into warp_v as we merge.

    __shared__ float cand_v[BLOCK * LOCAL];
    __shared__ int   cand_i[BLOCK * LOCAL];
    __shared__ float warp_v[N_WARPS * KK];
    __shared__ int   warp_i[N_WARPS * KK];

    const int chunk = (int)blockIdx.x;
    const int row = (int)blockIdx.y;
    const int chunk_sz = (nn + chunks - 1) / chunks;
    const int col0 = chunk * chunk_sz;
    int col1 = col0 + chunk_sz;
    if (col1 > nn) col1 = nn;
    if (col0 >= nn) return;

    const float* rowp = x + (size_t)row * nn;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int tid = threadIdx.x;

    float lk[LOCAL];
    int li[LOCAL];
    int nloc = 0;
    #pragma unroll
    for (int t = 0; t < LOCAL; ++t) { lk[t] = -FLT_MAX; li[t] = -1; }

    for (int idx = col0 + tid; idx < col1; idx += BLOCK) {
        if (nloc < LOCAL) {
            lk[nloc] = __ldg(rowp + idx);
            li[nloc] = idx;
            ++nloc;
        } else {
            insert_desc<LOCAL>(lk, li, __ldg(rowp + idx), idx);
        }
    }
    // Local sort of the populated prefix (always sort all 16; pads are -FLT_MAX).
    #pragma unroll
    for (int i = 1; i < LOCAL; ++i) {
        const float ck = lk[i];
        const int cv = li[i];
        int j = i - 1;
        while (j >= 0 && lk[j] < ck) {
            lk[j + 1] = lk[j];
            li[j + 1] = li[j];
            --j;
        }
        lk[j + 1] = ck;
        li[j + 1] = cv;
    }

    #pragma unroll
    for (int t = 0; t < LOCAL; ++t) {
        cand_v[tid * LOCAL + t] = lk[t];
        cand_i[tid * LOCAL + t] = li[t];
    }
    __syncwarp();

    // Pairwise: 32 lists of 16.
    // off=1: merge two 16 -> 16 (not enough!) we need to keep more.
    // So after first merge we must store 32 values. Use warp_v as extra space
    // plus the 16-wide cand slots: each even lane writes 32 results into
    // a 32-wide packing in cand (overwrite two 16-slots).
    //
    // Layout after off=1: even lanes hold 32 values in cand[tid*16 : tid*16+32]
    // which overlaps the next (odd) lane's slot — that's OK because odd lanes
    // are done.

    // off=1, keep 32: write 32 values starting at cand[tid*16]
    if ((lane & 1) == 0) {
        float tv[32];
        int ti[32];
        const float* a = cand_v + tid * LOCAL;
        const float* b = cand_v + (tid + 1) * LOCAL;
        const int* ai = cand_i + tid * LOCAL;
        const int* bi = cand_i + (tid + 1) * LOCAL;
        int ia = 0, ib = 0;
        #pragma unroll
        for (int t = 0; t < 32; ++t) {
            const bool take_a = (ia < 16) && (ib >= 16 || a[ia] >= b[ib]);
            if (take_a) { tv[t] = a[ia]; ti[t] = ai[ia]; ++ia; }
            else        { tv[t] = b[ib]; ti[t] = bi[ib]; ++ib; }
        }
        #pragma unroll
        for (int t = 0; t < 32; ++t) {
            cand_v[tid * LOCAL + t] = tv[t];
            cand_i[tid * LOCAL + t] = ti[t];
        }
    }
    __syncwarp();

    // off=2, keep 64: even-even lanes (0,4,8,...) merge two 32-lists.
    // 32 values sit at cand[tid*16 : tid*16+32]. The neighbor is tid+2,
    // whose 32 values sit at cand[(tid+2)*16 : ].
    // 32+32=64, write 64 values starting at cand[tid*16] — occupies 4
    // original slots (tid..tid+3). Those odd/skipped lanes are idle.
    if ((lane & 3) == 0) {
        float tv[64];
        int ti[64];
        const float* a = cand_v + tid * LOCAL;
        const float* b = cand_v + (tid + 2) * LOCAL;
        const int* ai = cand_i + tid * LOCAL;
        const int* bi = cand_i + (tid + 2) * LOCAL;
        int ia = 0, ib = 0;
        #pragma unroll
        for (int t = 0; t < 64; ++t) {
            const bool take_a = (ia < 32) && (ib >= 32 || a[ia] >= b[ib]);
            if (take_a) { tv[t] = a[ia]; ti[t] = ai[ia]; ++ia; }
            else        { tv[t] = b[ib]; ti[t] = bi[ib]; ++ib; }
        }
        #pragma unroll
        for (int t = 0; t < 64; ++t) {
            cand_v[tid * LOCAL + t] = tv[t];
            cand_i[tid * LOCAL + t] = ti[t];
        }
    }
    __syncwarp();

    // Remaining merges keep 64. off=4,8,16.
    for (int off = 4; off < 32; off <<= 1) {
        if ((lane & ((off << 1) - 1)) == 0) {
            float tv[64];
            int ti[64];
            const float* a = cand_v + tid * LOCAL;
            const float* b = cand_v + (tid + off) * LOCAL;
            const int* ai = cand_i + tid * LOCAL;
            const int* bi = cand_i + (tid + off) * LOCAL;
            int ia = 0, ib = 0;
            #pragma unroll
            for (int t = 0; t < 64; ++t) {
                const bool take_a = (ia < 64) && (ib >= 64 || a[ia] >= b[ib]);
                if (take_a) { tv[t] = a[ia]; ti[t] = ai[ia]; ++ia; }
                else        { tv[t] = b[ib]; ti[t] = bi[ib]; ++ib; }
            }
            #pragma unroll
            for (int t = 0; t < 64; ++t) {
                cand_v[tid * LOCAL + t] = tv[t];
                cand_i[tid * LOCAL + t] = ti[t];
            }
        }
        __syncwarp();
    }

    // Lane 0 holds the warp top-64 in cand[tid*16 : tid*16+64].
    // tid for lane 0 is warp*32. cand[(warp*32)*16] = cand[warp*512],
    // and we need 64 values — that's fine, 512>64.
    if (lane == 0) {
        #pragma unroll
        for (int t = 0; t < KK; ++t) {
            warp_v[warp * KK + t] = cand_v[(warp * 32) * LOCAL + t];
            warp_i[warp * KK + t] = cand_i[(warp * 32) * LOCAL + t];
        }
    }
    __syncthreads();

    for (int off = 1; off < N_WARPS; off <<= 1) {
        if ((warp & ((off << 1) - 1)) == 0 && lane == 0) {
            merge_two_desc<KK>(
                warp_v + warp * KK, warp_i + warp * KK,
                warp_v + (warp + off) * KK, warp_i + (warp + off) * KK);
        }
        __syncthreads();
    }

    if (tid < KK) {
        const int out_base = WRITE_I64 ? (row * KK + tid)
                                       : ((row * chunks + chunk) * KK + tid);
        out_v[out_base] = warp_v[tid];
        if constexpr (WRITE_I64) out_i64[out_base] = (long long)warp_i[tid];
        else out_i32[out_base] = warp_i[tid];
    }
}

static void launch_slice(const float* x, float* vals, long long* i64, int* i32,
                         int batch, int nn, int kk, int chunks, bool write_i64,
                         cudaStream_t stream) {
    dim3 grid(chunks, batch);
    if (kk <= 8) {
        if (write_i64) topk_insert_kernel<256, 8, true><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
        else           topk_insert_kernel<256, 8, false><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
    } else if (kk <= 16) {
        if (write_i64) topk_insert_kernel<256, 16, true><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
        else           topk_insert_kernel<256, 16, false><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
    } else if (kk <= 32) {
        if (write_i64) topk_wsel_kernel<256, 32, true><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
        else           topk_wsel_kernel<256, 32, false><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
    } else {
        if (write_i64) topk_k64_kernel<true><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
        else           topk_k64_kernel<false><<<grid, 256, 0, stream>>>(x, vals, i64, i32, nn, chunks);
    }
}

struct StreamEnv {
    cudaStream_t s;
    cuda::stream_ref query(cuda::get_stream_t) const noexcept {
        return cuda::stream_ref{s};
    }
};

__global__ void sort_small_desc(float* v, const int* i32, long long* i64, int kk) {
    float k0 = -FLT_MAX, k1 = -FLT_MAX;
    int v0 = -1, v1 = -1;
    const int lane = threadIdx.x & 31;
    if (lane < kk) { k0 = v[lane]; v0 = i32[lane]; }
    if (32 + lane < kk) { k1 = v[32 + lane]; v1 = i32[32 + lane]; }
    if (kk <= 32) {
        warp_sort32_desc(k0, v0);
        if (lane < kk) {
            v[lane] = k0;
            i64[lane] = (long long)v0;
        }
    } else {
        warp_sort64_desc(k0, k1, v0, v1);
        v[lane] = k0;
        i64[lane] = (long long)v0;
        if (32 + lane < kk) {
            v[32 + lane] = k1;
            i64[32 + lane] = (long long)v1;
        }
    }
}

static auto topk_env(cudaStream_t stream) {
    auto req = cuda::execution::require(
        cuda::execution::determinism::not_guaranteed,
        cuda::execution::output_ordering::unsorted);
    return cuda::std::execution::env{req, StreamEnv{stream}};
}

extern "C" size_t query_devicetopk_temp(int nn, int kk) {
    size_t tb = 0;
    auto env = topk_env(0);
    thrust::counting_iterator<int> ids(0);
    float* dummy = nullptr;
    int* dummy_i = nullptr;
    cub::DeviceTopK::MaxPairs(nullptr, tb, dummy, dummy, ids, dummy_i, nn, kk, env);
    return tb;
}

extern "C" void launch_devicetopk(
        const float* x, float* vals, long long* idxs,
        void* temp, size_t temp_bytes, int* idx32,
        int nn, int kk, cudaStream_t stream) {
    auto env = topk_env(stream);
    thrust::counting_iterator<int> ids(0);
    cub::DeviceTopK::MaxPairs(temp, temp_bytes, x, vals, ids, idx32, nn, kk, env);
    sort_small_desc<<<1, 32, 0, stream>>>(vals, idx32, idxs, kk);
}

extern "C" void launch_argmax(
        const float* x, float* vals, long long* idxs,
        int batch, int nn, cudaStream_t stream) {
    argmax_kernel<256><<<batch, 256, 0, stream>>>(x, vals, idxs, nn);
}

extern "C" void launch_topk_slice_i64(
        const float* x, float* vals, long long* idxs,
        int batch, int nn, int kk, int chunks, cudaStream_t stream) {
    launch_slice(x, vals, idxs, nullptr, batch, nn, kk, chunks, true, stream);
}

extern "C" void launch_topk_slice_i32(
        const float* x, float* vals, int* idxs,
        int batch, int nn, int kk, int chunks, cudaStream_t stream) {
    launch_slice(x, vals, nullptr, idxs, batch, nn, kk, chunks, false, stream);
}

extern "C" void launch_topk_merge(
        const float* in_v, const int* in_i,
        float* vals, long long* idxs,
        int batch, int n_cand, int kk, cudaStream_t stream) {
    if (kk <= 8)       topk_merge_insert_kernel<256, 8><<<batch, 256, 0, stream>>>(in_v, in_i, vals, idxs, n_cand);
    else if (kk <= 16) topk_merge_insert_kernel<256, 16><<<batch, 256, 0, stream>>>(in_v, in_i, vals, idxs, n_cand);
    else if (kk <= 32) topk_merge_insert_kernel<128, 32><<<batch, 128, 0, stream>>>(in_v, in_i, vals, idxs, n_cand);
    else               topk_merge_insert_kernel<64, 64><<<batch, 64, 0, stream>>>(in_v, in_i, vals, idxs, n_cand);
}
"""

_lib = None


def _nvcc_bin() -> str:
    for key in ("CUDACXX", "NVCC"):
        p = os.environ.get(key)
        if p and Path(p).exists():
            return p
    home = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
    if home:
        cand = Path(home) / "bin" / "nvcc"
        if cand.exists():
            return str(cand)
    return "nvcc"


def _ext():
    global _lib
    if _lib is not None:
        return _lib
    build = Path("/tmp/topk_bitonic_build")
    build.mkdir(parents=True, exist_ok=True)
    digest = hashlib.sha1(_CUDA_SRC.encode()).hexdigest()[:16]
    src = build / f"topk_kernels_{digest}.cu"
    so = build / f"topk_kernels_{digest}.so"
    if not so.exists():
        src.write_text(_CUDA_SRC)
        incs = []
        seen = set()
        for base in (
            os.environ.get("CUDA_HOME"),
            os.environ.get("CUDA_PATH"),
            "/usr/local/cuda-host",
            "/usr/local/cuda",
        ):
            if not base:
                continue
            cand = Path(base) / "include" / "cccl"
            key = str(cand)
            if cand.is_dir() and key not in seen:
                seen.add(key)
                incs.append(f"-I{key}")
        cmd = [
            _nvcc_bin(),
            "-O3",
            "--use_fast_math",
            "-std=c++17",
            "--expt-relaxed-constexpr",
            "--expt-extended-lambda",
            "--shared",
            "-Xcompiler",
            "-fPIC",
            "-gencode",
            "arch=compute_120,code=sm_120",
            *incs,
            "-o",
            str(so),
            str(src),
        ]
        proc = subprocess.run(cmd, capture_output=True, text=True)
        if proc.returncode != 0:
            sys.stderr.write(proc.stdout)
            sys.stderr.write(proc.stderr)
            raise RuntimeError("nvcc failed to build top-k kernel")
    _lib = ctypes.CDLL(str(so))
    c_void = ctypes.c_void_p
    c_int = ctypes.c_int
    for name, args in (
        ("launch_argmax", [c_void, c_void, c_void, c_int, c_int, c_void]),
        ("launch_topk_slice_i64", [c_void, c_void, c_void, c_int, c_int, c_int, c_int, c_void]),
        ("launch_topk_slice_i32", [c_void, c_void, c_void, c_int, c_int, c_int, c_int, c_void]),
        ("launch_topk_merge", [c_void, c_void, c_void, c_void, c_int, c_int, c_int, c_void]),
        (
            "launch_devicetopk",
            [c_void, c_void, c_void, c_void, ctypes.c_size_t, c_void, c_int, c_int, c_void],
        ),
    ):
        fn = getattr(_lib, name)
        fn.argtypes = args
        fn.restype = None
    q = _lib.query_devicetopk_temp
    q.argtypes = [c_int, c_int]
    q.restype = ctypes.c_size_t
    return _lib


def _choose_chunks(batch_: int, n_: int, k_: int) -> int:
    if k_ == 1:
        return 1
    if k_ >= 32 and batch_ == 1:
        return max(1, min(32, n_ // 4096))
    # k=32 mid-size rows: a few chunks raise occupancy.
    if k_ >= 32 and n_ >= 8192:
        return 8
    if batch_ >= 16:
        return 1
    target_blocks = 32
    chunks = (target_blocks + batch_ - 1) // batch_
    max_chunks = max(1, n_ // 2048)
    return max(1, min(chunks, max_chunks))


class Model(nn.Module):
    """Top-k over the last dim of a 2D tensor."""

    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._chunks = _choose_chunks(batch, n, k)
        self._use_devicetopk = batch == 1 and k >= 32 and n >= 65536
        self._ws_v = None
        self._ws_i = None
        self._cub_temp = None
        self._cub_idx32 = None
        self._cub_temp_bytes = 0
        self._out_v = None
        self._out_i = None

    def forward(self, x: torch.Tensor):
        x = x.contiguous()
        bsz = int(x.size(0))
        nn = int(x.size(1))
        if self._out_v is None or self._out_v.device != x.device or self._out_v.size(0) != bsz:
            self._out_v = torch.empty(bsz, self.k, device=x.device, dtype=torch.float32)
            self._out_i = torch.empty(bsz, self.k, device=x.device, dtype=torch.int64)
        vals = self._out_v
        idxs = self._out_i
        lib = _ext()
        stream = ctypes.c_void_p(torch.cuda.current_stream().cuda_stream)
        xp = ctypes.c_void_p(x.data_ptr())
        vp = ctypes.c_void_p(vals.data_ptr())
        ip = ctypes.c_void_p(idxs.data_ptr())
        if self._use_devicetopk:
            if self._cub_temp is None or self._cub_temp.device != x.device:
                nbytes = int(lib.query_devicetopk_temp(nn, int(self.k)))
                self._cub_temp_bytes = max(nbytes, 256)
                self._cub_temp = torch.empty(
                    self._cub_temp_bytes, device=x.device, dtype=torch.uint8
                )
                self._cub_idx32 = torch.empty(
                    self.k, device=x.device, dtype=torch.int32
                )
            lib.launch_devicetopk(
                xp,
                vp,
                ip,
                ctypes.c_void_p(self._cub_temp.data_ptr()),
                ctypes.c_size_t(self._cub_temp_bytes),
                ctypes.c_void_p(self._cub_idx32.data_ptr()),
                nn,
                int(self.k),
                stream,
            )
            return vals, idxs
        if self.k == 1:
            lib.launch_argmax(xp, vp, ip, bsz, nn, stream)
        elif self._chunks <= 1:
            lib.launch_topk_slice_i64(xp, vp, ip, bsz, nn, int(self.k), 1, stream)
        else:
            if self._ws_v is None or self._ws_v.device != x.device:
                self._ws_v = torch.empty(
                    bsz, self._chunks, self.k, device=x.device, dtype=torch.float32
                )
                self._ws_i = torch.empty(
                    bsz, self._chunks, self.k, device=x.device, dtype=torch.int32
                )
            lib.launch_topk_slice_i32(
                xp,
                ctypes.c_void_p(self._ws_v.data_ptr()),
                ctypes.c_void_p(self._ws_i.data_ptr()),
                bsz,
                nn,
                int(self.k),
                int(self._chunks),
                stream,
            )
            lib.launch_topk_merge(
                ctypes.c_void_p(self._ws_v.data_ptr()),
                ctypes.c_void_p(self._ws_i.data_ptr()),
                vp,
                ip,
                bsz,
                int(self._chunks) * int(self.k),
                int(self.k),
                stream,
            )
        return vals, idxs


def get_inputs():
    x = torch.randn(batch, n, dtype=torch.float32)
    return [x]


def get_init_inputs():
    return [batch, n, k]

20260813_072912_grok_grok-4.6_05_topk_bitonic