Skip to content
KernelIndex
Search⌘K

submission 836906

devichand579 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
4.79ms
#168 of 515
2026-06-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:492003ce0dac54934cdd8992d88a2cd2ebbd9f96ca1278df1f95ee9692bf78fc
license declaredunknown
license concludedunknown
authorsdevichand579
imported2026-08-26

Techniques

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

shared-memory__shared__ double sh[WARPS];

Kernel source

submission.py1359 lines
"""Batched square compact-Householder QR (geqrf-compatible (H, tau)).

Strategy
--------
The reference (torch.geqrf) factors each matrix in the batch *sequentially*
through cuSOLVER, so it collapses on high-batch shapes (e.g. 640 x 512 takes
~0.75 s). We replace that with a genuinely batched **blocked Householder
(compact-WY) QR**:

  * panel factorization runs in a single custom CUDA kernel per panel (one
    threadblock factors one matrix's panel, so the sequential column loop lives
    inside the kernel -- no per-column Python/dispatch overhead);
  * the trailing submatrix is updated with batched GEMMs (cuBLAS) via the WY
    representation  A_trail -= V (T^T (V^T A_trail)).

The kernel works on column-major data (A passed transposed) so sub-columns are
contiguous -> coalesced loads.  Reductions accumulate in fp64 inside the kernel
for robustness on the ill-conditioned / wide-dynamic-range stress cases.

Precision: the trailing update defaults to fp32.  Plain low precision (bf16/
1-pass tf32) *fails the factor gate* on the wide-dynamic-range cases (clustered
/ rowscale / mixed).  `_TRAILING_PRECISION = "3xtf32"` switches the trailing
GEMMs to a manual fp32-emulation (3 TF32 tensor-core matmuls recovering
~2^-20 relative precision) -- the primary untested B200 lever for the
compute-bound shapes; see the comment above `_TRAILING_PRECISION`.

Dispatch: shapes where the batched panel kernel cannot fill the GPU (large n
with tiny batch) or where launch overhead dominates (very small n) fall back to
cuSOLVER, which is already near-optimal there.  This guarantees we never regress
below the reference on any shape.
"""

import torch

# ---------------------------------------------------------------------------
# CUDA extension: batched Householder panel factorization (column-major).
# ---------------------------------------------------------------------------
_CPP = """
void panel_factor(torch::Tensor R, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads);
std::vector<torch::Tensor> qr_blocked(torch::Tensor Acm, int64_t block, int64_t threads);
std::vector<torch::Tensor> qr_blocked2(torch::Tensor Acm, int64_t superblk, int64_t inner, int64_t threads);
torch::Tensor build_T_from_G(torch::Tensor G, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads);
void set_cublas_emulation(int64_t strategy);
"""

_CUDA = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <ATen/cuda/CUDAContext.h>
#include <dlfcn.h>
#define MAX_WARPS 16   // up to 512 threads/block

// Enable cuBLAS fp32 tensor-core emulation (BF16x9). dlsym so the ext still builds
// on older cuBLAS lacking the symbol (CUDA 12.8 runtime) -> no-op there; on B200
// (cuBLAS 13) it resolves and enables emulation (verified neutral/safe on B200).
void set_cublas_emulation(int64_t strategy) {
    typedef cublasStatus_t (*set_emul_fn)(cublasHandle_t, int);
    static set_emul_fn fn = (set_emul_fn)dlsym(RTLD_DEFAULT, "cublasSetEmulationStrategy");
    if (fn) {
        cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
        fn(h, (int)strategy);
    }
}

// Hand-rolled grid-wide barrier (sense-reversing) over a 2-int global buffer
// bar = {count, sense}, zeroed by the host before each launch. Used by the
// cooperative multi-block panel kernel so all G*B co-resident blocks (guaranteed
// resident by cudaLaunchCooperativeKernel) can sync between reflector columns.
// Atomics + __threadfence only -> no -rdc / cooperative_groups device runtime.
__device__ __forceinline__ void grid_sync(unsigned int* bar, unsigned int gridSize, bool& sense) {
    __syncthreads();
    if (threadIdx.x == 0) {
        __threadfence();
        sense = !sense;
        unsigned int old = atomicAdd(&bar[0], 1u);
        if (old == gridSize - 1u) {
            atomicExch(&bar[0], 0u);
            __threadfence();
            atomicExch(&bar[1], sense ? 1u : 0u);
        } else {
            volatile unsigned int* vs = (volatile unsigned int*)&bar[1];
            while (((*vs) != 0u) != sense) { }
        }
    }
    __syncthreads();
}

__device__ __forceinline__ double warpReduceSum(double v) {
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        v += __shfl_down_sync(0xffffffffu, v, o);
    return v;
}

// Block reduction: shuffle within each warp (no sync), then combine the WARPS
// partials through a tiny shared array. Two __syncthreads per reduction vs ~9 for
// the old shared-memory tree -- this matters because the apply step calls it
// ~jb^2/2 times/panel. THREADS is a compile-time template parameter (one kernel
// instantiation per launch block size) so the strided loops and this reduction
// stay fully unrolled with constant strides -- a runtime blockDim.x costs ~10%.
template <int THREADS>
__device__ __forceinline__ double blockReduceSum(double v, double* sh) {
    constexpr int WARPS = THREADS / 32;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    v = warpReduceSum(v);
    if (lane == 0) sh[wid] = v;
    __syncthreads();
    double r = 0.0;
    #pragma unroll
    for (int i = 0; i < WARPS; ++i) r += sh[i];
    __syncthreads();
    return r;
}

// Single-sync block reduction: drops the trailing __syncthreads that protects sh
// for reuse. Safe ONLY when the caller does an explicit __syncthreads before the
// next reduction (true in the panel kernels: a sync follows after tau-broadcast,
// scale, and apply). Saves one sync per Householder step -- the panel is bound by
// the latency of its serial n-step chain + per-step syncs, so this is a direct cut.
template <int THREADS>
__device__ __forceinline__ double blockReduceSum1(double v, double* sh) {
    constexpr int WARPS = THREADS / 32;
    int lane = threadIdx.x & 31;
    int wid = threadIdx.x >> 5;
    v = warpReduceSum(v);
    if (lane == 0) sh[wid] = v;
    __syncthreads();
    double r = 0.0;
    #pragma unroll
    for (int i = 0; i < WARPS; ++i) r += sh[i];
    return r;
}

// Acm: (B,n,n) holding COLUMN-MAJOR n x n matrices (lda=n, batch stride n*n);
// the caller passes A^T as a contiguous row-major tensor.  Column 'col' of A is
// contiguous (stride 1) -> coalesced. tau: (B,n). Panel = cols [j, j+jb).
template <int THREADS>
__global__ void panel_factor_kernel(float* __restrict__ Acm, float* __restrict__ tau,
                                    int n, int j, int jb) {
    constexpr int WARPS = THREADS / 32;
    int b = blockIdx.x;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    __shared__ double sh[WARPS];
    __shared__ float s_tau, s_beta;
    float* Ab = Acm + (size_t)b * n * n;
    float* taub = tau + (size_t)b * n;

    for (int k = 0; k < jb; ++k) {
        int col = j + k;
        int m = n - col;                            // subcolumn length
        float* base = Ab + (size_t)col * n + col;   // A(col,col), entries stride 1
        float alpha = base[0];
        double part = 0.0;
        for (int i = t + 1; i < m; i += THREADS) {
            double val = base[i];
            part += val * val;
        }
        double sigma = blockReduceSum1<THREADS>(part, sh);
        if (t == 0) {
            float tau_k, beta;
            if (sigma <= 0.0) {                      // zero sub-column: no reflection
                tau_k = 0.f; beta = alpha;
            } else {
                double xnorm = sqrt((double)alpha * alpha + sigma);
                double sign = (alpha >= 0.f) ? 1.0 : -1.0;
                double betad = -sign * xnorm;
                beta = (float)betad;
                tau_k = (float)((betad - alpha) / betad);
            }
            s_tau = tau_k; s_beta = beta;
            taub[col] = tau_k;
        }
        __syncthreads();
        float tau_k = s_tau, beta = s_beta;

        if (sigma > 0.0) {
            double scale = 1.0 / ((double)alpha - beta);
            for (int i = t + 1; i < m; i += THREADS)
                base[i] = (float)(base[i] * scale);  // v_i = x_i / (alpha - beta)
            if (t == 0) base[0] = beta;              // R diagonal
        } else {
            for (int i = t + 1; i < m; i += THREADS)
                base[i] = 0.f;                       // clean reflector tail
        }
        __syncthreads();

        // apply H_k = I - tau v v^T (v[0]=1) to remaining panel columns:
        // one warp per column (warp-shuffle dot, no per-column block sync).
        // fp32 apply (norm/sigma above stays fp64): the intra-panel dot/update is
        // the same operation the fp32 cuBLAS trailing GEMM does and passes the
        // gate, so fp32 here is accurate too -- and ~2x the fp64 throughput.
        if (tau_k != 0.f) {
            for (int c = col + 1 + wid; c < j + jb; c += WARPS) {
                float* cptr = Ab + (size_t)c * n + col;   // column c, rows [col,n)
                float pivot = cptr[0];                    // read before any write
                float p = 0.f;
                for (int i = lane + 1; i < m; i += 32)
                    p += base[i] * cptr[i];
                #pragma unroll
                for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
                float full = __shfl_sync(0xffffffffu, p, 0) + pivot;
                float f = tau_k * full;
                if (lane == 0) cptr[0] = pivot - f;
                for (int i = lane + 1; i < m; i += 32)
                    cptr[i] = cptr[i] - f * base[i];
            }
        }
        __syncthreads();
    }
}

// Shared-memory panel factorization: load the panel cols [j,j+jb) rows [j,n)
// (jb*(n-j) floats) into shared ONCE, run the jb unblocked Householder steps
// with the warp-parallel apply reading/writing shared (not global), write the
// factored panel back ONCE. The plain panel_factor_kernel re-reads each panel
// column from global ~jb times during the apply (memory-bound, ~40-55% of CUDA
// time per the profiler); this cuts that to one read + one write per element.
// Trailing update stays in cuBLAS. Caller must ensure jb*(n-j)*4 <= smem capacity.
template <int THREADS>
__global__ void panel_factor_smem_kernel(float* __restrict__ Acm, float* __restrict__ tau,
                                         int n, int j, int jb) {
    extern __shared__ float sV[];               // jb*m, col-major tight: sV[c*m + r]
    constexpr int WARPS = THREADS / 32;
    int b = blockIdx.x;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    int m = n - j;
    float* Ab = Acm + (size_t)b * n * n;
    float* taub = tau + (size_t)b * n;
    __shared__ double sh[WARPS];
    __shared__ float s_tau, s_beta, s_alpha;

    for (int idx = t; idx < jb * m; idx += THREADS) {
        int c = idx / m, r = idx % m;
        sV[(size_t)c * m + r] = Ab[(size_t)(j + c) * n + (j + r)];
    }
    __syncthreads();

    for (int k = 0; k < jb; ++k) {
        int mk = m - k;                             // subcolumn length
        float* vk = sV + (size_t)k * m + k;
        float alpha = vk[0];
        double part = 0.0;
        for (int i = t + 1; i < mk; i += THREADS) { double v = vk[i]; part += v * v; }
        double sigma = blockReduceSum1<THREADS>(part, sh);
        if (t == 0) {
            float tau_k, beta;
            if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
            else {
                double xnorm = sqrt((double)alpha * alpha + sigma);
                double sign = (alpha >= 0.f) ? 1.0 : -1.0;
                double betad = -sign * xnorm;
                beta = (float)betad;
                tau_k = (float)((betad - alpha) / betad);
            }
            s_tau = tau_k; s_beta = beta; s_alpha = alpha;
            taub[j + k] = tau_k;
        }
        __syncthreads();
        float tau_k = s_tau, beta = s_beta, alpha2 = s_alpha;
        if (sigma > 0.0) {
            double scale = 1.0 / ((double)alpha2 - beta);
            for (int i = t + 1; i < mk; i += THREADS) vk[i] = (float)(vk[i] * scale);
            if (t == 0) vk[0] = beta;
        } else {
            for (int i = t + 1; i < mk; i += THREADS) vk[i] = 0.f;
        }
        __syncthreads();
        if (tau_k != 0.f) {
            for (int c = k + 1 + wid; c < jb; c += WARPS) {
                float* ac = sV + (size_t)c * m + k;
                float pivot = ac[0];
                float p = 0.f;
                for (int i = lane + 1; i < mk; i += 32) p += vk[i] * ac[i];
                #pragma unroll
                for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
                float tot = __shfl_sync(0xffffffffu, p, 0) + pivot;
                float f = tau_k * tot;
                if (lane == 0) ac[0] = pivot - f;
                for (int i = lane + 1; i < mk; i += 32)
                    ac[i] = ac[i] - f * vk[i];
            }
        }
        __syncthreads();
    }

    for (int idx = t; idx < jb * m; idx += THREADS) {
        int c = idx / m, r = idx % m;
        Ab[(size_t)(j + c) * n + (j + r)] = sV[(size_t)c * m + r];
    }
}

// Fully-fused unblocked Householder QR: one threadblock factors one entire
// n x n matrix, kept resident in shared memory (column-major, same layout as
// Acm). The whole factorization runs in a single kernel launch -- no Python
// panel loop, no trailing-GEMM round trips, no global traffic between steps.
//
// Key to speed: the reflector APPLY is parallelized across warps -- one warp
// per trailing column, with a warp-shuffle dot reduction (no __syncthreads).
// Only the n sequential reflector steps carry a couple of block syncs each, so
// we pay ~O(n) block syncs instead of the ~O(n^2) block reductions an unblocked
// in-block QR would otherwise need. Valid for n*n*4 bytes <= the opted-in
// dynamic shared capacity (see qr_full host fn).
template <int THREADS>
__global__ void qr_full_kernel(float* __restrict__ Acm, float* __restrict__ tau, int n) {
    extern __shared__ float sA[];                 // n*n floats, column-major
    constexpr int WARPS = THREADS / 32;
    int b = blockIdx.x;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    float* Ab = Acm + (size_t)b * n * n;
    float* taub = tau + (size_t)b * n;
    __shared__ double sh[WARPS];
    __shared__ float s_tau, s_beta, s_alpha;

    for (int idx = t; idx < n * n; idx += THREADS) sA[idx] = Ab[idx];
    __syncthreads();

    for (int k = 0; k < n; ++k) {
        int m = n - k;                            // subcolumn length
        float* vk = sA + (size_t)k * n + k;       // column k, rows [k,n)
        float alpha = vk[0];
        double part = 0.0;
        for (int i = t + 1; i < m; i += THREADS) { double v = vk[i]; part += v * v; }
        double sigma = blockReduceSum<THREADS>(part, sh);
        if (t == 0) {
            float tau_k, beta;
            if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
            else {
                double xnorm = sqrt((double)alpha * alpha + sigma);
                double sign = (alpha >= 0.f) ? 1.0 : -1.0;
                double betad = -sign * xnorm;
                beta = (float)betad;
                tau_k = (float)((betad - alpha) / betad);
            }
            s_tau = tau_k; s_beta = beta; s_alpha = alpha;
            taub[k] = tau_k;
        }
        __syncthreads();
        float tau_k = s_tau, beta = s_beta, alpha2 = s_alpha;
        if (sigma > 0.0) {
            double scale = 1.0 / ((double)alpha2 - beta);
            for (int i = t + 1; i < m; i += THREADS) vk[i] = (float)(vk[i] * scale);
            if (t == 0) vk[0] = beta;
        } else {
            for (int i = t + 1; i < m; i += THREADS) vk[i] = 0.f;
        }
        __syncthreads();

        // apply H_k = I - tau v v^T (v[0]=1): one warp per trailing column c
        if (tau_k != 0.f) {
            for (int c = k + 1 + wid; c < n; c += WARPS) {
                float* ac = sA + (size_t)c * n + k;   // column c, rows [k,n)
                float pivot = ac[0];
                double p = 0.0;
                for (int i = lane + 1; i < m; i += 32) p += (double)vk[i] * ac[i];
                #pragma unroll
                for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
                double tot = __shfl_sync(0xffffffffu, p, 0) + (double)pivot;
                float f = (float)((double)tau_k * tot);
                if (lane == 0) ac[0] = pivot - f;
                for (int i = lane + 1; i < m; i += 32)
                    ac[i] = (float)((double)ac[i] - (double)f * vk[i]);
            }
        }
        __syncthreads();
    }

    for (int idx = t; idx < n * n; idx += THREADS) Ab[idx] = sA[idx];
}

void qr_full(torch::Tensor R, torch::Tensor tau, int64_t threads) {
    TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
    TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
    int B = (int)R.size(0);
    int n = (int)R.size(1);
    float* rp = R.data_ptr<float>();
    float* tp = tau.data_ptr<float>();
    size_t shbytes = (size_t)n * n * sizeof(float);
    #define LAUNCH_FULL(TH) \
        cudaFuncSetAttribute(qr_full_kernel<TH>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shbytes); \
        qr_full_kernel<TH><<<B, TH, shbytes, 0>>>(rp, tp, n);
    switch ((int)threads) {
        case 128: LAUNCH_FULL(128); break;
        case 512: LAUNCH_FULL(512); break;
        default:  LAUNCH_FULL(256); break;
    }
    #undef LAUNCH_FULL
}

// Cooperative multi-block-per-matrix panel factorization. For starved shapes
// (large n, tiny batch) the one-block-per-matrix panel kernel uses only `batch`
// SMs; here G blocks cooperate on each matrix's panel so the work fills the GPU.
// Block bid -> matrix b = bid/G, sub-block g = bid%G. For each subcolumn the G
// blocks partition its rows, reduce the norm / apply-dots across blocks through
// global scratch, and sync via the hand-rolled grid barrier between the (still
// sequential) reflector columns. Math is identical to panel_factor_kernel.
template <int THREADS>
__global__ void panel_factor_coop_kernel(
        float* __restrict__ Acm, float* __restrict__ tau,
        float* __restrict__ normbuf, float* __restrict__ dotbuf,
        unsigned int* __restrict__ bar, int n, int j, int jb, int G) {
    int bid = blockIdx.x;
    int b = bid / G;
    int g = bid % G;
    int t = threadIdx.x;
    unsigned int gridSize = (unsigned int)(gridDim.x);
    bool sense = false;
    __shared__ double sh[THREADS / 32];
    float* Ab = Acm + (size_t)b * n * n;
    float* taub = tau + (size_t)b * n;
    float* nbrow = normbuf + (size_t)b * G;            // [G] partial norms
    float* dbrow = dotbuf + (size_t)(b * G + g) * jb;  // [jb] this block's dots

    for (int k = 0; k < jb; ++k) {
        int col = j + k;
        int M = n - col;                               // subcolumn length
        float* base = Ab + (size_t)col * n + col;      // base[0] = diagonal
        int chunk = (M + G - 1) / G;
        int gs = g * chunk;
        int ge = gs + chunk; if (ge > M) ge = M;
        int ns = gs < 1 ? 1 : gs;                      // tail start (skip pivot row 0)
        float alpha = base[0];

        double part = 0.0;
        for (int i = ns + t; i < ge; i += THREADS) { double v = base[i]; part += v * v; }
        double psum = blockReduceSum<THREADS>(part, sh);
        if (t == 0) nbrow[g] = (float)psum;
        grid_sync(bar, gridSize, sense);

        double sigma = 0.0;
        for (int gg = 0; gg < G; ++gg) sigma += (double)nbrow[gg];
        float tau_k, beta;
        if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
        else {
            double xnorm = sqrt((double)alpha * alpha + sigma);
            double sign = (alpha >= 0.f) ? 1.0 : -1.0;
            double betad = -sign * xnorm;
            beta = (float)betad;
            tau_k = (float)((betad - alpha) / betad);
        }
        if (sigma > 0.0) {
            double scale = 1.0 / ((double)alpha - beta);
            for (int i = ns + t; i < ge; i += THREADS) base[i] = (float)(base[i] * scale);
        } else {
            for (int i = ns + t; i < ge; i += THREADS) base[i] = 0.f;
        }
        if (g == 0 && t == 0) { base[0] = beta; taub[col] = tau_k; }  // block 0 owns pivot
        grid_sync(bar, gridSize, sense);

        // apply H_k = I - tau v v^T (v[0]=1): partial dots over owned rows
        for (int c = col + 1; c < j + jb; ++c) {
            float* cptr = Ab + (size_t)c * n + col;
            double p = 0.0;
            for (int i = gs + t; i < ge; i += THREADS) {
                double vi = (i == 0) ? 1.0 : (double)base[i];
                p += vi * (double)cptr[i];
            }
            double bp = blockReduceSum<THREADS>(p, sh);
            if (t == 0) dbrow[c - j] = (float)bp;
        }
        grid_sync(bar, gridSize, sense);

        for (int c = col + 1; c < j + jb; ++c) {
            double full = 0.0;
            for (int gg = 0; gg < G; ++gg) full += (double)dotbuf[(size_t)(b * G + gg) * jb + (c - j)];
            float f = (float)((double)tau_k * full);
            float* cptr = Ab + (size_t)c * n + col;
            for (int i = gs + t; i < ge; i += THREADS) {
                double vi = (i == 0) ? 1.0 : (double)base[i];
                cptr[i] = (float)((double)cptr[i] - (double)f * vi);
            }
        }
        grid_sync(bar, gridSize, sense);
    }
}

void panel_factor_coop(torch::Tensor R, torch::Tensor tau, int64_t j, int64_t jb, int64_t G) {
    TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
    TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
    int B = (int)R.size(0);
    int n = (int)R.size(1);
    int Gi = (int)G;
    float* rp = R.data_ptr<float>();
    float* tp = tau.data_ptr<float>();
    auto i32 = R.options().dtype(at::kInt);
    auto nbuf = torch::zeros({(long)(B * Gi)}, R.options());
    auto dbuf = torch::zeros({(long)(B * Gi * (int)jb)}, R.options());
    auto barb = torch::zeros({2}, i32);
    float* nbp = nbuf.data_ptr<float>();
    float* dbp = dbuf.data_ptr<float>();
    unsigned int* barp = (unsigned int*)barb.data_ptr<int>();
    int ni = n, ji = (int)j, jbi = (int)jb;
    void* args[] = { (void*)&rp, (void*)&tp, (void*)&nbp, (void*)&dbp,
                     (void*)&barp, (void*)&ni, (void*)&ji, (void*)&jbi, (void*)&Gi };
    dim3 grid(B * Gi), block(256);
    cudaError_t err = cudaLaunchCooperativeKernel(
        (const void*)panel_factor_coop_kernel<256>, grid, block, args, 0, 0);
    TORCH_CHECK(err == cudaSuccess, "coop panel launch failed: ", cudaGetErrorString(err));
}

// Blocked fully-fused Householder QR for larger n that does NOT fit in shared:
// one threadblock factors one entire matrix in a SINGLE launch (no Python panel
// loop, no cuBLAS, no T-build, no V-extraction copies). Only the current panel
// (jb columns) lives in shared; each trailing column is read from global
// EXACTLY ONCE per panel and has all jb reflectors applied to it in-place
// (kept in a per-warp shared scratch column), so global bandwidth stays at the
// blocked O(n/jb)-passes level rather than the O(n) of an unblocked global QR.
//
// Shared layout (dynamic): sV = jb*n floats (panel, col-major sV[c*n+r]);
//                          sC = WARPS*n floats (per-warp trailing column).
template <int THREADS>
__global__ void qr_block_fused_kernel(float* __restrict__ Acm, float* __restrict__ tau,
                                      int n, int blk) {
    extern __shared__ float smem[];
    constexpr int WARPS = THREADS / 32;
    float* sV = smem;                 // jb*n
    float* sC = smem + (size_t)blk * n;   // WARPS*n
    int b = blockIdx.x;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    float* Ab = Acm + (size_t)b * n * n;
    float* taub = tau + (size_t)b * n;
    __shared__ double sh[WARPS];
    __shared__ float s_tau, s_beta, s_alpha;

    for (int j = 0; j < n; j += blk) {
        int jb = blk < (n - j) ? blk : (n - j);
        int m = n - j;                          // panel/trailing row count

        // load panel cols [j,j+jb) rows [j,n) into sV (col-major, m rows)
        for (int idx = t; idx < jb * m; idx += THREADS) {
            int c = idx / m, r = idx % m;
            sV[(size_t)c * n + r] = Ab[(size_t)(j + c) * n + (j + r)];
        }
        __syncthreads();

        // factor panel in shared (unblocked Householder, warp-parallel apply)
        for (int k = 0; k < jb; ++k) {
            int mk = m - k;                     // subcolumn length
            float* vk = sV + (size_t)k * n + k;
            float alpha = vk[0];
            double part = 0.0;
            for (int i = t + 1; i < mk; i += THREADS) { double v = vk[i]; part += v * v; }
            double sigma = blockReduceSum<THREADS>(part, sh);
            if (t == 0) {
                float tau_k, beta;
                if (sigma <= 0.0) { tau_k = 0.f; beta = alpha; }
                else {
                    double xnorm = sqrt((double)alpha * alpha + sigma);
                    double sign = (alpha >= 0.f) ? 1.0 : -1.0;
                    double betad = -sign * xnorm;
                    beta = (float)betad;
                    tau_k = (float)((betad - alpha) / betad);
                }
                s_tau = tau_k; s_beta = beta; s_alpha = alpha;
                taub[j + k] = tau_k;
            }
            __syncthreads();
            float tau_k = s_tau, beta = s_beta, alpha2 = s_alpha;
            if (sigma > 0.0) {
                double scale = 1.0 / ((double)alpha2 - beta);
                for (int i = t + 1; i < mk; i += THREADS) vk[i] = (float)(vk[i] * scale);
                if (t == 0) vk[0] = beta;
            } else {
                for (int i = t + 1; i < mk; i += THREADS) vk[i] = 0.f;
            }
            __syncthreads();
            if (tau_k != 0.f) {
                for (int c = k + 1 + wid; c < jb; c += WARPS) {
                    float* ac = sV + (size_t)c * n + k;
                    float pivot = ac[0];
                    double p = 0.0;
                    for (int i = lane + 1; i < mk; i += 32) p += (double)vk[i] * ac[i];
                    #pragma unroll
                    for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
                    double tot = __shfl_sync(0xffffffffu, p, 0) + (double)pivot;
                    float f = (float)((double)tau_k * tot);
                    if (lane == 0) ac[0] = pivot - f;
                    for (int i = lane + 1; i < mk; i += 32)
                        ac[i] = (float)((double)ac[i] - (double)f * vk[i]);
                }
            }
            __syncthreads();
        }

        // write factored panel back (R rows + reflectors)
        for (int idx = t; idx < jb * m; idx += THREADS) {
            int c = idx / m, r = idx % m;
            Ab[(size_t)(j + c) * n + (j + r)] = sV[(size_t)c * n + r];
        }
        __syncthreads();

        // trailing update: cols [j+jb, n), one warp per column, read once
        float* col = sC + (size_t)wid * n;
        for (int c = j + jb + wid; c < n; c += WARPS) {
            for (int r = lane; r < m; r += 32) col[r] = Ab[(size_t)c * n + (j + r)];
            __syncwarp();
            for (int k = 0; k < jb; ++k) {
                int mk = m - k;
                float* vk = sV + (size_t)k * n + k;
                float tau_k = taub[j + k];
                if (tau_k != 0.f) {
                    float pivot = col[k];
                    double p = 0.0;
                    for (int i = lane + 1; i < mk; i += 32) p += (double)vk[i] * col[k + i];
                    #pragma unroll
                    for (int o = 16; o > 0; o >>= 1) p += __shfl_down_sync(0xffffffffu, p, o);
                    double tot = __shfl_sync(0xffffffffu, p, 0) + (double)pivot;
                    float f = (float)((double)tau_k * tot);
                    if (lane == 0) col[k] = pivot - f;
                    for (int i = lane + 1; i < mk; i += 32)
                        col[k + i] = (float)((double)col[k + i] - (double)f * vk[i]);
                    __syncwarp();
                }
            }
            for (int r = lane; r < m; r += 32) Ab[(size_t)c * n + (j + r)] = col[r];
        }
        __syncthreads();
    }
}

void qr_block_fused(torch::Tensor R, torch::Tensor tau, int64_t blk, int64_t threads) {
    TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
    TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
    int B = (int)R.size(0);
    int n = (int)R.size(1);
    float* rp = R.data_ptr<float>();
    float* tp = tau.data_ptr<float>();
    int W = (int)threads / 32;
    size_t shbytes = (size_t)(blk + W) * n * sizeof(float);
    #define LAUNCH_BF(TH) \
        cudaFuncSetAttribute(qr_block_fused_kernel<TH>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shbytes); \
        qr_block_fused_kernel<TH><<<B, TH, shbytes, 0>>>(rp, tp, n, (int)blk);
    switch ((int)threads) {
        case 128: LAUNCH_BF(128); break;
        case 512: LAUNCH_BF(512); break;
        default:  LAUNCH_BF(256); break;
    }
    #undef LAUNCH_BF
}

// Build the compact-WY block reflector T (jb x jb, upper triangular) for the
// panel [j, j+jb) directly from the reflectors already stored in Acm (col-major)
// and tau -- one block per matrix, one launch. Replaces the 5-op PyTorch T-build
// (G = bmm(V,V^T), triu, eye+tau*U, diag_embed, solve_triangular) and removes the
// cuSOLVER batched triangular solve entirely. Q_panel = I - V T V^T (LAPACK
// dlarft, DIRECT='F', STOREV='C'). Reflector p: v_p[p]=1, v_p[bb>p]=Acm tail,
// v_p[bb<p]=0; the R entries on/above the panel diagonal are never read.
template <int THREADS>
__global__ void build_T_kernel(const float* __restrict__ Acm, const float* __restrict__ tau,
                               float* __restrict__ Tout, int n, int j, int jb) {
    constexpr int WARPS = THREADS / 32;
    int b = blockIdx.x;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    const float* Ab = Acm + (size_t)b * n * n;
    const float* taub = tau + (size_t)b * n;
    float* Tb = Tout + (size_t)b * jb * jb;     // row-major jb x jb
    int m = n - j;

    extern __shared__ float smT[];
    float* sG = smT;                 // jb*jb (only upper+diag used)
    float* sT = smT + jb * jb;       // jb*jb
    float* stau = smT + 2 * jb * jb; // jb

    for (int p = t; p < jb; p += THREADS) stau[p] = taub[j + p];
    __syncthreads();

    // G[p][q] = v_p . v_q for p <= q : one warp per (p,q) pair
    for (int idx = wid; idx < jb * jb; idx += WARPS) {
        int p = idx / jb, q = idx % jb;
        if (p > q) continue;
        const float* vp = Ab + (size_t)(j + p) * n + j;   // vp[bb] = reflector p, local row bb
        const float* vq = Ab + (size_t)(j + q) * n + j;
        double s = 0.0;
        for (int bb = q + 1 + lane; bb < m; bb += 32) s += (double)vp[bb] * vq[bb];
        #pragma unroll
        for (int o = 16; o > 0; o >>= 1) s += __shfl_down_sync(0xffffffffu, s, o);
        if (lane == 0) {
            double vpq = (p == q) ? 1.0 : (double)vp[q];  // v_p[q]; v_q[q]=1
            sG[p * jb + q] = (float)(vpq + s);
        }
    }
    __syncthreads();

    // dlarft forward recurrence (sequential in q), done by warp 0
    if (wid == 0) {
        for (int q = 0; q < jb; ++q) {
            if (lane == 0) sT[q * jb + q] = stau[q];
            for (int p = lane; p < q; p += 32) {
                double acc = 0.0;
                for (int l = p; l < q; ++l) acc += (double)sT[p * jb + l] * (double)sG[l * jb + q];
                sT[p * jb + q] = (float)(-(double)stau[q] * acc);
            }
            __syncwarp();
        }
        for (int idx = lane; idx < jb * jb; idx += 32)
            if ((idx / jb) > (idx % jb)) sT[idx] = 0.f;   // zero strict lower
    }
    __syncthreads();
    for (int idx = t; idx < jb * jb; idx += THREADS) Tb[idx] = sT[idx];
}

// Build compact-WY T (jb x jb, row-major upper-tri) from the ALREADY-computed
// G = V^T V (jb x jb, from a cuBLAS tensor-core bmm) and tau. This replaces the
// 5-launch PyTorch tail (G.triu(1) + tau*U + eye+ + diag_embed + solve_triangular)
// with ONE launch, while KEEPING the fast TC bmm for G (the old build_T_kernel lost
// at jb=32 only because it recomputed G on CUDA cores warp-per-pair).
//   Amat = I + diag(tau)*striu(G,1)   (unit upper-tri);  Amat @ T = diag(tau).
// The jb columns of T are INDEPENDENT triangular solves, so one warp solves each
// column (back-substitution) -- critical depth ~jb, vs the dlarft recurrence's
// serial-over-columns ~jb^2/2 in a single warp. One block per matrix.
template <int THREADS>
__global__ void build_T_from_G_kernel(const float* __restrict__ Gin,
                                      const float* __restrict__ tau,
                                      float* __restrict__ Tout,
                                      int n, int j, int jb) {
    constexpr int WARPS = THREADS / 32;
    int b = blockIdx.x;
    int t = threadIdx.x;
    int lane = t & 31;
    int wid = t >> 5;
    const float* Gb = Gin + (size_t)b * jb * jb;     // row-major jb x jb
    const float* taub = tau + (size_t)b * n;
    float* Tb = Tout + (size_t)b * jb * jb;          // row-major jb x jb

    extern __shared__ float smTG[];
    float* sA = smTG;                 // jb*jb  Amat (unit upper-tri)
    float* sT = smTG + jb * jb;       // jb*jb  T (col q in sT[p*jb+q])
    float* st = smTG + 2 * jb * jb;   // jb     tau

    for (int p = t; p < jb; p += THREADS) st[p] = taub[j + p];
    __syncthreads();
    for (int idx = t; idx < jb * jb; idx += THREADS) {
        int p = idx / jb, q = idx % jb;
        sA[idx] = (p == q) ? 1.f : (p < q ? st[p] * Gb[idx] : 0.f);
        sT[idx] = 0.f;
    }
    __syncthreads();

    // one warp per column q: solve Amat @ x = tau[q] e_q  (x = T[:,q]); back-sub
    // x[q]=tau[q]; x[p<q] = -sum_{l=p+1..q} Amat[p][l]*x[l] (descending p).
    for (int q = wid; q < jb; q += WARPS) {
        if (lane == 0) sT[q * jb + q] = st[q];
        __syncwarp();
        for (int p = q - 1; p >= 0; --p) {
            float acc = 0.f;            // fp32 (B200 fp64 is ~1/30th); matches the
            for (int l = p + 1 + lane; l <= q; l += 32)   // fp32 solve_triangular baseline
                acc += sA[p * jb + l] * sT[l * jb + q];
            #pragma unroll
            for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffffu, acc, o);
            if (lane == 0) sT[p * jb + q] = -acc;
            __syncwarp();
        }
    }
    __syncthreads();
    for (int idx = t; idx < jb * jb; idx += THREADS) {
        int p = idx / jb, q = idx % jb;
        Tb[idx] = (p > q) ? 0.f : sT[idx];
    }
}

torch::Tensor build_T_from_G(torch::Tensor G, torch::Tensor tau, int64_t j,
                             int64_t jb, int64_t threads) {
    TORCH_CHECK(G.is_cuda() && tau.is_cuda(), "cuda required");
    TORCH_CHECK(G.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
    int B = (int)G.size(0);
    int n = (int)tau.size(1);
    auto T = torch::empty({B, (int)jb, (int)jb}, G.options());
    const float* gp = G.data_ptr<float>();
    const float* tp = tau.data_ptr<float>();
    float* Tp = T.data_ptr<float>();
    size_t shbytes = (size_t)(2 * jb * jb + jb) * sizeof(float);
    #define LAUNCH_TG(TH) build_T_from_G_kernel<TH><<<B, TH, shbytes, 0>>>(gp, tp, Tp, n, (int)j, (int)jb);
    switch ((int)threads) {
        case 128: LAUNCH_TG(128); break;
        case 512: LAUNCH_TG(512); break;
        default:  LAUNCH_TG(256); break;
    }
    #undef LAUNCH_TG
    return T;
}

// Emit a clean compact-WY V panel (B,jb,m) row-major from the reflectors stored
// in Acm (col-major (B,n,n) = A^T row-major): V[p][bb] = 1 (bb==p), Acm reflector
// (bb>p), 0 (bb<p). Replaces raw.triu(1).contiguous() (strided triu_tril_kernel +
// a separate diagonal.fill_(1) launch) with one coalesced pass: blockIdx=(p,b),
// threads stride over contiguous bb so both the Acm read and the Vt write coalesce.
__global__ void extract_V_kernel(const float* __restrict__ Acm, float* __restrict__ Vt,
                                 int n, int j, int jb, int m) {
    int b = blockIdx.y;
    int p = blockIdx.x;                       // reflector index within panel
    const float* Ab = Acm + (size_t)b * n * n + (size_t)(j + p) * n + j;  // row j+p, col j
    float* Vb = Vt + ((size_t)b * jb + p) * m;
    for (int bb = threadIdx.x; bb < m; bb += blockDim.x) {
        Vb[bb] = (bb < p) ? 0.f : (bb == p) ? 1.f : Ab[bb];
    }
}

torch::Tensor extract_V(torch::Tensor Acm, int64_t j, int64_t jb, int64_t m) {
    int B = (int)Acm.size(0);
    int n = (int)Acm.size(1);
    auto Vt = torch::empty({B, jb, m}, Acm.options());
    int th = (int)m < 256 ? (((int)m + 31) / 32) * 32 : 256;
    if (th < 32) th = 32;
    dim3 grid((unsigned)jb, (unsigned)B);
    extract_V_kernel<<<grid, th, 0, 0>>>(Acm.data_ptr<float>(), Vt.data_ptr<float>(),
                                         n, (int)j, (int)jb, (int)m);
    return Vt;
}

torch::Tensor build_T(torch::Tensor Acm, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads) {
    TORCH_CHECK(Acm.is_cuda() && tau.is_cuda(), "cuda required");
    TORCH_CHECK(Acm.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
    int B = (int)Acm.size(0);
    int n = (int)Acm.size(1);
    auto T = torch::empty({B, jb, jb}, Acm.options());
    const float* ap = Acm.data_ptr<float>();
    const float* tp = tau.data_ptr<float>();
    float* Tp = T.data_ptr<float>();
    size_t shbytes = (size_t)(2 * jb * jb + jb) * sizeof(float);
    #define LAUNCH_T(TH) build_T_kernel<TH><<<B, TH, shbytes, 0>>>(ap, tp, Tp, n, (int)j, (int)jb);
    switch ((int)threads) {
        case 128: LAUNCH_T(128); break;
        case 512: LAUNCH_T(512); break;
        default:  LAUNCH_T(256); break;
    }
    #undef LAUNCH_T
    return T;
}

// NEGATIVE RESULT (profiled on Ada, kept disabled): routing panel_factor to
// panel_factor_smem_kernel was a wash-to-loss (+0.5% @512/640, +1.3% @1024/60).
// Caching the jb*(n-j) panel needs ~64KB shared -> ~1 block/SM, which destroys
// the occupancy that makes the high-batch shapes fast. The panel is bound by its
// serial jb-step Householder dependency chain + per-step block syncs, not global
// bandwidth, so the staged-in-shared traffic savings don't pay for the occupancy
// loss. The global (near-zero-shared) kernel below stays the production path.
// Shared-memory cap for the staged panel (bytes). B200 allows up to ~227KB
// dynamic shared per block; stay under to keep >=1 block resident.
#define PANEL_SMEM_CAP (192 * 1024)

void panel_factor(torch::Tensor R, torch::Tensor tau, int64_t j, int64_t jb, int64_t threads) {
    TORCH_CHECK(R.is_cuda() && tau.is_cuda(), "cuda required");
    TORCH_CHECK(R.scalar_type() == at::kFloat && tau.scalar_type() == at::kFloat, "fp32 required");
    int B = (int)R.size(0);
    int n = (int)R.size(1);
    float* rp = R.data_ptr<float>();
    float* tp = tau.data_ptr<float>();
    // Stage the whole panel in shared (read once / write once instead of jb x
    // through global) when it fits. fp32 apply -> lower register pressure helps
    // the occupancy that previously made this a wash on Ada.
    size_t shbytes = (size_t)jb * (n - j) * sizeof(float);
    if (shbytes <= PANEL_SMEM_CAP) {
        #define LAUNCH_SMEM(TH) \
            cudaFuncSetAttribute(panel_factor_smem_kernel<TH>, \
                cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shbytes); \
            panel_factor_smem_kernel<TH><<<B, TH, shbytes, 0>>>(rp, tp, n, (int)j, (int)jb);
        switch ((int)threads) {
            case 64:   LAUNCH_SMEM(64); break;
            case 128:  LAUNCH_SMEM(128); break;
            case 512:  LAUNCH_SMEM(512); break;
            case 1024: LAUNCH_SMEM(1024); break;
            default:   LAUNCH_SMEM(256); break;
        }
        #undef LAUNCH_SMEM
        return;
    }
    switch ((int)threads) {
        case 64:   panel_factor_kernel<64><<<B, 64, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
        case 128:  panel_factor_kernel<128><<<B, 128, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
        case 512:  panel_factor_kernel<512><<<B, 512, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
        case 1024: panel_factor_kernel<1024><<<B, 1024, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
        default:   panel_factor_kernel<256><<<B, 256, 0, 0>>>(rp, tp, n, (int)j, (int)jb); break;
    }
}

// Blocked Householder QR host loop in C++: one Python entry, no per-panel GIL/dispatch.
// Acm is column-major (B,n,n); returns (H row-major, tau). Trailing GEMMs use ATen bmm
// (same cuBLAS path as Python torch.bmm). kernel_T path uses build_T for jb<=16.
std::vector<torch::Tensor> qr_blocked(torch::Tensor Acm, int64_t block, int64_t threads) {
    TORCH_CHECK(Acm.is_cuda(), "cuda required");
    TORCH_CHECK(Acm.scalar_type() == at::kFloat, "fp32 required");
    const int B = (int)Acm.size(0);
    const int n = (int)Acm.size(1);
    const int blk = (int)block;
    const int th = (int)threads;
    const bool kernel_T = blk <= 16;
    auto tau = torch::zeros({B, n}, Acm.options());
    int j = 0;
    while (j < n) {
        const int jb = std::min(blk, n - j);
        panel_factor(Acm, tau, j, jb, th);
        const int c2 = j + jb;
        if (c2 < n) {
            const int m = n - j;
            auto Vt = extract_V(Acm, j, jb, m);
            torch::Tensor T;
            if (kernel_T) {
                T = build_T(Acm, tau, j, jb, th);
            } else {
                // G = V^T V on tensor cores, then T via the one-launch parallel
                // column-solve kernel (replaces triu + 2 ew + diag_embed + trsm).
                auto G = at::bmm(Vt, Vt.transpose(1, 2));
                T = build_T_from_G(G, tau, j, jb, 512);
            }
            auto At = Acm.narrow(1, c2, n - c2).narrow(2, j, m);
            auto W1 = at::bmm(Vt, At.transpose(1, 2));
            auto W2 = at::bmm(T.transpose(1, 2), W1);
            // Fuse trailing subtract into the GEMM: At = 1*At + (-1)*(W2^T @ Vt)
            // in one cuBLAS call -- no materialized temp, no separate elementwise
            // pass over the trailing matrix (that aten::sub_ was 23-46% of CUDA time).
            At.baddbmm_(W2.transpose(1, 2), Vt, /*beta=*/1.0, /*alpha=*/-1.0);
        }
        j = c2;
    }
    auto H = Acm.transpose(-2, -1).contiguous();
    return {H, tau};
}

// Two-level blocked QR: super-panel SW, inner sub-panel IB. Decouples panel width
// (IB -> cheap O(IB^2 m) apply) from far-trailing GEMM width (SW). Factor the
// super-panel in IB chunks applying each chunk's WY only WITHIN [J,J+W), then apply
// the merged SW-wide WY to the FAR trailing [J+W,n) once (M=SW). vs single-level
// this fattens the far-GEMM M and halves (SW/IB x) its count. Stored reflectors are
// identical to a single SW-wide panel, so (H,tau) is unchanged.
std::vector<torch::Tensor> qr_blocked2(torch::Tensor Acm, int64_t superblk,
                                       int64_t inner, int64_t threads) {
    TORCH_CHECK(Acm.is_cuda(), "cuda required");
    TORCH_CHECK(Acm.scalar_type() == at::kFloat, "fp32 required");
    const int n = (int)Acm.size(1);
    const int SW = (int)superblk;
    const int IB = (int)inner;
    const int th = (int)threads;
    auto tau = torch::zeros({(int)Acm.size(0), n}, Acm.options());
    auto eye_full = torch::eye(SW, Acm.options());
    int J = 0;
    while (J < n) {
        const int W = std::min(SW, n - J);
        int j = J;
        while (j < J + W) {
            const int jb = std::min(IB, J + W - j);
            panel_factor(Acm, tau, j, jb, th);
            const int c2 = j + jb;
            if (c2 < J + W) {                        // within-super-panel update only
                const int m = n - j;
                auto Vt = extract_V(Acm, j, jb, m);
                auto G = at::bmm(Vt, Vt.transpose(1, 2));
                auto T = build_T_from_G(G, tau, j, jb, 512);
                auto At = Acm.narrow(1, c2, (J + W) - c2).narrow(2, j, m);
                auto W1 = at::bmm(Vt, At.transpose(1, 2));
                auto W2 = at::bmm(T.transpose(1, 2), W1);
                At.baddbmm_(W2.transpose(1, 2), Vt, 1.0, -1.0);
            }
            j = c2;
        }
        const int c3 = J + W;
        if (c3 < n) {                                // merged SW-wide far-trailing update
            const int m = n - J;
            auto Vt = extract_V(Acm, J, W, m);
            auto G = at::bmm(Vt, Vt.transpose(1, 2));
            // wide far-panel (W=64): cuBLAS batched trsm beats the in-kernel
            // back-sub at high batch (33KB smem caps occupancy; 63-step serial
            // chain). The narrow jb<=32 T-builds above use the fast kernel.
            torch::Tensor T;
            if (W <= 32) {
                T = build_T_from_G(G, tau, J, W, 512);
            } else {
                auto tau_blk = tau.narrow(1, J, W);
                auto U = G.triu(1);
                auto eye_jb = eye_full.narrow(0, 0, W).narrow(1, 0, W);
                auto Amat = eye_jb + tau_blk.unsqueeze(2) * U;
                auto Dmat = at::diag_embed(tau_blk);
                T = at::linalg_solve_triangular(Amat, Dmat, true, true, true);
            }
            auto At = Acm.narrow(1, c3, n - c3).narrow(2, J, m);
            auto W1 = at::bmm(Vt, At.transpose(1, 2));
            auto W2 = at::bmm(T.transpose(1, 2), W1);
            At.baddbmm_(W2.transpose(1, 2), Vt, 1.0, -1.0);
        }
        J += W;
    }
    auto H = Acm.transpose(-2, -1).contiguous();
    return {H, tau};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("qr_blocked2", &qr_blocked2, "two-level blocked QR (super-panel + inner)");
    m.def("panel_factor", &panel_factor, "batched householder panel factor (col-major)");
    m.def("qr_full", &qr_full, "fully-fused batched householder QR (shared-mem, col-major)");
    m.def("qr_block_fused", &qr_block_fused, "blocked fully-fused batched householder QR (one launch)");
    m.def("build_T", &build_T, "compact-WY block reflector T from panel reflectors (dlarft)");
    m.def("build_T_from_G", &build_T_from_G, "compact-WY T from G=V^T V (parallel column solve)");
    m.def("extract_V", &extract_V, "emit clean compact-WY V panel (unit diag, reflectors below)");
    m.def("panel_factor_coop", &panel_factor_coop, "cooperative multi-block-per-matrix panel factor");
    m.def("qr_blocked", &qr_blocked, "blocked QR host loop in C++ (panel+T+trailing)");
    m.def("set_cublas_emulation", &set_cublas_emulation, "enable cuBLAS fp32 TC emulation");
}
'''

# Trailing update precision for the WY trailing GEMMs (W1, W2, final update).
#
#   "fp32"   -- plain torch.bmm in fp32 (CUDA-core path off Hopper/Blackwell,
#                tensor-core fp32-emulation on newer cuBLAS where available).
#                Proven correct on all 70 (shape, case) combinations (Ada).
#
#   "3xtf32" -- manual fp32-emulation via three TF32 tensor-core matmuls
#                (split each operand into hi/lo TF32 parts and sum
#                 hi@hi + hi@lo + lo@hi, dropping the lo@lo term). Recovers
#                ~2^-20 relative precision (vs ~2^-10 for plain TF32), which is
#                well inside every factor/orthogonality gate (looser than
#                ~2^-17 at n=4096). On Ada, naive 3xTF32 measured 1.6-3x SLOWER
#                because TF32 tensor cores aren't fast enough relative to fp32
#                CUDA cores to amortize 3x the matmuls + the Dekker-split
#                overhead. On B200, TF32 tensor-core throughput is far higher
#                relative to fp32, so this is the primary untested B200 lever
#                for the dominant trailing GEMM (n=512 b=640 etc).
#
# Default is "fp32" (safe, validated). 3xtf32 measured 1.5-1.6x SLOWER on B200
# (512: 17.8->29.1ms) -- confirms plain fp32 bmm already uses fast hardware, so
# the trailing GEMM is not the bottleneck.
_TRAILING_PRECISION = "fp32"

# int32 bit-mask keeping sign(1) + exponent(8) + top 10 mantissa bits (TF32).
_TF32_MASK = -8192  # 0xFFFFE000 as a signed int32


def _round_tf32(x: torch.Tensor) -> torch.Tensor:
    bits = x.view(torch.int32)
    return (bits & _TF32_MASK).view(torch.float32)


def _bmm_3xtf32(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
    """fp32-accurate batched matmul via 3 TF32 tensor-core matmuls.

    a = a_hi + a_lo, b = b_hi + b_lo (both splits exact since a_lo/b_lo are the
    truncated low mantissa bits of a/b). a@b ~= a_hi@b_hi + a_hi@b_lo + a_lo@b_hi,
    dropping the a_lo@b_lo term (relative magnitude ~2^-20).
    """
    a_hi = _round_tf32(a)
    a_lo = a - a_hi
    b_hi = _round_tf32(b)
    b_lo = b - b_hi
    prev = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        out = torch.bmm(a_hi, b_hi)
        out += torch.bmm(a_hi, b_lo)
        out += torch.bmm(a_lo, b_hi)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev
    return out


def _trail_mm(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
    if _TRAILING_PRECISION == "3xtf32":
        return _bmm_3xtf32(a, b)
    return torch.bmm(a, b)


_ext = None
_ext_failed = False


def _get_ext():
    global _ext, _ext_failed
    if _ext is not None or _ext_failed:
        return _ext
    try:
        from torch.utils.cpp_extension import load_inline
        _ext = load_inline(
            name="qr2l_ext",
            cpp_sources=_CPP,
            cuda_sources=_CUDA,
            extra_cuda_cflags=["-O3"],
            extra_ldflags=["-lcublas"],
            verbose=False,
        )
    except Exception:
        _ext_failed = True
        _ext = None
    return _ext


_emul_set = False


def _enable_emulation():
    """Turn on cuBLAS BF16x9 fp32 emulation (tensor-core, full fp32 accuracy)
    once, so the trailing-update at::bmm/baddbmm run on tensor cores."""
    global _emul_set
    if _emul_set:
        return
    ext = _get_ext()
    if ext is not None:
        try:
            ext.set_cublas_emulation(1)  # 1 = PERFORMANT (emulate only when faster)
            _emul_set = True
        except Exception:
            pass


if torch.cuda.is_available():
    _get_ext()


def _qr_blocked(A, block):
    """Batched blocked Householder QR -> (H, tau) in geqrf-compact form.

    Step 1: entire panel loop runs in C++ (qr_blocked) -- one Python entry instead
    of ~n/jb iterations each dispatching panel_factor + build_T + 3x bmm.
    """
    ext = _get_ext()
    _enable_emulation()
    Acm = A.transpose(-2, -1).contiguous()
    n = A.shape[-1]
    threads = _panel_threads(n)
    # Run the trailing bmm/baddbmm on TF32 tensor cores (1-pass) only where the
    # 20*n*eps factor gate has slack: n in {1024,2048} pass 22/22 on B200 with a
    # big speedup (1024 8.4->6.6ms, 2048 18.2->15.5ms). n=512's tighter gate fails
    # the mixed-scale case under TF32, so it stays fp32. allow_tf32 is a global
    # cuBLAS-handle math-mode flag that at::bmm reads inside the C++ loop -> toggle
    # it per shape around the call. (Custom bf16x3 WMMA trailing was a dead end:
    # 2.3x slower than cuBLAS; see qr-v2-perf-findings.)
    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = (1024 <= n <= 2048)
    try:
        if block < 0:   # 2-level: SW-wide merged far-trailing GEMM (M=SW) + IB=32 panel.
            # SW=128 was rejected pre-fusion (within-super-panel T-build too costly);
            # re-testing now that build_T_from_G made the inner T-build cheap.
            # per-n: 1024 (launch-bound, b60) wins with SW=128 (fatter M=128 far-GEMM,
            # fewer panels: 6.29->6.06ms); 512 (occupancy-bound, b640) regresses at
            # SW=128 (10.8ms) so keeps SW=64. B200-measured.
            sw = 128 if n > 640 else 64
            H, tau = ext.qr_blocked2(Acm, sw, 32, threads)
        else:
            H, tau = ext.qr_blocked(Acm, block, threads)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32
    return H, tau


# Experimental cooperative-panel blocked QR for the GPU-starved low-batch large-n
# shapes (2048 b8, 4096 b2): G thread blocks cooperate on each matrix's panel so the
# panel phase fills the GPU instead of using only `batch` SMs. Off by default; set
# _COOP_N to the min n that routes here (e.g. 2048) for a B200 A/B. _COOP_G is the
# blocks-per-matrix (B*G must stay co-resident for cudaLaunchCooperativeKernel).
#
# MEASURED DEAD END (B200 benchmark, G=8): 2048 b8 = 63.7ms (vs ~18ms baseline),
# 4096 b2 = 152ms (vs 52ms geqrf). The panel is sequential over its reflector columns,
# so a multi-block panel needs a grid-sync PER column (4 here) -> O(n) barriers
# (~16k at 4096). That barrier cost dominates and scales the WRONG way with G (Ada:
# G=8 166ms, G=16 309ms). The serial column dependency can't be cheaply parallelized
# across blocks; cuSOLVER geqrf's structure wins. Kept off (_COOP_N=None).
_COOP_N = None
_COOP_G = 8
_COOP_BLK = 32


def _qr_blocked_coop(A, block=None, G=None):
    """Blocked Householder QR with a COOPERATIVE multi-block panel. Mirrors the C++
    qr_blocked loop (panel -> compact-WY T -> trailing baddbmm) but factors each
    panel with G blocks/matrix (ext.panel_factor_coop) to fill the GPU at low batch.
    Python panel loop (n/jb iters) -- fine for the large-n shapes (few panels)."""
    ext = _get_ext()
    _enable_emulation()
    if block is None:
        block = _COOP_BLK
    if G is None:
        G = _COOP_G
    Acm = A.transpose(-2, -1).contiguous()
    B, n, _ = A.shape
    th = _panel_threads(n)
    tau = torch.zeros(B, n, device=A.device, dtype=torch.float32)
    eye_full = torch.eye(block, device=A.device, dtype=torch.float32)
    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = (1024 <= n <= 2048)
    try:
        j = 0
        while j < n:
            jb = min(block, n - j)
            ext.panel_factor_coop(Acm, tau, j, jb, G)
            c2 = j + jb
            if c2 < n:
                m = n - j
                Vt = ext.extract_V(Acm, j, jb, m)
                tau_blk = tau.narrow(1, j, jb)
                Gmat = torch.bmm(Vt, Vt.transpose(1, 2))
                U = Gmat.triu(1)
                eye_jb = eye_full.narrow(0, 0, jb).narrow(1, 0, jb)
                Amat = eye_jb + tau_blk.unsqueeze(2) * U
                Dmat = torch.diag_embed(tau_blk)
                T = torch.linalg.solve_triangular(Amat, Dmat, upper=True,
                                                  left=True, unitriangular=True)
                At = Acm.narrow(1, c2, n - c2).narrow(2, j, m)
                W1 = torch.bmm(Vt, At.transpose(1, 2))
                W2 = torch.bmm(T.transpose(1, 2), W1)
                At.baddbmm_(W2.transpose(1, 2), Vt, beta=1.0, alpha=-1.0)
            j = c2
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32
    H = Acm.transpose(-2, -1).contiguous()
    return H, tau


_FUSED_MAX = 200  # n*n*4 bytes must fit the opted-in dynamic shared (B200 ~227KB)


def _fused_threads(n):
    if n <= 64:
        return 128
    return 512


def _qr_fused(A):
    """Whole-matrix Householder QR in a single fused shared-memory kernel: one
    threadblock factors one matrix end-to-end, no Python panel loop / trailing
    GEMM round trips. Only viable while the matrix fits in dynamic shared mem."""
    ext = _get_ext()
    B, n, _ = A.shape
    Acm = A.transpose(-2, -1).contiguous()
    tau = torch.zeros(B, n, device=A.device, dtype=torch.float32)
    ext.qr_full(Acm, tau, _fused_threads(n))
    H = Acm.transpose(-2, -1).contiguous()
    return H, tau


_BLOCKFUSED_MAX = 1024  # one block/matrix needs decent batch for SM occupancy


def _qr_block_fused(A, blk):
    """Blocked fully-fused QR: one kernel launch does the entire factorization
    (panel in shared, trailing columns read once each). No host panel loop,
    no cuBLAS, no T-build, no V-extraction copies."""
    ext = _get_ext()
    B, n, _ = A.shape
    Acm = A.transpose(-2, -1).contiguous()
    tau = torch.zeros(B, n, device=A.device, dtype=torch.float32)
    ext.qr_block_fused(Acm, tau, blk, 256)
    H = Acm.transpose(-2, -1).contiguous()
    return H, tau


def _panel_threads(n):
    """Block size (blockDim.x) for the panel kernel, per n. Reductions run over
    sub-columns of length <= n, so match the thread count to n: tiny fused QR
    (n<=64) wants a single/double warp (more matrices packed per SM, minimal
    reduction latency); the large compute-bound shapes want wide blocks so each
    matrix's panel finishes faster. Templated -> compile-time-constant strides."""
    if n <= 64:
        return 64
    if n <= 384:
        return 512                     # 352 (low batch 40): 16 warps cut the jb=32 apply waves;
                                       # occupancy is a non-issue at b40. B200: 1.91 -> 1.78ms.
    if n <= 512:
        return 256                     # 512 (high batch 640): 8 warps -> 8 blocks/SM beats the
                                       # 16-warp apply (512 threads measured ~5% SLOWER at b640).
    if n <= 1024:
        return 512                     # lever #4: 1024 gains from 16 warps (1024 thr measured
                                       # identical 10.8ms -> cuBLAS-bound, keep proven 512)
    return 1024                        # 2048 (b8): no occupancy pressure (8 blocks).
                                       # B200 A/B: 512thr=23.3ms -> 1024thr=22.8ms (~2% faster).


def _plan(batch, n):
    """Return (use_blocked, block).  Use the batched kernel only where it can
    fill the GPU and beat cuSOLVER; otherwise fall back to torch.geqrf."""
    if n <= 64:
        # Tiny matrices: geqrf wastes ~0.3ms of launch overhead per call. Do the
        # whole QR as a single fused panel (block = n -> one kernel launch, one
        # threadblock per matrix, no trailing GEMM) when there are enough
        # matrices to populate the SMs.
        if batch >= 8:
            return True, n
        return False, 0
    if n > 2048:
        return False, 0                # n=4096: blocked panel starves at batch 2 (2 blocks
                                       # -> 2 of ~148 SMs). Re-measured with baddbmm + jb=32:
                                       # blocked=88.4ms vs geqrf=52.4ms. geqrf still wins (the
                                       # cuBLAS trailing is already fast; only the serial panel
                                       # is starved, and only HW clusters could unstarve it).
    if batch < 8:
        return False, 0                # too few matrices to fill the GPU
    # Per-n block size (validated with the warp-parallel panel apply).
    if n <= 352:
        block = 32                     # 352: M=16->32 trailing-GEMM-fill fix.
                                       # B200 A/B: jb16=2.13ms -> jb32=1.91ms (~10% faster).
    elif n <= 640:
        block = -1                     # 512: 2-level (SW=64 merged far-GEMM, IB=32 panel).
                                       # plain jb=64 was slower (wide panel); 2-level keeps
                                       # the cheap jb=32 panel but fattens the far-trailing GEMM.
                                       # B200 A/B: SW=128 was slower (9.83->10.8ms) -- the extra
                                       # within-super-panel apply outweighs the far-GEMM savings.
    elif n <= 1280:
        block = -1                     # 1024: 2-level (SW=64 merged far-GEMM, IB=32 panel).
                                       # B200 A/B: single-level jb=32 8.92ms -> 2-level 8.55ms
                                       # (~4% faster); the M=64 far-trailing GEMM + halved
                                       # far-GEMM count beats the extra within-super-panel ops.
                                       # Re-confirmed post-TF32: single-level jb=32 = 7.22ms
                                       # vs 2-level 6.58ms -> 2-level still wins.
    else:
        block = 32                     # 2048: single-level jb=32. 2-level measured SLOWER
                                       # (18.4 -> 18.8ms): batch 8 is launch-sensitive, the
                                       # extra within-super-panel ops outweigh the M=64 win.
                                       # B200 A/B: jb16=25.0ms, jb32=23.4ms; jb64=31.0ms.
    return True, block


def custom_kernel(data):
    A = data
    if A.dtype != torch.float32:
        A = A.to(torch.float32)
    if not A.is_cuda:
        raise RuntimeError("custom_kernel requires a CUDA tensor")
    A = A.contiguous()
    batch, n, _ = A.shape

    ext = _get_ext()
    if ext is not None and 2 <= n <= _FUSED_MAX:
        try:
            H, tau = _qr_fused(A)
            return H.contiguous(), tau.contiguous()
        except Exception:
            pass  # fall through to blocked / reference on any failure

    # NOTE: a one-block-per-matrix blocked megakernel (qr_block_fused) was tried
    # for 352/512/1024 to kill host overhead. It is correct but 2-6x SLOWER: the
    # in-kernel trailing update (CUDA-core warps) cannot match cuBLAS tensor-core
    # GEMM, which dominates these shapes. So large-n keeps the cuBLAS trailing.

    # Experimental: cooperative-panel blocked QR for starved low-batch large-n.
    if _COOP_N is not None and ext is not None and n >= _COOP_N:
        try:
            H, tau = _qr_blocked_coop(A)
            return H.contiguous(), tau.contiguous()
        except Exception:
            pass

    use_blocked, block = _plan(batch, n)
    if use_blocked and ext is not None:
        try:
            H, tau = _qr_blocked(A, block)
            return H.contiguous(), tau.contiguous()
        except Exception:
            pass  # fall through to the reference path on any failure

    H, tau = torch.geqrf(A)
    return H.contiguous(), tau.contiguous()
scrolls · 1359 lines total

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

Best evidence level for this revision: reported

JSON