Skip to content
KernelIndex
Search⌘K

submission 808904

Leiko · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 786 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-808904?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
49.9ms
#396 of 515
2026-06-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:501385111e413dff823ae9731a334c64228fdc604cb49a6c617715a39705a4af
license declaredunknown
license concludedunknown
authorsLeiko
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 4static constexpr int NUM_WARPS = 4;
persistent-kernel__device__ __forceinline__ void persistent_barrier(int* counters, int* sense,
shared-memoryextern __shared__ int __shm[];
tile-k = 64static constexpr int BK = 64;
tile-m = 128static constexpr int BM = 128;
tile-n = 128static constexpr int BN = 128;

Kernel source

submission.py786 lines
"""General blocked Householder QR for the Popcorn `qr` leaderboard (B200 / sm_100a).

Design (see linalg-qr-b200.md):
  * Right-looking *blocked* Householder QR with compact-WY blocks.
  * Panel factorization, the compact-WY T factor, and the explicit reflector
    matrix V are computed in **fp32 on CUDA cores** (accuracy-critical work).
  * The trailing update  C <- C - V T^T (V^T C)  is expressed as three GEMMs and
    run on the **tcgen05 tensor cores via ThunderKittens**.

Precision note (important):
  The Blackwell design doc asks for a "tf32" trailing update, but ThunderKittens'
  tcgen05 MMA descriptor does NOT expose tf32 -- its `kind::f16` family only
  encodes `half` and `bf16` (see ThunderKittens/include/ops/thread/mma/tcgen05.cuh).
  We therefore realize the trailing GEMMs with **fp16 inputs + fp32 accumulation**.
  fp16 carries a 10-bit mantissa, the same as tf32, so this matches the intended
  accuracy target while staying on the verified tcgen05 path. If accuracy proves
  insufficient for the largest / most ill-conditioned cases, the upgrade path is
  bf16 3-split ("fp32 emulation") in the same kernel.

Status:
  The tcgen05 GEMM mirrors the verified `tk_tcgen05_probe_submission.py` handshake.
  Small matrices (n < 128) and the env var QR_FORCE_FALLBACK=1 use torch.geqrf
  (the LAPACK baseline) as a correctness oracle / fallback.

The previous all-custom-CUDA submission is preserved in v1.py.
"""

import os

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


_BLOCK = 128          # panel width == tcgen05 output-tile M/N
_TK_MIN_N = 128       # below this we just do a single panel (no trailing update)


def _force_fallback() -> bool:
    return os.environ.get("QR_FORCE_FALLBACK", "0") == "1"


# ---------------------------------------------------------------------------
# ThunderKittens setup (clone-or-find, identical to the verified probes).
# ---------------------------------------------------------------------------
def _thunderkittens_root() -> str:
    import subprocess

    local_root = os.path.abspath("ThunderKittens")
    if os.path.exists(os.path.join(local_root, "include", "kittens.cuh")):
        return local_root

    tmp_root = "/tmp/ThunderKittens"
    if not os.path.exists(os.path.join(tmp_root, "include", "kittens.cuh")):
        subprocess.check_call(
            [
                "git", "clone", "--depth", "1",
                "https://github.com/HazyResearch/ThunderKittens.git",
                tmp_root,
            ]
        )
    return tmp_root


# ---------------------------------------------------------------------------
# tcgen05 bf16-3-split GEMM (the only ThunderKittens kernel): D = A @ B, or
# D := D - A @ B when subtract != 0. Each operand is a batched 3D tensor; the
# tile the kernel reads/writes is selected by a per-operand (row,col) tile
# offset, so an operand can be a *strided sub-block* of a larger matrix (e.g.
# the trailing submatrix of `out`). Tile units: A rows / D rows / D cols / B
# cols are 128; A cols / B rows (the K dim) are 64. Mt,Nt,Kt are tile counts.
# fp32 in global memory, split to bf16 hi/lo for the MMA, fp32 accumulation.
# ---------------------------------------------------------------------------
_TK_CPP = r"""
void tk_gemm(torch::Tensor A, long ar, long ac, long a_ro, long a_co,
             torch::Tensor B, long br, long bc, long b_ro, long b_co,
             torch::Tensor D, long dr, long dc, long d_ro, long d_co,
             long batch, long Mt, long Nt, long Kt, long subtract);
"""

_TK_CUDA = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <stdexcept>
#include "kittens.cuh"

using namespace kittens;

static constexpr int BM = 128;
static constexpr int BN = 128;
static constexpr int BK = 64;
static constexpr int NUM_WARPS = 4;
static constexpr int NUM_THREADS = NUM_WARPS * WARP_THREADS;

// bf16 3-split ("fp32 emulation"): each fp32 operand is split into a hi and lo
// bf16, and the product is accumulated as hi*hi + hi*lo + lo*hi in fp32. bf16
// shares fp32's exponent range, so this is robust to row scaling.
using a_bf_t = st_bf<BM, BK>;     // 128 x 64
using b_bf_t = st_bf<BK, BN>;     // 64 x 128
using d_fl_t = st_fl<BM, BN>;     // 128 x 128

// Register tiles distribute the leading dim across the 4-warp group.
using a_rf_t = rt_fl<BM / NUM_WARPS, BK>;   // 32 x 64
using a_rb_t = rt_bf<BM / NUM_WARPS, BK>;
using b_rf_t = rt_fl<BK / NUM_WARPS, BN>;   // 16 x 128
using b_rb_t = rt_bf<BK / NUM_WARPS, BN>;
using d_rf_t = rt_fl<BM / NUM_WARPS, BN>;   // 32 x 128

// 4D global layouts: (batch, 1, rows, cols). No TMA tile type -- we move data
// with warpgroup::load/store (global<->register), not tma::load_async.
using a_gl = gl<float, -1, 1, -1, -1>;
using b_gl = gl<float, -1, 1, -1, -1>;
using d_gl = gl<float, -1, 1, -1, -1>;
using acc_tt = tt<float, BM, BN>;

__global__ __launch_bounds__(NUM_THREADS, 1)
void tk_gemm_kernel(const __grid_constant__ a_gl A,
                    const __grid_constant__ b_gl B,
                    const __grid_constant__ d_gl D,
                    int num_k,
                    int a_ro, int a_co, int b_ro, int b_co,
                    int d_ro, int d_co, int subtract) {
    const int bz = blockIdx.z;   // batch
    const int mt = blockIdx.y;   // output tile row
    const int nt = blockIdx.x;   // output tile col
    const int wg_lane = warpgroup::laneid();

    extern __shared__ int __shm[];
    tma_swizzle_allocator al((int*)&__shm[0]);
    a_bf_t (&a_hi) = al.allocate<a_bf_t>();
    a_bf_t (&a_lo) = al.allocate<a_bf_t>();
    b_bf_t (&b_hi) = al.allocate<b_bf_t>();
    b_bf_t (&b_lo) = al.allocate<b_bf_t>();
    d_fl_t (&d_smem) = al.allocate<d_fl_t>();

    __shared__ semaphore inputs_finished, scratch_sem, compute_done;
    if (threadIdx.x == 0) {
        init_semaphore(inputs_finished, 1, 0);   // pre-arrived; signalled by last MMA
        init_semaphore(scratch_sem, 1, 0);       // soaks up the non-final MMAs
        init_semaphore(compute_done, 0, 1);
    }
    __syncthreads();

    tensor_allocator<1, 1> tm_alloc{};
    acc_tt accum;
    if (wg_lane == 0) accum = tm_alloc.allocate<acc_tt>(0);
    warpgroup::sync(1);

    int phase = 0;
    for (int kt = 0; kt < num_k; ++kt) {
        // Wait until the previous tile's MMAs have consumed the shared operands.
        if (threadIdx.x == 0) wait(inputs_finished, phase ^ 1);
        warpgroup::sync(1);
        phase ^= 1;

        // global(fp32) -> register -> {hi,lo} bf16 -> shared. Scoped so the
        // fp32/bf16 register tiles for A are freed before B's are allocated.
        {
            a_rf_t f; warpgroup::load(f, A, {bz, 0, a_ro + mt, a_co + kt});
            a_rb_t hb; warp::copy(hb, f); warpgroup::store(a_hi, hb);
            a_rf_t t; warp::copy(t, hb); warp::sub(f, f, t);
            warp::copy(hb, f); warpgroup::store(a_lo, hb);
        }
        {
            b_rf_t f; warpgroup::load(f, B, {bz, 0, b_ro + kt, b_co + nt});
            b_rb_t hb; warp::copy(hb, f); warpgroup::store(b_hi, hb);
            b_rf_t t; warp::copy(t, hb); warp::sub(f, f, t);
            warp::copy(hb, f); warpgroup::store(b_lo, hb);
        }
        warpgroup::sync(1);

        if (wg_lane == 0) {
            if (kt == 0) mm_AB (accum, a_hi, b_hi, scratch_sem);
            else         mma_AB(accum, a_hi, b_hi, scratch_sem);
            mma_AB(accum, a_hi, b_lo, scratch_sem);
            mma_AB(accum, a_lo, b_hi, inputs_finished);   // last: signals reload-ok
        }
    }

    if (wg_lane == 0) kittens::detail::tcgen05::commit<1>(compute_done);
    wait(compute_done, 0);

    d_rf_t d_rf;
    warpgroup::load_async(d_rf, accum);
    tensor_load_wait();
    warpgroup::sync(1);
    if (subtract) {
        // In-place: D := D_existing - A @ B  (read the current D tile, subtract).
        d_rf_t cur;
        warpgroup::load(cur, D, {bz, 0, d_ro + mt, d_co + nt});
        warp::sub(d_rf, cur, d_rf);
    }
    warpgroup::store(d_smem, d_rf);
    warpgroup::sync(1);
    warpgroup::store(D, d_smem, {bz, 0, d_ro + mt, d_co + nt});
}

void tk_gemm(torch::Tensor A, long ar, long ac, long a_ro, long a_co,
             torch::Tensor B, long br, long bc, long b_ro, long b_co,
             torch::Tensor D, long dr, long dc, long d_ro, long d_co,
             long batch, long Mt, long Nt, long Kt, long subtract) {
    // The (rows, cols) are the *logical* per-call dims and set the gl strides
    // (col = row stride, rows*cols = batch stride). Scratch buffers are used as
    // per-panel compacted pools, so these must be passed explicitly, NOT read
    // from the (over-allocated) tensor shape.
    auto mkgl = [](torch::Tensor& T, long r, long c) {
        return gl<float, -1, 1, -1, -1>{
            reinterpret_cast<float*>(T.data_ptr()), (unsigned long)T.size(0),
            nullptr, (unsigned long)r, (unsigned long)c};
    };
    a_gl Agl = mkgl(A, ar, ac);
    b_gl Bgl = mkgl(B, br, bc);
    d_gl Dgl = mkgl(D, dr, dc);

    dim3 grid((unsigned)Nt, (unsigned)Mt, (unsigned)batch);
    int smem = MAX_SHARED_MEMORY - 1024;
    static bool attr_set = false;
    if (!attr_set) {
        cudaFuncSetAttribute(tk_gemm_kernel,
                             cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
        attr_set = true;
    }
    tk_gemm_kernel<<<grid, NUM_THREADS, smem>>>(
        Agl, Bgl, Dgl, (int)Kt,
        (int)a_ro, (int)a_co, (int)b_ro, (int)b_co,
        (int)d_ro, (int)d_co, (int)subtract);
    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) throw std::runtime_error(cudaGetErrorString(err));
}
"""

_tk_mod = None


def _get_tk():
    global _tk_mod
    if _tk_mod is None:
        tk_root = _thunderkittens_root()
        _tk_mod = load_inline(
            name="qr_tk_gemm_v2",
            cpp_sources=[_TK_CPP],
            cuda_sources=[_TK_CUDA],
            functions=["tk_gemm"],
            verbose=False,
            extra_include_paths=[
                os.path.join(tk_root, "include"),
                os.path.join(tk_root, "prototype"),
            ],
            extra_cuda_cflags=[
                "-std=c++20", "-O3", "--use_fast_math",
                "--expt-extended-lambda", "--expt-relaxed-constexpr",
                "-forward-unknown-to-host-compiler",
                "-Xcompiler=-Wno-psabi", "-Xcompiler=-fno-strict-aliasing",
                "-DKITTENS_SM100", "-DNDEBUG", "-lineinfo",
                "-ftemplate-backtrace-limit=0",
                "-gencode=arch=compute_100a,code=sm_100a",
            ],
            extra_ldflags=["-lcuda"],
        )
    return _tk_mod


# ---------------------------------------------------------------------------
# Plain fp32 CUDA-core kernels: the proven v1 path (fallback / small n) plus the
# blocked-QR helpers (panel factor, V/V^T build, compact-WY T, pad copy, axpy).
# Compiled with torch.cuda._compile_kernel (no ThunderKittens needed).
# ---------------------------------------------------------------------------
_CUDA_SRC = r"""
__device__ __forceinline__ float warp_reduce_sum(float v) {
    v += __shfl_down_sync(0xffffffff, v, 16);
    v += __shfl_down_sync(0xffffffff, v, 8);
    v += __shfl_down_sync(0xffffffff, v, 4);
    v += __shfl_down_sync(0xffffffff, v, 2);
    v += __shfl_down_sync(0xffffffff, v, 1);
    return v;
}

__device__ __forceinline__ float block_reduce_sum(float v, float* scratch) {
    int lane = threadIdx.x & 31;
    int warp = threadIdx.x >> 5;
    int num_warps = (blockDim.x + 31) >> 5;
    v = warp_reduce_sum(v);
    if (lane == 0) scratch[warp] = v;
    __syncthreads();
    float total = 0.0f;
    if (warp == 0) {
        total = (lane < num_warps) ? scratch[lane] : 0.0f;
        total = warp_reduce_sum(total);
        if (lane == 0) scratch[0] = total;
    }
    __syncthreads();
    return scratch[0];
}

// ---- blocked-QR helpers ---------------------------------------------------
// Grid-wide barrier across the `workers` blocks cooperating on one matrix
// (manual counters/sense spin, adapted from v1.py). Requires all `workers`
// blocks of a matrix to be co-resident (the launcher caps batch*workers).
__device__ __forceinline__ void persistent_barrier(int* counters, int* sense,
                                                    int matrix, int workers,
                                                    int& phase) {
    __syncthreads();
    if (threadIdx.x == 0) {
        __threadfence();
        int ticket = atomicAdd(counters + matrix, 1);
        if (ticket == workers - 1) {
            counters[matrix] = 0;
            __threadfence();
            atomicAdd(sense + matrix, 1);
        } else {
            while (atomicAdd(sense + matrix, 0) == phase) {}
        }
    }
    ++phase;
    __syncthreads();
}

// Cooperative panel factorization: `workers` blocks per matrix split the row
// reductions (column norms) and own disjoint panel columns for the reflector
// apply (split-Q style, no cross-worker reduction there). Factors columns
// [c0, c0+pw) over rows [c0, n) in place. workers==1 => single-block (degenerate
// barrier), identical to the old one-block path. counters/sense must be zeroed
// before launch.
extern "C" __global__
void panel_factor_coop(float* A, float* tau, float* partial,
                       int* counters, int* sense,
                       int n, int c0, int pw, int workers) {
    __shared__ float red[256];
    __shared__ float sa[4096];   // staged reflector column a[:,k] (max n=4096)
    int worker = blockIdx.x % workers;
    int b = blockIdx.x / workers;
    int tid = threadIdx.x;
    float* a = A + (long)b * n * n;
    float* p = partial + (long)b * workers;
    int stride = workers * blockDim.x;
    int phase = 0;

    for (int k = c0; k < c0 + pw; ++k) {
        // Column norm: rows split across all workers.
        float sum = 0.0f;
        for (int r = k + 1 + worker * blockDim.x + tid; r < n; r += stride) {
            float xi = a[r * n + k];
            sum += xi * xi;
        }
        float local = block_reduce_sum(sum, red);
        if (tid == 0) p[worker] = local;
        persistent_barrier(counters, sense, b, workers, phase);

        if (worker == 0) {
            float total = 0.0f;
            for (int i = tid; i < workers; i += blockDim.x) total += p[i];
            total = block_reduce_sum(total, red);
            if (tid == 0) {
                float alpha = a[k * n + k];
                float xnorm = sqrtf(total);
                float tau_k = 0.0f, inv = 0.0f;
                if (xnorm != 0.0f) {
                    float norm = sqrtf(alpha * alpha + xnorm * xnorm);
                    float beta = (alpha >= 0.0f) ? -norm : norm;
                    tau_k = (beta - alpha) / beta;
                    inv = 1.0f / (alpha - beta);
                    a[k * n + k] = beta;
                }
                tau[(long)b * n + k] = tau_k;
                p[0] = inv;
            }
        }
        persistent_barrier(counters, sense, b, workers, phase);

        float inv_scale = p[0];
        float tau_k = tau[(long)b * n + k];
        if (tau_k != 0.0f)
            for (int r = k + 1 + worker * blockDim.x + tid; r < n; r += stride)
                a[r * n + k] *= inv_scale;
        persistent_barrier(counters, sense, b, workers, phase);

        // Stage the (now scaled) reflector column into shared once, so the apply
        // loop reads it from SMEM instead of re-reading global ~pw times.
        for (int r = k + 1 + tid; r < n; r += blockDim.x) sa[r - (k + 1)] = a[r * n + k];
        __syncthreads();

        // Apply reflector to remaining panel columns; each worker owns a disjoint
        // set of columns and reduces over all rows within its own block.
        for (int j = k + 1 + worker; j < c0 + pw; j += workers) {
            float contrib = (tid == 0) ? a[k * n + j] : 0.0f;
            for (int r = k + 1 + tid; r < n; r += blockDim.x)
                contrib += sa[r - (k + 1)] * a[r * n + j];
            float dot = block_reduce_sum(contrib, red);
            float update = tau_k * dot;
            if (tid == 0) a[k * n + j] -= update;
            for (int r = k + 1 + tid; r < n; r += blockDim.x)
                a[r * n + j] -= sa[r - (k + 1)] * update;
        }
        persistent_barrier(counters, sense, b, workers, phase);
    }
}

// Materialize the explicit reflector matrix V (m x pw, unit-lower-trapezoidal)
// and its transpose V^T, each zero-padded. One block per matrix.
//   Vbuf : [batch, Mp, 128]   VTbuf : [batch, 128, Kp]   (m = n - c0)
extern "C" __global__
void build_V(const float* A, float* Vbuf, float* VTbuf,
             int n, int c0, int pw, int Mp, int Kp) {
    int b = blockIdx.x;
    const float* a = A + (long)b * n * n;
    float* V = Vbuf + (long)b * Mp * 128;
    float* VT = VTbuf + (long)b * 128 * Kp;
    int m = n - c0;
    long total = (long)Mp * 128;
    for (long idx = (long)threadIdx.x + (long)blockIdx.y * blockDim.x;
         idx < total; idx += (long)blockDim.x * gridDim.y) {
        int i = idx / 128;       // 0..Mp-1
        int j = idx % 128;       // 0..127
        float val = 0.0f;
        if (i < m && j < pw) {
            if (i == j) val = 1.0f;
            else if (i > j) val = a[(long)(c0 + i) * n + (c0 + j)];
        }
        V[(long)i * 128 + j] = val;
        if (i < Kp) VT[(long)j * Kp + i] = val;   // VT is 128 x Kp
    }
}

// Compact-WY T factor (pw x pw, upper triangular) and its transpose TT,
// both zero-padded to 128 x 128. One block per matrix; pw <= 128 threads active.
// T is held in global scratch (Tbuf) to avoid a >48KB static shared array.
//   Tbuf, TTbuf : [batch, 128, 128]
extern "C" __global__
void build_T(const float* Vbuf, const float* tau, float* Tbuf, float* TTbuf,
             int n, int c0, int pw, int Mp) {
    __shared__ float w[128];
    int b = blockIdx.x;
    const float* V = Vbuf + (long)b * Mp * 128;
    const float* tb = tau + (long)b * n;
    float* T = Tbuf + (long)b * 128 * 128;
    int p = threadIdx.x;

    for (int idx = p; idx < 128 * 128; idx += blockDim.x) T[idx] = 0.0f;
    __syncthreads();
    if (p < pw) T[p * 128 + p] = tb[c0 + p];
    __syncthreads();

    for (int j = 1; j < pw; ++j) {
        float tau_j = tb[c0 + j];
        // w[p] = V(:,p)^T V(:,j) for p < j
        if (p < j) {
            float s = 0.0f;
            for (int i = 0; i < Mp; ++i)
                s += V[(long)i * 128 + p] * V[(long)i * 128 + j];
            w[p] = s;
        }
        __syncthreads();
        // T(0:j, j) = -tau_j * ( T(0:j,0:j) @ w )   (T upper triangular)
        if (p < j) {
            float s = 0.0f;
            for (int q = p; q < j; ++q) s += T[p * 128 + q] * w[q];
            T[p * 128 + j] = -tau_j * s;
        }
        __syncthreads();
    }

    // Write TT = T^T, zero-padded, into global.
    float* TT = TTbuf + (long)b * 128 * 128;
    for (int idx = p; idx < 128 * 128; idx += blockDim.x) {
        int r = idx / 128, c = idx % 128;
        TT[idx] = T[c * 128 + r];
    }
}

// Same compact-WY T factor, but the V^T V Gram matrix (the O(Mp) reduction) is
// precomputed on the tensor cores; this kernel only runs the small O(pw^2)
// triangular recurrence. Gram : [batch, 128, 128] with Gram[p*128+j] = V(:,p).V(:,j).
extern "C" __global__
void build_T_gram(const float* Gram, const float* tau, float* Tbuf, float* TTbuf,
                  int n, int c0, int pw) {
    int b = blockIdx.x;
    const float* G = Gram + (long)b * 128 * 128;
    const float* tb = tau + (long)b * n;
    float* T = Tbuf + (long)b * 128 * 128;
    int p = threadIdx.x;

    for (int idx = p; idx < 128 * 128; idx += blockDim.x) T[idx] = 0.0f;
    __syncthreads();
    if (p < pw) T[p * 128 + p] = tb[c0 + p];
    __syncthreads();

    for (int j = 1; j < pw; ++j) {
        float tau_j = tb[c0 + j];
        if (p < j) {
            float s = 0.0f;
            for (int q = p; q < j; ++q) s += T[p * 128 + q] * G[q * 128 + j];
            T[p * 128 + j] = -tau_j * s;
        }
        __syncthreads();
    }

    float* TT = TTbuf + (long)b * 128 * 128;
    for (int idx = p; idx < 128 * 128; idx += blockDim.x) {
        int r = idx / 128, c = idx % 128;
        TT[idx] = T[c * 128 + r];
    }
}

// Copy trailing block C = A[c0:n, c0+pw:n] (m x t) into a zero-padded buffer.
//   Cbuf : [batch, Kp, Np]
extern "C" __global__
void copy_trailing(const float* A, float* Cbuf,
                   int n, int c0, int pw, int Kp, int Np) {
    int b = blockIdx.z;
    const float* a = A + (long)b * n * n;
    float* C = Cbuf + (long)b * Kp * Np;
    int m = n - c0;
    int t = n - c0 - pw;
    int i = blockIdx.y * blockDim.y + threadIdx.y;
    int j = blockIdx.x * blockDim.x + threadIdx.x;
    if (i >= Kp || j >= Np) return;
    float val = (i < m && j < t) ? a[(long)(c0 + i) * n + (c0 + pw + j)] : 0.0f;
    C[(long)i * Np + j] = val;
}

// A[c0+i, c0+pw+j] -= Cup[i, j]   for the real (unpadded) trailing region.
//   Cup : [batch, Mp, Np]
extern "C" __global__
void axpy_sub(float* A, const float* Cup,
              int n, int c0, int pw, int Mp, int Np) {
    int b = blockIdx.z;
    float* a = A + (long)b * n * n;
    const float* C = Cup + (long)b * Mp * Np;
    int m = n - c0;
    int t = n - c0 - pw;
    int i = blockIdx.y * blockDim.y + threadIdx.y;
    int j = blockIdx.x * blockDim.x + threadIdx.x;
    if (i >= m || j >= t) return;
    a[(long)(c0 + i) * n + (c0 + pw + j)] -= C[(long)i * Np + j];
}

// Generalized column-range versions for the recursive (within-panel) update:
// copy A[r0:n, lo:hi] into a zero-padded Cbuf; subtract Cup back into A[r0:n, lo:hi].
extern "C" __global__
void copy_block(const float* A, float* Cbuf,
                int n, int r0, int lo, int hi, int Kp, int Np) {
    int b = blockIdx.z;
    const float* a = A + (long)b * n * n;
    float* C = Cbuf + (long)b * Kp * Np;
    int m = n - r0;
    int t = hi - lo;
    int i = blockIdx.y * blockDim.y + threadIdx.y;
    int j = blockIdx.x * blockDim.x + threadIdx.x;
    if (i >= Kp || j >= Np) return;
    float val = (i < m && j < t) ? a[(long)(r0 + i) * n + (lo + j)] : 0.0f;
    C[(long)i * Np + j] = val;
}

extern "C" __global__
void axpy_block(float* A, const float* Cup,
                int n, int r0, int lo, int hi, int Mp, int Np) {
    int b = blockIdx.z;
    float* a = A + (long)b * n * n;
    const float* C = Cup + (long)b * Mp * Np;
    int m = n - r0;
    int t = hi - lo;
    int i = blockIdx.y * blockDim.y + threadIdx.y;
    int j = blockIdx.x * blockDim.x + threadIdx.x;
    if (i >= m || j >= t) return;
    a[(long)(r0 + i) * n + (lo + j)] -= C[(long)i * Np + j];
}
"""

_k = {}


def _kernels():
    if not _k:
        c = torch.cuda._compile_kernel
        for name in [
            "panel_factor_coop", "build_V", "build_T_gram", "copy_trailing",
            "axpy_sub", "copy_block", "axpy_block",
        ]:
            _k[name] = c(_CUDA_SRC, name)
    return _k


def _round_up(x, m):
    return (x + m - 1) // m * m


def _fallback_qr(data):
    """Baseline: torch's LAPACK geqrf -> (R+reflectors, tau), the expected format."""
    a, tau = torch.geqrf(data)
    return a.contiguous(), tau.contiguous()


# Scratch workspace cached per (batch, n, device). cudaMalloc of the multi-GB
# trailing-update scratch is slow and synchronous, so we allocate once per shape
# and reuse across calls. Every buffer is fully overwritten each panel, so reuse
# is safe. Only the returned `out`/`tau` are freshly allocated per call.
_ws = {}


def _workspace(batch, n, dev):
    key = (batch, n, str(dev))
    ws = _ws.get(key)
    if ws is None:
        Mp = _round_up(n, 128)
        Kp = _round_up(n, 64)
        Np = _round_up(n, 128)
        f = lambda *s: torch.empty(s, device=dev, dtype=torch.float32)
        ws = dict(
            V=f(batch, Mp, 128), VT=f(batch, 128, Kp),
            T=f(batch, 128, 128), TT=f(batch, 128, 128), G=f(batch, 128, 128),
            W=f(batch, 128, Np), W2=f(batch, 128, Np),
            # within-panel (recursive BLAS-3) update scratch: cols <= 128 wide
            Cb=f(batch, Kp, 128), Cupb=f(batch, Mp, 128),
            # cooperative-panel scratch (multi-block-per-matrix barrier)
            partial=f(batch, _TARGET_BLOCKS),
            counters=torch.zeros((batch,), device=dev, dtype=torch.int32),
            sense=torch.zeros((batch,), device=dev, dtype=torch.int32),
        )
        if n % 128:   # non-aligned: the scratch (copy_trailing / axpy) path
            ws["C"] = f(batch, Kp, Np)
            ws["Cup"] = f(batch, Mp, Np)
        _ws[key] = ws
    return ws


# Total cooperating blocks per launch is capped so every worker of a matrix is
# co-resident (the manual barrier spin-waits and would otherwise deadlock).
_TARGET_BLOCKS = 128

# Inner sub-block width for the recursive BLAS-3 panel (within-panel updates on
# the tensor cores). _PANEL_IB == _BLOCK disables recursion (one BLAS-2 panel).
_PANEL_IB = 32


def _panel_workers(batch, m):
    """Blocks per matrix for the cooperative panel: parallelize low-batch/tall
    panels, degenerate to 1 (single block) when there are already enough matrices.
    Capped so batch*workers <= _TARGET_BLOCKS (co-residency) and each worker still
    owns >= ~32 rows of the reduction."""
    return max(1, min(_TARGET_BLOCKS // batch, max(1, m // 32)))


def _blocked_qr_tk(data):
    """Blocked Householder QR with tcgen05 (bf16 3-split) trailing updates.

    For n % 128 == 0 the trailing block is tile-aligned, so the update reads C
    straight from the strided trailing submatrix of `out` and subtracts V W back
    in place (no copy_trailing / axpy / C / Cup scratch). Other n use the padded
    scratch path.
    """
    batch, n, _ = data.shape
    dev = data.device
    k = _kernels()
    tk = _get_tk()
    aligned = (n % 128 == 0)

    out = data.clone()
    tau = torch.zeros((batch, n), device=dev, dtype=torch.float32)

    ws = _workspace(batch, n, dev)
    Vbuf, VTbuf = ws["V"], ws["VT"]
    Tbuf, TTbuf, Gram = ws["T"], ws["TT"], ws["G"]
    Wbuf, W2buf = ws["W"], ws["W2"]
    Cb, Cupb = ws["Cb"], ws["Cupb"]
    partial, counters, sense = ws["partial"], ws["counters"], ws["sense"]

    def panel_coop(c0, pw):
        m = n - c0
        w = _panel_workers(batch, m)
        if w > 1:
            counters.zero_()
            sense.zero_()
        k["panel_factor_coop"](grid=(batch * w, 1, 1), block=(256, 1, 1),
                               args=[out, tau, partial, counters, sense, n, c0, pw, w])

    def within_panel(ic, ipw, lo, hi):
        # Apply sub-panel [ic, ic+ipw) reflectors to panel cols [lo, hi) on the
        # tensor cores (BLAS-3), so the panel is mostly GEMMs not rank-1 updates.
        m = n - ic
        Mp, Kp, Np = _round_up(m, 128), _round_up(m, 64), _round_up(hi - lo, 128)
        Mt, Nt, Kt = Mp // 128, Np // 128, Kp // 64
        k["build_V"](grid=(batch, 32, 1), block=(256, 1, 1),
                     args=[out, Vbuf, VTbuf, n, ic, ipw, Mp, Kp])
        tk.tk_gemm(VTbuf, 128, Kp, 0, 0, Vbuf, Mp, 128, 0, 0,
                   Gram, 128, 128, 0, 0, batch, 1, 1, Kp // 64, 0)
        k["build_T_gram"](grid=(batch, 1, 1), block=(128, 1, 1),
                          args=[Gram, tau, Tbuf, TTbuf, n, ic, ipw])
        bx, by = (Np + 15) // 16, (Kp + 15) // 16
        k["copy_block"](grid=(bx, by, batch), block=(16, 16, 1),
                        args=[out, Cb, n, ic, lo, hi, Kp, Np])
        tk.tk_gemm(VTbuf, 128, Kp, 0, 0, Cb, Kp, Np, 0, 0,
                   Wbuf, 128, Np, 0, 0, batch, 1, Nt, Kt, 0)
        tk.tk_gemm(TTbuf, 128, 128, 0, 0, Wbuf, 128, Np, 0, 0,
                   W2buf, 128, Np, 0, 0, batch, 1, Nt, 2, 0)
        tk.tk_gemm(Vbuf, Mp, 128, 0, 0, W2buf, 128, Np, 0, 0,
                   Cupb, Mp, Np, 0, 0, batch, Mt, Nt, 2, 0)
        bx, by = (hi - lo + 15) // 16, (m + 15) // 16
        k["axpy_block"](grid=(bx, by, batch), block=(16, 16, 1),
                        args=[out, Cupb, n, ic, lo, hi, Mp, Np])

    def factor_panel(c0, pw):
        # Recursive BLAS-3 panel: factor inner sub-blocks of width _PANEL_IB and
        # apply each to the rest of the panel via within_panel (tensor cores).
        for ic in range(c0, c0 + pw, _PANEL_IB):
            ipw = min(_PANEL_IB, c0 + pw - ic)
            panel_coop(ic, ipw)
            rem_lo = ic + ipw
            if rem_lo < c0 + pw:
                within_panel(ic, ipw, rem_lo, c0 + pw)

    for c0 in range(0, n, _BLOCK):
        pw = min(_BLOCK, n - c0)
        m = n - c0
        t = n - c0 - pw

        factor_panel(c0, pw)
        if t <= 0:
            continue

        Mp = _round_up(m, 128)
        Kp = _round_up(m, 64)
        Np = _round_up(t, 128)

        k["build_V"](grid=(batch, 32, 1), block=(256, 1, 1),
                     args=[out, Vbuf, VTbuf, n, c0, pw, Mp, Kp])
        # Gram = V^T V on the tensor cores (the O(Mp) reduction); then the small
        # O(pw^2) compact-WY recurrence. V is unit-triangular (|v|<=1), so this is
        # the LARFT inner-product, NOT a normal-equations data Gram.
        Gram = ws["G"]
        tk.tk_gemm(VTbuf, 128, Kp, 0, 0, Vbuf, Mp, 128, 0, 0,
                   Gram, 128, 128, 0, 0, batch, 1, 1, Kp // 64, 0)
        k["build_T_gram"](grid=(batch, 1, 1), block=(128, 1, 1),
                          args=[Gram, tau, Tbuf, TTbuf, n, c0, pw])

        Mt, Nt, Kt = Mp // 128, Np // 128, Kp // 64
        # tk_gemm(A, ar, ac, a_ro, a_co, B, br, bc, b_ro, b_co,
        #         D, dr, dc, d_ro, d_co, batch, Mt, Nt, Kt, subtract).
        # rows/cols are the per-panel *logical* dims (compacted buffer layout).
        if aligned:
            # Stage 1: W0 = V^T C, reading C from out's trailing block in place.
            tk.tk_gemm(VTbuf, 128, Kp, 0, 0, out, n, n, c0 // 64, (c0 + pw) // 128,
                       Wbuf, 128, Np, 0, 0, batch, 1, Nt, Kt, 0)
            # W = T^T W0  (small).
            tk.tk_gemm(TTbuf, 128, 128, 0, 0, Wbuf, 128, Np, 0, 0,
                       W2buf, 128, Np, 0, 0, batch, 1, Nt, 2, 0)
            # Stage 3: out_trailing -= V W, written in place.
            tk.tk_gemm(Vbuf, Mp, 128, 0, 0, W2buf, 128, Np, 0, 0,
                       out, n, n, c0 // 128, (c0 + pw) // 128, batch, Mt, Nt, 2, 1)
        else:
            Cbuf, Cupbuf = ws["C"], ws["Cup"]
            bx, by = (Np + 15) // 16, (Kp + 15) // 16
            k["copy_trailing"](grid=(bx, by, batch), block=(16, 16, 1),
                               args=[out, Cbuf, n, c0, pw, Kp, Np])
            tk.tk_gemm(VTbuf, 128, Kp, 0, 0, Cbuf, Kp, Np, 0, 0,
                       Wbuf, 128, Np, 0, 0, batch, 1, Nt, Kt, 0)
            tk.tk_gemm(TTbuf, 128, 128, 0, 0, Wbuf, 128, Np, 0, 0,
                       W2buf, 128, Np, 0, 0, batch, 1, Nt, 2, 0)
            tk.tk_gemm(Vbuf, Mp, 128, 0, 0, W2buf, 128, Np, 0, 0,
                       Cupbuf, Mp, Np, 0, 0, batch, Mt, Nt, 2, 0)
            bx, by = (t + 15) // 16, (m + 15) // 16
            k["axpy_sub"](grid=(bx, by, batch), block=(16, 16, 1),
                          args=[out, Cupbuf, n, c0, pw, Mp, Np])

    return out, tau


def custom_kernel(data: input_t) -> output_t:
    if not (
        data.is_cuda
        and data.dtype == torch.float32
        and data.is_contiguous()
        and data.ndim == 3
        and data.shape[-1] == data.shape[-2]
    ):
        raise RuntimeError("unsupported input for custom QR kernel")

    batch, n, _ = data.shape

    if _force_fallback() or n < _TK_MIN_N:
        return _fallback_qr(data)

    return _blocked_qr_tk(data)
scrolls · 786 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON