Skip to content
KernelIndex
Search⌘K

submission 834880

moogician · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-834880?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.16ms
#182 of 515
2026-06-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:46c90429a0d55dd1c5936858e331a10e334fed045061426ca2b32ce3c9175649
license declaredunknown
license concludedunknown
authorsmoogician
imported2026-08-26

Techniques

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

mmaG = tl.dot(tl.trans(Vc), Vc,

Kernel source

submission.py438 lines
import torch
from task import input_t, output_t

# Trailing-update GEMMs (V^T C, T^T Wt, V Wt) dominate the n=512/1024 FLOPs.
# TF32 tensor cores run them far faster; whether it stays inside the per-matrix
# factor-residual gate depends on n (the gate scales with n, so larger n has more
# slack). We toggle the global flag per call rather than globally.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False

try:
    import triton
    import triton.language as tl

    @triton.jit
    def _panel_kernel(H_ptr, tau_ptr, V_ptr, T_ptr, n, m, w, k,
                      sb, srow, scol, stb, vb, vrow, vcol, tb, trow, tcol,
                      BLOCK_M: tl.constexpr, BLOCK_W: tl.constexpr,
                      DOT_TF32: tl.constexpr, NO_VT: tl.constexpr,
                      LOG2W: tl.constexpr, TPREC: tl.constexpr = "ieee"):
        b = tl.program_id(0)
        rm = tl.arange(0, BLOCK_M)
        rw = tl.arange(0, BLOCK_W)
        rmask = rm < m
        wmask = rw < w
        base = H_ptr + b * sb + k * srow + k * scol
        ptrs = base + rm[:, None] * srow + rw[None, :] * scol
        mask2d = rmask[:, None] & wmask[None, :]
        P = tl.load(ptrs, mask=mask2d, other=0.0)
        tau_acc = tl.zeros([BLOCK_W], dtype=tl.float32)
        for c in range(0, BLOCK_W):
            if c < w:
                colc = tl.sum(tl.where((rw == c)[None, :], P, 0.0), axis=1)  # [BLOCK_M]
                ge = rm >= c
                # fuse the two M-dim reductions (alpha=colc[c], nf2=sum_{i>=c} colc[i]^2)
                # into one reduction over a [BLOCK_M, 2] stack to cut sequential
                # warp-reduction latency on the per-column critical path.
                two = tl.arange(0, 2)
                stk = tl.where(two[None, :] == 0,
                               tl.where(rm == c, colc, 0.0)[:, None],
                               tl.where(ge, colc * colc, 0.0)[:, None])      # [BLOCK_M, 2]
                red = tl.sum(stk, axis=0)                                    # [2]
                alpha = tl.sum(tl.where(two == 0, red, 0.0))
                nf2 = tl.sum(tl.where(two == 1, red, 0.0))
                nf = tl.sqrt(nf2)
                an = tl.abs(alpha)
                pos = nf > 0.0
                sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
                tauv = tl.where(pos, 1.0 + an / tl.where(pos, nf, 1.0), 0.0)
                inv = tl.where(pos, sgn / (nf + an), 0.0)
                beta = tl.where(pos, -sgn * nf, alpha)
                v = tl.where(rm == c, 1.0, tl.where(rm > c, colc * inv, 0.0))
                wvec = tl.sum(v[:, None] * P, axis=0)                        # [BLOCK_W]
                upd = (tauv * v)[:, None] * wvec[None, :]
                P = tl.where((rw > c)[None, :], P - upd, P)
                newcol = tl.where(rm < c, colc, tl.where(rm == c, beta, colc * inv))
                P = tl.where((rw == c)[None, :], newcol[:, None], P)
                tau_acc = tl.where(rw == c, tauv, tau_acc)
        tl.store(ptrs, P, mask=mask2d)
        tl.store(tau_ptr + b * stb + k + rw, tau_acc, mask=wmask)
        if NO_VT:
            return
        # clean unit-lower V (1 on diagonal, v below, 0 above) for the trailing update
        Vc = tl.where(rm[:, None] > rw[None, :], P,
                      tl.where(rm[:, None] == rw[None, :], 1.0, 0.0))
        vptrs = V_ptr + b * vb + rm[:, None] * vrow + rw[None, :] * vcol
        tl.store(vptrs, Vc, mask=mask2d)
        # block reflector T (upper-tri): Q = I - V T V^T, via LARFT recurrence on G = V^T V
        G = tl.dot(tl.trans(Vc), Vc,
                   input_precision=("tf32" if DOT_TF32 else TPREC))  # [BLOCK_W, BLOCK_W]
        # Block reflector  T = (I - N)^{-1} diag(tau),  N = -diag(tau) striu(G)
        # (strictly upper, nilpotent: N^BLOCK_W = 0). The 32-step sequential LARFT
        # recurrence is latency-bound on the warp-reduction hardware; instead build
        # (I-N)^{-1} = sum_{i<32} N^i = prod_{k} (I + N^{2^k}) with ~8 tl.dot's, which
        # run on the MMA units (off the shuffle-reduction critical path). fp32 dots
        # keep T as accurate as the exact recurrence. (BLOCK_W==32 -> 5 factors.)
        eye = rw[:, None] == rw[None, :]
        Imat = tl.where(eye, 1.0, 0.0)
        N = tl.where(rw[:, None] < rw[None, :], -tau_acc[:, None] * G, 0.0)
        # (I-N)^{-1} = sum_i N^i = prod_k (I + N^{2^k}); N is BLOCK_W-nilpotent, so
        # LOG2W-1 doubling stages (covering up to N^{BLOCK_W-1}) are exact.
        acc = Imat + N
        Np = N
        for _ in range(LOG2W - 1):
            Np = tl.dot(Np, Np, input_precision=TPREC)
            acc = tl.dot(acc, Imat + Np, input_precision=TPREC)
        T = acc * tau_acc[None, :]                                   # (I-N)^{-1} diag(tau)
        tptrs = T_ptr + b * tb + rw[:, None] * trow + rw[None, :] * tcol
        tl.store(tptrs, T, mask=wmask[:, None] & wmask[None, :])

    _HAS_TRITON = True
except Exception:
    _HAS_TRITON = False


def _next_pow2(x):
    p = 1
    while p < x:
        p <<= 1
    return p


@torch.jit.script
def _refl(colf: torch.Tensor, alpha: torch.Tensor):
    nf = torch.sqrt((colf * colf).sum(dim=1))            # full-column norm
    an = torch.abs(alpha)
    mask = nf > 0
    safen = torch.where(mask, nf, torch.ones_like(nf))
    tauj = torch.where(mask, 1.0 + an / safen, torch.zeros_like(nf))
    inv = torch.where(mask, torch.copysign(1.0 / (nf + an), alpha), torch.zeros_like(nf))
    newdiag = torch.where(mask, -torch.copysign(nf, alpha), alpha)
    return inv, tauj, newdiag


def _house_qr_blocked(A, nb=64):
    """Batched blocked Householder QR producing geqrf-compact (H, tau)."""
    B, n, _ = A.shape
    dev, dt = A.device, A.dtype
    H = A.clone()
    tau = torch.zeros(B, n, device=dev, dtype=dt)
    zeros = torch.zeros(B, device=dev, dtype=dt)
    for k in range(0, n, nb):
        kb = min(nb, n - k)
        jend = k + kb
        # ---- unblocked panel factorization (columns k..jend) ----
        for jj in range(kb):
            j = k + jj
            colf = H[:, j:, j]                                   # (B, m) view, full column
            alpha = H[:, j, j]                                   # (B,)
            inv, tauj, newdiag = _refl(colf, alpha)
            tau[:, j] = tauj
            if j + 1 < n:
                H[:, j + 1:, j].mul_(inv[:, None])               # scale v below diagonal
            # apply reflector within panel to cols j+1..jend-1
            if j + 1 < jend:
                colf[:, 0] = 1.0                                 # v[0]=1 (overwrites alpha)
                sub = H[:, j:, j + 1:jend]                        # (B, m, w) view of H
                w = torch.bmm(colf[:, None, :], sub).squeeze(1)   # (B, w) = v^T C
                w.mul_(tauj[:, None])
                torch.baddbmm(sub, colf[:, :, None], w[:, None, :],
                              beta=1.0, alpha=-1.0, out=sub)       # C -= v w
            H[:, j, j] = newdiag                                  # write R diagonal
        # ---- build block reflector V, T (compact WY) ----
        # T satisfies T^{-1} = diag(1/tau) + striu(V^T V), i.e.
        # (I + diag(tau) striu(G)) T = diag(tau)  -> one triangular solve.
        Vp = H[:, k:, k:jend]                                     # (B, m, kb)
        V = torch.tril(Vp, -1)
        di = torch.arange(kb, device=dev)
        V[:, di, di] = 1.0
        G = torch.bmm(V.transpose(1, 2), V)                       # (B, kb, kb) Gram
        taup = tau[:, k:jend]                                     # (B, kb)
        M = torch.triu(G, 1) * taup[:, :, None]                   # diag(tau) @ striu(G)
        M[:, di, di] = 1.0                                        # unit upper triangular
        RHS = torch.diag_embed(taup)
        T = torch.linalg.solve_triangular(M, RHS, upper=True, unitriangular=True)
        # ---- trailing update: C -= V (T^T (V^T C)) ----
        if jend < n:
            C = H[:, k:, jend:]
            Wt = torch.bmm(V.transpose(1, 2), C)
            Wt = torch.bmm(T.transpose(1, 2), Wt)
            torch.baddbmm(C, V, Wt, beta=1.0, alpha=-1.0, out=C)
    return H, tau


def _panel_warps(BLOCK_M, B, nwarps):
    # Only the BLOCK_M==1024 panels (n=1024 early panels) truly need 8 warps;
    # at 4 they are thread-starved (256 rows/warp).
    # The mid tiers (BLOCK_M==512/256) DIVERGE by shape, keyed on the nwarps hint
    # (n=1024 path passes nwarps=8, n=512 path passes 4):
    #  - n=1024 LATER panels (m<=512, nwarps>=8): need 4 -- at 2 they are
    #    thread-starved (CATASTROPHIC 6.8ms, the 512-tile reduction can't hide).
    #  - n=512 EARLY panels (the DOMINANT ones, narrow nb=16 tile, nwarps=4):
    #    reduce FASTER at 2 warps under heavy cross-call overlap (nslots=10) --
    #    more concurrent CTAs prefer fewer warps/CTA (fewer cross-warp combines).
    #    Re-swept @ nslots=10: 512=2,256=2 -> 4.84->4.80 (the old 4 was tuned at
    #    nslots=6, before every-iter-its-own-buffer max overlap shifted the optimum).
    if BLOCK_M >= 1024:
        return nwarps
    elif BLOCK_M >= 256:
        return 4 if nwarps >= 8 else 2
    elif BLOCK_M >= 64:
        return 2 if B >= 256 else 4
    return 2


def _house_qr_triton_agg(A, nb=32, ag=2, tf32=False, nwarps=8, dot_tf32=None,
                         inplace=False, wide_tf32=None, tprec="ieee"):
    """Blocked Householder QR, panel factored by the fused Triton kernel, but the
    WIDE trailing update is AGGREGATED across `ag` consecutive nb-panels into one
    rank-(ag*nb) update. The panel kernel stays at the fast nb=32 width; only the
    memory-bound wide trailing is widened, halving (for ag=2) the number of full
    passes over the trailing matrix C -- the dominant n=512 cost."""
    B, n, _ = A.shape
    dev, dt = A.device, A.dtype
    if dot_tf32 is None:
        dot_tf32 = tf32
    H = A if inplace else A.clone()
    tau = torch.empty(B, n, device=dev, dtype=dt)
    sb, srow, scol = H.stride()
    stb = tau.stride(0)
    BLOCK_W = nb
    LOG2W = max(1, _next_pow2(nb).bit_length() - 1)
    GW = ag * nb                                          # combined block width
    # combined-group V (unit lower-trapezoidal, GW wide) and combined T (GW x GW)
    Vg = torch.zeros(B, n, GW, device=dev, dtype=dt)
    Tg = torch.zeros(B, GW, GW, device=dev, dtype=dt)
    Wt1 = torch.empty(B, GW, n, device=dev, dtype=dt)
    Wt2 = torch.empty(B, GW, n, device=dev, dtype=dt)
    Gc = torch.empty(B, GW, nb, device=dev, dtype=dt)    # cross Gram / cross-T scratch
    Cc = torch.empty(B, GW, nb, device=dev, dtype=dt)
    torch.backends.cuda.matmul.allow_tf32 = tf32
    for kk in range(0, n, GW):
        gend = min(kk + GW, n)
        mg = n - kk                                      # rows in this group's frame
        gw = gend - kk                                   # actual group width (<=GW)
        wide = gend < n                                  # a wide trailing update follows
        subs = list(range(kk, gend, nb))
        # NB: the block-upper-triangular corners of Vg (rows above each sub-panel's
        # diagonal start) stay zero from the initial torch.zeros -- they are never
        # written by any kernel/bmm, and the initial zeros re-runs on each graph
        # replay -- so no per-group re-zeroing is needed (combined V stays unit
        # lower-trapezoidal; reads always slice to the valid mg rows).
        # ---- factor sub-panels; narrow-update later sub-panels; accumulate combined T ----
        for i, k in enumerate(subs):
            kb = min(nb, n - k)
            jend = k + kb
            m = n - k
            p = i * nb                                   # combined width before this sub-panel
            BLOCK_M = _next_pow2(m)
            pw = _panel_warps(BLOCK_M, B, nwarps)
            Vslice = Vg[:, k - kk:, p:p + kb]
            Tslice = Tg[:, p:p + kb, p:p + kb]
            vbi, vrowi, vcoli = Vslice.stride()
            tbi, trowi, tcoli = Tslice.stride()
            no_vt = (jend >= n)                          # only the final panel skips V/T
            _panel_kernel[(B,)](H, tau, Vslice, Tslice, n, m, kb, k,
                                sb, srow, scol, stb, vbi, vrowi, vcoli,
                                tbi, trowi, tcoli, BLOCK_M=BLOCK_M, BLOCK_W=BLOCK_W,
                                DOT_TF32=dot_tf32, NO_VT=no_vt, LOG2W=LOG2W,
                                TPREC=tprec, num_warps=pw)
            # accumulate the combined-WY cross block for the wide update:
            #   T[:p, p:p+kb] = -T[:p,:p] @ (V[:,:p]^T @ V[:,p:p+kb]) @ T_i
            if wide and p > 0:
                Vprev = Vg[:, :mg, :p]
                Vi = Vg[:, :mg, p:p + kb]
                Tprev = Tg[:, :p, :p]
                Ti = Tg[:, p:p + kb, p:p + kb]
                gram = torch.bmm(Vprev.transpose(1, 2), Vi, out=Gc[:, :p, :kb])
                cross = torch.bmm(Tprev, gram, out=Cc[:, :p, :kb])
                cross = torch.bmm(cross, Ti, out=Gc[:, :p, :kb]).neg_()
                Tg[:, :p, p:p + kb] = cross
            # narrow update: apply this block to the REMAINING sub-panels in the group
            if jend < gend:
                V = Vg[:, k - kk:k - kk + m, p:p + kb]
                T = Tslice
                ncols = gend - jend
                C = H[:, k:, jend:gend]
                w1 = Wt1[:, :kb, :ncols]
                w2 = Wt2[:, :kb, :ncols]
                torch.bmm(V.transpose(1, 2), C, out=w1)
                torch.bmm(T.transpose(1, 2), w1, out=w2)
                torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
        # ---- aggregated WIDE trailing update over cols [gend:n] (rank gw) ----
        if wide:
            V = Vg[:, :mg, :gw]
            T = Tg[:, :gw, :gw]
            cols = n - gend
            C = H[:, kk:, gend:]
            w1 = Wt1[:, :gw, :cols]
            w2 = Wt2[:, :gw, :cols]
            if wide_tf32 is None:
                torch.bmm(V.transpose(1, 2), C, out=w1)
                torch.bmm(T.transpose(1, 2), w1, out=w2)
                torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
            else:
                # PER-GEMM precision split (n=512): the two BIG wide-trailing GEMMs
                # (V^T C and C-=V w2, each m*gw*cols MACs) run in TF32 tensor cores
                # (~2x the CUDA-core fp32 path), while the tiny step2 (T^T w1) and
                # all narrow/cross-T work stay fp32. Only ONE of the two big GEMMs'
                # worth of tf32 error lands per pass and the aggregated path makes
                # only gw/nb-fewer passes, so mixed-n512 resid 1.12e-3 < gate
                # 1.22e-3 holds (full-tf32 was 1.65e-3, FAIL). Other n512 cases have
                # ~1000x slack. fp32-only schedules left the trailing CUDA-core-bound.
                torch.backends.cuda.matmul.allow_tf32 = wide_tf32
                torch.bmm(V.transpose(1, 2), C, out=w1)
                torch.backends.cuda.matmul.allow_tf32 = tf32
                torch.bmm(T.transpose(1, 2), w1, out=w2)
                torch.backends.cuda.matmul.allow_tf32 = wide_tf32
                torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
                torch.backends.cuda.matmul.allow_tf32 = tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    return H, tau


def _house_qr_triton(A, nb=32, tf32=False, nwarps=8, dot_tf32=None, inplace=False):
    """Blocked Householder QR with the panel factored by a fused Triton kernel."""
    B, n, _ = A.shape
    dev, dt = A.device, A.dtype
    if dot_tf32 is None:
        dot_tf32 = tf32
    # inplace=True factors A directly (caller owns a scratch buffer); avoids the
    # clone so a captured CUDA graph needs only ONE input copy per replay.
    H = A if inplace else A.clone()
    # every tau entry is written by the panel kernel (kb per panel, covering 0..n),
    # so an uninitialized buffer is safe and skips the zero-init launch.
    tau = torch.empty(B, n, device=dev, dtype=dt)
    sb, srow, scol = H.stride()
    stb = tau.stride(0)
    BLOCK_W = nb
    LOG2W = max(1, _next_pow2(nb).bit_length() - 1)  # log2(nb) for the T recurrence
    Vbuf = torch.empty(B, n, BLOCK_W, device=dev, dtype=dt)
    Tbuf = torch.empty(B, BLOCK_W, BLOCK_W, device=dev, dtype=dt)
    vb, vrow, vcol = Vbuf.stride()
    tbs, trow, tcol = Tbuf.stride()
    # reusable trailing-update scratch (V^T C and T^T(V^T C)); allocating once and
    # slicing per panel keeps the CUDA-graph capture pool from bloating (one buffer
    # instead of one fresh allocation per panel -> graphable even for n=512/1024).
    Wt1 = torch.empty(B, BLOCK_W, n, device=dev, dtype=dt)
    Wt2 = torch.empty(B, BLOCK_W, n, device=dev, dtype=dt)
    torch.backends.cuda.matmul.allow_tf32 = tf32
    for k in range(0, n, nb):
        kb = min(nb, n - k)
        jend = k + kb
        m = n - k
        BLOCK_M = _next_pow2(m)
        # Per-panel warp count: big early panels (BLOCK_M>=256) need all 8 warps
        # (fewer starves them -- memory: global nwarps=4 catastrophic), but the
        # small later panels (BLOCK_M<256) are oversubscribed at 8 warps (256
        # threads for <256 rows) -> fewer warps cuts cross-warp reduction overhead.
        # Optimal warp count falls with BLOCK_M (small tiles oversubscribe at 8
        # warps) but also depends on OCCUPANCY: a large batch (B>=256, e.g. n=512
        # b=640) saturates the GPU so its small tail panels tolerate just 2 warps,
        # while a small batch (n=176/352 b=40) needs more warps to hide latency.
        if BLOCK_M >= 512:
            pw = nwarps
        elif BLOCK_M >= 256:
            pw = 4
        elif BLOCK_M >= 64:
            pw = 2 if B >= 256 else 4
        else:
            pw = 2
        _panel_kernel[(B,)](H, tau, Vbuf, Tbuf, n, m, kb, k, sb, srow, scol, stb,
                            vb, vrow, vcol, tbs, trow, tcol,
                            BLOCK_M=BLOCK_M, BLOCK_W=BLOCK_W, DOT_TF32=dot_tf32,
                            NO_VT=(jend >= n), LOG2W=LOG2W, num_warps=pw)
        V = Vbuf[:, :m, :kb]                                   # clean unit-lower, from kernel
        T = Tbuf[:, :kb, :kb]                                  # block reflector, from kernel
        if jend < n:
            C = H[:, k:, jend:]
            cols = n - jend
            w1 = Wt1[:, :kb, :cols]
            w2 = Wt2[:, :kb, :cols]
            # fp32 trailing for n<=512 (tight gate); TF32 tensor cores for n>=1024
            # where the looser gate absorbs the rounding.
            torch.bmm(V.transpose(1, 2), C, out=w1)
            torch.bmm(T.transpose(1, 2), w1, out=w2)
            torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
    torch.backends.cuda.matmul.allow_tf32 = False
    return H, tau


# Small shapes (n<=352) are launch-overhead-bound (~4 launches/panel, several
# panels). Their working set is tiny, so a captured CUDA graph collapses all the
# per-panel launches into one replay without the cache-thrash that killed graphs
# for n>=512 (large per-panel trailing temporaries -> capture-pool bloat).
_GRAPH_CACHE = {}


def _graphed_triton(data, nb, tf32, dot_tf32, nwarps=8, agg=1, wide_tf32=None,
                    tprec="ieee", nslots=1):
    # Direct factorization on the default execution queue. CUDA graphs use capture
    # (an internal side queue), which the leaderboard disallows, so the graph
    # capture/replay path is removed. We clone the input since the factorization
    # runs in place.
    buf = data.clone()
    if agg >= 2:
        return _house_qr_triton_agg(buf, nb, ag=agg, tf32=tf32,
                                    dot_tf32=dot_tf32, nwarps=nwarps,
                                    inplace=True, wide_tf32=wide_tf32,
                                    tprec=tprec)
    return _house_qr_triton(buf, nb, tf32=tf32, dot_tf32=dot_tf32,
                            nwarps=nwarps, inplace=True)


def custom_kernel(data: input_t) -> output_t:
    B, n, _ = data.shape
    if _HAS_TRITON and 32 <= n <= 1024 and B >= 8:
        try:
            # n>=1024's factor-residual gate (20*n*eps) is loose enough to absorb
            # TF32 trailing updates; n<=512 stays fp32 to keep ill-conditioned
            # (mixed) batches inside the tighter gate.
            # n=512: nb=16 (NOT 32). The within-panel unblocked wvec work is
            # O(n*nb*m) -- it scales with the panel width -- so a NARROWER panel
            # halves the panel-kernel's dominant cross-warp v^T C reduction. The
            # memory's old "nb=16 is slower" finding was with PER-PANEL trailing
            # (which then makes 32 memory-bound passes over C); the aggregation path
            # DECOUPLES trailing rank from panel width (ag=4 keeps the rank-64 wide
            # trailing = 8 passes), so nb=16 wins ONLY in the agg context. nb must
            # be >=16 (tl.dot MMA minimum for the in-kernel G-dot). n=1024 is
            # occupancy-bound (60 progs, BLOCK_M=1024), NOT within-panel-bound, so
            # narrowing nb there just adds non-overlapped passes -> stays nb=32.
            nb = 16 if n == 512 else 32
            # n=512 stays fp32: a TF32 precision SCHEDULE was tried (TF32 early/late
            # windows) and the mixed-batch error stays ~1.6e-3 > 1.22e-3 gate no matter
            # which panels are TF32 -- a single extreme matrix (nearcollinear/rowscaled)
            # in the mixed batch is ruined by any TF32 trailing. fp32-locked confirmed.
            # All other Triton shapes (32/176/352 dense-only, 1024 looser gate) tolerate TF32.
            # TF32 trailing updates overran the leaderboard's correctness tolerance
            # on ill-conditioned public-test matrices (the locally-tuned gates were
            # looser than the grader's), so the Triton path runs fully fp32.
            use_tf32 = False
            # n=512: aggregate the memory-bound wide trailing update into rank-64
            # (nb=16 * ag=4) -- keeps trailing passes low while the narrow panel cuts
            # within-panel work. n=1024: rank-128 (nb=32 * ag=4).
            agg = 4 if n in (512, 1024) else 1
            # n=512: keep panel + step2 fp32 (tight mixed gate) but run the two big
            # wide-trailing GEMMs in TF32 -- per-GEMM split passes (1.12e-3<1.22e-3)
            # where full-TF32 fails, recovering the tensor-core trailing speedup.
            wide_tf32 = None
            # nwarps=4 for the BLOCK_M==512 panels (faster cross-warp reduction than
            # 8): wins for n=512 (nb=16 narrow tile) AND n=352 (plain-path BLOCK_M=512
            # early panels). ONLY n=1024 needs 8 (its BLOCK_M==1024 panels are
            # thread-starved at 4 -> 53ms). n=176/32 have no BLOCK_M>=512 panel so the
            # value is irrelevant for them. (The memory's "nwarps=4 catastrophic" was
            # for BLOCK_M==1024 / the wider nb=32 tile.)
            nwarps = 8 if n == 1024 else 4
            return _graphed_triton(data, nb, tf32=use_tf32, dot_tf32=use_tf32,
                                   agg=agg, wide_tf32=wide_tf32)
        except Exception:
            pass
    # Batched blocked Householder wins only when batch is large and n moderate;
    # otherwise cuSOLVER per-matrix (geqrf) is faster.
    if 352 < n <= 1024 and B >= 32:
        return _house_qr_blocked(data, 192 if n > 512 else 96)
    # Large-n small-batch and everything else: plain batched cuSOLVER geqrf on the
    # default execution queue.
    return torch.geqrf(data)
scrolls · 438 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