submission 833772
Lorenzo · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 748 lines, June 9 Researcher Reciprocity License v1.0.
submission_41.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-833772?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:357f509365aadafa45a004087fc7b1b2d23e7caee8b6ad22b841d0ce5ea648ba
license declaredunknown
license concludedunknown
authorsLorenzo
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float smem[];Kernel source
submission_41.py748 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
"""Batched compact-Householder QR, GPU Mode `qr_v2` (B200).
s31 = s30 + SMALLER-PANEL HIGH-OCCUPANCY bet for n=512 (the real footprint lever).
The panel data R*pb MUST be resident somewhere, so register-residency (B3) can't
break the occ~3 smem ceiling (smem<->reg tradeoff is 1:1) AND it cuts blocks/SM
(s8/s18/s17 say that regresses at b=640). The actual lever is making the panel
SMALLER: pb=16 -> smem 512*17*4=34.8KB -> occ 6 (vs pb=32 67KB occ 3). The s19
probe measured the panel phase pb16=4.78 vs pb32=6.83 ms (-30%) at threads=128.
The panel is 68% of n=512 runtime, so -30% there is the big swing; the only risk
is the trailing (2x the panels, K=16 vs 32 in A-=VY) eating it -- never tested
end-to-end with the cuBLAS pipeline (only the old fused s12). With s30's freed
registers, pb16+threads=128 should reach occ ~5. n=512 dispatch -> pb=16,
threads=128 (occ play needs low threads; threads=256 would be reg-capped to occ 2).
Other n keep pb=32 (1024 is b=60<SMs = latency-bound per block, not occ-bound;
smaller pb only adds skinnier panels there).
s30 = s27 + PANEL REGISTER-RESIDENCY lever (ideas Bet A / B1). The s19 probe found
the panel at threads=256 is REGISTER-bound to occ=2 (smem would allow occ=3), so
s20 dropped n=512 to threads=128 to reach occ=3 -- at the cost of half the threads.
The `accd[PBMAX]` apply-reduction accumulator was FP64 = 64 regs/thread, the
dominant register consumer. Each thread accumulates only ~R/nthreads (~2-4) terms
and the cross-warp combine stays FP64 (sWcol stays double), so dropping the
PER-THREAD accumulator to FP32 (32 regs) is ~accuracy-neutral but frees ~32
regs/thread. Hypothesis: that lets n=512 run threads=256 at occ=3 -- 2x the threads
of the threads=128 occ-3 config -> more latency hiding at the SAME occupancy.
Change: accd[]/vr -> float + warpReduceSumF (store to double sWcol); n=512 dispatch
-> threads=256, la=0. All other n unchanged. ISOLATES the register lever.
s20 = s11 + PANEL OCCUPANCY FIX. The s19 probe overturned s17: the panel is
occupancy-bound with a steep slope, and s11's threads=256 is REGISTER-bound to
occ=2 (124 regs/thread). Dropping to threads=128 lifts occ to 3 and cuts the
n=512 panel phase 9.0->6.85 ms (1.32x), same for n=1024. Requires making the two
cross-warp reduction loops use the ACTUAL warp count (nwarps), not hard-coded
WARPS=8 (else threads<256 reads stale sWarpD/sWcol slots). pb UNCHANGED (isolates
the threads lever). n=2048 stays threads=256 (R/128=16 > VREG_MAX=8).
s10 = s7's barrier-bound panel kernel (unchanged) + the trailing block-WY
update moved onto DIRECT cuBLAS strided-batched TF32 tensor-core GEMMs. s7's
fused FP32 larfb ran the O(n^3) trailing update on CUDA cores (~2.6 TFLOPS,
~3% of FP32 peak); the headline n=512 case spent ~22 ms there. cuBLAS TF32
tensor cores are ~2200 TFLOPS, so the trailing FLOPs are nearly free; the only
prior loss (s5) was torch's per-panel V materialization (clone/tril/mask) +
limb splits + many launches. Here the panel kernel emits a packed, GEMM-ready
V (unit diagonal, zeroed strict-upper, zeroed identity-reflector columns) so
cuBLAS consumes it with zero torch ops.
Trailing update per panel (cur=pb, R=n-k, M trailing cols):
W = V^T @ Atrail (cur x M, K=R) -> strided-batched tensor core
Y = Tg^T @ W (cur x M, K=cur) -> small, FP32
Atrail -= V @ Y (R x M, K=cur) -> strided-batched tensor core
Precision toggle `_CUBLAS_PREC`: "fp32" | "tf32" | "tf32x3". tf32x3 splits each
fp32 operand into a TF32-exact hi limb + residual lo limb (3 GEMMs:
AhBh+AhBl+AlBh, ~1e-5 rel err) entirely via a custom split kernel, well inside
the factor gate (~1.2e-3 at n=512).
Inherited from s7: per-n `_DISPATCH` (torch_trsm for small n, torch_tfac+bf16x3
for n=2048, geqrf for n=4096). Build/accuracy failure degrades to geqrf.
"""
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# FP32-safe globally: keep torch matmuls in true FP32 (no TF32). The bf16x3
# torch trailing path toggles allow_tf32 locally and restores it.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
_ENABLE_CUSTOM = os.environ.get("QR_ENABLE_CUSTOM", "1") == "1"
_IMPL_OVERRIDE = os.environ.get("QR_TRAILING_IMPL", "")
# Trailing GEMM precision for the cuBLAS path. "tf32x3" is the safe default
# (~1e-5 rel err); "tf32" is the fastest (1 GEMM) but may miss the factor gate
# on ill-conditioned cases; "fp32" is the safety net. Env override for A/B.
# Fallback only; per-n precision is set explicitly in _DISPATCH below.
_CUBLAS_PREC = os.environ.get("QR_CUBLAS_PREC", "tf32x3")
_PREC_CODE = {"fp32": 0, "tf32": 1, "tf32x3": 2, "mixed": 3}
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cublas_v2.h>
#include <cmath>
#define VREG_MAX 8
#define WARPS 8
#define PBMAX 32
__device__ __forceinline__ double warpReduceSumD(double val) {
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
int hi = __shfl_down_sync(0xffffffffu, __double2hiint(val), o);
int lo = __shfl_down_sync(0xffffffffu, __double2loint(val), o);
val += __hiloint2double(hi, lo);
}
return val;
}
__device__ __forceinline__ float warpReduceSumF(float val) {
#pragma unroll
for (int o = 16; o > 0; o >>= 1)
val += __shfl_down_sync(0xffffffffu, val, o);
return val;
}
#define LP 32
#define LNT 64
#define LRT 32
#define LTPB 256
// One CUDA block per matrix. Factors the panel H[:, k:n, k:k+pb] in place
// (unblocked geqr2 over pb sequential reflectors), writes tau[:, k:k+pb], the
// cur x cur compact-WY T-factor (Tg) when build_t, and (when write_v) a packed
// GEMM-ready V buffer (batch, R, pb): unit diagonal, zeroed strict-upper,
// zeroed identity-reflector columns -- ready for cuBLAS with no torch ops.
// LA (compile-time): when true, each reflector's sub-diagonal norm is folded
// into the PREVIOUS reflector's apply pass (look-ahead) so the per-reflector
// norm reduction + its __syncthreads are skipped. This helps LOW-occupancy
// launches (threads=256, occ=1: barriers exposed) but HURTS the occ=3 n=512
// family (barriers already hidden; the fold only adds apply-loop cost), so the
// dispatcher picks LA per case.
// NW = compile-time max warps the reduction buffers are sized for (>= blockDim/32).
// Templated so the occ-critical n=512 (thr=128, 4 warps) keeps NW=8 (small static
// smem -> occ 3) while n=2048 can run threads=512 (16 warps -> NW=16) for more
// per-block parallelism (it is occ=1 anyway, so the extra ~2KB static is free).
template <bool LA, int NW>
__global__ void panel_factor_kernel(float* __restrict__ H,
float* __restrict__ tau,
float* __restrict__ Tg,
float* __restrict__ V,
int n, int k, int pb, int R,
int build_t, int write_v) {
extern __shared__ float smem[];
const int PB1 = pb + 1;
float* sH = smem;
float* sT = sH + (size_t)R * PB1;
float* gz = sT + (size_t)pb * pb;
__shared__ double sWarpD[NW];
__shared__ double sWcol[NW * PBMAX];
__shared__ double sWfull[PBMAX];
__shared__ float stw[PBMAX];
__shared__ double bcast[4];
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int nthreads = blockDim.x;
const int nwarps = (nthreads + 31) >> 5; // actual warps (<= NW)
float* Hb = H + (size_t)b * n * n;
float* taub = tau + (size_t)b * n;
for (int idx = tid; idx < R * pb; idx += nthreads) {
int r = idx / pb, c = idx % pb;
sH[r * PB1 + c] = Hb[(size_t)(k + r) * n + (k + c)];
}
if (build_t) {
for (int idx = tid; idx < pb * pb; idx += nthreads) sT[idx] = 0.0f;
}
__syncthreads();
const int lane = tid & 31;
const int warp = tid >> 5;
for (int jj = 0; jj < pb; ++jj) {
float alpha = sH[jj * PB1 + jj];
// !LA: compute column jj's norm now (every reflector). LA: only jj==0;
// for jj>0 the norm is already in sWarpD from jj-1's apply look-ahead.
if (!LA || jj == 0) {
double local = 0.0;
for (int r = jj + 1 + tid; r < R; r += nthreads) {
float v = sH[r * PB1 + jj];
local += (double)v * (double)v;
}
local = warpReduceSumD(local);
if (lane == 0) sWarpD[warp] = local;
__syncthreads();
}
if (tid == 0) {
double s = 0.0;
for (int w = 0; w < nwarps; ++w) s += sWarpD[w];
double xnorm = sqrt(s);
double beta, tauj, vscale;
if (xnorm == 0.0) {
beta = alpha; tauj = 0.0; vscale = 0.0;
} else {
double a = alpha;
double nrm = hypot(a, xnorm);
beta = (a >= 0.0) ? -nrm : nrm;
tauj = (beta - a) / beta;
vscale = 1.0 / (a - beta);
}
bcast[0] = beta; bcast[1] = tauj; bcast[2] = vscale;
taub[k + jj] = (float)tauj;
sH[jj * PB1 + jj] = (float)beta;
}
__syncthreads();
float vscale = (float)bcast[2];
float tauj = (float)bcast[1];
// Fuse the v-scale into the reflector-load. Each thread scales only its
// OWN rows of column jj and immediately caches them in vreg; the next
// consumer (w-reduction) reads only the thread's own rows, so the old
// __syncthreads between scale and load is unnecessary. Free barrier cut.
float vreg[VREG_MAX];
int nv = 0;
for (int r = jj + 1 + tid; r < R; r += nthreads) {
float v = sH[r * PB1 + jj];
if (vscale != 0.0f) { v *= vscale; sH[r * PB1 + jj] = v; }
vreg[nv++] = v;
}
if (tauj != 0.0f) {
// FP32 per-thread accumulator (s30): each thread sums only ~R/nthreads
// (~2-4) products, so FP32 is accuracy-neutral here; the cross-warp
// combine below stays FP64 (sWcol is double). Halves accd registers
// (32 vs 64) -> targets occ at threads=256.
float accd[PBMAX];
#pragma unroll
for (int c = 0; c < PBMAX; ++c) accd[c] = 0.0f;
int i = 0;
for (int r = jj + 1 + tid; r < R; r += nthreads) {
float vr = vreg[i++];
#pragma unroll
for (int c = 0; c < PBMAX; ++c)
if (c < pb) accd[c] += vr * sH[r * PB1 + c];
}
#pragma unroll
for (int c = 0; c < PBMAX; ++c) {
if (c >= pb) break;
float t = warpReduceSumF(accd[c]);
if (lane == 0) sWcol[warp * PBMAX + c] = (double)t;
}
__syncthreads();
if (tid < pb) {
int c = tid;
double s = 0.0;
for (int w = 0; w < nwarps; ++w) s += sWcol[w * PBMAX + c];
s += (double)sH[jj * PB1 + c];
sWfull[c] = s;
if (c > jj) {
float tw = (float)((double)tauj * s);
stw[c] = tw;
sH[jj * PB1 + c] -= tw;
}
}
__syncthreads();
if (LA) {
// Apply within the panel AND fold the look-ahead norm of column
// jj+1 (rows jj+2..R-1) into the same pass; publish via sWarpD so
// the next reflector skips its own norm reduction + barrier.
double nrm_next = 0.0;
int i2 = 0;
for (int r = jj + 1 + tid; r < R; r += nthreads) {
float vr = vreg[i2++];
for (int c = jj + 1; c < pb; ++c) {
float val = sH[r * PB1 + c] - stw[c] * vr;
sH[r * PB1 + c] = val;
if (c == jj + 1 && r > jj + 1) nrm_next += (double)val * val;
}
}
nrm_next = warpReduceSumD(nrm_next);
if (lane == 0) sWarpD[warp] = nrm_next;
} else {
int i2 = 0;
for (int r = jj + 1 + tid; r < R; r += nthreads) {
float vr = vreg[i2++];
for (int c = jj + 1; c < pb; ++c)
sH[r * PB1 + c] -= stw[c] * vr;
}
}
__syncthreads();
} else if (LA) {
// Identity reflector (tau==0): column jj+1 is unchanged; still must
// hand the next reflector its norm look-ahead via sWarpD.
double nrm_next = 0.0;
if (jj + 1 < pb) {
for (int r = jj + 2 + tid; r < R; r += nthreads) {
float v = sH[r * PB1 + (jj + 1)];
nrm_next += (double)v * v;
}
}
nrm_next = warpReduceSumD(nrm_next);
if (lane == 0) sWarpD[warp] = nrm_next;
__syncthreads();
}
if (build_t) {
if (tauj != 0.0f && jj > 0) {
if (tid < jj) gz[tid] = (float)(-(double)tauj * sWfull[tid]);
__syncthreads();
for (int i = tid; i < jj; i += nthreads) {
double acc = 0.0;
for (int l = i; l < jj; ++l)
acc += (double)sT[i * pb + l] * (double)gz[l];
sT[i * pb + jj] = (float)acc;
}
__syncthreads();
}
if (tid == 0) sT[jj * pb + jj] = tauj;
__syncthreads();
}
}
for (int idx = tid; idx < R * pb; idx += nthreads) {
int r = idx / pb, c = idx % pb;
Hb[(size_t)(k + r) * n + (k + c)] = sH[r * PB1 + c];
}
if (build_t) {
float* Tgb = Tg + (size_t)b * pb * pb;
for (int idx = tid; idx < pb * pb; idx += nthreads) Tgb[idx] = sT[idx];
}
// Emit packed, GEMM-ready V (unit diag, zeroed upper, zeroed tau==0 cols).
if (write_v) {
float* Vb = V + (size_t)b * R * pb;
for (int idx = tid; idx < R * pb; idx += nthreads) {
int r = idx / pb, c = idx % pb;
float v;
if (taub[k + c] == 0.0f) v = 0.0f;
else if (r < c) v = 0.0f;
else if (r == c) v = 1.0f;
else v = sH[r * PB1 + c];
Vb[(size_t)r * pb + c] = v;
}
}
}
void panel_factor(torch::Tensor H, torch::Tensor tau, torch::Tensor Tg,
torch::Tensor V, int64_t k, int64_t pb,
int64_t build_t, int64_t write_v, int64_t threads_,
int64_t lookahead) {
const int n = (int)H.size(1);
const int R = n - (int)k;
const int threads = (int)threads_;
size_t smem = ((size_t)R * (pb + 1) + (size_t)pb * pb + (size_t)pb)
* sizeof(float);
// NW=16 only when blockDim>256 (>8 warps); else NW=8 (keeps n=512 occ=3).
const bool nw16 = threads > 256;
auto kern = lookahead
? (nw16 ? panel_factor_kernel<true, 16> : panel_factor_kernel<true, 8>)
: (nw16 ? panel_factor_kernel<false, 16> : panel_factor_kernel<false, 8>);
cudaFuncSetAttribute(kern,
cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
kern<<<(int)H.size(0), threads, smem>>>(
H.data_ptr<float>(), tau.data_ptr<float>(), Tg.data_ptr<float>(),
V.data_ptr<float>(), n, (int)k, (int)pb, R,
(int)build_t, (int)write_v);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess)
throw std::runtime_error(std::string("panel_factor launch: ")
+ cudaGetErrorString(err));
}
// ---- TF32 hi/lo split (for the tf32x3 trailing path) ----
// hi = fp32 with low 13 mantissa bits cleared (exactly TF32-representable),
// lo = x - hi. So hi survives a TF32 GEMM truncation losslessly and the lo
// residual carries the remaining bits; AhBh+AhBl+AlBh recovers ~18 bits.
__device__ __forceinline__ void split_tf32(float x, float& hi, float& lo) {
unsigned int b = __float_as_uint(x);
hi = __uint_as_float(b & 0xFFFFE000u);
lo = x - hi;
}
__global__ void split_contig_kernel(const float* __restrict__ X,
float* __restrict__ Hh,
float* __restrict__ Ll, long long N) {
for (long long i = (long long)blockIdx.x * blockDim.x + threadIdx.x;
i < N; i += (long long)gridDim.x * blockDim.x) {
float h, l; split_tf32(X[i], h, l); Hh[i] = h; Ll[i] = l;
}
}
// Gather the strided trailing block Atrail = H[:, k:n, k+cur:n] into contiguous
// (batch, R, M) hi/lo split buffers.
__global__ void gather_split_atrail_kernel(const float* __restrict__ H,
float* __restrict__ Ah,
float* __restrict__ Al,
int n, int k, int cur, int R, int M) {
const int b = blockIdx.z;
const long long tot = (long long)R * M;
const float* Hb = H + (size_t)b * n * n;
float* Ahb = Ah + (size_t)b * R * M;
float* Alb = Al + (size_t)b * R * M;
for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
idx < tot; idx += (long long)gridDim.x * blockDim.x) {
int r = idx / M, c = idx % M;
float x = Hb[(size_t)(k + r) * n + (k + cur + c)];
float h, l; split_tf32(x, h, l);
Ahb[idx] = h; Alb[idx] = l;
}
}
// Use a handle created by OUR linked cuBLAS instance. Mixing torch's handle
// (at::cuda::getCurrentCUDABlasHandle) with calls into a separately-linked
// libcublas yields CUBLAS_STATUS_NOT_INITIALIZED, since each library instance
// has its own global state. A freshly-created handle runs on the default
// execution context (no override), keeping ordering with our <<<>>> launches.
static cublasHandle_t qrHandle() {
static cublasHandle_t h = nullptr;
if (h == nullptr) {
cublasStatus_t st = cublasCreate(&h);
if (st != CUBLAS_STATUS_SUCCESS)
throw std::runtime_error("cublasCreate failed: " + std::to_string((int)st));
}
return h;
}
// Row-major batched GEMM: C(m x n) = alpha*opA(A)(m x k) @ opB(B)(k x n) + beta*C.
// Implemented via the standard operand-swap so H stays row-major.
static void gemm_rm(bool tA, bool tB, int m, int n, int k,
float alpha, const float* A, int lda, long long sA,
const float* B, int ldb, long long sB,
float beta, float* C, int ldc, long long sC,
int batch, cublasComputeType_t ct) {
cublasOperation_t opA = tA ? CUBLAS_OP_T : CUBLAS_OP_N;
cublasOperation_t opB = tB ? CUBLAS_OP_T : CUBLAS_OP_N;
cublasStatus_t st = cublasGemmStridedBatchedEx(
qrHandle(), opB, opA, n, m, k, &alpha,
B, CUDA_R_32F, ldb, sB,
A, CUDA_R_32F, lda, sA,
&beta, C, CUDA_R_32F, ldc, sC,
batch, ct, CUBLAS_GEMM_DEFAULT);
if (st != CUBLAS_STATUS_SUCCESS)
throw std::runtime_error("cublas gemm failed: " + std::to_string((int)st));
}
static void launch_split_contig(const float* X, float* Hh, float* Ll, long long N) {
int threads = 256;
long long blk = (N + threads - 1) / threads;
int blocks = (int)(blk > 65535 ? 65535 : blk);
split_contig_kernel<<<blocks, threads>>>(X, Hh, Ll, N);
}
// One panel's block-WY trailing update via direct cuBLAS strided-batched GEMMs.
// prec: 0=fp32, 1=tf32, 2=tf32x3.
void larfb_cublas(torch::Tensor H, torch::Tensor V, torch::Tensor Tg,
int64_t k_, int64_t cur_, int64_t prec) {
const int n = (int)H.size(1);
const int batch = (int)H.size(0);
const int k = (int)k_, cur = (int)cur_;
const int R = n - k;
const int M = n - k - cur;
if (M <= 0) return;
float* Hp = H.data_ptr<float>();
float* Vp = V.data_ptr<float>();
float* Tp = Tg.data_ptr<float>();
float* Ap = Hp + (size_t)k * n + (k + cur); // Atrail base (batch b: + b*n*n)
const long long sH = (long long)n * n;
const long long sV = (long long)R * cur; // V is (batch, R, cur)
const long long sT = (long long)cur * cur; // Tg is (batch, pb=cur, pb)
auto opt = H.options();
torch::Tensor W = torch::empty({batch, cur, M}, opt);
torch::Tensor Y = torch::empty({batch, cur, M}, opt);
float* Wp = W.data_ptr<float>();
float* Yp = Y.data_ptr<float>();
const long long sW = (long long)cur * M;
const long long sY = (long long)cur * M;
const cublasComputeType_t TC = CUBLAS_COMPUTE_32F_FAST_TF32;
const cublasComputeType_t F32 = CUBLAS_COMPUTE_32F;
if (prec == 2) {
// ---- tf32x3: split operands, 3 TF32 GEMMs per heavy product. ----
torch::Tensor Vh = torch::empty({batch, R, cur}, opt);
torch::Tensor Vl = torch::empty({batch, R, cur}, opt);
torch::Tensor Ah = torch::empty({batch, R, M}, opt);
torch::Tensor Al = torch::empty({batch, R, M}, opt);
float* Vhp = Vh.data_ptr<float>(); float* Vlp = Vl.data_ptr<float>();
float* Ahp = Ah.data_ptr<float>(); float* Alp = Al.data_ptr<float>();
launch_split_contig(Vp, Vhp, Vlp, (long long)batch * R * cur);
{
int threads = 256;
long long blk = ((long long)R * M + threads - 1) / threads;
int bx = (int)(blk > 65535 ? 65535 : blk);
dim3 grid(bx, 1, batch);
gather_split_atrail_kernel<<<grid, threads>>>(Hp, Ahp, Alp, n, k, cur, R, M);
}
// W = V^T @ Atrail (cur x M, K=R) = Vh^T Ah + Vh^T Al + Vl^T Ah
gemm_rm(true, false, cur, M, R, 1.f, Vhp, cur, sV, Ahp, M, (long long)R * M,
0.f, Wp, M, sW, batch, TC);
gemm_rm(true, false, cur, M, R, 1.f, Vhp, cur, sV, Alp, M, (long long)R * M,
1.f, Wp, M, sW, batch, TC);
gemm_rm(true, false, cur, M, R, 1.f, Vlp, cur, sV, Ahp, M, (long long)R * M,
1.f, Wp, M, sW, batch, TC);
// Y = Tg^T @ W (cur x M, K=cur), small -> FP32.
gemm_rm(true, false, cur, M, cur, 1.f, Tp, cur, sT, Wp, M, sW,
0.f, Yp, M, sY, batch, F32);
// Atrail -= V @ Y (R x M, K=cur) = Vh Yh + Vh Yl + Vl Yh
torch::Tensor Yh = torch::empty({batch, cur, M}, opt);
torch::Tensor Yl = torch::empty({batch, cur, M}, opt);
float* Yhp = Yh.data_ptr<float>(); float* Ylp = Yl.data_ptr<float>();
launch_split_contig(Yp, Yhp, Ylp, (long long)batch * cur * M);
gemm_rm(false, false, R, M, cur, -1.f, Vhp, cur, sV, Yhp, M, sY,
1.f, Ap, n, sH, batch, TC);
gemm_rm(false, false, R, M, cur, -1.f, Vhp, cur, sV, Ylp, M, sY,
1.f, Ap, n, sH, batch, TC);
gemm_rm(false, false, R, M, cur, -1.f, Vlp, cur, sV, Yhp, M, sY,
1.f, Ap, n, sH, batch, TC);
} else {
// prec: 1=tf32(both heavy GEMMs), 3=mixed (W=V^TA on tf32 since it is a
// K=R fat reduction with accuracy headroom; the trailing-mutating
// A-=V@Y stays FP32 to hold the factor gate on hard cases like band),
// else 0=fp32(both).
const cublasComputeType_t WC = (prec == 1 || prec == 3) ? TC : F32;
const cublasComputeType_t AC = (prec == 1) ? TC : F32;
// W = V^T @ Atrail (cur x M, K=R)
gemm_rm(true, false, cur, M, R, 1.f, Vp, cur, sV, Ap, n, sH,
0.f, Wp, M, sW, batch, WC);
// Y = Tg^T @ W (cur x M, K=cur), FP32
gemm_rm(true, false, cur, M, cur, 1.f, Tp, cur, sT, Wp, M, sW,
0.f, Yp, M, sY, batch, F32);
// Atrail -= V @ Y (R x M, K=cur)
gemm_rm(false, false, R, M, cur, -1.f, Vp, cur, sV, Yp, M, sY,
1.f, Ap, n, sH, batch, AC);
}
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess)
throw std::runtime_error(std::string("larfb_cublas: ")
+ cudaGetErrorString(err));
}
"""
_CPP_SRC = (
"void panel_factor(torch::Tensor H, torch::Tensor tau, torch::Tensor Tg, "
"torch::Tensor V, int64_t k, int64_t pb, int64_t build_t, int64_t write_v, "
"int64_t threads_, int64_t lookahead);\n"
"void larfb_cublas(torch::Tensor H, torch::Tensor V, torch::Tensor Tg, "
"int64_t k, int64_t cur, int64_t prec);"
)
_module = None
try:
_module = load_inline(
name="qr_v31_kernels",
cpp_sources=[_CPP_SRC],
cuda_sources=[_CUDA_SRC],
functions=["panel_factor", "larfb_cublas"],
extra_cuda_cflags=["-O3"],
extra_ldflags=["-lcublas"],
verbose=True,
)
print("[qr] load_inline build OK")
except Exception as exc:
print(f"[qr] load_inline build FAILED, using geqrf fallback: {exc!r}")
def _mm_tc_split(A: torch.Tensor, B: torch.Tensor, four: bool) -> torch.Tensor:
"""A@B at ~fp32 accuracy via a hi/lo limb split on TF32 tensor cores."""
prev = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
Ah = A.to(torch.bfloat16).float()
Bh = B.to(torch.bfloat16).float()
Al = A - Ah
Bl = B - Bh
out = torch.matmul(Ah, Bh)
out = out + torch.matmul(Ah, Bl)
out = out + torch.matmul(Al, Bh)
if four:
out = out + torch.matmul(Al, Bl)
finally:
torch.backends.cuda.matmul.allow_tf32 = prev
return out
def _trailing_torch_trsm(H, k, cur, n, tau, eye, Tg):
block = H[:, k:n, k:k + cur]
V = block.clone()
top = V[:, :cur, :cur]
V[:, :cur, :cur] = torch.tril(top, -1) + eye[:cur, :cur]
taup = tau[:, k:k + cur]
V = V * (taup != 0).to(V.dtype).unsqueeze(1)
Atrail = H[:, k:n, k + cur:n]
W = V.transpose(-1, -2) @ Atrail
S = V.transpose(-1, -2) @ V
d = torch.where(taup != 0, 1.0 / taup, torch.ones_like(taup))
Tinv = torch.triu(S, 1) + torch.diag_embed(d)
Y = torch.linalg.solve_triangular(Tinv.transpose(-1, -2), W, upper=False)
Atrail.sub_(V @ Y)
def _trailing_torch_tfac_tc(H, k, cur, n, tau, eye, Tg):
block = H[:, k:n, k:k + cur]
V = block.clone()
top = V[:, :cur, :cur]
V[:, :cur, :cur] = torch.tril(top, -1) + eye[:cur, :cur]
taup = tau[:, k:k + cur]
V = V * (taup != 0).to(V.dtype).unsqueeze(1)
Atrail = H[:, k:n, k + cur:n]
W = _mm_tc_split(V.transpose(-1, -2), Atrail, four=False)
Y = Tg[:, :cur, :cur].transpose(-1, -2) @ W
Atrail.sub_(_mm_tc_split(V, Y, four=False))
def _qr_blocked(data: torch.Tensor, pb: int, impl: str, prec: str,
threads: int, lookahead: int) -> output_t:
"""Blocked WY Householder QR. Custom panel kernel + selected trailing path."""
batch, n, _ = data.shape
H = data.contiguous().clone()
tau = torch.zeros(batch, n, device=data.device, dtype=torch.float32)
need_t = impl in ("torch_tfac", "fused", "cublas")
need_v = impl == "cublas"
if need_t:
Tg = torch.empty(batch, pb, pb, device=data.device, dtype=torch.float32)
else:
Tg = torch.empty(1, device=data.device, dtype=torch.float32)
eye = torch.eye(pb, device=data.device, dtype=torch.float32)
dummy_v = torch.empty(1, device=data.device, dtype=torch.float32)
prec_code = _PREC_CODE.get(prec, 2)
for k in range(0, n, pb):
cur = min(pb, n - k)
ntrail = n - k - cur
R = n - k
build_t = 1 if (need_t and ntrail > 0) else 0
write_v = 1 if (need_v and ntrail > 0) else 0
if write_v:
V = torch.empty(batch, R, cur, device=data.device, dtype=torch.float32)
else:
V = dummy_v
_module.panel_factor(H, tau, Tg, V, k, cur, build_t, write_v,
threads, lookahead)
if ntrail <= 0:
continue
if impl == "cublas":
_module.larfb_cublas(H, V, Tg, k, cur, prec_code)
elif impl == "torch_tfac":
_trailing_torch_tfac_tc(H, k, cur, n, tau, eye, Tg)
else:
_trailing_torch_trsm(H, k, cur, n, tau, eye, Tg)
return H, tau
def _qr_blocked_geqrf_tc(data: torch.Tensor, pb: int) -> output_t:
"""n=4096 single-large regime: WIDE-panel blocked QR. The panel (R x pb) is
factored by torch.geqrf (cuSOLVER, row-major + tall-skinny native), the
trailing update runs on TF32 tensor cores via s7's compact-WY trsm formula.
pb is WIDE (512) so only ~8 cuSOLVER panel calls (the s40 pb=64 regression
was 64-panel orchestration overhead). gate 9.8e-3 (loose) -> plain TF32 safe.
"""
batch, n, _ = data.shape
H = data.contiguous().clone()
tau = torch.zeros(batch, n, device=data.device, dtype=torch.float32)
eye = torch.eye(pb, device=data.device, dtype=torch.float32)
prev = torch.backends.cuda.matmul.allow_tf32
try:
for k in range(0, n, pb):
cur = min(pb, n - k)
ntrail = n - k - cur
block = H[:, k:n, k:k + cur].contiguous()
hb, tb = torch.geqrf(block) # FP32 cuSOLVER tall panel
H[:, k:n, k:k + cur] = hb
tau[:, k:k + cur] = tb
if ntrail <= 0:
continue
V = hb.clone()
V[:, :cur, :cur] = torch.tril(V[:, :cur, :cur], -1) + eye[:cur, :cur]
V = V * (tb != 0).to(V.dtype).unsqueeze(1)
Atrail = H[:, k:n, k + cur:n]
torch.backends.cuda.matmul.allow_tf32 = True
W = V.transpose(-1, -2) @ Atrail # TF32, K=R (fat reduction)
torch.backends.cuda.matmul.allow_tf32 = False
S = V.transpose(-1, -2) @ V # FP32, cur x cur (cheap)
d = torch.where(tb != 0, 1.0 / tb, torch.ones_like(tb))
Tinv = torch.triu(S, 1) + torch.diag_embed(d)
Y = torch.linalg.solve_triangular(Tinv.transpose(-1, -2), W, upper=False)
torch.backends.cuda.matmul.allow_tf32 = True
Atrail.sub_(V @ Y) # TF32, K=cur
torch.backends.cuda.matmul.allow_tf32 = False
finally:
torch.backends.cuda.matmul.allow_tf32 = prev
return H, tau
# Per-n dispatch with the direct-cuBLAS trailing path. Precision chosen per n by
# the factor-gate headroom, which scales ~4096/n for plain TF32: n=512 has too
# little headroom (band case fails at scaled 21.6 > 20) so it uses exact FP32;
# n>=1024 passes comfortably on plain TF32 (1 GEMM, tensor cores). n=2048 needs
# pb<=24 to keep the first-panel smem under the cap.
# threads=128 lifts panel occupancy 2->3 (s19 probe) ONLY when batch >> #SMs(148)
# so blocks compete per-SM. That is ONLY n=512 (b=640). For batch < 148 (1024 b=60,
# 2048 b=8, 176/352 b=40) each block owns an SM, so fewer threads just starves
# per-matrix parallelism -> KEEP threads=256 (the benchmark confirmed thr=128
# regressed n=1024 10.4->13.4 and 176/352). n=512 thr=128: 15.4->11.8 (1.30x).
# `la` (look-ahead norm fold): ON for the low-occupancy threads=256 cases (occ=1,
# barriers exposed -> fewer barriers help: s23 gave 1024 -4.5%, 2048 -6%, 176/352
# -3%); OFF for the occ=3 n=512 family (barriers already hidden, the fold only adds
# apply-loop cost -> s23 regressed 512 +10%).
_DISPATCH = {
32: {"pb": 32, "impl": "torch_trsm", "threads": 256, "la": 0}, # no trailing
176: {"pb": 32, "impl": "cublas", "prec": "fp32", "threads": 256, "la": 1},
352: {"pb": 32, "impl": "cublas", "prec": "fp32", "threads": 256, "la": 1},
512: {"pb": 16, "impl": "cublas", "prec": "mixed", "threads": 128, "la": 0}, # s31: pb16 -> smem 34.8KB -> occ 6 (panel -30%); thr=128 keeps occ high
1024: {"pb": 32, "impl": "cublas", "prec": "tf32", "threads": 512, "la": 1}, # s41: thr 256->512 (b=60<SMs, latency-bound -> shorten critical path; NW=16)
2048: {"pb": 24, "impl": "cublas", "prec": "tf32", "threads": 512, "la": 1}, # b=8<<SMs: max parallelism (NW=16)
4096: {"pb": 512, "impl": "geqrf_tc"}, # s41: WIDE-panel cuSOLVER + TF32 trailing (8 panels; was geqrf ~52ms)
}
_SMEM_CAP = 220 * 1024
def custom_kernel(data: input_t) -> output_t:
n = data.shape[-1]
cfg = _DISPATCH.get(n)
if cfg is not None:
pb = cfg["pb"]
impl = _IMPL_OVERRIDE or cfg["impl"]
prec = cfg.get("prec", "fp32")
threads = cfg.get("threads", 256)
la = cfg.get("la", 0)
if impl == "geqrf_tc":
if _ENABLE_CUSTOM:
try:
return _qr_blocked_geqrf_tc(data, pb)
except Exception as exc:
print(f"[qr] geqrf_tc n={n} raised, fallback: {exc!r}")
return torch.geqrf(data)
smem = (n * (pb + 1) + pb * pb + pb) * 4
if _ENABLE_CUSTOM and _module is not None and smem <= _SMEM_CAP:
try:
return _qr_blocked(data, pb, impl, prec, threads, la)
except Exception as exc:
print(f"[qr] custom path n={n} impl={impl} raised, fallback: {exc!r}")
return torch.geqrf(data)
scrolls · 748 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