submission 833233
deepestseeker_80318 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 403 lines, June 9 Researcher Reciprocity License v1.0.
submission_tsqr.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-833233?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:0c88c41398cfe7438202080a1541c27816f7f8da0f47ea13cec7a2007b8a0f8b
license declaredunknown
license concludedunknown
authorsdeepestseeker_80318
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float sm[];Kernel source
submission_tsqr.py403 lines
"""
GPU MODE #774 qr_v2 — TSQR-1: 2-level TSQR + Householder reconstruction for n>=2048. Custom local_qr (factor+orgqr) kernel; torch combine/assembly/recon (cheap for b=2/8). Validates the local_qr kernel + pipeline. Math verified CPU 5/5.
CUDA-3/6/7 showed panel impl & cheap-op dispatch don't move ~6,300. But nb=32
(2x more solve_triangular + GEMMs) WAS slower -> the per-panel cuSOLVER
triangular-solve and/or trailing GEMMs are the cost. cuSOLVER was catastrophic
for batched-small (geqrf). So CUDA-8 computes T (and Y) entirely in the CUDA
panel kernel (G=Y^T Y + WY T-recurrence in shared memory, verified 1e-7 vs
solve_triangular), eliminating cuSOLVER. Per panel: 1 CUDA kernel + 3 cuBLAS
GEMMs. n<=512 CUDA, n=1024 Triton(+torch T-solve), n>=2048 geqrf.
"""
import torch
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
try:
torch.set_float32_matmul_precision("highest")
except Exception:
pass
_NB = 64
_TF32_MIN_N = 512
_CUDA = None
if torch.cuda.is_available():
try:
from torch.utils.cpp_extension import load_inline
_src = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <math.h>
__device__ __forceinline__ float warpRed(float v){
for(int o=16;o>0;o>>=1) v += __shfl_down_sync(0xffffffffu, v, o);
return v;
}
__device__ __forceinline__ float blockRed1(float v, float* s){
int tid=threadIdx.x, lane=tid&31, wid=tid>>5;
v = warpRed(v);
if(lane==0) s[wid]=v;
__syncthreads();
int nw=(blockDim.x+31)>>5;
v = (tid<nw)? s[tid] : 0.f;
if(wid==0) v = warpRed(v);
if(tid==0) s[0]=v;
__syncthreads();
return s[0];
}
__global__ void panel_sh(float* __restrict__ H, float* __restrict__ TAU,
float* __restrict__ YB, float* __restrict__ TB,
int n, int j, int jb, int ydim){
int pid=blockIdx.x, tid=threadIdx.x, T=blockDim.x;
int m=n-j, do_trail=(j+jb)<n;
long base=(long)pid*n*n + (long)j*n + j;
float* Hp=H+base; float* taup=TAU+(long)pid*n+j;
extern __shared__ float sm[];
float* P=sm;
float* Gsh=P+(long)m*jb;
float* Tsh=Gsh+(long)jb*jb;
float* taush=Tsh+(long)jb*jb;
float* red=taush+jb;
__shared__ float beta_s, tau_s, inv_s;
for(long i=tid;i<(long)m*jb;i+=T){ int r=i/jb; int c=i-(long)r*jb; P[i]=Hp[(long)r*n+c]; }
__syncthreads();
for(int k=0;k<jb;k++){
float ps=0.f;
for(int r=k+1+tid;r<m;r+=T){ float x=P[(long)r*jb+k]; ps+=x*x; }
float sigma=blockRed1(ps,red);
if(tid==0){
float alpha=P[(long)k*jb+k];
float normf=sqrtf(alpha*alpha+sigma);
float sign=(alpha>=0.f)?1.f:-1.f;
float beta=-sign*normf, tau, inv;
if(sigma==0.f){ beta=alpha; tau=0.f; inv=0.f; }
else { tau=(beta-alpha)/beta; inv=1.f/(alpha-beta); }
beta_s=beta; tau_s=tau; inv_s=inv;
P[(long)k*jb+k]=beta; taup[k]=tau; taush[k]=tau;
}
__syncthreads();
float tau=tau_s, inv=inv_s;
for(int r=k+1+tid;r<m;r+=T){ P[(long)r*jb+k]*=inv; }
__syncthreads();
for(int c=k+1;c<jb;c++){
float pw=0.f;
for(int r=k+tid;r<m;r+=T){ float vr=(r==k)?1.f:P[(long)r*jb+k]; pw+=vr*P[(long)r*jb+c]; }
float w=blockRed1(pw,red);
float tw=tau*w;
for(int r=k+tid;r<m;r+=T){ float vr=(r==k)?1.f:P[(long)r*jb+k]; P[(long)r*jb+c]-=tw*vr; }
__syncthreads();
}
}
// writeback H + Y
float* Yp=YB+(long)pid*(long)n*ydim;
for(long i=tid;i<(long)m*jb;i+=T){ int r=i/jb; int c=i-(long)r*jb;
Hp[(long)r*n+c]=P[i];
if(do_trail){ float y=(r==c)?1.f:((r>c)?P[i]:0.f); Yp[(long)(j+r)*ydim+c]=y; }
}
if(!do_trail) return;
__syncthreads();
// G = Y^T Y (upper triangle a<=b)
for(int idx=tid;idx<jb*jb;idx+=T){
int a=idx/jb, b=idx-a*jb;
if(a>b){ Gsh[idx]=0.f; continue; }
float s=0.f;
for(int r=b;r<m;r++){
float ya=(r==a)?1.f:P[(long)r*jb+a];
float yb=(r==b)?1.f:P[(long)r*jb+b];
s+=ya*yb;
}
Gsh[idx]=s;
}
__syncthreads();
for(int idx=tid;idx<jb*jb;idx+=T) Tsh[idx]=0.f;
__syncthreads();
for(int k=0;k<jb;k++){
if(tid==0) Tsh[(long)k*jb+k]=taush[k];
__syncthreads();
for(int i=tid;i<k;i+=T){
float s=0.f;
for(int l=i;l<k;l++) s+=Tsh[(long)i*jb+l]*Gsh[(long)l*jb+k];
Tsh[(long)i*jb+k]=-taush[k]*s;
}
__syncthreads();
}
float* Tp=TB+(long)pid*jb*jb;
for(int idx=tid;idx<jb*jb;idx+=T) Tp[idx]=Tsh[idx];
}
// TSQR local block: factor BR-row x nb panel (Householder) + orgqr -> Q_local(nr x nb) + R_i(nb x nb).
// grid=(b,P), one block per (matrix,row-block). Reuses blockRed1.
__global__ void local_qr(const float* __restrict__ H, float* __restrict__ Qloc,
float* __restrict__ Rst, int n,int j,int nb,int BR,int P,int m){
int bi=blockIdx.x, pi=blockIdx.y, tid=threadIdx.x, T=blockDim.x;
int r0=pi*BR, r1=r0+BR; if(r1>m)r1=m; int nr=r1-r0; if(nr<=0) return;
extern __shared__ float sm[];
float* A=sm; float* Q=A+(long)nr*nb; float* taush=Q+(long)nr*nb; float* red=taush+nb;
__shared__ float ts,is_;
long mpad=(long)P*BR;
const float* Hp=H+(long)bi*n*n+(long)(j+r0)*n+j;
float* Qlp=Qloc+(long)bi*mpad*nb+(long)r0*nb;
float* Rp=Rst+((long)bi*P+pi)*nb*nb;
for(long i=tid;i<(long)nr*nb;i+=T){ int rr=i/nb,c=i-(long)rr*nb; A[i]=Hp[(long)rr*n+c]; }
__syncthreads();
for(int k=0;k<nb;k++){
float ps=0.f;
for(int r=k+1+tid;r<nr;r+=T){ float x=A[(long)r*nb+k]; ps+=x*x; }
float sg=blockRed1(ps,red);
if(tid==0){ float a=A[(long)k*nb+k]; float nf=sqrtf(a*a+sg); float sn=(a>=0.f)?1.f:-1.f;
float be=-sn*nf,tau,inv; if(sg==0.f){be=a;tau=0.f;inv=0.f;}else{tau=(be-a)/be;inv=1.f/(a-be);}
ts=tau;is_=inv; A[(long)k*nb+k]=be; taush[k]=tau; }
__syncthreads(); float tau=ts,inv=is_;
for(int r=k+1+tid;r<nr;r+=T) A[(long)r*nb+k]*=inv;
__syncthreads();
for(int c=k+1;c<nb;c++){
float pw=0.f;
for(int r=k+tid;r<nr;r+=T){ float vr=(r==k)?1.f:A[(long)r*nb+k]; pw+=vr*A[(long)r*nb+c]; }
float w=blockRed1(pw,red); float tw=tau*w;
for(int r=k+tid;r<nr;r+=T){ float vr=(r==k)?1.f:A[(long)r*nb+k]; A[(long)r*nb+c]-=tw*vr; }
__syncthreads();
}
}
for(long i=tid;i<(long)nb*nb;i+=T){ int rr=i/nb,c=i-(long)rr*nb;
Rp[i]=(rr<=c && rr<nr)? A[(long)rr*nb+c] : 0.f; }
for(long i=tid;i<(long)nr*nb;i+=T){ int rr=i/nb,c=i-(long)rr*nb; Q[i]=(rr==c)?1.f:0.f; }
__syncthreads();
for(int k=nb-1;k>=0;k--){
float tk=taush[k];
if(tk!=0.f){
for(int c=0;c<nb;c++){
float pw=0.f;
for(int r=k+tid;r<nr;r+=T){ float vr=(r==k)?1.f:A[(long)r*nb+k]; pw+=vr*Q[(long)r*nb+c]; }
float w=blockRed1(pw,red); float tw=tk*w;
for(int r=k+tid;r<nr;r+=T){ float vr=(r==k)?1.f:A[(long)r*nb+k]; Q[(long)r*nb+c]-=tw*vr; }
__syncthreads();
}
}
}
for(long i=tid;i<(long)nr*nb;i+=T){ Qlp[i]=Q[i]; }
}
void local_qr_launch(torch::Tensor H, torch::Tensor Qloc, torch::Tensor Rst,
int64_t n,int64_t j,int64_t nb,int64_t BR,int64_t P,int64_t m){
int b=H.size(0), T=256;
size_t shmem=((size_t)2*BR*nb + nb + 128)*sizeof(float);
static bool st=false; if(!st){ cudaFuncSetAttribute(local_qr,cudaFuncAttributeMaxDynamicSharedMemorySize,160*1024); st=true; }
dim3 grid(b,(unsigned)P);
local_qr<<<grid,T,shmem>>>(H.data_ptr<float>(),Qloc.data_ptr<float>(),Rst.data_ptr<float>(),
(int)n,(int)j,(int)nb,(int)BR,(int)P,(int)m);
}
void panel_sh_launch(torch::Tensor H, torch::Tensor TAU, torch::Tensor YB, torch::Tensor TB,
int64_t n, int64_t j, int64_t jb, int64_t ydim){
int b=H.size(0); int m=(int)(n-j); int T=256;
size_t shmem=((size_t)m*jb + 2*(size_t)jb*jb + jb + 128)*sizeof(float);
static bool set=false;
if(!set){ cudaFuncSetAttribute(panel_sh, cudaFuncAttributeMaxDynamicSharedMemorySize, 200*1024); set=true; }
panel_sh<<<b,T,shmem>>>(H.data_ptr<float>(),TAU.data_ptr<float>(),YB.data_ptr<float>(),
TB.data_ptr<float>(),(int)n,(int)j,(int)jb,(int)ydim);
}
'''
_CUDA = load_inline(name='qr_cuda8', cpp_sources='', cuda_sources=_src,
functions=['panel_sh_launch','local_qr_launch'], verbose=False)
except Exception:
_CUDA = None
try:
import triton
import triton.language as tl
_HAS_TRITON = torch.cuda.is_available()
except Exception:
_HAS_TRITON = False
def _tf32(on):
torch.backends.cuda.matmul.allow_tf32 = on
def _next_pow2(x):
return 1 << (x - 1).bit_length()
if _HAS_TRITON:
@triton.jit
def _panel_tri(H, TAU, n, j, jb, sb, si, sj, stb, stk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr):
pid = tl.program_id(0)
m = n - j
rows = tl.arange(0, BLOCK_M)
cols = tl.arange(0, BLOCK_N)
rmask = rows < m
cmask = cols < jb
pbase = H + pid * sb + j * si + j * sj
offs = rows[:, None] * si + cols[None, :] * sj
mask = rmask[:, None] & cmask[None, :]
tile = tl.load(pbase + offs, mask=mask, other=0.0)
for k in range(BLOCK_N):
kv = k < jb
col_k = tl.sum(tl.where(cols[None, :] == k, tile, 0.0), axis=1)
is_k = rows == k
below = (rows > k) & (rows < m)
alpha = tl.sum(tl.where(is_k, col_k, 0.0), axis=0)
sigma = tl.sum(tl.where(below, col_k * col_k, 0.0), axis=0)
norm_full = tl.sqrt(alpha * alpha + sigma)
sign = tl.where(alpha >= 0, 1.0, -1.0)
beta = -sign * norm_full
zero = sigma == 0.0
beta = tl.where(zero, alpha, beta)
denom = alpha - beta
safe = tl.where(zero, 1.0, denom)
tau_k = tl.where(zero, 0.0, (beta - alpha) / beta)
v_tail = col_k / safe
v = tl.where(is_k, 1.0, tl.where(below, v_tail, 0.0))
newcol = tl.where(is_k, beta, tl.where(below, v_tail, col_k))
tile = tl.where(cols[None, :] == k, newcol[:, None], tile)
w = tl.sum(v[:, None] * tile, axis=0)
update = tau_k * v[:, None] * w[None, :]
upd_mask = (cols[None, :] > k) & rmask[:, None]
tile = tile - tl.where(upd_mask, update, 0.0)
tl.store(TAU + pid * stb + (j + k) * stk, tau_k, mask=kv)
tl.store(pbase + offs, tile, mask=mask)
def _form_T_solve(Y, tau, eye):
G = Y.transpose(-1, -2) @ Y
M = eye + torch.triu(G, 1) * tau.unsqueeze(1)
Dl = torch.diag_embed(tau)
Tt = torch.linalg.solve_triangular(M.transpose(-1, -2), Dl, upper=False, unitriangular=True)
return Tt.transpose(-1, -2)
def _qr_cuda(A, tau_out, use_tf32):
b, n, _ = A.shape
H = A.contiguous()
YB = torch.empty(b, n, _NB, device=H.device, dtype=H.dtype)
TB = torch.empty(b, _NB, _NB, device=H.device, dtype=H.dtype)
j = 0
while j < n:
jb = min(_NB, n - j)
_CUDA.panel_sh_launch(H, tau_out, YB, TB, n, j, jb, _NB)
if j + jb < n:
Yp = YB[:, j:, :jb]
T = TB[:, :jb, :jb]
C = H[:, j:, j + jb:]
Yt = Yp.transpose(-1, -2)
Tt = T.transpose(-1, -2)
if use_tf32:
_tf32(True); W = Yt @ C
_tf32(False); W = Tt @ W
_tf32(True); YW = Yp @ W
_tf32(False); C -= YW
else:
W = Yt @ C; W = Tt @ W; C -= Yp @ W
j += jb
return H, tau_out
def _qr_tri(A, tau_out, use_tf32):
b, n, _ = A.shape
H = A.contiguous()
eye = torch.eye(_NB, device=H.device).unsqueeze(0).expand(b, _NB, _NB)
j = 0
while j < n:
jb = min(_NB, n - j)
m = n - j
BLOCK_M = _next_pow2(m)
nwarps = max(4, BLOCK_M // 64)
sb, si, sj = H.stride()
stb, stk = tau_out.stride()
_panel_tri[(b,)](H, tau_out, n, j, jb, sb, si, sj, stb, stk,
BLOCK_M=BLOCK_M, BLOCK_N=_NB, num_warps=nwarps)
if j + jb < n:
panel = H[:, j:, j:j + jb]
Yp = torch.tril(panel, diagonal=-1).clone()
idx = torch.arange(jb, device=H.device)
Yp[:, idx, idx] = 1.0
ev = eye if jb == _NB else eye[:, :jb, :jb]
T = _form_T_solve(Yp, tau_out[:, j:j + jb], ev)
C = H[:, j:, j + jb:]
Yt = Yp.transpose(-1, -2)
Tt = T.transpose(-1, -2)
if use_tf32:
_tf32(True); W = Yt @ C
_tf32(False); W = Tt @ W
_tf32(True); YW = Yp @ W
_tf32(False); C -= YW
else:
W = Yt @ C; W = Tt @ W; C -= Yp @ W
j += jb
return H, tau_out
def _batched_lu(negW, jb):
b, m, nb = negW.shape
Y = torch.zeros_like(negW); M = negW.clone()
for k in range(jb):
piv = M[:, k, k]; Y[:, k, k] = 1.0
mask = piv.abs() > 1e-30
safe = torch.where(mask, piv, torch.ones_like(piv))
col = M[:, k + 1:, k] / safe.unsqueeze(1)
col = torch.where(mask.unsqueeze(1), col, torch.zeros_like(col))
Y[:, k + 1:, k] = col
M[:, k + 1:, k:] = M[:, k + 1:, k:] - Y[:, k + 1:, k:k + 1] * M[:, k:k + 1, k:]
return Y
def _qr_tsqr(A, tau_out, use_tf32):
b, n, _ = A.shape
H = A.contiguous(); BR = 64
eyeN = torch.eye(_NB, device=H.device, dtype=H.dtype)
j = 0
while j < n:
jb = min(_NB, n - j); m = n - j
P = (m + BR - 1) // BR
Qloc = torch.zeros(b, P * BR, jb, device=H.device, dtype=H.dtype)
Rst = torch.zeros(b, P, jb, jb, device=H.device, dtype=H.dtype)
_CUDA.local_qr_launch(H, Qloc, Rst, n, j, jb, BR, P, m)
Rstack = Rst.reshape(b, P * jb, jb)
Qc, Rfin = torch.linalg.qr(Rstack) # combine (torch for now)
Qcb = Qc.reshape(b, P, jb, jb)
Qlocb = Qloc.reshape(b, P, BR, jb)
Qthin = torch.matmul(Qlocb, Qcb).reshape(b, P * BR, jb)[:, :m, :]
d = -torch.sign(torch.diagonal(Rfin, dim1=-2, dim2=-1))
d = torch.where(d == 0, torch.ones_like(d), d)
Qd = Qthin * d.unsqueeze(1)
W = Qd.clone(); W[:, :jb, :] = W[:, :jb, :] - eyeN[:jb, :jb].unsqueeze(0)
Y = _batched_lu(-W, jb)
vtv = 1.0 + (torch.tril(Y, -1) ** 2).sum(1)
tau = torch.where(vtv > 1e-20, 2.0 / vtv, torch.zeros_like(vtv))
Rg = d.unsqueeze(-1) * Rfin
newp = torch.tril(Y, -1)
newp[:, :jb, :jb] = newp[:, :jb, :jb] + torch.triu(Rg)
H[:, j:, j:j + jb] = newp
tau_out[:, j:j + jb] = tau
if j + jb < n:
Yp = torch.tril(Y, -1).clone(); idx = torch.arange(jb, device=H.device); Yp[:, idx, idx] = 1.0
ev = eyeN if jb == _NB else eyeN[:jb, :jb]
T = _form_T_solve(Yp, tau, ev.unsqueeze(0).expand(b, jb, jb))
C = H[:, j:, j + jb:]; Yt = Yp.transpose(-1, -2); Tt = T.transpose(-1, -2)
if use_tf32:
_tf32(True); Wm = Yt @ C; _tf32(False); Wm = Tt @ Wm
_tf32(True); YW = Yp @ Wm; _tf32(False); C -= YW
else:
Wm = Yt @ C; Wm = Tt @ Wm; C -= Yp @ Wm
j += jb
return H, tau_out
def custom_kernel(data):
A = data
b, n, _ = A.shape
if n <= 512 and _CUDA is not None:
tau = torch.zeros(b, n, device=A.device, dtype=A.dtype)
return _qr_cuda(A.clone(), tau, n >= _TF32_MIN_N)
if n <= 1024 and _HAS_TRITON:
tau = torch.zeros(b, n, device=A.device, dtype=A.dtype)
return _qr_tri(A.clone(), tau, n >= _TF32_MIN_N)
# n>=2048 (TSQR-1): 2-level TSQR + reconstruction
if _CUDA is not None:
tau = torch.zeros(b, n, device=A.device, dtype=A.dtype)
return _qr_tsqr(A.clone(), tau, n >= _TF32_MIN_N)
hs = []; ts = []
for i in range(b):
h, t = torch.geqrf(A[i]); hs.append(h); ts.append(t)
return torch.stack(hs, 0), torch.stack(ts, 0)
scrolls · 403 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