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
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 = 4
static constexpr int NUM_WARPS = 4;persistent-kernel
__device__ __forceinline__ void persistent_barrier(int* counters, int* sense,shared-memory
extern __shared__ int __shm[];tile-k = 64
static constexpr int BK = 64;tile-m = 128
static constexpr int BM = 128;tile-n = 128
static 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