Skip to content
KernelIndex
Search⌘K

submission 833772

Lorenzo · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_41.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-833772?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
6.31ms
#220 of 515
2026-06-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:357f509365aadafa45a004087fc7b1b2d23e7caee8b6ad22b841d0ce5ea648ba
license declaredunknown
license concludedunknown
authorsLorenzo
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float smem[];

Kernel source

submission_41.py748 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
"""Batched compact-Householder QR, GPU Mode `qr_v2` (B200).

s31 = s30 + SMALLER-PANEL HIGH-OCCUPANCY bet for n=512 (the real footprint lever).
The panel data R*pb MUST be resident somewhere, so register-residency (B3) can't
break the occ~3 smem ceiling (smem<->reg tradeoff is 1:1) AND it cuts blocks/SM
(s8/s18/s17 say that regresses at b=640). The actual lever is making the panel
SMALLER: pb=16 -> smem 512*17*4=34.8KB -> occ 6 (vs pb=32 67KB occ 3). The s19
probe measured the panel phase pb16=4.78 vs pb32=6.83 ms (-30%) at threads=128.
The panel is 68% of n=512 runtime, so -30% there is the big swing; the only risk
is the trailing (2x the panels, K=16 vs 32 in A-=VY) eating it -- never tested
end-to-end with the cuBLAS pipeline (only the old fused s12). With s30's freed
registers, pb16+threads=128 should reach occ ~5. n=512 dispatch -> pb=16,
threads=128 (occ play needs low threads; threads=256 would be reg-capped to occ 2).
Other n keep pb=32 (1024 is b=60<SMs = latency-bound per block, not occ-bound;
smaller pb only adds skinnier panels there).

s30 = s27 + PANEL REGISTER-RESIDENCY lever (ideas Bet A / B1). The s19 probe found
the panel at threads=256 is REGISTER-bound to occ=2 (smem would allow occ=3), so
s20 dropped n=512 to threads=128 to reach occ=3 -- at the cost of half the threads.
The `accd[PBMAX]` apply-reduction accumulator was FP64 = 64 regs/thread, the
dominant register consumer. Each thread accumulates only ~R/nthreads (~2-4) terms
and the cross-warp combine stays FP64 (sWcol stays double), so dropping the
PER-THREAD accumulator to FP32 (32 regs) is ~accuracy-neutral but frees ~32
regs/thread. Hypothesis: that lets n=512 run threads=256 at occ=3 -- 2x the threads
of the threads=128 occ-3 config -> more latency hiding at the SAME occupancy.
Change: accd[]/vr -> float + warpReduceSumF (store to double sWcol); n=512 dispatch
-> threads=256, la=0. All other n unchanged. ISOLATES the register lever.

s20 = s11 + PANEL OCCUPANCY FIX. The s19 probe overturned s17: the panel is
occupancy-bound with a steep slope, and s11's threads=256 is REGISTER-bound to
occ=2 (124 regs/thread). Dropping to threads=128 lifts occ to 3 and cuts the
n=512 panel phase 9.0->6.85 ms (1.32x), same for n=1024. Requires making the two
cross-warp reduction loops use the ACTUAL warp count (nwarps), not hard-coded
WARPS=8 (else threads<256 reads stale sWarpD/sWcol slots). pb UNCHANGED (isolates
the threads lever). n=2048 stays threads=256 (R/128=16 > VREG_MAX=8).

s10 = s7's barrier-bound panel kernel (unchanged) + the trailing block-WY
update moved onto DIRECT cuBLAS strided-batched TF32 tensor-core GEMMs. s7's
fused FP32 larfb ran the O(n^3) trailing update on CUDA cores (~2.6 TFLOPS,
~3% of FP32 peak); the headline n=512 case spent ~22 ms there. cuBLAS TF32
tensor cores are ~2200 TFLOPS, so the trailing FLOPs are nearly free; the only
prior loss (s5) was torch's per-panel V materialization (clone/tril/mask) +
limb splits + many launches. Here the panel kernel emits a packed, GEMM-ready
V (unit diagonal, zeroed strict-upper, zeroed identity-reflector columns) so
cuBLAS consumes it with zero torch ops.

Trailing update per panel (cur=pb, R=n-k, M trailing cols):
  W = V^T @ Atrail   (cur x M, K=R)   -> strided-batched tensor core
  Y = Tg^T @ W       (cur x M, K=cur) -> small, FP32
  Atrail -= V @ Y    (R x M,  K=cur)  -> strided-batched tensor core
Precision toggle `_CUBLAS_PREC`: "fp32" | "tf32" | "tf32x3". tf32x3 splits each
fp32 operand into a TF32-exact hi limb + residual lo limb (3 GEMMs:
AhBh+AhBl+AlBh, ~1e-5 rel err) entirely via a custom split kernel, well inside
the factor gate (~1.2e-3 at n=512).

Inherited from s7: per-n `_DISPATCH` (torch_trsm for small n, torch_tfac+bf16x3
for n=2048, geqrf for n=4096). Build/accuracy failure degrades to geqrf.
"""

import os

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

# FP32-safe globally: keep torch matmuls in true FP32 (no TF32). The bf16x3
# torch trailing path toggles allow_tf32 locally and restores it.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False

_ENABLE_CUSTOM = os.environ.get("QR_ENABLE_CUSTOM", "1") == "1"
_IMPL_OVERRIDE = os.environ.get("QR_TRAILING_IMPL", "")

# Trailing GEMM precision for the cuBLAS path. "tf32x3" is the safe default
# (~1e-5 rel err); "tf32" is the fastest (1 GEMM) but may miss the factor gate
# on ill-conditioned cases; "fp32" is the safety net. Env override for A/B.
# Fallback only; per-n precision is set explicitly in _DISPATCH below.
_CUBLAS_PREC = os.environ.get("QR_CUBLAS_PREC", "tf32x3")
_PREC_CODE = {"fp32": 0, "tf32": 1, "tf32x3": 2, "mixed": 3}


_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cublas_v2.h>
#include <cmath>

#define VREG_MAX 8
#define WARPS 8
#define PBMAX 32

__device__ __forceinline__ double warpReduceSumD(double val) {
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) {
        int hi = __shfl_down_sync(0xffffffffu, __double2hiint(val), o);
        int lo = __shfl_down_sync(0xffffffffu, __double2loint(val), o);
        val += __hiloint2double(hi, lo);
    }
    return val;
}

__device__ __forceinline__ float warpReduceSumF(float val) {
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        val += __shfl_down_sync(0xffffffffu, val, o);
    return val;
}

#define LP   32
#define LNT  64
#define LRT  32
#define LTPB 256

// One CUDA block per matrix. Factors the panel H[:, k:n, k:k+pb] in place
// (unblocked geqr2 over pb sequential reflectors), writes tau[:, k:k+pb], the
// cur x cur compact-WY T-factor (Tg) when build_t, and (when write_v) a packed
// GEMM-ready V buffer (batch, R, pb): unit diagonal, zeroed strict-upper,
// zeroed identity-reflector columns -- ready for cuBLAS with no torch ops.
// LA (compile-time): when true, each reflector's sub-diagonal norm is folded
// into the PREVIOUS reflector's apply pass (look-ahead) so the per-reflector
// norm reduction + its __syncthreads are skipped. This helps LOW-occupancy
// launches (threads=256, occ=1: barriers exposed) but HURTS the occ=3 n=512
// family (barriers already hidden; the fold only adds apply-loop cost), so the
// dispatcher picks LA per case.
// NW = compile-time max warps the reduction buffers are sized for (>= blockDim/32).
// Templated so the occ-critical n=512 (thr=128, 4 warps) keeps NW=8 (small static
// smem -> occ 3) while n=2048 can run threads=512 (16 warps -> NW=16) for more
// per-block parallelism (it is occ=1 anyway, so the extra ~2KB static is free).
template <bool LA, int NW>
__global__ void panel_factor_kernel(float* __restrict__ H,
                                    float* __restrict__ tau,
                                    float* __restrict__ Tg,
                                    float* __restrict__ V,
                                    int n, int k, int pb, int R,
                                    int build_t, int write_v) {
    extern __shared__ float smem[];
    const int PB1 = pb + 1;
    float* sH = smem;
    float* sT = sH + (size_t)R * PB1;
    float* gz = sT + (size_t)pb * pb;
    __shared__ double sWarpD[NW];
    __shared__ double sWcol[NW * PBMAX];
    __shared__ double sWfull[PBMAX];
    __shared__ float  stw[PBMAX];
    __shared__ double bcast[4];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int nthreads = blockDim.x;
    const int nwarps = (nthreads + 31) >> 5;   // actual warps (<= NW)

    float* Hb = H + (size_t)b * n * n;
    float* taub = tau + (size_t)b * n;

    for (int idx = tid; idx < R * pb; idx += nthreads) {
        int r = idx / pb, c = idx % pb;
        sH[r * PB1 + c] = Hb[(size_t)(k + r) * n + (k + c)];
    }
    if (build_t) {
        for (int idx = tid; idx < pb * pb; idx += nthreads) sT[idx] = 0.0f;
    }
    __syncthreads();

    const int lane = tid & 31;
    const int warp = tid >> 5;

    for (int jj = 0; jj < pb; ++jj) {
        float alpha = sH[jj * PB1 + jj];

        // !LA: compute column jj's norm now (every reflector). LA: only jj==0;
        // for jj>0 the norm is already in sWarpD from jj-1's apply look-ahead.
        if (!LA || jj == 0) {
            double local = 0.0;
            for (int r = jj + 1 + tid; r < R; r += nthreads) {
                float v = sH[r * PB1 + jj];
                local += (double)v * (double)v;
            }
            local = warpReduceSumD(local);
            if (lane == 0) sWarpD[warp] = local;
            __syncthreads();
        }

        if (tid == 0) {
            double s = 0.0;
            for (int w = 0; w < nwarps; ++w) s += sWarpD[w];
            double xnorm = sqrt(s);
            double beta, tauj, vscale;
            if (xnorm == 0.0) {
                beta = alpha; tauj = 0.0; vscale = 0.0;
            } else {
                double a = alpha;
                double nrm = hypot(a, xnorm);
                beta = (a >= 0.0) ? -nrm : nrm;
                tauj = (beta - a) / beta;
                vscale = 1.0 / (a - beta);
            }
            bcast[0] = beta; bcast[1] = tauj; bcast[2] = vscale;
            taub[k + jj] = (float)tauj;
            sH[jj * PB1 + jj] = (float)beta;
        }
        __syncthreads();

        float vscale = (float)bcast[2];
        float tauj = (float)bcast[1];

        // Fuse the v-scale into the reflector-load. Each thread scales only its
        // OWN rows of column jj and immediately caches them in vreg; the next
        // consumer (w-reduction) reads only the thread's own rows, so the old
        // __syncthreads between scale and load is unnecessary. Free barrier cut.
        float vreg[VREG_MAX];
        int nv = 0;
        for (int r = jj + 1 + tid; r < R; r += nthreads) {
            float v = sH[r * PB1 + jj];
            if (vscale != 0.0f) { v *= vscale; sH[r * PB1 + jj] = v; }
            vreg[nv++] = v;
        }

        if (tauj != 0.0f) {
            // FP32 per-thread accumulator (s30): each thread sums only ~R/nthreads
            // (~2-4) products, so FP32 is accuracy-neutral here; the cross-warp
            // combine below stays FP64 (sWcol is double). Halves accd registers
            // (32 vs 64) -> targets occ at threads=256.
            float accd[PBMAX];
            #pragma unroll
            for (int c = 0; c < PBMAX; ++c) accd[c] = 0.0f;
            int i = 0;
            for (int r = jj + 1 + tid; r < R; r += nthreads) {
                float vr = vreg[i++];
                #pragma unroll
                for (int c = 0; c < PBMAX; ++c)
                    if (c < pb) accd[c] += vr * sH[r * PB1 + c];
            }
            #pragma unroll
            for (int c = 0; c < PBMAX; ++c) {
                if (c >= pb) break;
                float t = warpReduceSumF(accd[c]);
                if (lane == 0) sWcol[warp * PBMAX + c] = (double)t;
            }
            __syncthreads();

            if (tid < pb) {
                int c = tid;
                double s = 0.0;
                for (int w = 0; w < nwarps; ++w) s += sWcol[w * PBMAX + c];
                s += (double)sH[jj * PB1 + c];
                sWfull[c] = s;
                if (c > jj) {
                    float tw = (float)((double)tauj * s);
                    stw[c] = tw;
                    sH[jj * PB1 + c] -= tw;
                }
            }
            __syncthreads();

            if (LA) {
                // Apply within the panel AND fold the look-ahead norm of column
                // jj+1 (rows jj+2..R-1) into the same pass; publish via sWarpD so
                // the next reflector skips its own norm reduction + barrier.
                double nrm_next = 0.0;
                int i2 = 0;
                for (int r = jj + 1 + tid; r < R; r += nthreads) {
                    float vr = vreg[i2++];
                    for (int c = jj + 1; c < pb; ++c) {
                        float val = sH[r * PB1 + c] - stw[c] * vr;
                        sH[r * PB1 + c] = val;
                        if (c == jj + 1 && r > jj + 1) nrm_next += (double)val * val;
                    }
                }
                nrm_next = warpReduceSumD(nrm_next);
                if (lane == 0) sWarpD[warp] = nrm_next;
            } else {
                int i2 = 0;
                for (int r = jj + 1 + tid; r < R; r += nthreads) {
                    float vr = vreg[i2++];
                    for (int c = jj + 1; c < pb; ++c)
                        sH[r * PB1 + c] -= stw[c] * vr;
                }
            }
            __syncthreads();
        } else if (LA) {
            // Identity reflector (tau==0): column jj+1 is unchanged; still must
            // hand the next reflector its norm look-ahead via sWarpD.
            double nrm_next = 0.0;
            if (jj + 1 < pb) {
                for (int r = jj + 2 + tid; r < R; r += nthreads) {
                    float v = sH[r * PB1 + (jj + 1)];
                    nrm_next += (double)v * v;
                }
            }
            nrm_next = warpReduceSumD(nrm_next);
            if (lane == 0) sWarpD[warp] = nrm_next;
            __syncthreads();
        }

        if (build_t) {
            if (tauj != 0.0f && jj > 0) {
                if (tid < jj) gz[tid] = (float)(-(double)tauj * sWfull[tid]);
                __syncthreads();
                for (int i = tid; i < jj; i += nthreads) {
                    double acc = 0.0;
                    for (int l = i; l < jj; ++l)
                        acc += (double)sT[i * pb + l] * (double)gz[l];
                    sT[i * pb + jj] = (float)acc;
                }
                __syncthreads();
            }
            if (tid == 0) sT[jj * pb + jj] = tauj;
            __syncthreads();
        }
    }

    for (int idx = tid; idx < R * pb; idx += nthreads) {
        int r = idx / pb, c = idx % pb;
        Hb[(size_t)(k + r) * n + (k + c)] = sH[r * PB1 + c];
    }
    if (build_t) {
        float* Tgb = Tg + (size_t)b * pb * pb;
        for (int idx = tid; idx < pb * pb; idx += nthreads) Tgb[idx] = sT[idx];
    }
    // Emit packed, GEMM-ready V (unit diag, zeroed upper, zeroed tau==0 cols).
    if (write_v) {
        float* Vb = V + (size_t)b * R * pb;
        for (int idx = tid; idx < R * pb; idx += nthreads) {
            int r = idx / pb, c = idx % pb;
            float v;
            if (taub[k + c] == 0.0f) v = 0.0f;
            else if (r < c)         v = 0.0f;
            else if (r == c)        v = 1.0f;
            else                    v = sH[r * PB1 + c];
            Vb[(size_t)r * pb + c] = v;
        }
    }
}

void panel_factor(torch::Tensor H, torch::Tensor tau, torch::Tensor Tg,
                  torch::Tensor V, int64_t k, int64_t pb,
                  int64_t build_t, int64_t write_v, int64_t threads_,
                  int64_t lookahead) {
    const int n = (int)H.size(1);
    const int R = n - (int)k;
    const int threads = (int)threads_;
    size_t smem = ((size_t)R * (pb + 1) + (size_t)pb * pb + (size_t)pb)
                  * sizeof(float);
    // NW=16 only when blockDim>256 (>8 warps); else NW=8 (keeps n=512 occ=3).
    const bool nw16 = threads > 256;
    auto kern = lookahead
        ? (nw16 ? panel_factor_kernel<true, 16>  : panel_factor_kernel<true, 8>)
        : (nw16 ? panel_factor_kernel<false, 16> : panel_factor_kernel<false, 8>);
    cudaFuncSetAttribute(kern,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
    kern<<<(int)H.size(0), threads, smem>>>(
        H.data_ptr<float>(), tau.data_ptr<float>(), Tg.data_ptr<float>(),
        V.data_ptr<float>(), n, (int)k, (int)pb, R,
        (int)build_t, (int)write_v);
    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess)
        throw std::runtime_error(std::string("panel_factor launch: ")
                                 + cudaGetErrorString(err));
}

// ---- TF32 hi/lo split (for the tf32x3 trailing path) ----
// hi = fp32 with low 13 mantissa bits cleared (exactly TF32-representable),
// lo = x - hi. So hi survives a TF32 GEMM truncation losslessly and the lo
// residual carries the remaining bits; AhBh+AhBl+AlBh recovers ~18 bits.
__device__ __forceinline__ void split_tf32(float x, float& hi, float& lo) {
    unsigned int b = __float_as_uint(x);
    hi = __uint_as_float(b & 0xFFFFE000u);
    lo = x - hi;
}

__global__ void split_contig_kernel(const float* __restrict__ X,
                                     float* __restrict__ Hh,
                                     float* __restrict__ Ll, long long N) {
    for (long long i = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         i < N; i += (long long)gridDim.x * blockDim.x) {
        float h, l; split_tf32(X[i], h, l); Hh[i] = h; Ll[i] = l;
    }
}

// Gather the strided trailing block Atrail = H[:, k:n, k+cur:n] into contiguous
// (batch, R, M) hi/lo split buffers.
__global__ void gather_split_atrail_kernel(const float* __restrict__ H,
                                           float* __restrict__ Ah,
                                           float* __restrict__ Al,
                                           int n, int k, int cur, int R, int M) {
    const int b = blockIdx.z;
    const long long tot = (long long)R * M;
    const float* Hb = H + (size_t)b * n * n;
    float* Ahb = Ah + (size_t)b * R * M;
    float* Alb = Al + (size_t)b * R * M;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < tot; idx += (long long)gridDim.x * blockDim.x) {
        int r = idx / M, c = idx % M;
        float x = Hb[(size_t)(k + r) * n + (k + cur + c)];
        float h, l; split_tf32(x, h, l);
        Ahb[idx] = h; Alb[idx] = l;
    }
}

// Use a handle created by OUR linked cuBLAS instance. Mixing torch's handle
// (at::cuda::getCurrentCUDABlasHandle) with calls into a separately-linked
// libcublas yields CUBLAS_STATUS_NOT_INITIALIZED, since each library instance
// has its own global state. A freshly-created handle runs on the default
// execution context (no override), keeping ordering with our <<<>>> launches.
static cublasHandle_t qrHandle() {
    static cublasHandle_t h = nullptr;
    if (h == nullptr) {
        cublasStatus_t st = cublasCreate(&h);
        if (st != CUBLAS_STATUS_SUCCESS)
            throw std::runtime_error("cublasCreate failed: " + std::to_string((int)st));
    }
    return h;
}

// Row-major batched GEMM: C(m x n) = alpha*opA(A)(m x k) @ opB(B)(k x n) + beta*C.
// Implemented via the standard operand-swap so H stays row-major.
static void gemm_rm(bool tA, bool tB, int m, int n, int k,
                    float alpha, const float* A, int lda, long long sA,
                    const float* B, int ldb, long long sB,
                    float beta, float* C, int ldc, long long sC,
                    int batch, cublasComputeType_t ct) {
    cublasOperation_t opA = tA ? CUBLAS_OP_T : CUBLAS_OP_N;
    cublasOperation_t opB = tB ? CUBLAS_OP_T : CUBLAS_OP_N;
    cublasStatus_t st = cublasGemmStridedBatchedEx(
        qrHandle(), opB, opA, n, m, k, &alpha,
        B, CUDA_R_32F, ldb, sB,
        A, CUDA_R_32F, lda, sA,
        &beta, C, CUDA_R_32F, ldc, sC,
        batch, ct, CUBLAS_GEMM_DEFAULT);
    if (st != CUBLAS_STATUS_SUCCESS)
        throw std::runtime_error("cublas gemm failed: " + std::to_string((int)st));
}

static void launch_split_contig(const float* X, float* Hh, float* Ll, long long N) {
    int threads = 256;
    long long blk = (N + threads - 1) / threads;
    int blocks = (int)(blk > 65535 ? 65535 : blk);
    split_contig_kernel<<<blocks, threads>>>(X, Hh, Ll, N);
}

// One panel's block-WY trailing update via direct cuBLAS strided-batched GEMMs.
//   prec: 0=fp32, 1=tf32, 2=tf32x3.
void larfb_cublas(torch::Tensor H, torch::Tensor V, torch::Tensor Tg,
                  int64_t k_, int64_t cur_, int64_t prec) {
    const int n = (int)H.size(1);
    const int batch = (int)H.size(0);
    const int k = (int)k_, cur = (int)cur_;
    const int R = n - k;
    const int M = n - k - cur;
    if (M <= 0) return;

    float* Hp = H.data_ptr<float>();
    float* Vp = V.data_ptr<float>();
    float* Tp = Tg.data_ptr<float>();
    float* Ap = Hp + (size_t)k * n + (k + cur);   // Atrail base (batch b: + b*n*n)

    const long long sH = (long long)n * n;
    const long long sV = (long long)R * cur;      // V is (batch, R, cur)
    const long long sT = (long long)cur * cur;    // Tg is (batch, pb=cur, pb)

    auto opt = H.options();
    torch::Tensor W = torch::empty({batch, cur, M}, opt);
    torch::Tensor Y = torch::empty({batch, cur, M}, opt);
    float* Wp = W.data_ptr<float>();
    float* Yp = Y.data_ptr<float>();
    const long long sW = (long long)cur * M;
    const long long sY = (long long)cur * M;

    const cublasComputeType_t TC = CUBLAS_COMPUTE_32F_FAST_TF32;
    const cublasComputeType_t F32 = CUBLAS_COMPUTE_32F;

    if (prec == 2) {
        // ---- tf32x3: split operands, 3 TF32 GEMMs per heavy product. ----
        torch::Tensor Vh = torch::empty({batch, R, cur}, opt);
        torch::Tensor Vl = torch::empty({batch, R, cur}, opt);
        torch::Tensor Ah = torch::empty({batch, R, M}, opt);
        torch::Tensor Al = torch::empty({batch, R, M}, opt);
        float* Vhp = Vh.data_ptr<float>(); float* Vlp = Vl.data_ptr<float>();
        float* Ahp = Ah.data_ptr<float>(); float* Alp = Al.data_ptr<float>();

        launch_split_contig(Vp, Vhp, Vlp, (long long)batch * R * cur);
        {
            int threads = 256;
            long long blk = ((long long)R * M + threads - 1) / threads;
            int bx = (int)(blk > 65535 ? 65535 : blk);
            dim3 grid(bx, 1, batch);
            gather_split_atrail_kernel<<<grid, threads>>>(Hp, Ahp, Alp, n, k, cur, R, M);
        }

        // W = V^T @ Atrail  (cur x M, K=R)  =  Vh^T Ah + Vh^T Al + Vl^T Ah
        gemm_rm(true, false, cur, M, R, 1.f, Vhp, cur, sV, Ahp, M, (long long)R * M,
                0.f, Wp, M, sW, batch, TC);
        gemm_rm(true, false, cur, M, R, 1.f, Vhp, cur, sV, Alp, M, (long long)R * M,
                1.f, Wp, M, sW, batch, TC);
        gemm_rm(true, false, cur, M, R, 1.f, Vlp, cur, sV, Ahp, M, (long long)R * M,
                1.f, Wp, M, sW, batch, TC);

        // Y = Tg^T @ W  (cur x M, K=cur), small -> FP32.
        gemm_rm(true, false, cur, M, cur, 1.f, Tp, cur, sT, Wp, M, sW,
                0.f, Yp, M, sY, batch, F32);

        // Atrail -= V @ Y  (R x M, K=cur)  =  Vh Yh + Vh Yl + Vl Yh
        torch::Tensor Yh = torch::empty({batch, cur, M}, opt);
        torch::Tensor Yl = torch::empty({batch, cur, M}, opt);
        float* Yhp = Yh.data_ptr<float>(); float* Ylp = Yl.data_ptr<float>();
        launch_split_contig(Yp, Yhp, Ylp, (long long)batch * cur * M);

        gemm_rm(false, false, R, M, cur, -1.f, Vhp, cur, sV, Yhp, M, sY,
                1.f, Ap, n, sH, batch, TC);
        gemm_rm(false, false, R, M, cur, -1.f, Vhp, cur, sV, Ylp, M, sY,
                1.f, Ap, n, sH, batch, TC);
        gemm_rm(false, false, R, M, cur, -1.f, Vlp, cur, sV, Yhp, M, sY,
                1.f, Ap, n, sH, batch, TC);
    } else {
        // prec: 1=tf32(both heavy GEMMs), 3=mixed (W=V^TA on tf32 since it is a
        // K=R fat reduction with accuracy headroom; the trailing-mutating
        // A-=V@Y stays FP32 to hold the factor gate on hard cases like band),
        // else 0=fp32(both).
        const cublasComputeType_t WC = (prec == 1 || prec == 3) ? TC : F32;
        const cublasComputeType_t AC = (prec == 1) ? TC : F32;
        // W = V^T @ Atrail  (cur x M, K=R)
        gemm_rm(true, false, cur, M, R, 1.f, Vp, cur, sV, Ap, n, sH,
                0.f, Wp, M, sW, batch, WC);
        // Y = Tg^T @ W  (cur x M, K=cur), FP32
        gemm_rm(true, false, cur, M, cur, 1.f, Tp, cur, sT, Wp, M, sW,
                0.f, Yp, M, sY, batch, F32);
        // Atrail -= V @ Y  (R x M, K=cur)
        gemm_rm(false, false, R, M, cur, -1.f, Vp, cur, sV, Yp, M, sY,
                1.f, Ap, n, sH, batch, AC);
    }
    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess)
        throw std::runtime_error(std::string("larfb_cublas: ")
                                 + cudaGetErrorString(err));
}
"""

_CPP_SRC = (
    "void panel_factor(torch::Tensor H, torch::Tensor tau, torch::Tensor Tg, "
    "torch::Tensor V, int64_t k, int64_t pb, int64_t build_t, int64_t write_v, "
    "int64_t threads_, int64_t lookahead);\n"
    "void larfb_cublas(torch::Tensor H, torch::Tensor V, torch::Tensor Tg, "
    "int64_t k, int64_t cur, int64_t prec);"
)

_module = None
try:
    _module = load_inline(
        name="qr_v31_kernels",
        cpp_sources=[_CPP_SRC],
        cuda_sources=[_CUDA_SRC],
        functions=["panel_factor", "larfb_cublas"],
        extra_cuda_cflags=["-O3"],
        extra_ldflags=["-lcublas"],
        verbose=True,
    )
    print("[qr] load_inline build OK")
except Exception as exc:
    print(f"[qr] load_inline build FAILED, using geqrf fallback: {exc!r}")


def _mm_tc_split(A: torch.Tensor, B: torch.Tensor, four: bool) -> torch.Tensor:
    """A@B at ~fp32 accuracy via a hi/lo limb split on TF32 tensor cores."""
    prev = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        Ah = A.to(torch.bfloat16).float()
        Bh = B.to(torch.bfloat16).float()
        Al = A - Ah
        Bl = B - Bh
        out = torch.matmul(Ah, Bh)
        out = out + torch.matmul(Ah, Bl)
        out = out + torch.matmul(Al, Bh)
        if four:
            out = out + torch.matmul(Al, Bl)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev
    return out


def _trailing_torch_trsm(H, k, cur, n, tau, eye, Tg):
    block = H[:, k:n, k:k + cur]
    V = block.clone()
    top = V[:, :cur, :cur]
    V[:, :cur, :cur] = torch.tril(top, -1) + eye[:cur, :cur]
    taup = tau[:, k:k + cur]
    V = V * (taup != 0).to(V.dtype).unsqueeze(1)

    Atrail = H[:, k:n, k + cur:n]
    W = V.transpose(-1, -2) @ Atrail
    S = V.transpose(-1, -2) @ V
    d = torch.where(taup != 0, 1.0 / taup, torch.ones_like(taup))
    Tinv = torch.triu(S, 1) + torch.diag_embed(d)
    Y = torch.linalg.solve_triangular(Tinv.transpose(-1, -2), W, upper=False)
    Atrail.sub_(V @ Y)


def _trailing_torch_tfac_tc(H, k, cur, n, tau, eye, Tg):
    block = H[:, k:n, k:k + cur]
    V = block.clone()
    top = V[:, :cur, :cur]
    V[:, :cur, :cur] = torch.tril(top, -1) + eye[:cur, :cur]
    taup = tau[:, k:k + cur]
    V = V * (taup != 0).to(V.dtype).unsqueeze(1)

    Atrail = H[:, k:n, k + cur:n]
    W = _mm_tc_split(V.transpose(-1, -2), Atrail, four=False)
    Y = Tg[:, :cur, :cur].transpose(-1, -2) @ W
    Atrail.sub_(_mm_tc_split(V, Y, four=False))


def _qr_blocked(data: torch.Tensor, pb: int, impl: str, prec: str,
                threads: int, lookahead: int) -> output_t:
    """Blocked WY Householder QR. Custom panel kernel + selected trailing path."""
    batch, n, _ = data.shape
    H = data.contiguous().clone()
    tau = torch.zeros(batch, n, device=data.device, dtype=torch.float32)
    need_t = impl in ("torch_tfac", "fused", "cublas")
    need_v = impl == "cublas"
    if need_t:
        Tg = torch.empty(batch, pb, pb, device=data.device, dtype=torch.float32)
    else:
        Tg = torch.empty(1, device=data.device, dtype=torch.float32)
    eye = torch.eye(pb, device=data.device, dtype=torch.float32)
    dummy_v = torch.empty(1, device=data.device, dtype=torch.float32)
    prec_code = _PREC_CODE.get(prec, 2)

    for k in range(0, n, pb):
        cur = min(pb, n - k)
        ntrail = n - k - cur
        R = n - k
        build_t = 1 if (need_t and ntrail > 0) else 0
        write_v = 1 if (need_v and ntrail > 0) else 0
        if write_v:
            V = torch.empty(batch, R, cur, device=data.device, dtype=torch.float32)
        else:
            V = dummy_v
        _module.panel_factor(H, tau, Tg, V, k, cur, build_t, write_v,
                              threads, lookahead)

        if ntrail <= 0:
            continue

        if impl == "cublas":
            _module.larfb_cublas(H, V, Tg, k, cur, prec_code)
        elif impl == "torch_tfac":
            _trailing_torch_tfac_tc(H, k, cur, n, tau, eye, Tg)
        else:
            _trailing_torch_trsm(H, k, cur, n, tau, eye, Tg)

    return H, tau


def _qr_blocked_geqrf_tc(data: torch.Tensor, pb: int) -> output_t:
    """n=4096 single-large regime: WIDE-panel blocked QR. The panel (R x pb) is
    factored by torch.geqrf (cuSOLVER, row-major + tall-skinny native), the
    trailing update runs on TF32 tensor cores via s7's compact-WY trsm formula.
    pb is WIDE (512) so only ~8 cuSOLVER panel calls (the s40 pb=64 regression
    was 64-panel orchestration overhead). gate 9.8e-3 (loose) -> plain TF32 safe.
    """
    batch, n, _ = data.shape
    H = data.contiguous().clone()
    tau = torch.zeros(batch, n, device=data.device, dtype=torch.float32)
    eye = torch.eye(pb, device=data.device, dtype=torch.float32)
    prev = torch.backends.cuda.matmul.allow_tf32
    try:
        for k in range(0, n, pb):
            cur = min(pb, n - k)
            ntrail = n - k - cur
            block = H[:, k:n, k:k + cur].contiguous()
            hb, tb = torch.geqrf(block)              # FP32 cuSOLVER tall panel
            H[:, k:n, k:k + cur] = hb
            tau[:, k:k + cur] = tb
            if ntrail <= 0:
                continue
            V = hb.clone()
            V[:, :cur, :cur] = torch.tril(V[:, :cur, :cur], -1) + eye[:cur, :cur]
            V = V * (tb != 0).to(V.dtype).unsqueeze(1)
            Atrail = H[:, k:n, k + cur:n]
            torch.backends.cuda.matmul.allow_tf32 = True
            W = V.transpose(-1, -2) @ Atrail         # TF32, K=R (fat reduction)
            torch.backends.cuda.matmul.allow_tf32 = False
            S = V.transpose(-1, -2) @ V              # FP32, cur x cur (cheap)
            d = torch.where(tb != 0, 1.0 / tb, torch.ones_like(tb))
            Tinv = torch.triu(S, 1) + torch.diag_embed(d)
            Y = torch.linalg.solve_triangular(Tinv.transpose(-1, -2), W, upper=False)
            torch.backends.cuda.matmul.allow_tf32 = True
            Atrail.sub_(V @ Y)                       # TF32, K=cur
            torch.backends.cuda.matmul.allow_tf32 = False
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev
    return H, tau


# Per-n dispatch with the direct-cuBLAS trailing path. Precision chosen per n by
# the factor-gate headroom, which scales ~4096/n for plain TF32: n=512 has too
# little headroom (band case fails at scaled 21.6 > 20) so it uses exact FP32;
# n>=1024 passes comfortably on plain TF32 (1 GEMM, tensor cores). n=2048 needs
# pb<=24 to keep the first-panel smem under the cap.
# threads=128 lifts panel occupancy 2->3 (s19 probe) ONLY when batch >> #SMs(148)
# so blocks compete per-SM. That is ONLY n=512 (b=640). For batch < 148 (1024 b=60,
# 2048 b=8, 176/352 b=40) each block owns an SM, so fewer threads just starves
# per-matrix parallelism -> KEEP threads=256 (the benchmark confirmed thr=128
# regressed n=1024 10.4->13.4 and 176/352). n=512 thr=128: 15.4->11.8 (1.30x).
# `la` (look-ahead norm fold): ON for the low-occupancy threads=256 cases (occ=1,
# barriers exposed -> fewer barriers help: s23 gave 1024 -4.5%, 2048 -6%, 176/352
# -3%); OFF for the occ=3 n=512 family (barriers already hidden, the fold only adds
# apply-loop cost -> s23 regressed 512 +10%).
_DISPATCH = {
    32:   {"pb": 32, "impl": "torch_trsm", "threads": 256, "la": 0},  # no trailing
    176:  {"pb": 32, "impl": "cublas", "prec": "fp32", "threads": 256, "la": 1},
    352:  {"pb": 32, "impl": "cublas", "prec": "fp32", "threads": 256, "la": 1},
    512:  {"pb": 16, "impl": "cublas", "prec": "mixed", "threads": 128, "la": 0},  # s31: pb16 -> smem 34.8KB -> occ 6 (panel -30%); thr=128 keeps occ high
    1024: {"pb": 32, "impl": "cublas", "prec": "tf32", "threads": 512, "la": 1},  # s41: thr 256->512 (b=60<SMs, latency-bound -> shorten critical path; NW=16)
    2048: {"pb": 24, "impl": "cublas", "prec": "tf32", "threads": 512, "la": 1},  # b=8<<SMs: max parallelism (NW=16)
    4096: {"pb": 512, "impl": "geqrf_tc"},  # s41: WIDE-panel cuSOLVER + TF32 trailing (8 panels; was geqrf ~52ms)
}

_SMEM_CAP = 220 * 1024


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    cfg = _DISPATCH.get(n)
    if cfg is not None:
        pb = cfg["pb"]
        impl = _IMPL_OVERRIDE or cfg["impl"]
        prec = cfg.get("prec", "fp32")
        threads = cfg.get("threads", 256)
        la = cfg.get("la", 0)
        if impl == "geqrf_tc":
            if _ENABLE_CUSTOM:
                try:
                    return _qr_blocked_geqrf_tc(data, pb)
                except Exception as exc:
                    print(f"[qr] geqrf_tc n={n} raised, fallback: {exc!r}")
            return torch.geqrf(data)
        smem = (n * (pb + 1) + pb * pb + pb) * 4
        if _ENABLE_CUSTOM and _module is not None and smem <= _SMEM_CAP:
            try:
                return _qr_blocked(data, pb, impl, prec, threads, la)
            except Exception as exc:
                print(f"[qr] custom path n={n} impl={impl} raised, fallback: {exc!r}")
    return torch.geqrf(data)
scrolls · 748 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON