submission 836906
devichand579 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1359 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-836906?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:492003ce0dac54934cdd8992d88a2cd2ebbd9f96ca1278df1f95ee9692bf78fc
license declaredunknown
license concludedunknown
authorsdevichand579
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ double sh[WARPS];Kernel source
submission.py1359 lines
"""Batched square compact-Householder QR (geqrf-compatible (H, tau)).
Strategy
--------
The reference (torch.geqrf) factors each matrix in the batch *sequentially*
through cuSOLVER, so it collapses on high-batch shapes (e.g. 640 x 512 takes
~0.75 s). We replace that with a genuinely batched **blocked Householder
(compact-WY) QR**:
* panel factorization runs in a single custom CUDA kernel per panel (one
threadblock factors one matrix's panel, so the sequential column loop lives
inside the kernel -- no per-column Python/dispatch overhead);
* the trailing submatrix is updated with batched GEMMs (cuBLAS) via the WY
representation A_trail -= V (T^T (V^T A_trail)).
The kernel works on column-major data (A passed transposed) so sub-columns are
contiguous -> coalesced loads. Reductions accumulate in fp64 inside the kernel
for robustness on the ill-conditioned / wide-dynamic-range stress cases.
Precision: the trailing update defaults to fp32. Plain low precision (bf16/
1-pass tf32) *fails the factor gate* on the wide-dynamic-range cases (clustered
/ rowscale / mixed). `_TRAILING_PRECISION = "3xtf32"` switches the trailing
GEMMs to a manual fp32-emulation (3 TF32 tensor-core matmuls recovering
~2^-20 relative precision) -- the primary untested B200 lever for the
compute-bound shapes; see the comment above `_TRAILING_PRECISION`.
Dispatch: shapes where the batched panel kernel cannot fill the GPU (large n
with tiny batch) or where launch overhead dominates (very small n) fall back to
cuSOLVER, which is already near-optimal there. This guarantees we never regress
below the reference on any shape.
"""
import torch
# ---------------------------------------------------------------------------
# CUDA extension: batched Householder panel factorization (column-major).
# ---------------------------------------------------------------------------
_CPP = """
void panel_factor(torch::Tensor R, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads);
std::vector<torch::Tensor> qr_blocked(torch::Tensor Acm, int64_t block, int64_t threads);
std::vector<torch::Tensor> qr_blocked2(torch::Tensor Acm, int64_t superblk, int64_t inner, int64_t threads);
torch::Tensor build_T_from_G(torch::Tensor G, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads);
void set_cublas_emulation(int64_t strategy);
"""
_CUDA = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <ATen/cuda/CUDAContext.h>
#include <dlfcn.h>
#define MAX_WARPS 16 // up to 512 threads/block
// Enable cuBLAS fp32 tensor-core emulation (BF16x9). dlsym so the ext still builds
// on older cuBLAS lacking the symbol (CUDA 12.8 runtime) -> no-op there; on B200
// (cuBLAS 13) it resolves and enables emulation (verified neutral/safe on B200).
void set_cublas_emulation(int64_t strategy) {
typedef cublasStatus_t (*set_emul_fn)(cublasHandle_t, int);
static set_emul_fn fn = (set_emul_fn)dlsym(RTLD_DEFAULT, "cublasSetEmulationStrategy");
if (fn) {
cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
fn(h, (int)strategy);
}
}
// Hand-rolled grid-wide barrier (sense-reversing) over a 2-int global buffer
// bar = {count, sense}, zeroed by the host before each launch. Used by the
// cooperative multi-block panel kernel so all G*B co-resident blocks (guaranteed
// resident by cudaLaunchCooperativeKernel) can sync between reflector columns.
// Atomics + __threadfence only -> no -rdc / cooperative_groups device runtime.
__device__ __forceinline__ void grid_sync(unsigned int* bar, unsigned int gridSize, bool& sense) {
__syncthreads();
if (threadIdx.x == 0) {
__threadfence();
sense = !sense;
unsigned int old = atomicAdd(&bar[0], 1u);
if (old == gridSize - 1u) {
atomicExch(&bar[0], 0u);
__threadfence();
atomicExch(&bar[1], sense ? 1u : 0u);
} else {
volatile unsigned int* vs = (volatile unsigned int*)&bar[1];
while (((*vs) != 0u) != sense) { }
}
}
__syncthreads();
}
__device__ __forceinline__ double warpReduceSum(double v) {
#pragma unroll
for (int o = 16; o > 0; o >>= 1)
v += __shfl_down_sync(0xffffffffu, v, o);
return v;
}
// Block reduction: shuffle within each warp (no sync), then combine the WARPS
// partials through a tiny shared array. Two __syncthreads per reduction vs ~9 for
// the old shared-memory tree -- this matters because the apply step calls it
// ~jb^2/2 times/panel. THREADS is a compile-time template parameter (one kernel
// instantiation per launch block size) so the strided loops and this reduction
// stay fully unrolled with constant strides -- a runtime blockDim.x costs ~10%.
template <int THREADS>
__device__ __forceinline__ double blockReduceSum(double v, double* sh) {
constexpr int WARPS = THREADS / 32;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
v = warpReduceSum(v);
if (lane == 0) sh[wid] = v;
__syncthreads();
double r = 0.0;
#pragma unroll
for (int i = 0; i < WARPS; ++i) r += sh[i];
__syncthreads();
return r;
}
// Single-sync block reduction: drops the trailing __syncthreads that protects sh
// for reuse. Safe ONLY when the caller does an explicit __syncthreads before the
// next reduction (true in the panel kernels: a sync follows after tau-broadcast,
// scale, and apply). Saves one sync per Householder step -- the panel is bound by
// the latency of its serial n-step chain + per-step syncs, so this is a direct cut.
template <int THREADS>
__device__ __forceinline__ double blockReduceSum1(double v, double* sh) {
constexpr int WARPS = THREADS / 32;
int lane = threadIdx.x & 31;
int wid = threadIdx.x >> 5;
v = warpReduceSum(v);
if (lane == 0) sh[wid] = v;
__syncthreads();
double r = 0.0;
#pragma unroll
for (int i = 0; i < WARPS; ++i) r += sh[i];
return r;
}
// Acm: (B,n,n) holding COLUMN-MAJOR n x n matrices (lda=n, batch stride n*n);
// the caller passes A^T as a contiguous row-major tensor. Column 'col' of A is
// contiguous (stride 1) -> coalesced. tau: (B,n). Panel = cols [j, j+jb).
template <int THREADS>
__global__ void panel_factor_kernel(float* __restrict__ Acm, float* __restrict__ tau,
int n, int j, int jb) {
constexpr int WARPS = THREADS / 32;
int b = blockIdx.x;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
__shared__ double sh[WARPS];
__shared__ float s_tau, s_beta;
float* Ab = Acm + (size_t)b * n * n;
float* taub = tau + (size_t)b * n;
for (int k = 0; k < jb; ++k) {
int col = j + k;
int m = n - col; // subcolumn length
float* base = Ab + (size_t)col * n + col; // A(col,col), entries stride 1
float alpha = base[0];
double part = 0.0;
for (int i = t + 1; i < m; i += THREADS) {
double val = base[i];
part += val * val;
}
double sigma = blockReduceSum1<THREADS>(part, sh);
if (t == 0) {
float tau_k, beta;
if (sigma <= 0.0) { // zero sub-column: no reflection
tau_k = 0.f; beta = alpha;
} else {
double xnorm = sqrt((double)alpha * alpha + sigma);
double sign = (alpha >= 0.f) ? 1.0 : -1.0;
double betad = -sign * xnorm;
beta = (float)betad;
tau_k = (float)((betad - alpha) / betad);
}
s_tau = tau_k; s_beta = beta;
taub[col] = tau_k;
}
__syncthreads();
float tau_k = s_tau, beta = s_beta;
if (sigma > 0.0) {
double scale = 1.0 / ((double)alpha - beta);
for (int i = t + 1; i < m; i += THREADS)
base[i] = (float)(base[i] * scale); // v_i = x_i / (alpha - beta)
if (t == 0) base[0] = beta; // R diagonal
} else {
for (int i = t + 1; i < m; i += THREADS)
base[i] = 0.f; // clean reflector tail
}
__syncthreads();
// apply H_k = I - tau v v^T (v[0]=1) to remaining panel columns:
// one warp per column (warp-shuffle dot, no per-column block sync).
// fp32 apply (norm/sigma above stays fp64): the intra-panel dot/update is
// the same operation the fp32 cuBLAS trailing GEMM does and passes the
// gate, so fp32 here is accurate too -- and ~2x the fp64 throughput.
if (tau_k != 0.f) {
for (int c = col + 1 + wid; c < j + jb; c += WARPS) {
float* cptr = Ab + (size_t)c * n + col; // column c, rows [col,n)
float pivot = cptr[0]; // read before any write
float p = 0.f;
for (int i = lane + 1; i < m; i += 32)
p += base[i] * cptr[i];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
float full = __shfl_sync(0xffffffffu, p, 0) + pivot;
float f = tau_k * full;
if (lane == 0) cptr[0] = pivot - f;
for (int i = lane + 1; i < m; i += 32)
cptr[i] = cptr[i] - f * base[i];
}
}
__syncthreads();
}
}
// Shared-memory panel factorization: load the panel cols [j,j+jb) rows [j,n)
// (jb*(n-j) floats) into shared ONCE, run the jb unblocked Householder steps
// with the warp-parallel apply reading/writing shared (not global), write the
// factored panel back ONCE. The plain panel_factor_kernel re-reads each panel
// column from global ~jb times during the apply (memory-bound, ~40-55% of CUDA
// time per the profiler); this cuts that to one read + one write per element.
// Trailing update stays in cuBLAS. Caller must ensure jb*(n-j)*4 <= smem capacity.
template <int THREADS>
__global__ void panel_factor_smem_kernel(float* __restrict__ Acm, float* __restrict__ tau,
int n, int j, int jb) {
extern __shared__ float sV[]; // jb*m, col-major tight: sV[c*m + r]
constexpr int WARPS = THREADS / 32;
int b = blockIdx.x;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
int m = n - j;
float* Ab = Acm + (size_t)b * n * n;
float* taub = tau + (size_t)b * n;
__shared__ double sh[WARPS];
__shared__ float s_tau, s_beta, s_alpha;
for (int idx = t; idx < jb * m; idx += THREADS) {
int c = idx / m, r = idx % m;
sV[(size_t)c * m + r] = Ab[(size_t)(j + c) * n + (j + r)];
}
__syncthreads();
for (int k = 0; k < jb; ++k) {
int mk = m - k; // subcolumn length
float* vk = sV + (size_t)k * m + k;
float alpha = vk[0];
double part = 0.0;
for (int i = t + 1; i < mk; i += THREADS) { double v = vk[i]; part += v * v; }
double sigma = blockReduceSum1<THREADS>(part, sh);
if (t == 0) {
float tau_k, beta;
if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
else {
double xnorm = sqrt((double)alpha * alpha + sigma);
double sign = (alpha >= 0.f) ? 1.0 : -1.0;
double betad = -sign * xnorm;
beta = (float)betad;
tau_k = (float)((betad - alpha) / betad);
}
s_tau = tau_k; s_beta = beta; s_alpha = alpha;
taub[j + k] = tau_k;
}
__syncthreads();
float tau_k = s_tau, beta = s_beta, alpha2 = s_alpha;
if (sigma > 0.0) {
double scale = 1.0 / ((double)alpha2 - beta);
for (int i = t + 1; i < mk; i += THREADS) vk[i] = (float)(vk[i] * scale);
if (t == 0) vk[0] = beta;
} else {
for (int i = t + 1; i < mk; i += THREADS) vk[i] = 0.f;
}
__syncthreads();
if (tau_k != 0.f) {
for (int c = k + 1 + wid; c < jb; c += WARPS) {
float* ac = sV + (size_t)c * m + k;
float pivot = ac[0];
float p = 0.f;
for (int i = lane + 1; i < mk; i += 32) p += vk[i] * ac[i];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
float tot = __shfl_sync(0xffffffffu, p, 0) + pivot;
float f = tau_k * tot;
if (lane == 0) ac[0] = pivot - f;
for (int i = lane + 1; i < mk; i += 32)
ac[i] = ac[i] - f * vk[i];
}
}
__syncthreads();
}
for (int idx = t; idx < jb * m; idx += THREADS) {
int c = idx / m, r = idx % m;
Ab[(size_t)(j + c) * n + (j + r)] = sV[(size_t)c * m + r];
}
}
// Fully-fused unblocked Householder QR: one threadblock factors one entire
// n x n matrix, kept resident in shared memory (column-major, same layout as
// Acm). The whole factorization runs in a single kernel launch -- no Python
// panel loop, no trailing-GEMM round trips, no global traffic between steps.
//
// Key to speed: the reflector APPLY is parallelized across warps -- one warp
// per trailing column, with a warp-shuffle dot reduction (no __syncthreads).
// Only the n sequential reflector steps carry a couple of block syncs each, so
// we pay ~O(n) block syncs instead of the ~O(n^2) block reductions an unblocked
// in-block QR would otherwise need. Valid for n*n*4 bytes <= the opted-in
// dynamic shared capacity (see qr_full host fn).
template <int THREADS>
__global__ void qr_full_kernel(float* __restrict__ Acm, float* __restrict__ tau, int n) {
extern __shared__ float sA[]; // n*n floats, column-major
constexpr int WARPS = THREADS / 32;
int b = blockIdx.x;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
float* Ab = Acm + (size_t)b * n * n;
float* taub = tau + (size_t)b * n;
__shared__ double sh[WARPS];
__shared__ float s_tau, s_beta, s_alpha;
for (int idx = t; idx < n * n; idx += THREADS) sA[idx] = Ab[idx];
__syncthreads();
for (int k = 0; k < n; ++k) {
int m = n - k; // subcolumn length
float* vk = sA + (size_t)k * n + k; // column k, rows [k,n)
float alpha = vk[0];
double part = 0.0;
for (int i = t + 1; i < m; i += THREADS) { double v = vk[i]; part += v * v; }
double sigma = blockReduceSum<THREADS>(part, sh);
if (t == 0) {
float tau_k, beta;
if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
else {
double xnorm = sqrt((double)alpha * alpha + sigma);
double sign = (alpha >= 0.f) ? 1.0 : -1.0;
double betad = -sign * xnorm;
beta = (float)betad;
tau_k = (float)((betad - alpha) / betad);
}
s_tau = tau_k; s_beta = beta; s_alpha = alpha;
taub[k] = tau_k;
}
__syncthreads();
float tau_k = s_tau, beta = s_beta, alpha2 = s_alpha;
if (sigma > 0.0) {
double scale = 1.0 / ((double)alpha2 - beta);
for (int i = t + 1; i < m; i += THREADS) vk[i] = (float)(vk[i] * scale);
if (t == 0) vk[0] = beta;
} else {
for (int i = t + 1; i < m; i += THREADS) vk[i] = 0.f;
}
__syncthreads();
// apply H_k = I - tau v v^T (v[0]=1): one warp per trailing column c
if (tau_k != 0.f) {
for (int c = k + 1 + wid; c < n; c += WARPS) {
float* ac = sA + (size_t)c * n + k; // column c, rows [k,n)
float pivot = ac[0];
double p = 0.0;
for (int i = lane + 1; i < m; i += 32) p += (double)vk[i] * ac[i];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
double tot = __shfl_sync(0xffffffffu, p, 0) + (double)pivot;
float f = (float)((double)tau_k * tot);
if (lane == 0) ac[0] = pivot - f;
for (int i = lane + 1; i < m; i += 32)
ac[i] = (float)((double)ac[i] - (double)f * vk[i]);
}
}
__syncthreads();
}
for (int idx = t; idx < n * n; idx += THREADS) Ab[idx] = sA[idx];
}
void qr_full(torch::Tensor R, torch::Tensor tau, int64_t threads) {
TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
int B = (int)R.size(0);
int n = (int)R.size(1);
float* rp = R.data_ptr<float>();
float* tp = tau.data_ptr<float>();
size_t shbytes = (size_t)n * n * sizeof(float);
#define LAUNCH_FULL(TH) \
cudaFuncSetAttribute(qr_full_kernel<TH>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shbytes); \
qr_full_kernel<TH><<<B, TH, shbytes, 0>>>(rp, tp, n);
switch ((int)threads) {
case 128: LAUNCH_FULL(128); break;
case 512: LAUNCH_FULL(512); break;
default: LAUNCH_FULL(256); break;
}
#undef LAUNCH_FULL
}
// Cooperative multi-block-per-matrix panel factorization. For starved shapes
// (large n, tiny batch) the one-block-per-matrix panel kernel uses only `batch`
// SMs; here G blocks cooperate on each matrix's panel so the work fills the GPU.
// Block bid -> matrix b = bid/G, sub-block g = bid%G. For each subcolumn the G
// blocks partition its rows, reduce the norm / apply-dots across blocks through
// global scratch, and sync via the hand-rolled grid barrier between the (still
// sequential) reflector columns. Math is identical to panel_factor_kernel.
template <int THREADS>
__global__ void panel_factor_coop_kernel(
float* __restrict__ Acm, float* __restrict__ tau,
float* __restrict__ normbuf, float* __restrict__ dotbuf,
unsigned int* __restrict__ bar, int n, int j, int jb, int G) {
int bid = blockIdx.x;
int b = bid / G;
int g = bid % G;
int t = threadIdx.x;
unsigned int gridSize = (unsigned int)(gridDim.x);
bool sense = false;
__shared__ double sh[THREADS / 32];
float* Ab = Acm + (size_t)b * n * n;
float* taub = tau + (size_t)b * n;
float* nbrow = normbuf + (size_t)b * G; // [G] partial norms
float* dbrow = dotbuf + (size_t)(b * G + g) * jb; // [jb] this block's dots
for (int k = 0; k < jb; ++k) {
int col = j + k;
int M = n - col; // subcolumn length
float* base = Ab + (size_t)col * n + col; // base[0] = diagonal
int chunk = (M + G - 1) / G;
int gs = g * chunk;
int ge = gs + chunk; if (ge > M) ge = M;
int ns = gs < 1 ? 1 : gs; // tail start (skip pivot row 0)
float alpha = base[0];
double part = 0.0;
for (int i = ns + t; i < ge; i += THREADS) { double v = base[i]; part += v * v; }
double psum = blockReduceSum<THREADS>(part, sh);
if (t == 0) nbrow[g] = (float)psum;
grid_sync(bar, gridSize, sense);
double sigma = 0.0;
for (int gg = 0; gg < G; ++gg) sigma += (double)nbrow[gg];
float tau_k, beta;
if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
else {
double xnorm = sqrt((double)alpha * alpha + sigma);
double sign = (alpha >= 0.f) ? 1.0 : -1.0;
double betad = -sign * xnorm;
beta = (float)betad;
tau_k = (float)((betad - alpha) / betad);
}
if (sigma > 0.0) {
double scale = 1.0 / ((double)alpha - beta);
for (int i = ns + t; i < ge; i += THREADS) base[i] = (float)(base[i] * scale);
} else {
for (int i = ns + t; i < ge; i += THREADS) base[i] = 0.f;
}
if (g == 0 && t == 0) { base[0] = beta; taub[col] = tau_k; } // block 0 owns pivot
grid_sync(bar, gridSize, sense);
// apply H_k = I - tau v v^T (v[0]=1): partial dots over owned rows
for (int c = col + 1; c < j + jb; ++c) {
float* cptr = Ab + (size_t)c * n + col;
double p = 0.0;
for (int i = gs + t; i < ge; i += THREADS) {
double vi = (i == 0) ? 1.0 : (double)base[i];
p += vi * (double)cptr[i];
}
double bp = blockReduceSum<THREADS>(p, sh);
if (t == 0) dbrow[c - j] = (float)bp;
}
grid_sync(bar, gridSize, sense);
for (int c = col + 1; c < j + jb; ++c) {
double full = 0.0;
for (int gg = 0; gg < G; ++gg) full += (double)dotbuf[(size_t)(b * G + gg) * jb + (c - j)];
float f = (float)((double)tau_k * full);
float* cptr = Ab + (size_t)c * n + col;
for (int i = gs + t; i < ge; i += THREADS) {
double vi = (i == 0) ? 1.0 : (double)base[i];
cptr[i] = (float)((double)cptr[i] - (double)f * vi);
}
}
grid_sync(bar, gridSize, sense);
}
}
void panel_factor_coop(torch::Tensor R, torch::Tensor tau, int64_t j, int64_t jb, int64_t G) {
TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
int B = (int)R.size(0);
int n = (int)R.size(1);
int Gi = (int)G;
float* rp = R.data_ptr<float>();
float* tp = tau.data_ptr<float>();
auto i32 = R.options().dtype(at::kInt);
auto nbuf = torch::zeros({(long)(B * Gi)}, R.options());
auto dbuf = torch::zeros({(long)(B * Gi * (int)jb)}, R.options());
auto barb = torch::zeros({2}, i32);
float* nbp = nbuf.data_ptr<float>();
float* dbp = dbuf.data_ptr<float>();
unsigned int* barp = (unsigned int*)barb.data_ptr<int>();
int ni = n, ji = (int)j, jbi = (int)jb;
void* args[] = { (void*)&rp, (void*)&tp, (void*)&nbp, (void*)&dbp,
(void*)&barp, (void*)&ni, (void*)&ji, (void*)&jbi, (void*)&Gi };
dim3 grid(B * Gi), block(256);
cudaError_t err = cudaLaunchCooperativeKernel(
(const void*)panel_factor_coop_kernel<256>, grid, block, args, 0, 0);
TORCH_CHECK(err == cudaSuccess, "coop panel launch failed: ", cudaGetErrorString(err));
}
// Blocked fully-fused Householder QR for larger n that does NOT fit in shared:
// one threadblock factors one entire matrix in a SINGLE launch (no Python panel
// loop, no cuBLAS, no T-build, no V-extraction copies). Only the current panel
// (jb columns) lives in shared; each trailing column is read from global
// EXACTLY ONCE per panel and has all jb reflectors applied to it in-place
// (kept in a per-warp shared scratch column), so global bandwidth stays at the
// blocked O(n/jb)-passes level rather than the O(n) of an unblocked global QR.
//
// Shared layout (dynamic): sV = jb*n floats (panel, col-major sV[c*n+r]);
// sC = WARPS*n floats (per-warp trailing column).
template <int THREADS>
__global__ void qr_block_fused_kernel(float* __restrict__ Acm, float* __restrict__ tau,
int n, int blk) {
extern __shared__ float smem[];
constexpr int WARPS = THREADS / 32;
float* sV = smem; // jb*n
float* sC = smem + (size_t)blk * n; // WARPS*n
int b = blockIdx.x;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
float* Ab = Acm + (size_t)b * n * n;
float* taub = tau + (size_t)b * n;
__shared__ double sh[WARPS];
__shared__ float s_tau, s_beta, s_alpha;
for (int j = 0; j < n; j += blk) {
int jb = blk < (n - j) ? blk : (n - j);
int m = n - j; // panel/trailing row count
// load panel cols [j,j+jb) rows [j,n) into sV (col-major, m rows)
for (int idx = t; idx < jb * m; idx += THREADS) {
int c = idx / m, r = idx % m;
sV[(size_t)c * n + r] = Ab[(size_t)(j + c) * n + (j + r)];
}
__syncthreads();
// factor panel in shared (unblocked Householder, warp-parallel apply)
for (int k = 0; k < jb; ++k) {
int mk = m - k; // subcolumn length
float* vk = sV + (size_t)k * n + k;
float alpha = vk[0];
double part = 0.0;
for (int i = t + 1; i < mk; i += THREADS) { double v = vk[i]; part += v * v; }
double sigma = blockReduceSum<THREADS>(part, sh);
if (t == 0) {
float tau_k, beta;
if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
else {
double xnorm = sqrt((double)alpha * alpha + sigma);
double sign = (alpha >= 0.f) ? 1.0 : -1.0;
double betad = -sign * xnorm;
beta = (float)betad;
tau_k = (float)((betad - alpha) / betad);
}
s_tau = tau_k; s_beta = beta; s_alpha = alpha;
taub[j + k] = tau_k;
}
__syncthreads();
float tau_k = s_tau, beta = s_beta, alpha2 = s_alpha;
if (sigma > 0.0) {
double scale = 1.0 / ((double)alpha2 - beta);
for (int i = t + 1; i < mk; i += THREADS) vk[i] = (float)(vk[i] * scale);
if (t == 0) vk[0] = beta;
} else {
for (int i = t + 1; i < mk; i += THREADS) vk[i] = 0.f;
}
__syncthreads();
if (tau_k != 0.f) {
for (int c = k + 1 + wid; c < jb; c += WARPS) {
float* ac = sV + (size_t)c * n + k;
float pivot = ac[0];
double p = 0.0;
for (int i = lane + 1; i < mk; i += 32) p += (double)vk[i] * ac[i];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
double tot = __shfl_sync(0xffffffffu, p, 0) + (double)pivot;
float f = (float)((double)tau_k * tot);
if (lane == 0) ac[0] = pivot - f;
for (int i = lane + 1; i < mk; i += 32)
ac[i] = (float)((double)ac[i] - (double)f * vk[i]);
}
}
__syncthreads();
}
// write factored panel back (R rows + reflectors)
for (int idx = t; idx < jb * m; idx += THREADS) {
int c = idx / m, r = idx % m;
Ab[(size_t)(j + c) * n + (j + r)] = sV[(size_t)c * n + r];
}
__syncthreads();
// trailing update: cols [j+jb, n), one warp per column, read once
float* col = sC + (size_t)wid * n;
for (int c = j + jb + wid; c < n; c += WARPS) {
for (int r = lane; r < m; r += 32) col[r] = Ab[(size_t)c * n + (j + r)];
__syncwarp();
for (int k = 0; k < jb; ++k) {
int mk = m - k;
float* vk = sV + (size_t)k * n + k;
float tau_k = taub[j + k];
if (tau_k != 0.f) {
float pivot = col[k];
double p = 0.0;
for (int i = lane + 1; i < mk; i += 32) p += (double)vk[i] * col[k + i];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
double tot = __shfl_sync(0xffffffffu, p, 0) + (double)pivot;
float f = (float)((double)tau_k * tot);
if (lane == 0) col[k] = pivot - f;
for (int i = lane + 1; i < mk; i += 32)
col[k + i] = (float)((double)col[k + i] - (double)f * vk[i]);
__syncwarp();
}
}
for (int r = lane; r < m; r += 32) Ab[(size_t)c * n + (j + r)] = col[r];
}
__syncthreads();
}
}
void qr_block_fused(torch::Tensor R, torch::Tensor tau, int64_t blk, int64_t threads) {
TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
int B = (int)R.size(0);
int n = (int)R.size(1);
float* rp = R.data_ptr<float>();
float* tp = tau.data_ptr<float>();
int W = (int)threads / 32;
size_t shbytes = (size_t)(blk + W) * n * sizeof(float);
#define LAUNCH_BF(TH) \
cudaFuncSetAttribute(qr_block_fused_kernel<TH>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shbytes); \
qr_block_fused_kernel<TH><<<B, TH, shbytes, 0>>>(rp, tp, n, (int)blk);
switch ((int)threads) {
case 128: LAUNCH_BF(128); break;
case 512: LAUNCH_BF(512); break;
default: LAUNCH_BF(256); break;
}
#undef LAUNCH_BF
}
// Build the compact-WY block reflector T (jb x jb, upper triangular) for the
// panel [j, j+jb) directly from the reflectors already stored in Acm (col-major)
// and tau -- one block per matrix, one launch. Replaces the 5-op PyTorch T-build
// (G = bmm(V,V^T), triu, eye+tau*U, diag_embed, solve_triangular) and removes the
// cuSOLVER batched triangular solve entirely. Q_panel = I - V T V^T (LAPACK
// dlarft, DIRECT='F', STOREV='C'). Reflector p: v_p[p]=1, v_p[bb>p]=Acm tail,
// v_p[bb<p]=0; the R entries on/above the panel diagonal are never read.
template <int THREADS>
__global__ void build_T_kernel(const float* __restrict__ Acm, const float* __restrict__ tau,
float* __restrict__ Tout, int n, int j, int jb) {
constexpr int WARPS = THREADS / 32;
int b = blockIdx.x;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
const float* Ab = Acm + (size_t)b * n * n;
const float* taub = tau + (size_t)b * n;
float* Tb = Tout + (size_t)b * jb * jb; // row-major jb x jb
int m = n - j;
extern __shared__ float smT[];
float* sG = smT; // jb*jb (only upper+diag used)
float* sT = smT + jb * jb; // jb*jb
float* stau = smT + 2 * jb * jb; // jb
for (int p = t; p < jb; p += THREADS) stau[p] = taub[j + p];
__syncthreads();
// G[p][q] = v_p . v_q for p <= q : one warp per (p,q) pair
for (int idx = wid; idx < jb * jb; idx += WARPS) {
int p = idx / jb, q = idx % jb;
if (p > q) continue;
const float* vp = Ab + (size_t)(j + p) * n + j; // vp[bb] = reflector p, local row bb
const float* vq = Ab + (size_t)(j + q) * n + j;
double s = 0.0;
for (int bb = q + 1 + lane; bb < m; bb += 32) s += (double)vp[bb] * vq[bb];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
if (lane == 0) {
double vpq = (p == q) ? 1.0 : (double)vp[q]; // v_p[q]; v_q[q]=1
sG[p * jb + q] = (float)(vpq + s);
}
}
__syncthreads();
// dlarft forward recurrence (sequential in q), done by warp 0
if (wid == 0) {
for (int q = 0; q < jb; ++q) {
if (lane == 0) sT[q * jb + q] = stau[q];
for (int p = lane; p < q; p += 32) {
double acc = 0.0;
for (int l = p; l < q; ++l) acc += (double)sT[p * jb + l] * (double)sG[l * jb + q];
sT[p * jb + q] = (float)(-(double)stau[q] * acc);
}
__syncwarp();
}
for (int idx = lane; idx < jb * jb; idx += 32)
if ((idx / jb) > (idx % jb)) sT[idx] = 0.f; // zero strict lower
}
__syncthreads();
for (int idx = t; idx < jb * jb; idx += THREADS) Tb[idx] = sT[idx];
}
// Build compact-WY T (jb x jb, row-major upper-tri) from the ALREADY-computed
// G = V^T V (jb x jb, from a cuBLAS tensor-core bmm) and tau. This replaces the
// 5-launch PyTorch tail (G.triu(1) + tau*U + eye+ + diag_embed + solve_triangular)
// with ONE launch, while KEEPING the fast TC bmm for G (the old build_T_kernel lost
// at jb=32 only because it recomputed G on CUDA cores warp-per-pair).
// Amat = I + diag(tau)*striu(G,1) (unit upper-tri); Amat @ T = diag(tau).
// The jb columns of T are INDEPENDENT triangular solves, so one warp solves each
// column (back-substitution) -- critical depth ~jb, vs the dlarft recurrence's
// serial-over-columns ~jb^2/2 in a single warp. One block per matrix.
template <int THREADS>
__global__ void build_T_from_G_kernel(const float* __restrict__ Gin,
const float* __restrict__ tau,
float* __restrict__ Tout,
int n, int j, int jb) {
constexpr int WARPS = THREADS / 32;
int b = blockIdx.x;
int t = threadIdx.x;
int lane = t & 31;
int wid = t >> 5;
const float* Gb = Gin + (size_t)b * jb * jb; // row-major jb x jb
const float* taub = tau + (size_t)b * n;
float* Tb = Tout + (size_t)b * jb * jb; // row-major jb x jb
extern __shared__ float smTG[];
float* sA = smTG; // jb*jb Amat (unit upper-tri)
float* sT = smTG + jb * jb; // jb*jb T (col q in sT[p*jb+q])
float* st = smTG + 2 * jb * jb; // jb tau
for (int p = t; p < jb; p += THREADS) st[p] = taub[j + p];
__syncthreads();
for (int idx = t; idx < jb * jb; idx += THREADS) {
int p = idx / jb, q = idx % jb;
sA[idx] = (p == q) ? 1.f : (p < q ? st[p] * Gb[idx] : 0.f);
sT[idx] = 0.f;
}
__syncthreads();
// one warp per column q: solve Amat @ x = tau[q] e_q (x = T[:,q]); back-sub
// x[q]=tau[q]; x[p<q] = -sum_{l=p+1..q} Amat[p][l]*x[l] (descending p).
for (int q = wid; q < jb; q += WARPS) {
if (lane == 0) sT[q * jb + q] = st[q];
__syncwarp();
for (int p = q - 1; p >= 0; --p) {
float acc = 0.f; // fp32 (B200 fp64 is ~1/30th); matches the
for (int l = p + 1 + lane; l <= q; l += 32) // fp32 solve_triangular baseline
acc += sA[p * jb + l] * sT[l * jb + q];
#pragma unroll
for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
if (lane == 0) sT[p * jb + q] = -acc;
__syncwarp();
}
}
__syncthreads();
for (int idx = t; idx < jb * jb; idx += THREADS) {
int p = idx / jb, q = idx % jb;
Tb[idx] = (p > q) ? 0.f : sT[idx];
}
}
torch::Tensor build_T_from_G(torch::Tensor G, torch::Tensor tau, int64_t j,
int64_t jb, int64_t threads) {
TORCH_CHECK(G.is_cuda() && tau.is_cuda(), "cuda required");
TORCH_CHECK(G.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
int B = (int)G.size(0);
int n = (int)tau.size(1);
auto T = torch::empty({B, (int)jb, (int)jb}, G.options());
const float* gp = G.data_ptr<float>();
const float* tp = tau.data_ptr<float>();
float* Tp = T.data_ptr<float>();
size_t shbytes = (size_t)(2 * jb * jb + jb) * sizeof(float);
#define LAUNCH_TG(TH) build_T_from_G_kernel<TH><<<B, TH, shbytes, 0>>>(gp, tp, Tp, n, (int)j, (int)jb);
switch ((int)threads) {
case 128: LAUNCH_TG(128); break;
case 512: LAUNCH_TG(512); break;
default: LAUNCH_TG(256); break;
}
#undef LAUNCH_TG
return T;
}
// Emit a clean compact-WY V panel (B,jb,m) row-major from the reflectors stored
// in Acm (col-major (B,n,n) = A^T row-major): V[p][bb] = 1 (bb==p), Acm reflector
// (bb>p), 0 (bb<p). Replaces raw.triu(1).contiguous() (strided triu_tril_kernel +
// a separate diagonal.fill_(1) launch) with one coalesced pass: blockIdx=(p,b),
// threads stride over contiguous bb so both the Acm read and the Vt write coalesce.
__global__ void extract_V_kernel(const float* __restrict__ Acm, float* __restrict__ Vt,
int n, int j, int jb, int m) {
int b = blockIdx.y;
int p = blockIdx.x; // reflector index within panel
const float* Ab = Acm + (size_t)b * n * n + (size_t)(j + p) * n + j; // row j+p, col j
float* Vb = Vt + ((size_t)b * jb + p) * m;
for (int bb = threadIdx.x; bb < m; bb += blockDim.x) {
Vb[bb] = (bb < p) ? 0.f : (bb == p) ? 1.f : Ab[bb];
}
}
torch::Tensor extract_V(torch::Tensor Acm, int64_t j, int64_t jb, int64_t m) {
int B = (int)Acm.size(0);
int n = (int)Acm.size(1);
auto Vt = torch::empty({B, jb, m}, Acm.options());
int th = (int)m < 256 ? (((int)m + 31) / 32) * 32 : 256;
if (th < 32) th = 32;
dim3 grid((unsigned)jb, (unsigned)B);
extract_V_kernel<<<grid, th, 0, 0>>>(Acm.data_ptr<float>(), Vt.data_ptr<float>(),
n, (int)j, (int)jb, (int)m);
return Vt;
}
torch::Tensor build_T(torch::Tensor Acm, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads) {
TORCH_CHECK(Acm.is_cuda() && tau.is_cuda(), "cuda required");
TORCH_CHECK(Acm.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
int B = (int)Acm.size(0);
int n = (int)Acm.size(1);
auto T = torch::empty({B, jb, jb}, Acm.options());
const float* ap = Acm.data_ptr<float>();
const float* tp = tau.data_ptr<float>();
float* Tp = T.data_ptr<float>();
size_t shbytes = (size_t)(2 * jb * jb + jb) * sizeof(float);
#define LAUNCH_T(TH) build_T_kernel<TH><<<B, TH, shbytes, 0>>>(ap, tp, Tp, n, (int)j, (int)jb);
switch ((int)threads) {
case 128: LAUNCH_T(128); break;
case 512: LAUNCH_T(512); break;
default: LAUNCH_T(256); break;
}
#undef LAUNCH_T
return T;
}
// NEGATIVE RESULT (profiled on Ada, kept disabled): routing panel_factor to
// panel_factor_smem_kernel was a wash-to-loss (+0.5% @512/640, +1.3% @1024/60).
// Caching the jb*(n-j) panel needs ~64KB shared -> ~1 block/SM, which destroys
// the occupancy that makes the high-batch shapes fast. The panel is bound by its
// serial jb-step Householder dependency chain + per-step block syncs, not global
// bandwidth, so the staged-in-shared traffic savings don't pay for the occupancy
// loss. The global (near-zero-shared) kernel below stays the production path.
// Shared-memory cap for the staged panel (bytes). B200 allows up to ~227KB
// dynamic shared per block; stay under to keep >=1 block resident.
#define PANEL_SMEM_CAP (192 * 1024)
void panel_factor(torch::Tensor R, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads) {
TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
int B = (int)R.size(0);
int n = (int)R.size(1);
float* rp = R.data_ptr<float>();
float* tp = tau.data_ptr<float>();
// Stage the whole panel in shared (read once / write once instead of jb x
// through global) when it fits. fp32 apply -> lower register pressure helps
// the occupancy that previously made this a wash on Ada.
size_t shbytes = (size_t)jb * (n - j) * sizeof(float);
if (shbytes <= PANEL_SMEM_CAP) {
#define LAUNCH_SMEM(TH) \
cudaFuncSetAttribute(panel_factor_smem_kernel<TH>, \
cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shbytes); \
panel_factor_smem_kernel<TH><<<B, TH, shbytes, 0>>>(rp, tp, n, (int)j, (int)jb);
switch ((int)threads) {
case 64: LAUNCH_SMEM(64); break;
case 128: LAUNCH_SMEM(128); break;
case 512: LAUNCH_SMEM(512); break;
case 1024: LAUNCH_SMEM(1024); break;
default: LAUNCH_SMEM(256); break;
}
#undef LAUNCH_SMEM
return;
}
switch ((int)threads) {
case 64: panel_factor_kernel<64><<<B, 64, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
case 128: panel_factor_kernel<128><<<B, 128, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
case 512: panel_factor_kernel<512><<<B, 512, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
case 1024: panel_factor_kernel<1024><<<B, 1024, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
default: panel_factor_kernel<256><<<B, 256, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
}
}
// Blocked Householder QR host loop in C++: one Python entry, no per-panel GIL/dispatch.
// Acm is column-major (B,n,n); returns (H row-major, tau). Trailing GEMMs use ATen bmm
// (same cuBLAS path as Python torch.bmm). kernel_T path uses build_T for jb<=16.
std::vector<torch::Tensor> qr_blocked(torch::Tensor Acm, int64_t block, int64_t threads) {
TORCH_CHECK(Acm.is_cuda(), "cuda required");
TORCH_CHECK(Acm.scalar_type() == at::kFloat, "fp32 required");
const int B = (int)Acm.size(0);
const int n = (int)Acm.size(1);
const int blk = (int)block;
const int th = (int)threads;
const bool kernel_T = blk <= 16;
auto tau = torch::zeros({B, n}, Acm.options());
int j = 0;
while (j < n) {
const int jb = std::min(blk, n - j);
panel_factor(Acm, tau, j, jb, th);
const int c2 = j + jb;
if (c2 < n) {
const int m = n - j;
auto Vt = extract_V(Acm, j, jb, m);
torch::Tensor T;
if (kernel_T) {
T = build_T(Acm, tau, j, jb, th);
} else {
// G = V^T V on tensor cores, then T via the one-launch parallel
// column-solve kernel (replaces triu + 2 ew + diag_embed + trsm).
auto G = at::bmm(Vt, Vt.transpose(1, 2));
T = build_T_from_G(G, tau, j, jb, 512);
}
auto At = Acm.narrow(1, c2, n - c2).narrow(2, j, m);
auto W1 = at::bmm(Vt, At.transpose(1, 2));
auto W2 = at::bmm(T.transpose(1, 2), W1);
// Fuse trailing subtract into the GEMM: At = 1*At + (-1)*(W2^T @ Vt)
// in one cuBLAS call -- no materialized temp, no separate elementwise
// pass over the trailing matrix (that aten::sub_ was 23-46% of CUDA time).
At.baddbmm_(W2.transpose(1, 2), Vt, /*beta=*/1.0, /*alpha=*/-1.0);
}
j = c2;
}
auto H = Acm.transpose(-2, -1).contiguous();
return {H, tau};
}
// Two-level blocked QR: super-panel SW, inner sub-panel IB. Decouples panel width
// (IB -> cheap O(IB^2 m) apply) from far-trailing GEMM width (SW). Factor the
// super-panel in IB chunks applying each chunk's WY only WITHIN [J,J+W), then apply
// the merged SW-wide WY to the FAR trailing [J+W,n) once (M=SW). vs single-level
// this fattens the far-GEMM M and halves (SW/IB x) its count. Stored reflectors are
// identical to a single SW-wide panel, so (H,tau) is unchanged.
std::vector<torch::Tensor> qr_blocked2(torch::Tensor Acm, int64_t superblk,
int64_t inner, int64_t threads) {
TORCH_CHECK(Acm.is_cuda(), "cuda required");
TORCH_CHECK(Acm.scalar_type() == at::kFloat, "fp32 required");
const int n = (int)Acm.size(1);
const int SW = (int)superblk;
const int IB = (int)inner;
const int th = (int)threads;
auto tau = torch::zeros({(int)Acm.size(0), n}, Acm.options());
auto eye_full = torch::eye(SW, Acm.options());
int J = 0;
while (J < n) {
const int W = std::min(SW, n - J);
int j = J;
while (j < J + W) {
const int jb = std::min(IB, J + W - j);
panel_factor(Acm, tau, j, jb, th);
const int c2 = j + jb;
if (c2 < J + W) { // within-super-panel update only
const int m = n - j;
auto Vt = extract_V(Acm, j, jb, m);
auto G = at::bmm(Vt, Vt.transpose(1, 2));
auto T = build_T_from_G(G, tau, j, jb, 512);
auto At = Acm.narrow(1, c2, (J + W) - c2).narrow(2, j, m);
auto W1 = at::bmm(Vt, At.transpose(1, 2));
auto W2 = at::bmm(T.transpose(1, 2), W1);
At.baddbmm_(W2.transpose(1, 2), Vt, 1.0, -1.0);
}
j = c2;
}
const int c3 = J + W;
if (c3 < n) { // merged SW-wide far-trailing update
const int m = n - J;
auto Vt = extract_V(Acm, J, W, m);
auto G = at::bmm(Vt, Vt.transpose(1, 2));
// wide far-panel (W=64): cuBLAS batched trsm beats the in-kernel
// back-sub at high batch (33KB smem caps occupancy; 63-step serial
// chain). The narrow jb<=32 T-builds above use the fast kernel.
torch::Tensor T;
if (W <= 32) {
T = build_T_from_G(G, tau, J, W, 512);
} else {
auto tau_blk = tau.narrow(1, J, W);
auto U = G.triu(1);
auto eye_jb = eye_full.narrow(0, 0, W).narrow(1, 0, W);
auto Amat = eye_jb + tau_blk.unsqueeze(2) * U;
auto Dmat = at::diag_embed(tau_blk);
T = at::linalg_solve_triangular(Amat, Dmat, true, true, true);
}
auto At = Acm.narrow(1, c3, n - c3).narrow(2, J, m);
auto W1 = at::bmm(Vt, At.transpose(1, 2));
auto W2 = at::bmm(T.transpose(1, 2), W1);
At.baddbmm_(W2.transpose(1, 2), Vt, 1.0, -1.0);
}
J += W;
}
auto H = Acm.transpose(-2, -1).contiguous();
return {H, tau};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("qr_blocked2", &qr_blocked2, "two-level blocked QR (super-panel + inner)");
m.def("panel_factor", &panel_factor, "batched householder panel factor (col-major)");
m.def("qr_full", &qr_full, "fully-fused batched householder QR (shared-mem, col-major)");
m.def("qr_block_fused", &qr_block_fused, "blocked fully-fused batched householder QR (one launch)");
m.def("build_T", &build_T, "compact-WY block reflector T from panel reflectors (dlarft)");
m.def("build_T_from_G", &build_T_from_G, "compact-WY T from G=V^T V (parallel column solve)");
m.def("extract_V", &extract_V, "emit clean compact-WY V panel (unit diag, reflectors below)");
m.def("panel_factor_coop", &panel_factor_coop, "cooperative multi-block-per-matrix panel factor");
m.def("qr_blocked", &qr_blocked, "blocked QR host loop in C++ (panel+T+trailing)");
m.def("set_cublas_emulation", &set_cublas_emulation, "enable cuBLAS fp32 TC emulation");
}
'''
# Trailing update precision for the WY trailing GEMMs (W1, W2, final update).
#
# "fp32" -- plain torch.bmm in fp32 (CUDA-core path off Hopper/Blackwell,
# tensor-core fp32-emulation on newer cuBLAS where available).
# Proven correct on all 70 (shape, case) combinations (Ada).
#
# "3xtf32" -- manual fp32-emulation via three TF32 tensor-core matmuls
# (split each operand into hi/lo TF32 parts and sum
# hi@hi + hi@lo + lo@hi, dropping the lo@lo term). Recovers
# ~2^-20 relative precision (vs ~2^-10 for plain TF32), which is
# well inside every factor/orthogonality gate (looser than
# ~2^-17 at n=4096). On Ada, naive 3xTF32 measured 1.6-3x SLOWER
# because TF32 tensor cores aren't fast enough relative to fp32
# CUDA cores to amortize 3x the matmuls + the Dekker-split
# overhead. On B200, TF32 tensor-core throughput is far higher
# relative to fp32, so this is the primary untested B200 lever
# for the dominant trailing GEMM (n=512 b=640 etc).
#
# Default is "fp32" (safe, validated). 3xtf32 measured 1.5-1.6x SLOWER on B200
# (512: 17.8->29.1ms) -- confirms plain fp32 bmm already uses fast hardware, so
# the trailing GEMM is not the bottleneck.
_TRAILING_PRECISION = "fp32"
# int32 bit-mask keeping sign(1) + exponent(8) + top 10 mantissa bits (TF32).
_TF32_MASK = -8192 # 0xFFFFE000 as a signed int32
def _round_tf32(x: torch.Tensor) -> torch.Tensor:
bits = x.view(torch.int32)
return (bits & _TF32_MASK).view(torch.float32)
def _bmm_3xtf32(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
"""fp32-accurate batched matmul via 3 TF32 tensor-core matmuls.
a = a_hi + a_lo, b = b_hi + b_lo (both splits exact since a_lo/b_lo are the
truncated low mantissa bits of a/b). a@b ~= a_hi@b_hi + a_hi@b_lo + a_lo@b_hi,
dropping the a_lo@b_lo term (relative magnitude ~2^-20).
"""
a_hi = _round_tf32(a)
a_lo = a - a_hi
b_hi = _round_tf32(b)
b_lo = b - b_hi
prev = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
out = torch.bmm(a_hi, b_hi)
out += torch.bmm(a_hi, b_lo)
out += torch.bmm(a_lo, b_hi)
finally:
torch.backends.cuda.matmul.allow_tf32 = prev
return out
def _trail_mm(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
if _TRAILING_PRECISION == "3xtf32":
return _bmm_3xtf32(a, b)
return torch.bmm(a, b)
_ext = None
_ext_failed = False
def _get_ext():
global _ext, _ext_failed
if _ext is not None or _ext_failed:
return _ext
try:
from torch.utils.cpp_extension import load_inline
_ext = load_inline(
name="qr2l_ext",
cpp_sources=_CPP,
cuda_sources=_CUDA,
extra_cuda_cflags=["-O3"],
extra_ldflags=["-lcublas"],
verbose=False,
)
except Exception:
_ext_failed = True
_ext = None
return _ext
_emul_set = False
def _enable_emulation():
"""Turn on cuBLAS BF16x9 fp32 emulation (tensor-core, full fp32 accuracy)
once, so the trailing-update at::bmm/baddbmm run on tensor cores."""
global _emul_set
if _emul_set:
return
ext = _get_ext()
if ext is not None:
try:
ext.set_cublas_emulation(1) # 1 = PERFORMANT (emulate only when faster)
_emul_set = True
except Exception:
pass
if torch.cuda.is_available():
_get_ext()
def _qr_blocked(A, block):
"""Batched blocked Householder QR -> (H, tau) in geqrf-compact form.
Step 1: entire panel loop runs in C++ (qr_blocked) -- one Python entry instead
of ~n/jb iterations each dispatching panel_factor + build_T + 3x bmm.
"""
ext = _get_ext()
_enable_emulation()
Acm = A.transpose(-2, -1).contiguous()
n = A.shape[-1]
threads = _panel_threads(n)
# Run the trailing bmm/baddbmm on TF32 tensor cores (1-pass) only where the
# 20*n*eps factor gate has slack: n in {1024,2048} pass 22/22 on B200 with a
# big speedup (1024 8.4->6.6ms, 2048 18.2->15.5ms). n=512's tighter gate fails
# the mixed-scale case under TF32, so it stays fp32. allow_tf32 is a global
# cuBLAS-handle math-mode flag that at::bmm reads inside the C++ loop -> toggle
# it per shape around the call. (Custom bf16x3 WMMA trailing was a dead end:
# 2.3x slower than cuBLAS; see qr-v2-perf-findings.)
prev_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = (1024 <= n <= 2048)
try:
if block < 0: # 2-level: SW-wide merged far-trailing GEMM (M=SW) + IB=32 panel.
# SW=128 was rejected pre-fusion (within-super-panel T-build too costly);
# re-testing now that build_T_from_G made the inner T-build cheap.
# per-n: 1024 (launch-bound, b60) wins with SW=128 (fatter M=128 far-GEMM,
# fewer panels: 6.29->6.06ms); 512 (occupancy-bound, b640) regresses at
# SW=128 (10.8ms) so keeps SW=64. B200-measured.
sw = 128 if n > 640 else 64
H, tau = ext.qr_blocked2(Acm, sw, 32, threads)
else:
H, tau = ext.qr_blocked(Acm, block, threads)
finally:
torch.backends.cuda.matmul.allow_tf32 = prev_tf32
return H, tau
# Experimental cooperative-panel blocked QR for the GPU-starved low-batch large-n
# shapes (2048 b8, 4096 b2): G thread blocks cooperate on each matrix's panel so the
# panel phase fills the GPU instead of using only `batch` SMs. Off by default; set
# _COOP_N to the min n that routes here (e.g. 2048) for a B200 A/B. _COOP_G is the
# blocks-per-matrix (B*G must stay co-resident for cudaLaunchCooperativeKernel).
#
# MEASURED DEAD END (B200 benchmark, G=8): 2048 b8 = 63.7ms (vs ~18ms baseline),
# 4096 b2 = 152ms (vs 52ms geqrf). The panel is sequential over its reflector columns,
# so a multi-block panel needs a grid-sync PER column (4 here) -> O(n) barriers
# (~16k at 4096). That barrier cost dominates and scales the WRONG way with G (Ada:
# G=8 166ms, G=16 309ms). The serial column dependency can't be cheaply parallelized
# across blocks; cuSOLVER geqrf's structure wins. Kept off (_COOP_N=None).
_COOP_N = None
_COOP_G = 8
_COOP_BLK = 32
def _qr_blocked_coop(A, block=None, G=None):
"""Blocked Householder QR with a COOPERATIVE multi-block panel. Mirrors the C++
qr_blocked loop (panel -> compact-WY T -> trailing baddbmm) but factors each
panel with G blocks/matrix (ext.panel_factor_coop) to fill the GPU at low batch.
Python panel loop (n/jb iters) -- fine for the large-n shapes (few panels)."""
ext = _get_ext()
_enable_emulation()
if block is None:
block = _COOP_BLK
if G is None:
G = _COOP_G
Acm = A.transpose(-2, -1).contiguous()
B, n, _ = A.shape
th = _panel_threads(n)
tau = torch.zeros(B, n, device=A.device, dtype=torch.float32)
eye_full = torch.eye(block, device=A.device, dtype=torch.float32)
prev_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = (1024 <= n <= 2048)
try:
j = 0
while j < n:
jb = min(block, n - j)
ext.panel_factor_coop(Acm, tau, j, jb, G)
c2 = j + jb
if c2 < n:
m = n - j
Vt = ext.extract_V(Acm, j, jb, m)
tau_blk = tau.narrow(1, j, jb)
Gmat = torch.bmm(Vt, Vt.transpose(1, 2))
U = Gmat.triu(1)
eye_jb = eye_full.narrow(0, 0, jb).narrow(1, 0, jb)
Amat = eye_jb + tau_blk.unsqueeze(2) * U
Dmat = torch.diag_embed(tau_blk)
T = torch.linalg.solve_triangular(Amat, Dmat, upper=True,
left=True, unitriangular=True)
At = Acm.narrow(1, c2, n - c2).narrow(2, j, m)
W1 = torch.bmm(Vt, At.transpose(1, 2))
W2 = torch.bmm(T.transpose(1, 2), W1)
At.baddbmm_(W2.transpose(1, 2), Vt, beta=1.0, alpha=-1.0)
j = c2
finally:
torch.backends.cuda.matmul.allow_tf32 = prev_tf32
H = Acm.transpose(-2, -1).contiguous()
return H, tau
_FUSED_MAX = 200 # n*n*4 bytes must fit the opted-in dynamic shared (B200 ~227KB)
def _fused_threads(n):
if n <= 64:
return 128
return 512
def _qr_fused(A):
"""Whole-matrix Householder QR in a single fused shared-memory kernel: one
threadblock factors one matrix end-to-end, no Python panel loop / trailing
GEMM round trips. Only viable while the matrix fits in dynamic shared mem."""
ext = _get_ext()
B, n, _ = A.shape
Acm = A.transpose(-2, -1).contiguous()
tau = torch.zeros(B, n, device=A.device, dtype=torch.float32)
ext.qr_full(Acm, tau, _fused_threads(n))
H = Acm.transpose(-2, -1).contiguous()
return H, tau
_BLOCKFUSED_MAX = 1024 # one block/matrix needs decent batch for SM occupancy
def _qr_block_fused(A, blk):
"""Blocked fully-fused QR: one kernel launch does the entire factorization
(panel in shared, trailing columns read once each). No host panel loop,
no cuBLAS, no T-build, no V-extraction copies."""
ext = _get_ext()
B, n, _ = A.shape
Acm = A.transpose(-2, -1).contiguous()
tau = torch.zeros(B, n, device=A.device, dtype=torch.float32)
ext.qr_block_fused(Acm, tau, blk, 256)
H = Acm.transpose(-2, -1).contiguous()
return H, tau
def _panel_threads(n):
"""Block size (blockDim.x) for the panel kernel, per n. Reductions run over
sub-columns of length <= n, so match the thread count to n: tiny fused QR
(n<=64) wants a single/double warp (more matrices packed per SM, minimal
reduction latency); the large compute-bound shapes want wide blocks so each
matrix's panel finishes faster. Templated -> compile-time-constant strides."""
if n <= 64:
return 64
if n <= 384:
return 512 # 352 (low batch 40): 16 warps cut the jb=32 apply waves;
# occupancy is a non-issue at b40. B200: 1.91 -> 1.78ms.
if n <= 512:
return 256 # 512 (high batch 640): 8 warps -> 8 blocks/SM beats the
# 16-warp apply (512 threads measured ~5% SLOWER at b640).
if n <= 1024:
return 512 # lever #4: 1024 gains from 16 warps (1024 thr measured
# identical 10.8ms -> cuBLAS-bound, keep proven 512)
return 1024 # 2048 (b8): no occupancy pressure (8 blocks).
# B200 A/B: 512thr=23.3ms -> 1024thr=22.8ms (~2% faster).
def _plan(batch, n):
"""Return (use_blocked, block). Use the batched kernel only where it can
fill the GPU and beat cuSOLVER; otherwise fall back to torch.geqrf."""
if n <= 64:
# Tiny matrices: geqrf wastes ~0.3ms of launch overhead per call. Do the
# whole QR as a single fused panel (block = n -> one kernel launch, one
# threadblock per matrix, no trailing GEMM) when there are enough
# matrices to populate the SMs.
if batch >= 8:
return True, n
return False, 0
if n > 2048:
return False, 0 # n=4096: blocked panel starves at batch 2 (2 blocks
# -> 2 of ~148 SMs). Re-measured with baddbmm + jb=32:
# blocked=88.4ms vs geqrf=52.4ms. geqrf still wins (the
# cuBLAS trailing is already fast; only the serial panel
# is starved, and only HW clusters could unstarve it).
if batch < 8:
return False, 0 # too few matrices to fill the GPU
# Per-n block size (validated with the warp-parallel panel apply).
if n <= 352:
block = 32 # 352: M=16->32 trailing-GEMM-fill fix.
# B200 A/B: jb16=2.13ms -> jb32=1.91ms (~10% faster).
elif n <= 640:
block = -1 # 512: 2-level (SW=64 merged far-GEMM, IB=32 panel).
# plain jb=64 was slower (wide panel); 2-level keeps
# the cheap jb=32 panel but fattens the far-trailing GEMM.
# B200 A/B: SW=128 was slower (9.83->10.8ms) -- the extra
# within-super-panel apply outweighs the far-GEMM savings.
elif n <= 1280:
block = -1 # 1024: 2-level (SW=64 merged far-GEMM, IB=32 panel).
# B200 A/B: single-level jb=32 8.92ms -> 2-level 8.55ms
# (~4% faster); the M=64 far-trailing GEMM + halved
# far-GEMM count beats the extra within-super-panel ops.
# Re-confirmed post-TF32: single-level jb=32 = 7.22ms
# vs 2-level 6.58ms -> 2-level still wins.
else:
block = 32 # 2048: single-level jb=32. 2-level measured SLOWER
# (18.4 -> 18.8ms): batch 8 is launch-sensitive, the
# extra within-super-panel ops outweigh the M=64 win.
# B200 A/B: jb16=25.0ms, jb32=23.4ms; jb64=31.0ms.
return True, block
def custom_kernel(data):
A = data
if A.dtype != torch.float32:
A = A.to(torch.float32)
if not A.is_cuda:
raise RuntimeError("custom_kernel requires a CUDA tensor")
A = A.contiguous()
batch, n, _ = A.shape
ext = _get_ext()
if ext is not None and 2 <= n <= _FUSED_MAX:
try:
H, tau = _qr_fused(A)
return H.contiguous(), tau.contiguous()
except Exception:
pass # fall through to blocked / reference on any failure
# NOTE: a one-block-per-matrix blocked megakernel (qr_block_fused) was tried
# for 352/512/1024 to kill host overhead. It is correct but 2-6x SLOWER: the
# in-kernel trailing update (CUDA-core warps) cannot match cuBLAS tensor-core
# GEMM, which dominates these shapes. So large-n keeps the cuBLAS trailing.
# Experimental: cooperative-panel blocked QR for starved low-batch large-n.
if _COOP_N is not None and ext is not None and n >= _COOP_N:
try:
H, tau = _qr_blocked_coop(A)
return H.contiguous(), tau.contiguous()
except Exception:
pass
use_blocked, block = _plan(batch, n)
if use_blocked and ext is not None:
try:
H, tau = _qr_blocked(A, block)
return H.contiguous(), tau.contiguous()
except Exception:
pass # fall through to the reference path on any failure
H, tau = torch.geqrf(A)
return H.contiguous(), tau.contiguous()
scrolls · 1359 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