Skip to content
KernelIndex
Search⌘K

submission 831293

Jesus · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 709 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-831293?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
9.53ms
#281 of 515
2026-06-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c727bf1a0f78ef0ec4cada15da1a9d2654fc07c32d0d48c08545de3277aa2b12
license declaredunknown
license concludedunknown
authorsJesus
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memoryextern __shared__ float smem[];

Kernel source

submission.py709 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = False   # trailing is exact FP32 (the bf16/TF32 split was
                                                # accurate but slower for the WY trailing -- see _split_bmm)

# V4 -- hybrid:
#   n <= _CUDA_MAX_N : custom CUDA one-block-per-matrix unblocked Householder.
#   n  > _CUDA_MAX_N : blocked Householder where the *panel factorization AND the
#     WY matrix T are built inside a CUDA kernel* (one block per matrix), and only
#     the big trailing update is left to cuBLAS (torch.bmm). This removes the
#     ~2n-long Python loop that made the pure-torch blocked version overhead-bound
#     for small-batch large-n (4096 b2 was 790ms, ~99% Python). The trailing GEMM
#     is one large batched matmul -> uses all SMs regardless of batch.
#
# Output = geqrf compact convention (H = R upper + reflectors below, tau coeffs).

# Dispatch (from B200 V4 benchmark):
#   n <= 384         -> CUDA one-block          (32/176/352)
#   384 < n <= 1024  -> blocked-kernel          (512: 40.6ms, 1024: 48.2ms -- big wins)
#   n  > 1024        -> geqrf                    (2048/4096: blocked-kernel's one-block panel
#                                                 starves on b8/b2 -> 476/2143ms; geqrf is 77/52)
_CUDA_MAX_N = 256              # 352 now -> blocked-kernel (trailing on all SMs vs one-block b40)
_BLOCKED_MAX_N = 1024
_NB = 32                       # panel width for the flat blocked path (keeps shared < 48KB)
_USE_RECURSIVE = False         # recursive blocked QR (Idea #1): TESTED, doesn't help 2048/4096
                               # (bottleneck is the one-block panel kernel on 2/8 SMs, not the
                               # within-panel updates/T that recursion GEMM-ifies). Kept for reuse.
_REC_BASE = 16                 # recursion base: small -> cheap leaves, GEMM buildup (all SMs)
_USE_COOP = True               # grid-cooperative panel (beats geqrf on 2048 b8: 70 vs 77ms)
_COOP_MAX_N = 2048             # cluster path for 2048 (45.4ms < geqrf 77); 4096 stays geqrf
                               # (b2 -> only 2 clusters/~32 SMs vs geqrf's 148; coop 110-136 > 52)
_TRAILING_SPLIT = False        # bf16/TF32 3x split (in _split_bmm) is accurate (6/6) but SLOWER
                               # for the WY trailing: 3x GEMMs + tiny K=pb -> tensor cores don't pay.
                               # Plain FP32 bmm wins here. (Kept off; building block reusable later.)
_USE_2LEVEL = False            # Fase A: two-level blocked QR (wide outer block -> K=NB_OUTER large
                               # so the 3xBF16 trailing could use tensor cores). TESTED on B200 (K=64):
                               # REGRESSED 512 15.1->24.6ms, 1024 20.7->27.9ms. torch.bmm 3xBF16 (3
                               # GEMMs + bf16-rounding traffic) + the Gram/larft/internal-update
                               # overhead beat the TF32 benefit. Confirms (K=32 in V8b, K=64 here)
                               # that FP32 cuBLAS bmm is the trailing ceiling with pure-torch prims;
_NB_OUTER = 64                 # a real tensor-core trailing needs a fused bf16-in/fp32-out GEMM
                               # (cublasGemmEx/CUTLASS). Kept dormant for reference (route off).


_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;

// ---- one-block-per-matrix unblocked Householder (small n) ----
__global__ void householder_qr_kernel(float* __restrict__ A,
                                      float* __restrict__ tau, int n) {
    const int b = blockIdx.x, tid = threadIdx.x, nt = blockDim.x;
    float* M = A + (size_t)b * n * n;
    float* T = tau + (size_t)b * n;
    extern __shared__ float smem[];
    float* v = smem; float* red = smem + n;
    __shared__ float s_tau, s_beta, s_denom; __shared__ int s_active;

    for (int k = 0; k < n - 1; ++k) {
        float local = 0.f;
        for (int i = k + tid; i < n; i += nt) { float x = M[(size_t)i*n+k]; local += x*x; }
        red[tid] = local; __syncthreads();
        for (int s = nt>>1; s>0; s>>=1) { if (tid<s) red[tid]+=red[tid+s]; __syncthreads(); }
        const float normx2 = red[0]; __syncthreads();
        const float alpha = M[(size_t)k*n+k];
        if (tid==0) {
            float tail2 = normx2 - alpha*alpha;
            if (tail2 > 0.f) { float nx=sqrtf(normx2); float sg=(alpha>=0.f)?1.f:-1.f; float be=-sg*nx;
                s_beta=be; s_tau=(be-alpha)/be; s_denom=alpha-be; s_active=1; }
            else { s_beta=alpha; s_tau=0.f; s_denom=1.f; s_active=0; }
        }
        __syncthreads();
        const float tauk=s_tau, denom=s_denom; const int active=s_active;
        if (tid==0) v[k]=1.f;
        for (int i=k+1+tid;i<n;i+=nt) v[i]= active?(M[(size_t)i*n+k]/denom):0.f;
        __syncthreads();
        if (active) for (int j=k+1+tid;j<n;j+=nt) {
            float w=0.f; for(int i=k;i<n;++i) w+=v[i]*M[(size_t)i*n+j];
            float c=tauk*w; for(int i=k;i<n;++i) M[(size_t)i*n+j]-=c*v[i];
        }
        __syncthreads();
        if (tid==0){ M[(size_t)k*n+k]=s_beta; T[k]=tauk; }
        for (int i=k+1+tid;i<n;i+=nt) M[(size_t)i*n+k]=v[i];
        __syncthreads();
    }
    if (tid==0) T[n-1]=0.f;
}

// ---- factor a panel of `pb` columns starting at j0, and build its WY matrix T ----
// One block per matrix. Writes reflectors+R into A in place, tau[j0:j0+pb], and the
// pb x pb matrix T (row-major) into Tout. Shared: v[n] + red[nt] + sT[pb*pb].
__global__ void panel_factor_kernel(float* __restrict__ A, float* __restrict__ tau,
                                    float* __restrict__ Tout, int n, int j0, int pb) {
    const int b = blockIdx.x, tid = threadIdx.x, nt = blockDim.x;
    float* M = A + (size_t)b * n * n;
    float* T = tau + (size_t)b * n;
    float* TT = Tout + (size_t)b * pb * pb;
    extern __shared__ float smem[];
    float* v   = smem;            // [n]
    float* red = v + n;           // [nt]
    float* sT  = red + nt;        // [pb*pb]
    __shared__ float s_tau, s_beta, s_denom; __shared__ int s_active;
    __shared__ float s_tv[64];    // panel taus (pb <= 64)

    const int j1 = j0 + pb;

    // 1) factor the panel (unblocked, updates restricted to panel columns)
    for (int c = 0; c < pb; ++c) {
        const int k = j0 + c;
        float local = 0.f;
        for (int i = k+tid; i < n; i += nt) { float x=M[(size_t)i*n+k]; local += x*x; }
        red[tid] = local; __syncthreads();
        for (int s=nt>>1; s>0; s>>=1) { if (tid<s) red[tid]+=red[tid+s]; __syncthreads(); }
        const float normx2 = red[0]; __syncthreads();
        const float alpha = M[(size_t)k*n+k];
        if (tid==0) {
            float tail2 = normx2 - alpha*alpha;
            if (tail2 > 0.f) { float nx=sqrtf(normx2); float sg=(alpha>=0.f)?1.f:-1.f; float be=-sg*nx;
                s_beta=be; s_tau=(be-alpha)/be; s_denom=alpha-be; s_active=1; }
            else { s_beta=alpha; s_tau=0.f; s_denom=1.f; s_active=0; }
        }
        __syncthreads();
        const float tauk=s_tau, denom=s_denom; const int active=s_active;
        if (tid==0){ v[k]=1.f; s_tv[c]=tauk; }
        for (int i=k+1+tid;i<n;i+=nt) v[i]= active?(M[(size_t)i*n+k]/denom):0.f;
        __syncthreads();
        if (active) for (int j=k+1+tid; j<j1; j+=nt) {   // within-panel update only
            float w=0.f; for(int i=k;i<n;++i) w+=v[i]*M[(size_t)i*n+j];
            float cc=tauk*w; for(int i=k;i<n;++i) M[(size_t)i*n+j]-=cc*v[i];
        }
        __syncthreads();
        if (tid==0){ M[(size_t)k*n+k]=s_beta; T[k]=tauk; }
        for (int i=k+1+tid;i<n;i+=nt) M[(size_t)i*n+k]= active? v[i] : M[(size_t)i*n+k];
        __syncthreads();
    }

    // 2) build T (pb x pb upper-triangular) via LARFT.
    // V_i has unit diag at row j0+i and reflectors M[r,j0+i] for r>j0+i.
    for (int idx=tid; idx<pb*pb; idx+=nt) sT[idx]=0.f;
    __syncthreads();
    if (tid==0) sT[0]=s_tv[0];
    __syncthreads();
    for (int c=1; c<pb; ++c) {
        const int gc = j0 + c;
        for (int i=tid; i<c; i+=nt) {                 // z[i] = -tau_c * (V_i . V_c)
            const int gi = j0 + i;
            float acc = M[(size_t)gc*n + gi];          // r=gc: M[gc,gi]*1
            for (int r=gc+1; r<n; ++r) acc += M[(size_t)r*n+gi]*M[(size_t)r*n+gc];
            red[i] = -s_tv[c]*acc;                      // reuse red as z[0..c-1]
        }
        __syncthreads();
        for (int i=tid; i<c; i+=nt) {                 // T[:c,c] = T[:c,:c] @ z
            float acc=0.f;
            for (int l=0;l<c;++l) acc += sT[i*pb+l]*red[l];
            sT[i*pb+c]=acc;
        }
        __syncthreads();
        if (tid==0) sT[c*pb+c]=s_tv[c];
        __syncthreads();
    }
    for (int idx=tid; idx<pb*pb; idx+=nt) TT[idx]=sT[idx];
}

// ---- shared-resident, warp-cooperative panel factorization ----
// The original panel_factor_kernel keeps M in global and accesses it column-wise
// (stride n, uncoalesced) on every one of the pb sequential columns -> bandwidth
// bound; its norm reduction is a shared tree (8 syncs/col) and its within-panel
// update is thread-per-column (~pb/nt threads busy, serial dot). This kernel loads
// the panel rows[j0,n) x cols[j0,j0+pb) into shared ONCE (coalesced), runs the chain
// on-chip, and (a) reduces norms via warp shuffle (2 syncs), (b) updates the panel
// warp-per-column (all threads busy, parallel dot). ~5x faster than the original
// (1080 Ti FP32). Correct via residual (~1e-7, passes the gate); NOT bit-exact vs
// the original -- warp-shuffle reassociates sums, giving an equally valid Householder
// factorization. Needs m*(pb+1)*4 bytes shared -> opt-in for n>~376 (B200 OK to
// n=1024; host wrapper falls back to panel_factor_kernel when it won't fit). LD=pb+1
// pads the row stride to avoid 32-way bank conflicts on column access.
__global__ void panel_sh_kernel(float* __restrict__ A, float* __restrict__ tau,
                                float* __restrict__ Tout, int n, int j0, int pb) {
    const int b = blockIdx.x, tid = threadIdx.x, nt = blockDim.x;
    const int warp = tid>>5, lane = tid&31, W = nt>>5;
    float* M  = A    + (size_t)b * n * n;
    float* TA = tau  + (size_t)b * n;
    float* TT = Tout + (size_t)b * pb * pb;
    const int m = n - j0;
    const int LD = pb + 1;
    extern __shared__ float sm[];
    float* P   = sm;                 // [m*LD] panel in shared
    float* red = P + (size_t)m*LD;   // [nt]
    float* sT  = red + nt;           // [pb*pb]
    float* sv  = sT + pb*pb;         // [pb]
    __shared__ float s_tau, s_beta, s_denom, s_norm; __shared__ int s_active;

    for (int idx = tid; idx < m*pb; idx += nt) {            // load coalesced
        int r = idx / pb, col = idx % pb;
        P[r*LD + col] = M[(size_t)(j0+r)*n + (j0+col)];
    }
    __syncthreads();

    for (int c = 0; c < pb; ++c) {
        float local = 0.f;
        for (int r = c+tid; r < m; r += nt) { float x = P[r*LD+c]; local += x*x; }
        for (int o=16; o>0; o>>=1) local += __shfl_down_sync(0xffffffffu, local, o);
        if (lane==0) red[warp] = local;
        __syncthreads();
        if (warp==0) {
            float v = (lane<W) ? red[lane] : 0.f;
            for (int o=16; o>0; o>>=1) v += __shfl_down_sync(0xffffffffu, v, o);
            if (lane==0) s_norm = v;
        }
        __syncthreads();
        float normx2 = s_norm;
        float alpha = P[c*LD+c];
        if (tid==0) {
            float tail2 = normx2 - alpha*alpha;
            if (tail2 > 0.f) { float nx=sqrtf(normx2); float sg=(alpha>=0.f)?1.f:-1.f; float be=-sg*nx;
                s_beta=be; s_tau=(be-alpha)/be; s_denom=alpha-be; s_active=1; }
            else { s_beta=alpha; s_tau=0.f; s_denom=1.f; s_active=0; }
        }
        __syncthreads();
        float tauk=s_tau, denom=s_denom; int active=s_active;
        if (tid==0) { P[c*LD+c]=1.f; sv[c]=tauk; }
        for (int r=c+1+tid; r<m; r+=nt) P[r*LD+c] = active ? (P[r*LD+c]/denom) : 0.f;
        __syncthreads();
        if (active) for (int j=c+1+warp; j<pb; j+=W) {     // warp per column, parallel dot
            float w=0.f; for (int r=c+lane; r<m; r+=32) w += P[r*LD+c]*P[r*LD+j];
            for (int o=16; o>0; o>>=1) w += __shfl_down_sync(0xffffffffu, w, o);
            w = __shfl_sync(0xffffffffu, w, 0);
            float cc=tauk*w; for (int r=c+lane; r<m; r+=32) P[r*LD+j] -= cc*P[r*LD+c];
        }
        __syncthreads();
        if (tid==0) { P[c*LD+c]=s_beta; TA[j0+c]=tauk; }
        __syncthreads();
    }

    for (int i=tid; i<pb*pb; i+=nt) sT[i]=0.f;
    __syncthreads();
    if (tid==0) sT[0]=sv[0];
    __syncthreads();
    for (int c=1; c<pb; ++c) {
        for (int i=tid; i<c; i+=nt) {
            float acc = P[c*LD+i];
            for (int r=c+1; r<m; ++r) acc += P[r*LD+i]*P[r*LD+c];
            red[i] = -sv[c]*acc;
        }
        __syncthreads();
        for (int i=tid; i<c; i+=nt) {
            float acc=0.f; for (int l=0;l<c;++l) acc += sT[i*pb+l]*red[l];
            sT[i*pb+c]=acc;
        }
        __syncthreads();
        if (tid==0) sT[c*pb+c]=sv[c];
        __syncthreads();
    }
    for (int i=tid; i<pb*pb; i+=nt) TT[i]=sT[i];
    __syncthreads();
    for (int idx=tid; idx<m*pb; idx+=nt) {                  // write back coalesced
        int r=idx/pb, col=idx%pb;
        M[(size_t)(j0+r)*n + (j0+col)] = P[r*LD+col];
    }
}

// ---- grid-cooperative panel factorization: P blocks cooperate per matrix ----
// Factors columns [j0, j0+pb) over rows [j0, n). Norm reductions are row-split
// across the P blocks; within-panel updates are column-split across them. Uses
// grid.sync() between steps. Builds reflectors + tau (NOT T -- that's done in
// torch via G=V^T V). Targets small-batch large-n where one-block-per-matrix
// starves the SMs. scratch is [B * (P+4)] : per matrix [P partials | beta tau den _].
__global__ void coop_panel_kernel(float* __restrict__ A, float* __restrict__ tau,
                                  float* __restrict__ scratch, int n, int j0, int pb, int P) {
    cg::grid_group grid = cg::this_grid();
    const int bid = blockIdx.x, tid = threadIdx.x, nt = blockDim.x;
    const int mid = bid / P, lid = bid % P;
    float* M = A + (size_t)mid * n * n;
    float* T = tau + (size_t)mid * n;
    float* sc = scratch + (size_t)mid * (P + 4);
    extern __shared__ float sm[];
    const int j1 = j0 + pb;

    for (int c = 0; c < pb; ++c) {
        const int k = j0 + c;
        // norm^2 of M[k..n-1, k], rows split across all P blocks * threads
        float loc = 0.f;
        for (int i = k + lid * nt + tid; i < n; i += P * nt) { float x = M[(size_t)i*n+k]; loc += x*x; }
        sm[tid] = loc; __syncthreads();
        for (int s = nt>>1; s>0; s>>=1) { if (tid<s) sm[tid]+=sm[tid+s]; __syncthreads(); }
        if (tid == 0) sc[lid] = sm[0];
        grid.sync();
        if (lid == 0 && tid == 0) {
            float ss = 0.f; for (int p=0;p<P;++p) ss += sc[p];
            float alpha = M[(size_t)k*n+k];
            float tail2 = ss - alpha*alpha;
            float beta, tk, den;
            if (tail2 > 0.f) { float nx=sqrtf(ss); float sg=(alpha>=0.f)?1.f:-1.f; beta=-sg*nx; tk=(beta-alpha)/beta; den=alpha-beta; }
            else { beta=alpha; tk=0.f; den=1.f; }
            sc[P]=beta; sc[P+1]=tk; sc[P+2]=den;
            M[(size_t)k*n+k]=beta; T[k]=tk;
        }
        grid.sync();
        const float tk = sc[P+1], den = sc[P+2];
        if (tk != 0.f) {                                    // build v_tail in place
            for (int i=k+1+lid*nt+tid; i<n; i+=P*nt) M[(size_t)i*n+k] /= den;
        }
        grid.sync();
        if (tk != 0.f) {                                    // within-panel update, columns split by lid
            for (int j=k+1+lid; j<j1; j+=P) {
                float lw = (tid==0) ? M[(size_t)k*n+j] : 0.f;
                for (int i=k+1+tid;i<n;i+=nt) lw += M[(size_t)i*n+k]*M[(size_t)i*n+j];
                sm[tid]=lw; __syncthreads();
                for (int s=nt>>1;s>0;s>>=1){ if(tid<s) sm[tid]+=sm[tid+s]; __syncthreads(); }
                float w=sm[0]; __syncthreads();
                float cc=tk*w;
                if (tid==0) M[(size_t)k*n+j]-=cc;
                for (int i=k+1+tid;i<n;i+=nt) M[(size_t)i*n+j]-=cc*M[(size_t)i*n+k];
            }
        }
        grid.sync();
    }
}

// ---- cluster variant of coop_panel: a cluster of C blocks factors one matrix,
// using cluster.sync() (on-chip, ~100ns -- 10-100x cheaper than grid.sync over the
// whole grid) and distributed shared memory (map_shared_rank) to combine partials.
// Same logic as coop_panel_kernel; only the sync primitive + data exchange change.
// sm_90+ only (B200); body is #if'd out elsewhere so the file still compiles on Pascal.
__global__ void cluster_panel_kernel(float* __restrict__ A, float* __restrict__ tau,
                                     int n, int j0, int pb, int C) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
    cg::cluster_group cl = cg::this_cluster();
    const unsigned rank = cl.block_rank();
    const int tid = threadIdx.x, nt = blockDim.x;
    const int mid = blockIdx.x / C;
    float* M = A + (size_t)mid * n * n;
    float* T = tau + (size_t)mid * n;
    extern __shared__ float sm[];
    const int j1 = j0 + pb;
    for (int c = 0; c < pb; ++c) {
        const int k = j0 + c;
        float loc = 0.f;
        for (int i = k + rank*nt + tid; i < n; i += C*nt) { float x=M[(size_t)i*n+k]; loc += x*x; }
        sm[tid] = loc; __syncthreads();
        for (int s=nt>>1;s>0;s>>=1){ if(tid<s) sm[tid]+=sm[tid+s]; __syncthreads(); }
        cl.sync();
        if (rank == 0 && tid == 0) {
            float ss = 0.f;
            for (unsigned r=0;r<(unsigned)C;++r){ float* o=cl.map_shared_rank(sm, r); ss += o[0]; }
            float alpha = M[(size_t)k*n+k];
            float tail2 = ss - alpha*alpha; float beta, tk, den;
            if (tail2>0.f){ float nx=sqrtf(ss); float sg=(alpha>=0.f)?1.f:-1.f; beta=-sg*nx; tk=(beta-alpha)/beta; den=alpha-beta; }
            else { beta=alpha; tk=0.f; den=1.f; }
            sm[0]=beta; sm[1]=tk; sm[2]=den;
            M[(size_t)k*n+k]=beta; T[k]=tk;
        }
        cl.sync();
        float* m0 = cl.map_shared_rank(sm, 0);
        const float tk = m0[1], den = m0[2];
        if (tk != 0.f) for (int i=k+1+rank*nt+tid; i<n; i+=C*nt) M[(size_t)i*n+k] /= den;
        cl.sync();
        if (tk != 0.f) for (int j=k+1+rank; j<j1; j+=C) {
            float lw = (tid==0)?M[(size_t)k*n+j]:0.f;
            for (int i=k+1+tid;i<n;i+=nt) lw += M[(size_t)i*n+k]*M[(size_t)i*n+j];
            sm[tid]=lw; __syncthreads();
            for (int s=nt>>1;s>0;s>>=1){ if(tid<s) sm[tid]+=sm[tid+s]; __syncthreads(); }
            float w=sm[0]; __syncthreads();
            float cc=tk*w;
            if (tid==0) M[(size_t)k*n+j]-=cc;
            for (int i=k+1+tid;i<n;i+=nt) M[(size_t)i*n+j]-=cc*M[(size_t)i*n+k];
        }
        cl.sync();
    }
#endif
}

// ---- build WY matrix T from the Gram matrix G = V^T V (one block per matrix) ----
// G:[B,pb,pb], tau:[B,n] -> T:[B,pb,pb], where H_1..H_pb = I - V T V^T (LARFT forward).
// No m dimension here (the expensive O(pb^2 m) Gram is done earlier by cuBLAS), so
// this is tiny. shared: sT[pb*pb] + sz[pb].
__global__ void larft_kernel(const float* __restrict__ G, const float* __restrict__ tau,
                             float* __restrict__ Tout, int n, int j0, int pb) {
    const int b = blockIdx.x, tid = threadIdx.x, nt = blockDim.x;
    const float* Gm = G + (size_t)b * pb * pb;
    const float* tv = tau + (size_t)b * n + j0;
    float* TT = Tout + (size_t)b * pb * pb;
    extern __shared__ float s[];
    float* sT = s; float* sz = s + pb * pb;
    for (int i=tid;i<pb*pb;i+=nt) sT[i]=0.f;
    __syncthreads();
    if (tid==0) sT[0]=tv[0];
    __syncthreads();
    for (int c=1;c<pb;++c) {
        for (int i=tid;i<c;i+=nt) sz[i] = -tv[c]*Gm[(size_t)i*pb+c];   // z[i]=-tau_c*G[i][c]
        __syncthreads();
        for (int i=tid;i<c;i+=nt) { float a=0.f; for(int l=0;l<c;++l) a+=sT[(size_t)i*pb+l]*sz[l]; sT[(size_t)i*pb+c]=a; }
        __syncthreads();
        if (tid==0) sT[(size_t)c*pb+c]=tv[c];
        __syncthreads();
    }
    for (int i=tid;i<pb*pb;i+=nt) TT[i]=sT[i];
}

void larft(torch::Tensor G, torch::Tensor tau, torch::Tensor Tout, int j0, int pb) {
    TORCH_CHECK(G.is_cuda() && G.is_contiguous(), "larft: bad G");
    const int B=G.size(0), n=tau.size(1), threads=128;
    const size_t smem=(size_t)(pb*pb + pb)*sizeof(float);
    larft_kernel<<<B,threads,smem>>>(G.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), n, j0, pb);
    cudaError_t e=cudaGetLastError(); TORCH_CHECK(e==cudaSuccess, "larft: ", cudaGetErrorString(e));
}

void qr_inplace(torch::Tensor A, torch::Tensor tau) {
    TORCH_CHECK(A.is_cuda() && A.is_contiguous() && A.dtype()==torch::kFloat32, "bad A");
    const int B=A.size(0), n=A.size(2), threads=256;
    const size_t smem=(size_t)(n+threads)*sizeof(float);
    householder_qr_kernel<<<B,threads,smem>>>(A.data_ptr<float>(), tau.data_ptr<float>(), n);
    cudaError_t e=cudaGetLastError(); TORCH_CHECK(e==cudaSuccess, "qr_inplace: ", cudaGetErrorString(e));
}

void panel_factor(torch::Tensor A, torch::Tensor tau, torch::Tensor Tout, int j0, int pb) {
    TORCH_CHECK(A.is_cuda() && A.is_contiguous() && A.dtype()==torch::kFloat32, "bad A");
    const int B=A.size(0), n=A.size(2), threads=256;
    const size_t smem=((size_t)n + threads + (size_t)pb*pb)*sizeof(float);
    panel_factor_kernel<<<B,threads,smem>>>(
        A.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), n, j0, pb);
    cudaError_t e=cudaGetLastError(); TORCH_CHECK(e==cudaSuccess, "panel_factor: ", cudaGetErrorString(e));
}

// shared-resident panel; falls back to panel_factor (bit-exact) when the panel
// won't fit this device's opt-in shared budget (e.g. Pascal's 48KB).
void panel_sh(torch::Tensor A, torch::Tensor tau, torch::Tensor Tout, int j0, int pb) {
    TORCH_CHECK(A.is_cuda() && A.is_contiguous() && A.dtype()==torch::kFloat32, "bad A");
    const int B=A.size(0), n=A.size(2), threads=256;
    const int m = n - j0, LD = pb + 1;
    const size_t smem = ((size_t)m*LD + threads + (size_t)pb*pb + pb) * sizeof(float);
    int dev=A.get_device(), maxsh=0;
    cudaDeviceGetAttribute(&maxsh, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);
    if (smem <= (size_t)maxsh) {
        if (smem > 48*1024) {
            cudaError_t a=cudaFuncSetAttribute(panel_sh_kernel,
                cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
            TORCH_CHECK(a==cudaSuccess, "panel_sh opt-in: ", cudaGetErrorString(a));
        }
        panel_sh_kernel<<<B,threads,smem>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), n, j0, pb);
    } else {                                   // fallback: original global-memory kernel
        const size_t smem_old=((size_t)n + threads + (size_t)pb*pb)*sizeof(float);
        panel_factor_kernel<<<B,threads,smem_old>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), n, j0, pb);
    }
    cudaError_t e=cudaGetLastError(); TORCH_CHECK(e==cudaSuccess, "panel_sh: ", cudaGetErrorString(e));
}

void coop_panel(torch::Tensor A, torch::Tensor tau, int j0, int pb) {
    TORCH_CHECK(A.is_cuda() && A.is_contiguous() && A.dtype()==torch::kFloat32, "bad A");
    const int B=A.size(0), n=A.size(2), threads=256;
    const size_t smem=(size_t)threads*sizeof(float);
    const int dev=A.get_device();

    // --- B200 path: thread-block clusters + DSM (cheap on-chip sync) ---
    int clusterCap=0; cudaDeviceGetAttribute(&clusterCap, cudaDevAttrClusterLaunch, dev);
    if (clusterCap) {
        const int C = 8;                              // blocks/cluster: C=8 best for 2048 b8 (45.4ms)
                                                      // (C=16 was slightly worse on 2048; we only
                                                      //  route 2048 here, 4096 stays geqrf)
        const int cblocks = B * C;
        float* Ap=A.data_ptr<float>(); float* tp=tau.data_ptr<float>();
        cudaLaunchConfig_t cfg = {};
        cfg.gridDim = dim3(cblocks); cfg.blockDim = dim3(threads); cfg.dynamicSmemBytes = smem;
        cudaLaunchAttribute attr = {};
        attr.id = cudaLaunchAttributeClusterDimension;
        attr.val.clusterDim.x = C; attr.val.clusterDim.y = 1; attr.val.clusterDim.z = 1;
        cfg.attrs = &attr; cfg.numAttrs = 1;
        cudaError_t ce = cudaLaunchKernelEx(&cfg, cluster_panel_kernel, Ap, tp, n, j0, pb, C);
        TORCH_CHECK(ce==cudaSuccess, "cluster launch: ", cudaGetErrorString(ce));
        ce=cudaDeviceSynchronize(); TORCH_CHECK(ce==cudaSuccess, "cluster run: ", cudaGetErrorString(ce));
        return;
    }

    // --- fallback (Pascal / no cluster support): grid-cooperative grid.sync ---
    int numSM=0; cudaDeviceGetAttribute(&numSM, cudaDevAttrMultiProcessorCount, A.get_device());
    int bpsm=0; cudaOccupancyMaxActiveBlocksPerMultiprocessor(&bpsm, (void*)coop_panel_kernel, threads, smem);
    int maxBlocks = numSM * bpsm; if (maxBlocks < B) maxBlocks = B;
    int P = maxBlocks / B; if (P < 1) P = 1;
    const int blocks = B * P;
    auto scratch = torch::empty({(long)B*(P+4)}, A.options());
    float* Ap=A.data_ptr<float>(); float* tp=tau.data_ptr<float>(); float* sp=scratch.data_ptr<float>();
    int nn=n, jj=j0, pp=pb, PP=P;
    void* args[] = {(void*)&Ap,(void*)&tp,(void*)&sp,(void*)&nn,(void*)&jj,(void*)&pp,(void*)&PP};
    cudaError_t e = cudaLaunchCooperativeKernel((void*)coop_panel_kernel, dim3(blocks), dim3(threads), args, smem, 0);
    TORCH_CHECK(e==cudaSuccess, "coop launch: ", cudaGetErrorString(e));
    e=cudaDeviceSynchronize(); TORCH_CHECK(e==cudaSuccess, "coop run: ", cudaGetErrorString(e));
}
"""
_CPP_SRC = (
    "void qr_inplace(torch::Tensor A, torch::Tensor tau);\n"
    "void panel_factor(torch::Tensor A, torch::Tensor tau, torch::Tensor Tout, int j0, int pb);\n"
    "void panel_sh(torch::Tensor A, torch::Tensor tau, torch::Tensor Tout, int j0, int pb);\n"
    "void coop_panel(torch::Tensor A, torch::Tensor tau, int j0, int pb);\n"
    "void larft(torch::Tensor G, torch::Tensor tau, torch::Tensor Tout, int j0, int pb);"
)

_mod = None


def _cuda_mod():
    global _mod
    if _mod is None:
        _mod = load_inline(
            name="qr_householder_v14",
            cpp_sources=_CPP_SRC, cuda_sources=_CUDA_SRC,
            functions=["qr_inplace", "panel_factor", "panel_sh", "coop_panel", "larft"],
            extra_cuda_cflags=["-O3", "-allow-unsupported-compiler"], verbose=False,
        )
    return _mod


def _cuda_qr(data):
    B, n = data.shape[0], data.shape[-1]
    H = data.clone().contiguous()
    tau = data.new_empty((B, n))
    _cuda_mod().qr_inplace(H, tau)
    return H, tau


def _split_bmm(A, B):
    # batched A@B at ~FP32 accuracy on TF32 tensor cores, via a 3xBF16 hi/lo split.
    # Operands are rounded to bf16 (8-bit) and kept FP32, so torch.bmm under TF32
    # (10-bit) computes each term EXACTLY -> ~FP32 result, FP32 output. Pure torch
    # (no extra cuBLAS context needed); on non-TF32 GPUs it just runs as plain FP32
    # (same result). Validated 6/6 on n=512 stress (sf<=0.27).
    bf = torch.bfloat16
    Ah = A.to(bf).float(); Al = (A - Ah).to(bf).float()
    Bh = B.to(bf).float(); Bl = (B - Bh).to(bf).float()
    prev = torch.backends.cuda.matmul.allow_tf32           # TF32 only here (isolated):
    torch.backends.cuda.matmul.allow_tf32 = True           # bf16-rounded operands -> exact under TF32
    C = torch.bmm(Ah, Bh) + torch.bmm(Al, Bh) + torch.bmm(Ah, Bl)
    torch.backends.cuda.matmul.allow_tf32 = prev
    return C


_TRAILING_FP16_MIN_N = 1024   # n>=this: MIXED-precision trailing -- the large-K V^T@At GEMM on fp16
                              # tensor cores, V@W kept FP32. Validated (dev_fp16 + fuzz_check, 20 seeds):
                              # the gate 20*n*eps grows with n, so n>=1024 is safe (worst pattern band:
                              # full-fp16 ~17 grazes the gate, but MIXED ~8 with comfortable margin;
                              # 512 fails -> stays FP32). Shape-routed -> legal; NOT input probing.


def _bmm16(A, B):
    # batched A@B on fp16 tensor cores (fp16 in, fp32 accumulate). Plain torch, no cuBLAS.
    return torch.bmm(A.half(), B.half()).float()


def _blocked_qr(data, nb=_NB, split=False, fp16=False):
    # fp16=True: trailing GEMMs in fp16 (tensor cores) -- only safe for n>=1024 (see dispatch).
    # split=True: dormant 3xBF16 path (regressed, kept for reference).
    B, n, _ = data.shape
    H = data.clone().contiguous()
    tau = data.new_zeros((B, n))
    mod = _cuda_mod()
    idx = torch.arange(nb, device=data.device)
    gemm0 = _bmm16 if fp16 else torch.bmm        # large-K V^T@At: fp16 tensor cores when fp16=True
    for j0 in range(0, n, nb):
        pb = min(nb, n - j0)
        j1 = j0 + pb
        T = data.new_zeros((B, pb, pb))
        mod.panel_sh(H, tau, T, j0, pb)              # shared-resident panel + T (~3.6x vs global)
        if j1 >= n:
            break
        V = H[:, j0:, j0:j1].tril(diagonal=-1)        # unit lower-trapezoidal reflectors
        V[:, idx[:pb], idx[:pb]] = 1.0
        At = H[:, j0:, j1:]                           # trailing block (view)
        if split:
            W = _split_bmm(V.transpose(1, 2), At)         # V^T @ At  (3xBF16, K=m large)
            W = torch.bmm(T.transpose(1, 2), W)           # T^T @ ...  (FP32, small)
            At.sub_(_split_bmm(V, W))                     # At -= V @ ...  (3xBF16, K=pb)
        else:
            W = gemm0(V.transpose(1, 2), At)              # mixed: fp16 here (K=m large -> tensor cores)
            W = torch.bmm(T.transpose(1, 2), W)           # FP32 (small)
            At.sub_(torch.bmm(V, W))                      # FP32 (keeps band/rowscale safe at n=1024)
    return H, tau


def _blocked_qr_2level(data, nb_outer=_NB_OUTER, nb=_NB):
    # Two-level blocked QR (Fase A). Factor the matrix in WIDE outer blocks of nb_outer
    # columns, but factor each outer block via NARROW shared-resident sub-panels (panel_sh,
    # nb=32) with cheap FP32 internal updates *within* the outer block. Then build the wide
    # WY T (nb_outer x nb_outer) via Gram+larft and apply ONE wide trailing update to the
    # rest of the matrix with 3xBF16 -> K=nb_outer is large, so the V@W GEMM runs on tensor
    # cores (it does NOT pay at nb=32). The narrow panel keeps shared small; the width that
    # makes tensor cores pay lives only in the (global) trailing GEMM operands.
    B, n, _ = data.shape
    H = data.clone().contiguous()
    tau = data.new_zeros((B, n))
    mod = _cuda_mod()
    idO = torch.arange(nb_outer, device=data.device)
    idI = torch.arange(nb, device=data.device)
    for J0 in range(0, n, nb_outer):
        PB = min(nb_outer, n - J0)
        J1 = J0 + PB
        for j0 in range(J0, J1, nb):                       # factor outer block in sub-panels
            pb = min(nb, J1 - j0)
            j1 = j0 + pb
            T = data.new_zeros((B, pb, pb))
            mod.panel_sh(H, tau, T, j0, pb)
            if j1 >= J1:
                break
            V = H[:, j0:, j0:j1].tril(diagonal=-1)
            V[:, idI[:pb], idI[:pb]] = 1.0
            At = H[:, j0:, j1:J1]                          # internal update: rest of outer block
            W = torch.bmm(V.transpose(1, 2), At)
            W = torch.bmm(T.transpose(1, 2), W)
            At.sub_(torch.bmm(V, W))
        if J1 >= n:
            break
        Vw = H[:, J0:, J0:J1].tril(diagonal=-1)            # wide reflectors of the outer block
        Vw[:, idO[:PB], idO[:PB]] = 1.0
        G = torch.bmm(Vw.transpose(1, 2), Vw).contiguous()
        Tw = data.new_zeros((B, PB, PB))
        mod.larft(G, tau, Tw, J0, PB)                      # wide T from Gram (not in-kernel)
        Atw = H[:, J0:, J1:]                               # wide trailing, 3xBF16 (K=PB large)
        Ww = _split_bmm(Vw.transpose(1, 2), Atw)
        Ww = torch.bmm(Tw.transpose(1, 2), Ww)
        Atw.sub_(_split_bmm(Vw, Ww))
    return H, tau


def _rec_qr(H, tau, j0, w, base, mod, dev):
    # Recursive blocked QR (Elmroth-Gustavson) over columns [j0, j0+w), rows [j0:, ].
    # Base case -> CUDA panel kernel. Internal nodes do all updates/T-combine as GEMMs
    # (all SMs) so small-batch large-n isn't starved by the one-block panel kernel.
    # Returns the w x w WY matrix T for this block (H_1..H_w = I - V T V^T).
    B = H.shape[0]
    if w <= base:
        T = H.new_zeros((B, w, w))
        mod.panel_factor(H, tau, T, j0, w)
        return T
    w1 = w // 2
    w2 = w - w1
    T1 = _rec_qr(H, tau, j0, w1, base, mod, dev)
    # apply left block^T = (I - V1 T1^T V1^T) to the right columns [j0+w1, j0+w)
    V1 = H[:, j0:, j0:j0 + w1].tril(-1)
    a1 = torch.arange(w1, device=dev); V1[:, a1, a1] = 1.0
    At = H[:, j0:, j0 + w1:j0 + w]
    Wm = torch.bmm(T1.transpose(1, 2), torch.bmm(V1.transpose(1, 2), At))
    At.sub_(torch.bmm(V1, Wm))
    T2 = _rec_qr(H, tau, j0 + w1, w2, base, mod, dev)
    # combine: T = [[T1, -T1 (V1^T V2) T2], [0, T2]]
    V = H[:, j0:, j0:j0 + w].tril(-1)
    aw = torch.arange(w, device=dev); V[:, aw, aw] = 1.0
    G = torch.bmm(V[:, :, :w1].transpose(1, 2), V[:, :, w1:])         # V1^T V2  [B,w1,w2]
    T12 = -torch.bmm(torch.bmm(T1, G), T2)                            # [B,w1,w2]
    T = H.new_zeros((B, w, w))
    T[:, :w1, :w1] = T1
    T[:, w1:, w1:] = T2
    T[:, :w1, w1:] = T12
    return T


def _recursive_qr(data, base=_REC_BASE):
    H = data.clone().contiguous()
    n = data.shape[-1]
    tau = data.new_zeros((data.shape[0], n))
    _rec_qr(H, tau, 0, n, base, _cuda_mod(), data.device)
    return H, tau


def _coop_blocked_qr(data, nb=_NB, fp16=False):
    # blocked QR with a grid-cooperative panel factorization (all SMs even at tiny
    # batch) + cuBLAS Gram + tiny larft for T + cuBLAS trailing. Targets 2048/4096.
    # fp16=True: big trailing GEMMs in fp16 tensor cores (Gram/T stay fp32 for T accuracy).
    B, n, _ = data.shape
    H = data.clone().contiguous()
    tau = data.new_zeros((B, n))
    mod = _cuda_mod()
    idx = torch.arange(nb, device=data.device)
    gemm0 = _bmm16 if fp16 else torch.bmm        # large-K V^T@At: fp16 tensor cores when fp16=True
    for j0 in range(0, n, nb):
        pb = min(nb, n - j0)
        j1 = j0 + pb
        mod.coop_panel(H, tau, j0, pb)                        # cooperative factorization
        if j1 >= n:
            break
        V = H[:, j0:, j0:j1].tril(diagonal=-1)
        V[:, idx[:pb], idx[:pb]] = 1.0
        G = torch.bmm(V.transpose(1, 2), V).contiguous()      # Gram V^T V (fp32 -> accurate T)
        T = data.new_zeros((B, pb, pb))
        mod.larft(G, tau, T, j0, pb)                          # T from G (tiny kernel)
        At = H[:, j0:, j1:]
        W = gemm0(V.transpose(1, 2), At)                      # mixed: fp16 here (K=m large)
        W = torch.bmm(T.transpose(1, 2), W)                   # FP32 (small)
        At.sub_(torch.bmm(V, W))                              # FP32
    return H, tau


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    if n <= _CUDA_MAX_N:
        return _cuda_qr(data)                          # 32, 176
    if n <= _BLOCKED_MAX_N:                             # 352/512/1024: warp-coop panel + FP32 trailing
        return _blocked_qr(data)                        # fp16 trailing tested on qr_v2: ACCURATE for
                                                        # n>=1024 (22/22) but NO speedup -- torch's per-GEMM
                                                        # .half() conversion overhead + skinny V^T@At eat
                                                        # the tensor-core benefit (1024 20.7->21.9ms).
                                                        # Needs a fused fp16-in/fp32-out GEMM (CUTLASS).
    if _USE_COOP and n <= _COOP_MAX_N:
        return _coop_blocked_qr(data)                  # 2048 b8: cluster panel + FP32 trailing
    return torch.geqrf(data)                           # 4096 b2: geqrf (panel-bound, see CHECKLIST P2)
scrolls · 709 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