Skip to content
KernelIndex
Search⌘K

submission 798644

switchtovim · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-798644?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
5.06ms
#177 of 515
2026-06-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cfd33e46ce59a58098e9f6e2805dc6d35dbb3ac842decf1d1d9e9c566fe1918c
license declaredunknown
license concludedunknown
authorsswitchtovim
imported2026-08-26

Techniques

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

num-warps = 4def _panel_factor_mb(H, k, bb, tau, BM=128, num_warps=4):
shared-memoryextern __shared__ float s[]; // m x bb, column-major: s[c*m + r], m = n-k
tile-m = 128def _panel_factor_mb(H, k, bb, tau, BM=128, num_warps=4):

Kernel source

submission.py984 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline

try:
    from task import input_t, output_t
except ImportError:
    input_t = torch.Tensor
    output_t = tuple[torch.Tensor, torch.Tensor]


# ----------------------------------------------------------------------------
# Shared-memory-resident panel factorization (CUDA, load_inline). One thread
# block per matrix factors the m x bb panel H[k:n, k:k+bb] entirely in shared
# memory: the panel is read from global ONCE (column-major) rather than re-read
# per column as the Triton panel does. Householder vectors / beta / tau written
# back in place. Used where the panel fits B200 smem (n=512: 512x64x4 = 128KB).
# Each trailing-column rank-1 update is handled by one warp (shuffle-reduced dot)
# so all 256 threads stay busy through the dominant in-panel apply.
# ----------------------------------------------------------------------------
_PANEL_CUDA = r"""
#include <ATen/cuda/CUDAContext.h>

// Single-block smem-resident panel. In one launch it (1) factors the m x bb panel,
// (2) writes the in-place H result (beta on diag, v below), (3) emits the dense
// unit-lower V (B,n,bcap) and (4) the bb x bb WY T-factor (B,bcap,bcap) — the latter
// two were previously separate Python ops (_extract_V + _form_T = bmm + triangular
// solve). T via closed form T = inv(diag(1/tau) + striu(V^T V, 1)) with V resident.
__global__ void qr_panel_k(float* __restrict__ H, float* __restrict__ tau,
                           float* __restrict__ Vd, float* __restrict__ Tout,
                           int n, int k, int bb, int bcap, int do_wy) {
    extern __shared__ float s[];        // m x bb, column-major: s[c*m + r], m = n-k
    int m = n - k;
    int t = threadIdx.x, nt = blockDim.x;
    int lane = t & 31, warp = t >> 5, nwarps = nt >> 5;
    float* Gm = s + (size_t)m * bb;     // bb x bb scratch (G then T-inverse)
    __shared__ float red[32];           // up to 32 warps (nt<=1024, hw threads/block cap)
    __shared__ float sh_tau, sh_vscale;
    __shared__ float taus[64];          // bb <= 64
    float* Hb = H + (size_t)blockIdx.x * n * n;
    for (int idx = t; idx < m * bb; idx += nt) {
        int c = idx / m, r = idx - c * m;
        s[idx] = Hb[(size_t)(k + r) * n + (k + c)];
    }
    __syncthreads();
    for (int c = 0; c < bb; c++) {
        float* sc = s + (size_t)c * m;
        float loc = 0.f;
        for (int r = c + 1 + t; r < m; r += nt) { float v = sc[r]; loc += v * v; }
        for (int o = 16; o > 0; o >>= 1) loc += __shfl_down_sync(0xffffffff, loc, o);
        if (lane == 0) red[warp] = loc;
        __syncthreads();
        if (t == 0) {
            float sigma = 0.f;
            for (int i = 0; i < nwarps; i++) sigma += red[i];
            float alpha = sc[c], norm = sqrtf(alpha * alpha + sigma);
            float beta = (alpha >= 0.f) ? -norm : norm;
            bool nz = sigma > 0.f;
            sh_tau    = nz ? (beta - alpha) / beta : 0.f;
            sh_vscale = nz ? 1.f / (alpha - beta) : 0.f;
            sc[c] = nz ? beta : alpha;
            if (do_wy) taus[c] = sh_tau;   // taus[64] only read on do_wy path; guard lets bb>64 (fused whole-matrix factor)
            tau[(size_t)blockIdx.x * n + (k + c)] = sh_tau;
        }
        __syncthreads();
        float tj = sh_tau, vs = sh_vscale;
        if (tj != 0.f) {
            for (int r = c + 1 + t; r < m; r += nt) sc[r] *= vs;
            __syncthreads();
            for (int cc = c + 1 + warp; cc < bb; cc += nwarps) {
                float* scc = s + (size_t)cc * m;
                float w = 0.f;
                for (int r = c + 1 + lane; r < m; r += 32) w += sc[r] * scc[r];
                for (int o = 16; o > 0; o >>= 1) w += __shfl_down_sync(0xffffffff, w, o);
                w = __shfl_sync(0xffffffff, w, 0);
                w = tj * (w + scc[c]);
                if (lane == 0) scc[c] -= w;
                for (int r = c + 1 + lane; r < m; r += 32) scc[r] -= sc[r] * w;
            }
        }
        __syncthreads();
    }
    // write H in place (+ dense unit-lower V if WY factors requested)
    float* Vb = Vd + (size_t)blockIdx.x * n * bcap;
    for (int idx = t; idx < m * bb; idx += nt) {
        int c = idx / m, r = idx - c * m;
        float val = s[idx];
        Hb[(size_t)(k + r) * n + (k + c)] = val;
        if (do_wy) Vb[(size_t)r * bcap + c] = (r < c) ? 0.f : (r == c) ? 1.f : val;
    }
    if (!do_wy) return;
    // G[i][j] = (V^T V)[i][j] for i<j  (one warp per pair, shfl-reduced dot)
    for (int idx = warp; idx < bb * bb; idx += nwarps) {
        int i = idx / bb, j = idx - i * bb;
        if (i < j) {
            float* si = s + (size_t)i * m;
            float* sj = s + (size_t)j * m;
            float acc = (lane == 0) ? si[j] : 0.f;   // r==j term: v_j[j]=1
            for (int r = j + 1 + lane; r < m; r += 32) acc += si[r] * sj[r];
            for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffff, acc, o);
            if (lane == 0) Gm[(size_t)i * bb + j] = acc;
        }
    }
    __syncthreads();
    // T = inv(M), M upper-tri with M[i][i]=1/tau_i, M[i][j>i]=G[i][j]. Solve M X = I.
    // Columns of X are independent -> one thread per column (smem back-substitution),
    // O(bb^2) depth instead of an O(bb^3) single-thread serial tail.
    float* Tsm = Gm + (size_t)bb * bb;
    for (int j = t; j < bb; j += nt) {
        Tsm[(size_t)j * bb + j] = taus[j];        // 1/M[j][j] = tau_j (0 -> ~identity)
        for (int i = j - 1; i >= 0; i--) {
            float sum = 0.f;
            for (int l = i + 1; l <= j; l++)
                sum += Gm[(size_t)i * bb + l] * Tsm[(size_t)l * bb + j];
            Tsm[(size_t)i * bb + j] = -taus[i] * sum;
        }
        for (int i = j + 1; i < bb; i++) Tsm[(size_t)i * bb + j] = 0.f;
    }
    __syncthreads();
    float* Tb = Tout + (size_t)blockIdx.x * bcap * bcap;
    for (int idx = t; idx < bb * bb; idx += nt) {
        int i = idx / bb, j = idx - i * bb;
        Tb[(size_t)i * bcap + j] = Tsm[(size_t)i * bb + j];
    }
}

// ---- Multi-block cooperative panel (tiny-batch huge-n). G blocks per matrix each
// own a BM-row slice resident in smem; cross-block reductions via global scratch +
// a hand-rolled device barrier (same structure as the Triton mb path, + residency).
__device__ __forceinline__ void grid_bar(int* counter, volatile int* sense_arr,
                                          int cb, int G, int* msh, int t) {
    __syncthreads();
    if (t == 0) {
        int ms = *msh ^ 1; *msh = ms;
        __threadfence();
        int old = atomicAdd(&counter[cb], 1);
        if (old == G - 1) {
            atomicExch(&counter[cb], 0);
            __threadfence();
            atomicExch((int*)&sense_arr[cb], ms);
        } else {
            while (sense_arr[cb] != ms) { }
        }
    }
    __syncthreads();
}

__global__ void qr_panel_mb_k(float* __restrict__ H, float* __restrict__ tau,
                              float* __restrict__ Vd,
                              int* counter, int* sense_arr, float* sigp,
                              float* alphasc, float* wp,
                              int n, int k, int bb, int G, int BM, int bcap, int do_wy) {
    extern __shared__ float s[];        // bb x BM, column-major: s[c*BM + rl]
    int t = threadIdx.x, nt = blockDim.x;
    int lane = t & 31, warp = t >> 5, nwarps = nt >> 5;
    int mid = blockIdx.x / G, g = blockIdx.x % G, m = n - k;
    float* Hb = H + (size_t)mid * n * n;
    int gbase = mid * G + g, go = g * BM;
    __shared__ int msh;
    __shared__ float red[32], sh_alpha, sh_sigma;
    if (t == 0) msh = 0;
    for (int idx = t; idx < bb * BM; idx += nt) {     // load slice into smem (once)
        int c = idx / BM, rl = idx - c * BM, pr = go + rl;
        s[idx] = (pr < m) ? Hb[(size_t)(k + pr) * n + (k + c)] : 0.f;
    }
    __syncthreads();
    int BG = gridDim.x;                               // B*G; double-buffer stride
    for (int c = 0; c < bb; c++) {
        int g_piv = c / BM, p = c & 1;                // ping-pong scratch by column parity
        float* sigp_p = sigp + p * BG;                // -> removes the 3rd grid barrier
        float* alp_p = alphasc + p * (BG / G);
        float* wp_p = wp + (size_t)p * BG * bb;
        float loc = 0.f;                              // partial sigma over local rows pr>c
        for (int rl = t; rl < BM; rl += nt) {
            int pr = go + rl;
            if (pr > c && pr < m) { float v = s[c * BM + rl]; loc += v * v; }
        }
        for (int o = 16; o > 0; o >>= 1) loc += __shfl_down_sync(0xffffffff, loc, o);
        if (lane == 0) red[warp] = loc;
        __syncthreads();
        if (t == 0) { float ss = 0.f; for (int i = 0; i < nwarps; i++) ss += red[i]; sigp_p[gbase] = ss; }
        if (g == g_piv && t == 0) alp_p[mid] = s[c * BM + (c - go)];
        grid_bar(counter, sense_arr, mid, G, &msh, t);
        if (t == 0) {
            float sg = 0.f; for (int i = 0; i < G; i++) sg += sigp_p[mid * G + i];
            sh_sigma = sg; sh_alpha = alp_p[mid];
        }
        __syncthreads();
        float sigma = sh_sigma, alpha = sh_alpha;
        float norm = sqrtf(alpha * alpha + sigma);
        float beta = (alpha >= 0.f) ? -norm : norm;
        bool nz = sigma > 0.f;
        float tj = nz ? (beta - alpha) / beta : 0.f, vs = nz ? 1.f / (alpha - beta) : 0.f;
        if (g == g_piv && t == 0) {
            s[c * BM + (c - go)] = nz ? beta : alpha;
            tau[(size_t)mid * n + (k + c)] = tj;
        }
        if (nz) {
            for (int rl = t; rl < BM; rl += nt) { int pr = go + rl; if (pr > c && pr < m) s[c * BM + rl] *= vs; }
            __syncthreads();
            for (int cc = c + 1 + warp; cc < bb; cc += nwarps) {   // partial w[cc]
                float wsum = 0.f;
                for (int rl = lane; rl < BM; rl += 32) {
                    int pr = go + rl;
                    if (pr >= c && pr < m) { float v = (pr == c) ? 1.f : s[c * BM + rl]; wsum += v * s[cc * BM + rl]; }
                }
                for (int o = 16; o > 0; o >>= 1) wsum += __shfl_down_sync(0xffffffff, wsum, o);
                if (lane == 0) wp_p[(size_t)gbase * bb + cc] = wsum;
            }
        }
        grid_bar(counter, sense_arr, mid, G, &msh, t);
        if (nz) {
            for (int cc = c + 1 + warp; cc < bb; cc += nwarps) {
                float wv = 0.f; for (int i = 0; i < G; i++) wv += wp_p[(size_t)(mid * G + i) * bb + cc];
                wv *= tj;
                for (int rl = lane; rl < BM; rl += 32) {
                    int pr = go + rl;
                    if (pr >= c && pr < m) { float v = (pr == c) ? 1.f : s[c * BM + rl]; s[cc * BM + rl] -= v * wv; }
                }
            }
        }
        __syncthreads();   // block-local only: order this column's smem writes before next col
    }
    float* Vb = Vd + (size_t)mid * n * bcap;
    for (int idx = t; idx < bb * BM; idx += nt) {      // write slice back (+ dense V)
        int c = idx / BM, rl = idx - c * BM, pr = go + rl;
        if (pr < m) {
            float val = s[idx];
            Hb[(size_t)(k + pr) * n + (k + c)] = val;
            if (do_wy) Vb[(size_t)pr * bcap + c] = (pr < c) ? 0.f : (pr == c) ? 1.f : val;
        }
    }
}

// Standalone WY T-factor from a dense unit-lower V (used by the multi-block panel,
// which can't cheaply reduce V^T V across its row-blocks). One block per matrix:
// G = striu(V^T V), then T = inv(diag(1/tau)+G) by one-thread-per-column back-sub.
// Replaces the Python _form_T (V^T V bmm + triangular solve + triu + diag) launches.
__global__ void form_t_k(const float* __restrict__ Vd, const float* __restrict__ tau,
                         float* __restrict__ Tout, int n, int k, int bb, int bcap) {
    extern __shared__ float sm[];            // Gm[bb*bb] then Tsm[bb*bb]
    int m = n - k;
    int t = threadIdx.x, nt = blockDim.x, lane = t & 31, warp = t >> 5, nwarps = nt >> 5;
    int mid = blockIdx.x;
    const float* Vb = Vd + (size_t)mid * n * bcap;
    float* Gm = sm;
    float* Tsm = sm + (size_t)bb * bb;
    __shared__ float taus[64];
    for (int i = t; i < bb; i += nt) taus[i] = tau[(size_t)mid * n + (k + i)];
    __syncthreads();
    for (int idx = warp; idx < bb * bb; idx += nwarps) {
        int i = idx / bb, j = idx - i * bb;
        if (i < j) {
            float acc = 0.f;
            for (int r = j + lane; r < m; r += 32)
                acc += Vb[(size_t)r * bcap + i] * Vb[(size_t)r * bcap + j];
            for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffff, acc, o);
            if (lane == 0) Gm[(size_t)i * bb + j] = acc;
        }
    }
    __syncthreads();
    for (int j = t; j < bb; j += nt) {
        Tsm[(size_t)j * bb + j] = taus[j];
        for (int i = j - 1; i >= 0; i--) {
            float sum = 0.f;
            for (int l = i + 1; l <= j; l++)
                sum += Gm[(size_t)i * bb + l] * Tsm[(size_t)l * bb + j];
            Tsm[(size_t)i * bb + j] = -taus[i] * sum;
        }
        for (int i = j + 1; i < bb; i++) Tsm[(size_t)i * bb + j] = 0.f;
    }
    __syncthreads();
    float* Tb = Tout + (size_t)mid * bcap * bcap;
    for (int idx = t; idx < bb * bb; idx += nt) {
        int i = idx / bb, j = idx - i * bb;
        Tb[(size_t)i * bcap + j] = Tsm[(size_t)i * bb + j];
    }
}

void qr_panel_mb(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,
                 torch::Tensor counter, torch::Tensor sense_arr, torch::Tensor sigp,
                 torch::Tensor alphasc, torch::Tensor wp, int k, int bb, int G,
                 int BM, int do_wy) {
    int B = H.size(0), n = H.size(1), nt = 1024, bcap = Vd.size(2);   // 32 warps (red[32]):
    size_t smem = (size_t)bb * BM * sizeof(float);                    // fill idle SMs (B*G
    qr_panel_mb_k<<<B * G, nt, smem>>>(                              // ~16 blocks) + hide barrier spin
        H.data_ptr<float>(), tau.data_ptr<float>(), Vd.data_ptr<float>(),
        counter.data_ptr<int>(), sense_arr.data_ptr<int>(), sigp.data_ptr<float>(),
        alphasc.data_ptr<float>(), wp.data_ptr<float>(), n, k, bb, G, BM, bcap, do_wy);
}

void form_t(torch::Tensor Vd, torch::Tensor tau, torch::Tensor Tout, int k, int bb) {
    int B = Vd.size(0), n = Vd.size(1), bcap = Vd.size(2), nt = 256;
    size_t smem = 2 * (size_t)bb * bb * sizeof(float);
    form_t_k<<<B, nt, smem>>>(
        Vd.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), n, k, bb, bcap);
}

void qr_panel_init(int max_smem) {
    cudaFuncSetAttribute(qr_panel_k,
        cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
    cudaFuncSetAttribute(qr_panel_mb_k,
        cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
}

void qr_panel(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,
              torch::Tensor Tout, int k, int bb, int do_wy, int nt) {
    int B = H.size(0), n = H.size(1), bcap = Vd.size(2);
    int m = n - k;   // nt (<=1024, red[32]): caller fills the SM when batch is small,
                     // 512 when the batch already saturates the GPU (else reduce-tree waste)
    size_t smem = ((size_t)m * bb + (do_wy ? 2 * (size_t)bb * bb : 0)) * sizeof(float);
    qr_panel_k<<<B, nt, smem>>>(
        H.data_ptr<float>(), tau.data_ptr<float>(), Vd.data_ptr<float>(),
        Tout.data_ptr<float>(), n, k, bb, bcap, do_wy);
}
"""
_PANEL_CPP = (
    "void qr_panel(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,"
    " torch::Tensor Tout, int k, int bb, int do_wy, int nt);\n"
    "void qr_panel_mb(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,"
    " torch::Tensor counter, torch::Tensor sense_arr, torch::Tensor sigp,"
    " torch::Tensor alphasc, torch::Tensor wp, int k, int bb, int G, int BM,"
    " int do_wy);\n"
    "void form_t(torch::Tensor Vd, torch::Tensor tau, torch::Tensor Tout, int k, int bb);\n"
    "void qr_panel_init(int max_smem);"
)

_PANEL = None
try:
    _PANEL = load_inline(
        name="qr_panel_mod",
        cpp_sources=[_PANEL_CPP],
        cuda_sources=[_PANEL_CUDA],
        functions=["qr_panel", "qr_panel_mb", "form_t", "qr_panel_init"],
        verbose=False,
    )
    _PANEL.qr_panel_init(227000)
except Exception as _e:
    _PANEL = None


# ----------------------------------------------------------------------------
# CholeskyQR1 + Householder-reconstruction panel (chol_bs) for the n=4096 wall.
# The two ops that would sync via cuSOLVER (32x32 Cholesky, unpivoted LU) are custom
# one-block-per-matrix kernels; the rest (G=A^TA, Q1=A R^-1, Y2=-Q1bot S^-1) are cuBLAS
# GEMM/trsm. This trades the sequential bb-column Householder critical path (the panel
# wall at batch 2) for parallel GEMMs + a trivial 32x32 factorization.
# ----------------------------------------------------------------------------
_TSQR_CUDA = r"""
#include <ATen/cuda/CUDAContext.h>

// Upper-Cholesky R (R^T R = G) of bb x bb SPD G, one block/matrix. info[b]=0 if PD else
// (first non-PD row + 1) so the caller routes that panel to the direct Householder fallback.
__global__ void chol_up_k(const float* __restrict__ G, float* __restrict__ R,
                          int* __restrict__ info, int bb, float rel_tol) {
    extern __shared__ float sm[];                 // Gs[bb*bb] | Rs[bb*bb]
    float* Gs = sm; float* Rs = sm + (size_t)bb * bb;
    int t = threadIdx.x;
    const float* Gb = G + (size_t)blockIdx.x * bb * bb;
    for (int idx = t; idx < bb * bb; idx += blockDim.x) { Gs[idx] = Gb[idx]; Rs[idx] = 0.f; }
    __shared__ int bad; if (t == 0) bad = 0;
    __syncthreads();
    for (int i = 0; i < bb; i++) {
        if (t >= i && t < bb) {
            float s = 0.f;
            for (int k = 0; k < i; k++) s += Rs[(size_t)k * bb + i] * Rs[(size_t)k * bb + t];
            if (t == i) {
                float gii = Gs[(size_t)i * bb + i], d = gii - s;
                // relative pivot guard: tiny d/gii => column near-dependent (cond(A^TA) past
                // fp32) => flag this matrix bad so the caller redoes it with direct geqrf.
                if (d <= rel_tol * gii) { atomicExch(&bad, i + 1); d = (d > 0.f) ? d : 1.f; }
                Rs[(size_t)i * bb + i] = sqrtf(d);
            } else Rs[(size_t)i * bb + t] = Gs[(size_t)i * bb + t] - s;
        }
        __syncthreads();
        if (t > i && t < bb) Rs[(size_t)i * bb + t] /= Rs[(size_t)i * bb + i];
        __syncthreads();
    }
    float* Rb = R + (size_t)blockIdx.x * bb * bb;
    for (int idx = t; idx < bb * bb; idx += blockDim.x) Rb[idx] = Rs[idx];
    if (t == 0) info[blockIdx.x] = bad;
}

// Unpivoted LU of M0 (bb x bb), one block/matrix: in place strict-lower = Y1 multipliers,
// upper(incl diag) = S so M0 = Y1 @ S. |pivot|<=tol -> degenerate col (identity reflector):
// S row = e_i, Y1 col below = 0, degen[col]=1.
__global__ void lu_unpiv_k(const float* __restrict__ M0, float* __restrict__ LU,
                           int* __restrict__ degen, int bb, float tol) {
    extern __shared__ float sm[];                 // Ws[bb*bb]
    float* Ws = sm; int t = threadIdx.x;
    const float* Mb = M0 + (size_t)blockIdx.x * bb * bb;
    for (int idx = t; idx < bb * bb; idx += blockDim.x) Ws[idx] = Mb[idx];
    int* db = degen + (size_t)blockIdx.x * bb;
    for (int idx = t; idx < bb; idx += blockDim.x) db[idx] = 0;
    __syncthreads();
    for (int i = 0; i < bb; i++) {
        float piv = Ws[(size_t)i * bb + i];
        bool deg = fabsf(piv) <= tol;
        if (deg) {
            if (t == i) { Ws[(size_t)i * bb + i] = 1.f; db[i] = 1; }
            if (t > i && t < bb) { Ws[(size_t)t * bb + i] = 0.f; Ws[(size_t)i * bb + t] = 0.f; }
            __syncthreads();
        } else {
            if (t > i && t < bb) Ws[(size_t)t * bb + i] /= piv;
            __syncthreads();
            if (t > i && t < bb) {
                float u = Ws[(size_t)i * bb + t];
                for (int r = i + 1; r < bb; r++) Ws[(size_t)r * bb + t] -= Ws[(size_t)r * bb + i] * u;
            }
            __syncthreads();
        }
    }
    float* Lb = LU + (size_t)blockIdx.x * bb * bb;
    for (int idx = t; idx < bb * bb; idx += blockDim.x) Lb[idx] = Ws[idx];
}

void chol_up(torch::Tensor G, torch::Tensor R, torch::Tensor info, int bb, double rel_tol) {
    int B = G.size(0), nt = bb < 32 ? 32 : bb;
    chol_up_k<<<B, nt, 2 * (size_t)bb * bb * sizeof(float)>>>(
        G.data_ptr<float>(), R.data_ptr<float>(), info.data_ptr<int>(), bb, (float)rel_tol);
}
void lu_unpiv(torch::Tensor M0, torch::Tensor LU, torch::Tensor degen, int bb, double tol) {
    int B = M0.size(0), nt = bb < 32 ? 32 : bb;
    lu_unpiv_k<<<B, nt, (size_t)bb * bb * sizeof(float)>>>(
        M0.data_ptr<float>(), LU.data_ptr<float>(), degen.data_ptr<int>(), bb, (float)tol);
}

void tsqr_init(int max_smem) {   // opt-in dynamic smem so wide panels (bb=128: chol 128KB) fit
    cudaFuncSetAttribute(chol_up_k, cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
    cudaFuncSetAttribute(lu_unpiv_k, cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
}
"""
_TSQR_CPP = (
    "void chol_up(torch::Tensor G, torch::Tensor R, torch::Tensor info, int bb, double rel_tol);\n"
    "void lu_unpiv(torch::Tensor M0, torch::Tensor LU, torch::Tensor degen, int bb, double tol);\n"
    "void tsqr_init(int max_smem);"
)
_TSQR = None
try:
    _TSQR = load_inline(
        name="qr_tsqr_mod",
        cpp_sources=[_TSQR_CPP],
        cuda_sources=[_TSQR_CUDA],
        functions=["chol_up", "lu_unpiv", "tsqr_init"],
        verbose=False,
    )
    _TSQR.tsqr_init(227000)
except Exception as _e:
    _TSQR = None


_CUDA_MB_SCRATCH = {}


def _panel_factor_cuda_mb(H, k, bb, tau, BM, Vd, do_wy):
    """Multi-block smem-resident CUDA panel. G blocks/matrix cooperate via grid barrier.
    Emits dense unit-lower V into Vd (each block writes its own row-slice) when do_wy."""
    B, n, _ = H.shape
    Gmax = (n + BM - 1) // BM
    key = (B, Gmax, bb, H.device)
    sc = _CUDA_MB_SCRATCH.get(key)
    if sc is None:
        z = lambda *s: torch.zeros(*s, dtype=torch.int32, device=H.device)
        f = lambda *s: torch.zeros(*s, dtype=torch.float32, device=H.device)
        # sigp/alphasc/wp doubled for column-parity ping-pong (drops the 3rd barrier)
        sc = (z(B), z(B), f(2 * B * Gmax), f(2 * B), f(2 * B * Gmax * bb))
        _CUDA_MB_SCRATCH[key] = sc
    counter, sense, sigp, alphasc, wp = sc
    G = (n - k + BM - 1) // BM
    _PANEL.qr_panel_mb(
        H, tau, Vd, counter, sense, sigp, alphasc, wp, k, bb, G, BM, do_wy
    )


# TF32 tensor cores for the WY trailing-update bmms (~30% of runtime). TF32 keeps
# an FP32 accumulator; multiplicands round to 19 bits. Toggled per-call: safe for
# n>=1024 (rtol scales with n) but breaks the tightest n=512 wide-dynamic-range
# cases (band/rowscale lose their small entries to the 10-bit TF32 mantissa).
def _set_tf32(on):
    torch.backends.cuda.matmul.allow_tf32 = on


# (Shelved: a 3xTF32 / fp16x3 split recovers FP32 accuracy for the trailing GEMM
# but is a net SLOWDOWN on B200 — the GEMM isn't the bottleneck. See qr_notes.md.)


# ----------------------------------------------------------------------------
# Triton panel factorization: factor columns [k, k+bb) of each matrix in the
# batch in place (LARFG convention), writing Householder vectors below the
# diagonal, beta on the diagonal, and tau.  grid = (batch,).  Sequential over
# the bb columns (runtime loop) with row tiling of size BM.
# ----------------------------------------------------------------------------
@triton.jit
def _panel_kernel(
    Hptr,
    tauptr,
    n,
    k,
    bb,
    sH0,
    sH1,
    sH2,
    stau0,
    stau1,
    BB: tl.constexpr,
    BM: tl.constexpr,
    PADDED: tl.constexpr,
):
    # BB = tile width (power of 2) >= bb (actual panel columns). When bb < BB the
    # extra columns are masked out (PADDED).  Benchmark shapes have bb == BB (no mask).
    pid = tl.program_id(0)
    Hb = Hptr + pid * sH0
    coff = tl.arange(0, BB)
    cols = k + coff  # global panel column indices
    for c in range(0, bb):
        j = k + c
        # ---- pass 1: alpha, sigma = sum_{i>j} H[i,j]^2 ----
        alpha = tl.load(Hb + j * sH1 + j * sH2)
        sigma = tl.zeros((), dtype=tl.float32)
        for r0 in range(0, n, BM):
            rows = r0 + tl.arange(0, BM)
            m = (rows > j) & (rows < n)
            col = tl.load(Hb + rows * sH1 + j * sH2, mask=m, other=0.0)
            sigma += tl.sum(col * col)
        nrm = tl.sqrt(alpha * alpha + sigma)
        beta = tl.where(alpha >= 0, -nrm, nrm)
        nz = sigma > 0.0
        tau_c = tl.where(nz, (beta - alpha) / beta, 0.0)
        inv = tl.where(nz, 1.0 / (alpha - beta), 0.0)
        tl.store(Hb + j * sH1 + j * sH2, tl.where(nz, beta, alpha))
        tl.store(tauptr + pid * stau0 + j * stau1, tau_c)
        if nz:
            # scale v below the diagonal
            for r0 in range(0, n, BM):
                rows = r0 + tl.arange(0, BM)
                m = (rows > j) & (rows < n)
                col = tl.load(Hb + rows * sH1 + j * sH2, mask=m, other=0.0)
                tl.store(Hb + rows * sH1 + j * sH2, col * inv, mask=m)
            tl.debug_barrier()
            # ---- apply reflector c to panel cols (vectorized over BB) ----
            # w[cc] = tau_c * sum_{rows>=j} v_row * H[row, k+cc]
            w = tl.zeros((BB,), dtype=tl.float32)
            for r0 in range(0, n, BM):
                rows = r0 + tl.arange(0, BM)
                mr = (rows > j) & (rows < n)
                v = tl.load(Hb + rows * sH1 + j * sH2, mask=mr, other=0.0)
                v = tl.where(rows == j, 1.0, v)
                mt = (rows[:, None] >= j) & (rows[:, None] < n)
                if PADDED:
                    mt = mt & (coff[None, :] < bb)
                ct = tl.load(
                    Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
                )
                w += tl.sum(v[:, None] * ct, axis=0)
            w = tl.where(coff > c, tau_c * w, 0.0)  # only update cols cc>c
            for r0 in range(0, n, BM):
                rows = r0 + tl.arange(0, BM)
                mr = (rows > j) & (rows < n)
                v = tl.load(Hb + rows * sH1 + j * sH2, mask=mr, other=0.0)
                v = tl.where(rows == j, 1.0, v)
                mt = (rows[:, None] >= j) & (rows[:, None] < n)
                if PADDED:
                    mt = mt & (coff[None, :] < bb)
                ct = tl.load(
                    Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
                )
                ct = ct - v[:, None] * w[None, :]
                tl.store(Hb + rows[:, None] * sH1 + cols[None, :] * sH2, ct, mask=mt)
        tl.debug_barrier()


@triton.jit
def _bar(counter, sense, cb, G, my_sense):
    # Device-scope grid barrier across the G blocks of one matrix. acq_rel/gpu makes
    # each block's partial-data writes visible to the others. Returns flipped sense.
    my_sense = my_sense ^ 1
    old = tl.atomic_add(counter + cb, 1, sem="acq_rel", scope="gpu")
    if old == G - 1:
        tl.atomic_xchg(counter + cb, 0, sem="relaxed", scope="gpu")
        tl.atomic_xchg(sense + cb, my_sense, sem="release", scope="gpu")
    else:
        while tl.load(sense + cb, volatile=True) != my_sense:
            pass
    return my_sense


@triton.jit
def _panel_kernel_mb(
    Hptr,
    tauptr,
    n,
    k,
    bb,
    G,
    sH0,
    sH1,
    sH2,
    stau0,
    stau1,
    counter,
    sense,
    sigp,
    wp,
    alphasc,
    BB: tl.constexpr,
    BM: tl.constexpr,
    PADDED: tl.constexpr,
):
    # Multi-block panel factorization: G blocks cooperate on one matrix via a grid
    # barrier, each owning a BM-row slice. Used for small-batch large-n (panel-bound).
    pid = tl.program_id(0)
    mid = pid // G
    g = pid % G
    Hb = Hptr + mid * sH0
    coff = tl.arange(0, BB)
    cols = k + coff
    rows = k + g * BM + tl.arange(0, BM)  # this block's row slice
    rv = rows < n
    # local sense starts at 0 for all blocks; global sense/counter start at 0 and return
    # to 0 after the panel (4*bb barriers = even), so reused scratch stays clean.
    ms = 0
    for c in range(0, bb):
        j = k + c
        # block 0 owns row j -> broadcast alpha = H[j,j] (cross-block read otherwise stale)
        if g == 0:
            tl.atomic_xchg(
                alphasc + mid,
                tl.load(Hb + j * sH1 + j * sH2),
                sem="release",
                scope="gpu",
            )
        # ---- partial sigma over rows>j ----
        below = rv & (rows > j)
        col = tl.load(Hb + rows * sH1 + j * sH2, mask=below, other=0.0)
        tl.atomic_xchg(
            sigp + mid * G + g, tl.sum(col * col), sem="release", scope="gpu"
        )
        ms = _bar(counter, sense, mid, G, ms)
        sigma = tl.zeros((), dtype=tl.float32)
        for i in range(0, G):
            sigma += tl.atomic_add(sigp + mid * G + i, 0.0, sem="acquire", scope="gpu")
        alpha = tl.atomic_add(alphasc + mid, 0.0, sem="acquire", scope="gpu")
        nrm = tl.sqrt(alpha * alpha + sigma)
        beta = tl.where(alpha >= 0, -nrm, nrm)
        nz = sigma > 0.0
        tau_c = tl.where(nz, (beta - alpha) / beta, 0.0)
        inv = tl.where(nz, 1.0 / (alpha - beta), 0.0)
        if g == 0:
            tl.store(Hb + j * sH1 + j * sH2, tl.where(nz, beta, alpha))
            tl.store(tauptr + mid * stau0 + j * stau1, tau_c)
        if nz:
            tl.store(Hb + rows * sH1 + j * sH2, col * inv, mask=below)
            # ---- partial w[cc] = sum rows>=j v*H[row,cc] (own rows -> no barrier vs scale) ----
            v = tl.load(Hb + rows * sH1 + j * sH2, mask=below, other=0.0)
            v = tl.where(rows == j, 1.0, v)
            mt = (rows[:, None] >= j) & rv[:, None]
            if PADDED:
                mt = mt & (coff[None, :] < bb)
            ct = tl.load(
                Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
            )
            tl.atomic_xchg(
                wp + (mid * G + g) * BB + coff,
                tl.sum(v[:, None] * ct, axis=0),
                sem="release",
                scope="gpu",
            )
        ms = _bar(counter, sense, mid, G, ms)
        if nz:
            w = tl.zeros((BB,), dtype=tl.float32)
            for i in range(0, G):
                w += tl.atomic_add(
                    wp + (mid * G + i) * BB + coff, 0.0, sem="acquire", scope="gpu"
                )
            w = tl.where(coff > c, tau_c * w, 0.0)
            v = tl.load(Hb + rows * sH1 + j * sH2, mask=below, other=0.0)
            v = tl.where(rows == j, 1.0, v)
            mt = (rows[:, None] >= j) & rv[:, None]
            if PADDED:
                mt = mt & (coff[None, :] < bb)
            ct = tl.load(
                Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
            )
            ct = ct - v[:, None] * w[None, :]
            tl.store(Hb + rows[:, None] * sH1 + cols[None, :] * sH2, ct, mask=mt)
        ms = _bar(counter, sense, mid, G, ms)


_MB_SCRATCH = {}


def _panel_factor_mb(H, k, bb, tau, BM=128, num_warps=4):
    B, n, _ = H.shape
    BB = _next_pow2(bb)
    G = (n - k + BM - 1) // BM  # blocks per matrix
    key = (B, G, BB, H.device)
    sc = _MB_SCRATCH.get(key)
    if sc is None:
        counter = torch.zeros(B, dtype=torch.int32, device=H.device)
        sense = torch.zeros(B, dtype=torch.int32, device=H.device)
        sigp = torch.zeros(B * G, dtype=torch.float32, device=H.device)
        wp = torch.zeros(B * G * BB, dtype=torch.float32, device=H.device)
        alphasc = torch.zeros(B, dtype=torch.float32, device=H.device)
        sc = (counter, sense, sigp, wp, alphasc)
        _MB_SCRATCH[key] = sc
    counter, sense, sigp, wp, alphasc = sc
    _panel_kernel_mb[(B * G,)](
        H,
        tau,
        n,
        k,
        bb,
        G,
        H.stride(0),
        H.stride(1),
        H.stride(2),
        tau.stride(0),
        tau.stride(1),
        counter,
        sense,
        sigp,
        wp,
        alphasc,
        BB=BB,
        BM=BM,
        PADDED=(BB != bb),
        num_warps=num_warps,
    )


def _next_pow2(x):
    return 1 << (x - 1).bit_length()


def _panel_threads(B, width):
    """Threads/block for the single-block panel kernel (qr_panel_k).

    Once the batch already saturates the GPU (>~num_SMs blocks) 16 warps is best — extra
    warps only add reduction-tree overhead (n=512 B=640 measurably regresses with more).
    Otherwise the SM is under-utilized (smem caps it at 1 block/SM), so fill it with the
    max 32 warps = 1024 threads (the hardware threads/block limit; can't go to 64 warps).
    `width` is accepted for API symmetry but the 1024 cap makes it moot here."""
    if B >= 256:
        return 512
    return 1024


def _panel_factor(H, k, bb, tau, BM=128, num_warps=4):
    B, n, _ = H.shape
    BB = _next_pow2(bb)  # tile width must be a power of 2
    _panel_kernel[(B,)](
        H,
        tau,
        n,
        k,
        bb,
        H.stride(0),
        H.stride(1),
        H.stride(2),
        tau.stride(0),
        tau.stride(1),
        BB=BB,
        BM=BM,
        PADDED=(BB != bb),
        num_warps=num_warps,
    )


def _form_T(V, tau):
    """V: (B,m,b) unit-lower; tau: (B,b) -> T (B,b,b) upper, Q=I-V T V^T.

    Closed form (no Python loop): T = inv(diag(1/tau) + striu(V^T V, 1)).
    tau==0 (identity reflector) handled via a huge diagonal -> T entry ~0."""
    B, m, b = V.shape
    M = torch.triu(torch.bmm(V.transpose(1, 2), V), 1)  # strictly upper
    inv_tau = torch.where(tau != 0, 1.0 / tau, tau.new_full((), 1e30))
    M.diagonal(dim1=-2, dim2=-1).copy_(inv_tau)
    eye = torch.eye(b, device=V.device, dtype=V.dtype).expand(B, b, b)
    return torch.linalg.solve_triangular(M, eye, upper=True)


def _extract_V(H, k, bb):
    """Unit-lower-trapezoidal V (B,m,bb) from factored panel."""
    B, n, _ = H.shape
    V = H[:, k:, k : k + bb].clone()
    ii = torch.arange(bb, device=H.device)
    top = V[:, :bb, :]
    top.masked_fill_(ii[:, None] < ii[None, :], 0.0)
    top[:, ii, ii] = 1.0
    return V


def _recon_cuda_panel(H, tau, k, bb, tol=1e-7, rel_tol=1e-3):
    """chol_bs reconstruction of panel H[:,k:,k:k+bb]: writes packed H (R above diag, v's below)
    + tau, returns (V, info) where V is the dense unit-lower (B,m,bb) for the trailing WY update
    and info[b]!=0 flags an ill-conditioned chol (caller redoes that matrix via reference geqrf)."""
    B, n, _ = H.shape
    dev = H.device
    m = n - k
    panel = H[:, k:, k : k + bb]
    G = (panel.transpose(1, 2) @ panel).contiguous()  # cuBLAS
    R = torch.empty(B, bb, bb, device=dev)
    info = torch.empty(B, dtype=torch.int32, device=dev)
    _TSQR.chol_up(G, R, info, bb, rel_tol)  # R upper, R^T R = G
    # Q1 = panel @ R^-1   via   R^T Q1^T = panel^T
    Q1 = torch.linalg.solve_triangular(
        R.transpose(1, 2), panel.transpose(1, 2), upper=False
    ).transpose(1, 2)
    M0 = (torch.eye(bb, device=dev) - Q1[:, :bb, :]).contiguous()
    LU = torch.empty(B, bb, bb, device=dev)
    degen = torch.empty(B, bb, dtype=torch.int32, device=dev)
    _TSQR.lu_unpiv(M0, LU, degen, bb, tol)  # M0 = Y1 @ S
    Y1 = torch.tril(LU, -1) + torch.eye(bb, device=dev)
    if m > bb:
        S = torch.triu(LU)
        Y2 = torch.linalg.solve_triangular(
            S.transpose(1, 2), (-Q1[:, bb:, :]).transpose(1, 2), upper=False
        ).transpose(1, 2)  # Y2 = -Q1bot S^-1
        V = torch.cat([Y1, Y2], dim=1)
    else:
        V = Y1
    degb = degen.bool()
    if degb.any():
        ii = torch.arange(bb, device=dev)
        cid = torch.arange(m, device=dev)[None, :, None] == ii[None, None, :]
        V = torch.where(
            degb[:, None, :].expand(B, m, bb), cid.to(V.dtype).expand(B, m, bb), V
        )
    tk = 2.0 / (V * V).sum(dim=1)
    tk = torch.where(degb, torch.zeros_like(tk), tk)
    ar = torch.arange(m, device=dev)
    ac = torch.arange(bb, device=dev)
    upper = (ar[:, None] <= ac[None, :])[None]
    lower = (ar[:, None] > ac[None, :])[None]
    Rfull = torch.zeros_like(V)
    Rfull[:, :bb, :] = R
    H[:, k:, k : k + bb] = torch.where(
        upper, Rfull, torch.where(lower, V, torch.zeros_like(V))
    )
    tau[:, k : k + bb] = tk
    return V, info


def _qr_tsqr_cholbs(A, b=32, recon_min_m=64, tol=1e-7, rel_tol=1e-3):
    """Block QR for n>=4096 via the chol_bs panel; near-square trailing panels (m<recon_min_m)
    use the direct CUDA Householder panel. If any panel's chol is ill-conditioned (cond(A^TA)
    past fp32: upper/rankdef/nearcollinear), the whole matrix is redone with reference geqrf -
    a rare correctness-only path; dense (the timed case) never trips it (one sync at the end)."""
    B, n, _ = A.shape
    H = A.clone()
    tau = torch.zeros(B, n, device=A.device, dtype=A.dtype)
    nt = _panel_threads(B, b)
    _dummy = torch.empty(1, 1, 1, device=A.device, dtype=A.dtype)
    any_bad = torch.zeros((), device=A.device)
    for k in range(0, n, b):
        bb = min(b, n - k)
        if (n - k) >= recon_min_m:
            V, info = _recon_cuda_panel(H, tau, k, bb, tol, rel_tol)
            any_bad = any_bad + info.sum()
        else:
            _PANEL.qr_panel(
                H, tau, _dummy, _dummy, k, bb, 0, nt
            )  # direct HH, writes H+tau
            V = _extract_V(H, k, bb)
        if k + bb < n:
            C = H[:, k:, k + bb :]
            Tm = _form_T(V, tau[:, k : k + bb])
            W1 = torch.bmm(V.transpose(1, 2), C)
            W2 = torch.bmm(Tm.transpose(1, 2), W1)
            C.sub_(torch.bmm(V, W2))
    if any_bad.item() > 0:  # ill-conditioned -> geqrf
        return torch.geqrf(A)
    return H, tau


def custom_kernel(data: input_t) -> output_t:
    A = data
    B, n, _ = A.shape
    # Precision: single-pass TF32 tensor cores ok for n>=1024 (factor rtol scales
    # with n, loose enough); n<=512 stays fp32 (single TF32 fails band/rowscale, and
    # 3xTF32 recovery is correct but slower since the GEMM isn't the bottleneck).
    _set_tf32(n >= 1024)
    # n>=4096: the panel wall. chol_bs (CholeskyQR1 + Householder reconstruction) trades the
    # sequential bb-column Householder critical path for parallel GEMMs + tiny 32x32 chol/LU.
    if (_TSQR is not None) and (_PANEL is not None) and (n >= 4096):
        return _qr_tsqr_cholbs(
            A, b=64, recon_min_m=128
        )  # wide panel cuts per-panel launches;
        #                                                    bb=128 is within noise but needs 128KB smem
    H = A.clone()  # one copy: checker reads original A for the residual
    tau = torch.zeros(B, n, device=A.device, dtype=A.dtype)
    # Fully-fused small-n path: whole matrix fits B200 smem (n*n*4 <= ~227KB, n<=238).
    # One launch does the entire unblocked Householder factor in smem (panel == full
    # matrix, bb=n), so the in-kernel rank-1 trailing apply replaces ALL the serialized
    # per-panel relaunches + extract_V/form_T/bmm glue that dominate small/mid-n runtime.
    if (_PANEL is not None) and (n * n * 4 <= 227000):
        # Fused whole-matrix factor (panel == full matrix): width = n.
        _dummy = torch.empty(1, 1, 1, device=A.device, dtype=A.dtype)
        _PANEL.qr_panel(H, tau, _dummy, _dummy, 0, n, 0, _panel_threads(B, n))
        return H, tau
    b = 64 if n <= 1024 else 32
    bm = 256 if n <= 1024 else 512
    nw = 4 if n <= 1024 else 8
    # smem-resident CUDA panel when the largest panel (m=n) fits B200 smem (~195KB).
    # n=512: b=64 -> 128KB. n=1024: need b=32 -> 128KB (b=64 would be 256KB, too big).
    # Multi-block cooperative panel for small-batch large-n: one block/matrix starves
    # the GPU (n=2048 B=8 -> 8 of 148 SMs; n=4096 B=2 -> 2). G blocks/matrix cooperate
    # via grid barriers so the panel fills the GPU. BM trades occupancy vs barrier cost
    # (barriers scale with G participants): n>=4096 -> mbm=512 (G=8); n=2048 -> mbm=256
    # (G=8 -> 64 blocks). [Re-testing n=2048 mb: the prior "B=8 fills enough" was wrong
    # for 148 SMs and predates the current smem-resident CUDA cooperative kernel.]
    use_mb = (B <= 8) and (n >= 4096)
    if use_mb:
        mbm = 256  # n=4096: G=16 (32 blocks) is the sweet spot at nt=1024 (~53ms): smaller
        # mbm=128 (G=32) hits the 32-way grid-barrier wall (60ms), larger 512
        # (G=8) underfills the SMs (56ms). n=2048 stays single-block — its grid
        # barriers cost more than filling 140 idle SMs (mb 25.8 vs sb 17.6ms).
    use_cuda_panel = (_PANEL is not None) and (n <= 2048) and not use_mb
    if use_cuda_panel and n == 1024:
        b = 32  # 1024x32x4 = 128KB fits smem (b=64 would be 256KB)
    if use_cuda_panel and n == 2048:
        b = 24  # widest that fits ~227KB; low batch -> wide panel wins (b sweep: 24 best)
    # CUDA multi-block panel: smem-resident row-slice + grid barrier. b even (barrier
    # parity) and bb*BM*4 must fit smem (32*256*4=32KB @2048; 32*512*4=64KB @4096).
    use_cuda_mb = (_PANEL is not None) and use_mb
    if use_cuda_mb:
        b = 32
    # The panel kernel emits the WY V/T directly (no _extract_V/_form_T launches) only
    # where it beats batched cuBLAS form_T: narrow-panel many-panel cases (n=1024/2048).
    # At n<=512 the batch is large and the panel wide (b=64) -> cuBLAS V^TV wins; keep it.
    # Single-block panel emits V/T in-kernel where it beats batched cuBLAS form_T
    # (narrow-panel many-panel n=1024/2048); n<=512 keeps cuBLAS. The multi-block panel
    # (n=4096) emits dense V per row-block + a standalone form_t kernel for T -> kills the
    # per-panel Python _extract_V/_form_T launches (was 34% of the 4096 time).
    use_emit_sb = use_cuda_panel and (
        n >= 1024
    )  # n=512 (B=640): cuBLAS form_T (tensor)
    # beats in-kernel G=VtV (scalar) by far
    use_emit_mb = use_cuda_mb
    use_emit = use_emit_sb or use_emit_mb
    panel_nt = _panel_threads(
        B, b
    )  # single-block panel threads (width = block width b)
    if use_emit:
        Vd_buf = torch.empty(B, n, b, device=A.device, dtype=A.dtype)
        T_buf = torch.empty(B, b, b, device=A.device, dtype=A.dtype)
    else:
        Vd_buf = torch.empty(1, 1, 1, device=A.device, dtype=A.dtype)
        T_buf = Vd_buf
    for k in range(0, n, b):
        bb = min(b, n - k)
        if use_cuda_panel:
            _PANEL.qr_panel(
                H, tau, Vd_buf, T_buf, k, bb, 1 if use_emit_sb else 0, panel_nt
            )
        elif use_cuda_mb:
            _panel_factor_cuda_mb(
                H, k, bb, tau, BM=mbm, Vd=Vd_buf, do_wy=1 if use_emit_mb else 0
            )
        elif use_mb:
            _panel_factor_mb(H, k, bb, tau, BM=mbm, num_warps=4)
        else:
            _panel_factor(H, k, bb, tau, BM=bm, num_warps=nw)
        if k + bb < n:
            C = H[:, k:, k + bb :]
            if use_emit_sb:
                V = Vd_buf[:, : n - k, :bb]
                Tm = T_buf[:, :bb, :bb]
            elif use_emit_mb:
                # emitted V skips extract_V; T via cuBLAS form_T (big-K V^TV, cuBLAS wins)
                V = Vd_buf[:, : n - k, :bb]
                Tm = _form_T(V, tau[:, k : k + bb])
            else:
                V = _extract_V(H, k, bb)
                Tm = _form_T(V, tau[:, k : k + bb])
            W1 = torch.bmm(V.transpose(1, 2), C)
            W2 = torch.bmm(Tm.transpose(1, 2), W1)
            C.sub_(torch.bmm(V, W2))
    return H, tau
scrolls · 984 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