Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
6.25ms
#216 of 515
2026-06-24

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-memoryextern __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