Skip to content
KernelIndex
Search⌘K

submission 837855

gct · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-837855?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.24ms
#146 of 515
2026-06-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:749ab52f4856ca86df86337ea271129a4c9efeb1a137951cc5f7003e7514d086
license declaredunknown
license concludedunknown
authorsgct
imported2026-08-26

Techniques

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

mmaVtC = tl.dot(tl.trans(V), C, input_precision="tf32x3")
num-warps = 4num_warps=4,
persistent-kernelPanel factorization + trailing update ALL inside one persistent kernel per batch element.
shared-memoryextern __shared__ float sh[];

Kernel source

submission.py423 lines
"""Shape-routed hybrid QR — Fully-fused Triton: single kernel launch for entire QR.
Panel factorization + trailing update ALL inside one persistent kernel per batch element.
Eliminates Python loop overhead (31 launches → 1 launch for n=512).
Panel is computed column-by-column in registers (BM×BN tile).
Trailing is updated tile-by-tile within the same kernel.
For n>BM_MAX (panel tile limit), falls back to multi-launch or geqrf.
"""
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

NB = 32
import os
# Co-residence guardrail: B*SPLIT_M must fit resident CTAs. B200 ~148 SMs → SPLIT_M<=16 at B=8.
# 3060 has 28 SMs → use SPLIT_M=2 locally (B*2 co-resident at small B). Override via env.
_SPLIT_M_N2048 = int(os.environ.get("SPLIT_M_N2048", "16"))

# ---- CUDA whole-matrix Householder QR (one CTA per matrix) ----
_CUDA = r"""
#include <torch/extension.h>
#include <vector>
__global__ void hh_qr(float* __restrict__ A, float* __restrict__ tau, int n) {
    extern __shared__ float sh[];
    float* As = sh; float* v = sh + n * n;
    __shared__ float red[256]; __shared__ float s_tau, s_scale, s_diag;
    int b = blockIdx.x, t = threadIdx.x, T = blockDim.x;
    float* Ab = A + (size_t)b * n * n;
    for (int idx = t; idx < n * n; idx += T) As[idx] = Ab[idx];
    __syncthreads();
    for (int k = 0; k < n; k++) {
        float loc = 0.f;
        for (int i = k + 1 + t; i < n; i += T) { float x = As[i * n + k]; loc += x * x; }
        red[t] = loc; __syncthreads();
        for (int s = T / 2; s > 0; s >>= 1) { if (t < s) red[t] += red[t + s]; __syncthreads(); }
        if (t == 0) {
            float a = As[k * n + k], xn = red[0];
            float sg = (a >= 0.f) ? 1.f : -1.f, be = -sg * sqrtf(a * a + xn);
            bool ac = xn > 0.f;
            s_tau = ac ? (be - a) / be : 0.f; s_scale = ac ? 1.f / (a - be) : 0.f; s_diag = ac ? be : a;
        }
        __syncthreads();
        float tk = s_tau, sc = s_scale;
        for (int i = t; i < n; i += T) v[i] = (i < k) ? 0.f : (i == k ? 1.f : As[i * n + k] * sc);
        __syncthreads();
        if (t == 0) { As[k * n + k] = s_diag; tau[(size_t)b * n + k] = tk; }
        for (int i = k + 1 + t; i < n; i += T) As[i * n + k] = v[i];
        __syncthreads();
        for (int j = k + 1 + t; j < n; j += T) {
            float w = 0.f;
            for (int i = k; i < n; i++) w += v[i] * As[i * n + j];
            w *= tk;
            for (int i = k; i < n; i++) As[i * n + j] -= v[i] * w;
        }
        __syncthreads();
    }
    for (int idx = t; idx < n * n; idx += T) Ab[idx] = As[idx];
}
std::vector<torch::Tensor> cuda_qr(torch::Tensor A) {
    int B = A.size(0), n = A.size(1);
    auto H = A.clone().contiguous();
    auto tau = torch::zeros({B, n}, A.options());
    size_t shmem = ((size_t)n * n + n) * sizeof(float);
    cudaFuncSetAttribute(hh_qr, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shmem);
    hh_qr<<<B, 256, shmem>>>(H.data_ptr<float>(), tau.data_ptr<float>(), n);
    return {H, tau};
}
"""
_cu = load_inline(name="cuda_hh_qr_fused", cpp_sources="std::vector<torch::Tensor> cuda_qr(torch::Tensor A);",
                  cuda_sources=_CUDA, functions=["cuda_qr"], verbose=False)

_lim = torch.cuda.get_device_properties(0).shared_memory_per_block_optin // 4
_NMAX = int(_lim ** 0.5)
while (_NMAX * _NMAX + _NMAX) > _lim:
    _NMAX -= 1


@triton.jit
def _fused_qr_kernel(A, TAU, N: tl.constexpr, NB_val: tl.constexpr,
                     sab, sai, saj, stb, stj,
                     BM: tl.constexpr, BN: tl.constexpr):
    """Fully-fused QR: one CTA per batch element, loops over all panels internally.

    For each panel k0=0,NB,2*NB,...:
      1. Load panel columns [k0:N, k0:k0+NB] into registers (BM×BN tile)
      2. Factor NB columns (Householder reflections) → V, T, tau
      3. For each trailing tile [k0:N, hi:hi+BN]:
         - Load C tile
         - VtC = V^T @ C
         - TtVtC = T^T @ VtC
         - C -= V @ TtVtC
         - Store C tile
      4. Store factored panel back
    """
    b = tl.program_id(0)
    row = tl.arange(0, BM)
    col = tl.arange(0, BN)

    num_panels = (N + NB_val - 1) // NB_val

    for panel_idx in range(num_panels):
        k0 = panel_idx * NB_val
        nb = tl.minimum(NB_val, N - k0)
        M = N - k0

        grow = k0 + row
        rmask = grow < N
        cmask = col < nb

        # Load panel
        m2 = rmask[:, None] & cmask[None, :]
        ptr = A + b * sab + grow[:, None] * sai + (k0 + col)[None, :] * saj
        P = tl.load(ptr, mask=m2, other=0.0)

        tauv = tl.zeros((BN,), dtype=tl.float32)
        Tm = tl.zeros((BN, BN), dtype=tl.float32)

        # Panel factorization (same as _panel_qr)
        for j in tl.static_range(BN):
            cj = tl.sum(tl.where(col[None, :] == j, P, 0.0), axis=1)
            diag = row == j; below = row > j
            alpha = tl.sum(tl.where(diag, cj, 0.0))
            xnsq = tl.sum(tl.where(below & rmask, cj * cj, 0.0))
            s = tl.where(alpha >= 0.0, 1.0, -1.0)
            beta = -s * tl.sqrt(alpha * alpha + xnsq)
            active = xnsq > 0.0
            tau_j = tl.where(active, (beta - alpha) / beta, 0.0)
            scale = tl.where(active, 1.0 / (alpha - beta), 0.0)
            v = tl.where(diag, 1.0, tl.where(below, cj * scale, 0.0))
            w = tl.sum(v[:, None] * P, axis=0)
            tcp = tl.where(col < j, -tau_j * w, 0.0)
            mvec = tl.sum(Tm * tcp[None, :], axis=1)
            newT = tl.where(col == j, tau_j, tl.where(col < j, mvec, 0.0))
            Tm = tl.where(col[None, :] == j, newT[:, None], Tm)
            upd = (col[None, :] > j) & (row[:, None] >= j)
            P = P - tl.where(upd, tau_j * v[:, None] * w[None, :], 0.0)
            dval = tl.where(active, beta, alpha)
            newcol = tl.where(diag, dval, tl.where(below, cj * scale, 0.0))
            P = tl.where(col[None, :] == j, tl.where(row[:, None] < j, P, newcol[:, None]), P)
            tauv = tl.where(col == j, tau_j, tauv)

        # Store factored panel
        tl.store(ptr, P, mask=m2)
        tl.store(TAU + b * stb + (k0 + col) * stj, tauv, mask=cmask)

        # Build V from factored panel (V[i,j] = 1 if i==j, P[i,j] if i>j, 0 if i<j)
        V = tl.where(row[:, None] > col[None, :], P,
                     tl.where(row[:, None] == col[None, :], 1.0, 0.0))

        # T transpose for trailing update
        Tt = tl.trans(Tm)

        # Trailing update: for each BN-wide tile of columns after the panel
        hi = k0 + NB_val
        Nc = N - hi
        num_tiles = (Nc + BN - 1) // BN if Nc > 0 else 0

        for tile_idx in range(num_tiles):
            jcol = tile_idx * BN + tl.arange(0, BN)
            jmask = jcol < Nc

            # Load C tile
            c_ptr = A + b * sab + grow[:, None] * sai + (hi + jcol)[None, :] * saj
            C = tl.load(c_ptr, mask=rmask[:, None] & jmask[None, :], other=0.0)

            # VtC = V^T @ C  (BN × BN)
            VtC = tl.dot(tl.trans(V), C, input_precision="tf32x3")

            # TtVtC = T^T @ VtC
            TtVtC = tl.dot(Tt, VtC, input_precision="tf32x3")

            # C -= V @ TtVtC
            update = tl.dot(V, TtVtC, input_precision="tf32x3")
            C = C - update

            tl.store(c_ptr, C, mask=rmask[:, None] & jmask[None, :])


def _fused_triton_qr(A):
    B, N, _ = A.shape
    H = A.clone()
    tau = torch.zeros((B, N), device=A.device, dtype=A.dtype)
    BM = triton.next_power_of_2(N)
    nw = 16 if BM >= 1024 else 8
    _fused_qr_kernel[(B,)](
        H, tau, N, NB,
        H.stride(0), H.stride(1), H.stride(2),
        tau.stride(0), tau.stride(1),
        BM=BM, BN=NB, num_warps=nw,
    )
    return H, tau


# Multi-launch Triton path for sizes where BM would be too large for fused
@triton.jit
def _panel_qr(A, TAU, TT, N, K0, NBw, sab, sai, saj, stb, stj, ttb, tti, ttj,
              BM: tl.constexpr, BN: tl.constexpr):
    b = tl.program_id(0)
    row = tl.arange(0, BM); col = tl.arange(0, BN)
    grow = K0 + row; rmask = grow < N; cmask = col < NBw
    m2 = rmask[:, None] & cmask[None, :]
    ptr = A + b * sab + grow[:, None] * sai + (K0 + col)[None, :] * saj
    P = tl.load(ptr, mask=m2, other=0.0)
    tauv = tl.zeros((BN,), dtype=tl.float32)
    Tm = tl.zeros((BN, BN), dtype=tl.float32)
    for j in tl.static_range(BN):
        cj = tl.sum(tl.where(col[None, :] == j, P, 0.0), axis=1)
        diag = row == j; below = row > j
        alpha = tl.sum(tl.where(diag, cj, 0.0))
        xnsq = tl.sum(tl.where(below & rmask, cj * cj, 0.0))
        s = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = -s * tl.sqrt(alpha * alpha + xnsq)
        active = xnsq > 0.0
        tau_j = tl.where(active, (beta - alpha) / beta, 0.0)
        scale = tl.where(active, 1.0 / (alpha - beta), 0.0)
        v = tl.where(diag, 1.0, tl.where(below, cj * scale, 0.0))
        w = tl.sum(v[:, None] * P, axis=0)
        tcp = tl.where(col < j, -tau_j * w, 0.0)
        mvec = tl.sum(Tm * tcp[None, :], axis=1)
        newT = tl.where(col == j, tau_j, tl.where(col < j, mvec, 0.0))
        Tm = tl.where(col[None, :] == j, newT[:, None], Tm)
        upd = (col[None, :] > j) & (row[:, None] >= j)
        P = P - tl.where(upd, tau_j * v[:, None] * w[None, :], 0.0)
        dval = tl.where(active, beta, alpha)
        newcol = tl.where(diag, dval, tl.where(below, cj * scale, 0.0))
        P = tl.where(col[None, :] == j, tl.where(row[:, None] < j, P, newcol[:, None]), P)
        tauv = tl.where(col == j, tau_j, tauv)
    tl.store(ptr, P, mask=m2)
    tl.store(TAU + b * stb + (K0 + col) * stj, tauv, mask=cmask)
    tptr = TT + b * ttb + col[:, None] * tti + col[None, :] * ttj
    tl.store(tptr, Tm, mask=cmask[:, None] & cmask[None, :])


@triton.jit
def _fused_trailing(
    H, TT, N, K0, NB_actual, Nc,
    sab, sai, saj, ttb, tti, ttj,
    VBUF, vb_batch, vb_row, vb_col,
    BM_TILE: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
):
    b = tl.program_id(0)
    tile_j = tl.program_id(1)
    hi = K0 + NB_actual
    M = N - K0
    kcol = tl.arange(0, BK)
    kmask = kcol < NB_actual
    trow = tl.arange(0, BK)
    tcol_t = tl.arange(0, BK)
    t_ptr = TT + b * ttb + trow[:, None] * tti + tcol_t[None, :] * ttj
    T = tl.load(t_ptr, mask=kmask[:, None] & kmask[None, :], other=0.0)
    Tt = tl.trans(T)
    jcol = tile_j * BN + tl.arange(0, BN)
    jmask = jcol < Nc
    VtC = tl.zeros((BK, BN), dtype=tl.float32)
    for m_start in range(0, M, BM_TILE):
        row = m_start + tl.arange(0, BM_TILE)
        grow = K0 + row
        rmask = (row < M)
        h_ptr = H + b * sab + grow[:, None] * sai + (K0 + kcol)[None, :] * saj
        Hpanel = tl.load(h_ptr, mask=rmask[:, None] & kmask[None, :], other=0.0, eviction_policy="evict_last")
        V_chunk = tl.where(row[:, None] > kcol[None, :], Hpanel,
                           tl.where(row[:, None] == kcol[None, :], 1.0, 0.0))
        c_ptr = H + b * sab + grow[:, None] * sai + (hi + jcol)[None, :] * saj
        C_chunk = tl.load(c_ptr, mask=rmask[:, None] & jmask[None, :], other=0.0)
        VtC += tl.dot(tl.trans(V_chunk), C_chunk, input_precision="tf32x3")
        v_ptr = VBUF + b * vb_batch + row[:, None] * vb_row + kcol[None, :] * vb_col
        tl.store(v_ptr, V_chunk, mask=rmask[:, None] & kmask[None, :])
    TtVtC = tl.dot(Tt, VtC, input_precision="tf32x3")
    for m_start in range(0, M, BM_TILE):
        row = m_start + tl.arange(0, BM_TILE)
        grow = K0 + row
        rmask = (row < M)
        v_ptr = VBUF + b * vb_batch + row[:, None] * vb_row + kcol[None, :] * vb_col
        V_chunk = tl.load(v_ptr, mask=rmask[:, None] & kmask[None, :], other=0.0)
        c_ptr = H + b * sab + grow[:, None] * sai + (hi + jcol)[None, :] * saj
        C_chunk = tl.load(c_ptr, mask=rmask[:, None] & jmask[None, :], other=0.0)
        update = tl.dot(V_chunk, TtVtC, input_precision="tf32x3")
        C_chunk = C_chunk - update
        tl.store(c_ptr, C_chunk, mask=rmask[:, None] & jmask[None, :])


@triton.jit
def _splitm_panel_qr(A, TAU, TT, VAL, CNT, N, K0, NBw, SPLIT_M: tl.constexpr,
                     sab, sai, saj, stb, stj, ttb, tti, ttj,
                     M, ROWS_PER, BM: tl.constexpr, BN: tl.constexpr, W: tl.constexpr):
    # Cross-CTA split-M Householder panel factorization (std Triton, atomic spin-barrier).
    # grid = (B, SPLIT_M). program_id(1) owns a row-slice of the panel; per reflector j
    # two cross-CTA reductions (phase0: [alpha,xnsq], phase1: w = v^T P) via per-(b,j,phase)
    # scratch slots used exactly once (release/acquire arrival counter, no reset).
    b = tl.program_id(0)
    pm = tl.program_id(1)
    base = pm * ROWS_PER                        # panel-local start row for this CTA
    hi_local = tl.minimum(base + ROWS_PER, M)   # exclusive upper bound for this slice
    prow = base + tl.arange(0, BM)              # panel-local row indices
    rmask = prow < hi_local                     # own-slice guard (no double counting)
    col = tl.arange(0, BN); cmask = col < NBw
    grow = K0 + prow
    m2 = rmask[:, None] & cmask[None, :]
    ptr = A + b * sab + grow[:, None] * sai + (K0 + col)[None, :] * saj
    P = tl.load(ptr, mask=m2, other=0.0)
    wcol = tl.arange(0, W)
    tauv = tl.zeros((BN,), dtype=tl.float32)
    Tm = tl.zeros((BN, BN), dtype=tl.float32)
    for j in range(BN):
        cj = tl.sum(tl.where(col[None, :] == j, P, 0.0), axis=1)
        diag = prow == j; below = prow > j
        alpha_p = tl.sum(tl.where(diag & rmask, cj, 0.0))
        xnsq_p = tl.sum(tl.where(below & rmask, cj * cj, 0.0))
        # ---- reduction phase 0 : [alpha, xnsq] ----
        slot = b * (BN * 2) + j * 2 + 0
        vptr = VAL + slot * W + wcol
        pvals = tl.where(wcol == 0, alpha_p, tl.where(wcol == 1, xnsq_p, 0.0))
        tl.atomic_add(vptr, pvals, sem="release")
        cptr = CNT + slot
        tl.atomic_add(cptr, 1, sem="acq_rel")
        done = tl.atomic_add(cptr, 0, sem="acquire")
        while done < SPLIT_M:
            done = tl.atomic_add(cptr, 0, sem="acquire")
        red = tl.load(vptr)
        alpha = tl.sum(tl.where(wcol == 0, red, 0.0))
        xnsq = tl.sum(tl.where(wcol == 1, red, 0.0))
        # ---- reflector (identical on all CTAs) ----
        s = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = -s * tl.sqrt(alpha * alpha + xnsq)
        active = xnsq > 0.0
        tau_j = tl.where(active, (beta - alpha) / beta, 0.0)
        scale = tl.where(active, 1.0 / (alpha - beta), 0.0)
        v = tl.where(diag, 1.0, tl.where(below, cj * scale, 0.0))
        v = tl.where(rmask, v, 0.0)
        # ---- reduction phase 1 : w = v^T P  (length NB) ----
        w_p = tl.sum(v[:, None] * P, axis=0)
        slot1 = b * (BN * 2) + j * 2 + 1
        vptr1 = VAL + slot1 * W + wcol
        wp_full = tl.where(wcol < BN,
                           tl.sum(tl.where(col[None, :] == wcol[:, None], w_p[None, :], 0.0), axis=1),
                           0.0)
        tl.atomic_add(vptr1, wp_full, sem="release")
        cptr1 = CNT + slot1
        tl.atomic_add(cptr1, 1, sem="acq_rel")
        done1 = tl.atomic_add(cptr1, 0, sem="acquire")
        while done1 < SPLIT_M:
            done1 = tl.atomic_add(cptr1, 0, sem="acquire")
        wred = tl.load(vptr1)
        w = tl.sum(tl.where(wcol[:, None] == col[None, :], wred[:, None], 0.0), axis=0)  # (NB,)
        # ---- T matrix (compact-WY), identical on all CTAs ----
        tcp = tl.where(col < j, -tau_j * w, 0.0)
        mvec = tl.sum(Tm * tcp[None, :], axis=1)
        newT = tl.where(col == j, tau_j, tl.where(col < j, mvec, 0.0))
        Tm = tl.where(col[None, :] == j, newT[:, None], Tm)
        tauv = tl.where(col == j, tau_j, tauv)
        # ---- local P update ----
        upd = (col[None, :] > j) & (prow[:, None] >= j)
        P = P - tl.where(upd & rmask[:, None], tau_j * v[:, None] * w[None, :], 0.0)
        dval = tl.where(active, beta, alpha)
        newcol = tl.where(diag, dval, tl.where(below, cj * scale, 0.0))
        P = tl.where(col[None, :] == j, tl.where(prow[:, None] < j, P, newcol[:, None]), P)
    tl.store(ptr, P, mask=m2)
    if pm == 0:
        tl.store(TAU + b * stb + (K0 + col) * stj, tauv, mask=cmask)
        tptr = TT + b * ttb + col[:, None] * tti + col[None, :] * ttj
        tl.store(tptr, Tm, mask=cmask[:, None] & cmask[None, :])


def _triton_qr(A, SPLIT_M=1):
    B, N, _ = A.shape; dev, dt = A.device, A.dtype
    H = A.clone(); tau = torch.zeros((B, N), device=dev, dtype=dt)
    Tt = torch.empty((B, NB, NB), device=dev, dtype=dt)
    Vbuf = torch.empty((B, N, NB), device=dev, dtype=dt)
    W = max(triton.next_power_of_2(NB), 2)
    if SPLIT_M > 1:
        VAL = torch.empty((B * NB * 2 * W,), device=dev, dtype=torch.float32)
        CNT = torch.empty((B * NB * 2,), device=dev, dtype=torch.int32)
    for k0 in range(0, N, NB):
        nb = min(NB, N - k0); M = N - k0; BM = triton.next_power_of_2(M)
        nw = 16 if BM >= 256 else 8
        if SPLIT_M > 1:
            rows_per = (M + SPLIT_M - 1) // SPLIT_M
            BMs = triton.next_power_of_2(rows_per)
            nws = 16 if BMs >= 256 else 8
            VAL.zero_(); CNT.zero_()
            _splitm_panel_qr[(B, SPLIT_M)](
                H, tau, Tt, VAL, CNT, N, k0, nb, SPLIT_M,
                H.stride(0), H.stride(1), H.stride(2),
                tau.stride(0), tau.stride(1), Tt.stride(0), Tt.stride(1), Tt.stride(2),
                M, rows_per, BM=BMs, BN=NB, W=W, num_warps=nws)
        else:
            _panel_qr[(B,)](H, tau, Tt, N, k0, nb, H.stride(0), H.stride(1), H.stride(2),
                            tau.stride(0), tau.stride(1), Tt.stride(0), Tt.stride(1), Tt.stride(2),
                            BM=BM, BN=NB, num_warps=nw)
        hi = k0 + nb
        if hi < N:
            Nc = N - hi
            BN_tile = min(triton.next_power_of_2(Nc), 128)
            grid = (B, triton.cdiv(Nc, BN_tile))
            _fused_trailing[grid](
                H, Tt, N, k0, nb, Nc,
                H.stride(0), H.stride(1), H.stride(2),
                Tt.stride(0), Tt.stride(1), Tt.stride(2),
                Vbuf, Vbuf.stride(0), Vbuf.stride(1), Vbuf.stride(2),
                BM_TILE=32, BN=BN_tile, BK=NB,
                num_warps=4,
            )
    return H, tau


def custom_kernel(data: input_t) -> output_t:
    B, n, _ = data.shape
    if 16 <= n <= _NMAX:
        H, tau = _cu.cuda_qr(data)
        return H, tau
    # Fused single-kernel path: single launch for entire QR
    # BM=next_pow2(n), shmem ~= BM*32*4*several. B200 has 228KB, fits n<=512.
    _shmem_limit = torch.cuda.get_device_properties(0).shared_memory_per_block_optin
    _fused_max = 256 if _shmem_limit >= 228 * 1024 else 128
    if n <= _fused_max:
        return _fused_triton_qr(data)
    if 128 <= n <= 2048:
        # split-M panel ONLY at n=2048 (B=8 starves SMs); all other n keep single-CTA panel.
        split_m = _SPLIT_M_N2048 if n == 2048 else 1
        return _triton_qr(data, SPLIT_M=split_m)
    return torch.geqrf(data)
scrolls · 423 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