Skip to content
KernelIndex
Search⌘K

submission 825389

umbrella___ · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_2level.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-825389?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
3.35ms
#100 of 515
2026-06-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:96696595111f46de12b8147f9f56951a92bffa85a65bc5d0165108c6f763c339
license declaredunknown
license concludedunknown
authorsumbrella___
imported2026-08-26

Techniques

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

num-warps = 8def _qr(A, block, num_warps=8):

Kernel source

submission_2level.py277 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t

# ===========================================================================
# Batched compact-Householder QR (Triton blocked compact-WY), FP32.
#
# A Triton panel kernel does the unblocked geqr2 + in-kernel compact-WY T per
# panel; the FLOP-heavy trailing update C -= V (T^T (V^T C)) is batched cuBLAS
# FP32 GEMM (bmm/baddbmm). n>2048 falls back to cuSolver (torch.geqrf) -- a
# custom batched panel there is far slower than cuSolver.
#
# Precision: the trailing GEMMs run on TF32 tensor cores when it is SAFE, else
# FP32. Plain TF32 fails the band/rowscale gate and is too tight for n<512
# (rankdef/nearrank fail at n=384), so TF32 is gated on n>=512 AND a structural
# detector that flags rowscale (row-norm spread) and band/diagonal (both
# off-corner blocks ~0). The panel (reflectors) is always FP32. Validated by a
# 10-bit-mantissa TF32 simulation: every TF32-failing case is caught -> FP32.
# ===========================================================================

torch.backends.cuda.matmul.allow_tf32 = False


def _tf32_unsafe(A):
    # True -> the trailing update must stay FP32 (TF32 would risk the gate).
    # Cheap (~20us even on n=512 b=640): sub-sample columns for the row-norm
    # spread, touch only the n/4 corners for band/diagonal. Avoids a full
    # A.abs() materialization (which cost ~560us on the big batches).
    n = A.shape[-1]
    step = max(1, n // 32)
    s = A[:, :, ::step]
    rn2 = (s * s).sum(dim=2)                        # (B, n) approx squared row norms
    mx = rn2.amax(dim=1)
    mn = rn2.amin(dim=1).clamp_min(1e-37)
    if (mx / mn).amax() > 1e6:                      # rowscale: norm spread (1e4)^2
        return True
    q = max(1, n // 4)
    scale = A[:, ::step, ::step].abs().amax().clamp_min(1e-30)
    tr = A[:, :q, n - q:].abs().amax()             # top-right corner
    bl = A[:, n - q:, :q].abs().amax()             # bottom-left corner
    if (tr < 1e-6 * scale) and (bl < 1e-6 * scale):  # band / diagonal
        return True
    return False


@triton.jit
def _panel_kernel(P, TAU, T, VOUT, M, IB,
                  spb, spr, spc, stb, sti, sTb, sTr, sTc, svb, svr, svc,
                  BM: tl.constexpr, BNB: tl.constexpr):
    b = tl.program_id(0)
    r = tl.arange(0, BM)
    c = tl.arange(0, BNB)
    rm = r < M
    cm = c < IB
    p = P + b * spb + r[:, None] * spr + c[None, :] * spc
    tile = tl.load(p, mask=rm[:, None] & cm[None, :], other=0.0)
    tau_vec = tl.zeros((BNB,), dtype=tl.float32)
    for j in range(BNB):
        colj = tl.sum(tl.where(c[None, :] == j, tile, 0.0), axis=1)
        alpha = tl.sum(tl.where(r == j, colj, 0.0))
        xn2 = tl.sum(tl.where(r > j, colj * colj, 0.0))
        reflect = xn2 > 0.0
        sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = tl.where(reflect, -sgn * tl.sqrt(alpha * alpha + xn2), alpha)
        tau_j = tl.where(reflect, (beta - alpha) / tl.where(reflect, beta, 1.0), 0.0)
        denom = tl.where(reflect, alpha - beta, 1.0)
        vb = colj / denom
        v = tl.where(r == j, 1.0, tl.where(r > j, vb, 0.0))
        vmask = tl.where(r >= j, v, 0.0)
        w = tl.sum(tl.where(c[None, :] > j, vmask[:, None] * tile, 0.0), axis=0)
        tile = tile - tau_j * vmask[:, None] * w[None, :]
        newcol = tl.where(r < j, colj, tl.where(r == j, beta, vb))
        tile = tl.where(c[None, :] == j, newcol[:, None], tile)
        tau_vec = tl.where(c == j, tau_j, tau_vec)
    V = tl.where(r[:, None] == c[None, :], 1.0, tl.where(r[:, None] > c[None, :], tile, 0.0))
    tl.store(VOUT + b * svb + r[:, None] * svr + c[None, :] * svc, V, mask=rm[:, None] & cm[None, :])
    Tt = tl.zeros((BNB, BNB), dtype=tl.float32)
    tau0 = tl.sum(tl.where(c == 0, tau_vec, 0.0))
    Tt = tl.where((c[:, None] == 0) & (c[None, :] == 0), tau0, Tt)
    for i in range(1, BNB):
        tau_i = tl.sum(tl.where(c == i, tau_vec, 0.0))
        Vi = tl.sum(tl.where(c[None, :] == i, V, 0.0), axis=1)
        dots = tl.sum(V * Vi[:, None], axis=0)
        z = tl.where(c < i, -tau_i * dots, 0.0)
        Tz = tl.sum(tl.where(c[None, :] < i, Tt * z[None, :], 0.0), axis=1)
        newTcol = tl.where(c < i, Tz, tl.where(c == i, tau_i, 0.0))
        Tt = tl.where(c[None, :] == i, newTcol[:, None], Tt)
    tl.store(T + b * sTb + c[:, None] * sTr + c[None, :] * sTc, Tt, mask=cm[:, None] & cm[None, :])
    tl.store(P + b * spb + r[:, None] * spr + c[None, :] * spc, tile, mask=rm[:, None] & cm[None, :])
    tl.store(TAU + b * stb + c * sti, tau_vec, mask=cm)


_WS = {}


def _ws(B, m, block, dev):
    # reused (un-zeroed) workspaces: the panel fully overwrites the used region.
    key = (B, m, block, str(dev))
    ws = _WS.get(key)
    if ws is None:
        BNB = triton.next_power_of_2(block)
        Vbuf = torch.empty((B, m, block), device=dev, dtype=torch.float32)
        Tbuf = torch.empty((B, BNB, BNB), device=dev, dtype=torch.float32)
        _WS[key] = ws = (Vbuf, Tbuf)
    return ws


def _qr(A, block, num_warps=8):
    B, m, n = A.shape
    bs = int(block)
    BNB = triton.next_power_of_2(bs)
    H = A.clone()
    tau = A.new_empty(B, n)                          # panel writes every column
    Vbuf, Tbuf = _ws(B, m, bs, A.device)
    for k in range(0, n, bs):
        ib = min(bs, n - k)
        BM = triton.next_power_of_2(m - k)
        Hv = H[:, k:, k:k + ib]
        Vb = Vbuf[:, :m - k, :ib]
        tv = tau[:, k:]                              # panel writes tau in place
        _panel_kernel[(B,)](Hv, tv, Tbuf, Vb, m - k, ib,
                            Hv.stride(0), Hv.stride(1), Hv.stride(2),
                            tv.stride(0), tv.stride(1),
                            Tbuf.stride(0), Tbuf.stride(1), Tbuf.stride(2),
                            Vb.stride(0), Vb.stride(1), Vb.stride(2),
                            BM=BM, BNB=BNB, num_warps=num_warps)
        hi = k + ib
        if hi < n:
            V = Vb
            T = Tbuf[:, :ib, :ib]
            C = H[:, k:, hi:]
            W = V.transpose(-1, -2) @ C
            W = T.transpose(-1, -2) @ W
            C.baddbmm_(V, W, beta=1, alpha=-1)
    return H, tau


@triton.jit
def _larft_kernel(S, TAU, T, sSb, sSr, sSc, stb, sti, sTb, sTr, sTc,
                  OBW: tl.constexpr):
    # build the compact-WY T (OBW x OBW upper-tri) from S = V^T V and tau.
    b = tl.program_id(0)
    rr = tl.arange(0, OBW)
    cc = tl.arange(0, OBW)
    Sm = tl.load(S + b * sSb + rr[:, None] * sSr + cc[None, :] * sSc)
    tauv = tl.load(TAU + b * stb + cc * sti)
    Tt = tl.zeros((OBW, OBW), dtype=tl.float32)
    tau0 = tl.sum(tl.where(cc == 0, tauv, 0.0))
    Tt = tl.where((rr[:, None] == 0) & (cc[None, :] == 0), tau0, Tt)
    for i in range(1, OBW):
        tau_i = tl.sum(tl.where(cc == i, tauv, 0.0))
        Sci = tl.sum(tl.where(cc[None, :] == i, Sm, 0.0), axis=1)
        z = tl.where(rr < i, -tau_i * Sci, 0.0)
        Tz = tl.sum(tl.where(cc[None, :] < i, Tt * z[None, :], 0.0), axis=1)
        newcol = tl.where(rr < i, Tz, tl.where(rr == i, tau_i, 0.0))
        Tt = tl.where(cc[None, :] == i, newcol[:, None], Tt)
    tl.store(T + b * sTb + rr[:, None] * sTr + cc[None, :] * sTc, Tt)


def _qr_blk(A, block):
    # Blocked QR for large n with tiny batch: cuSolver factors each tall panel
    # (FP32, well-utilized on tall-skinny), the O(n^3) trailing update runs on
    # TF32 tensor cores. Beats a single big FP32 cuSolver geqrf when TF32-safe.
    B, m, n = A.shape
    OBW = triton.next_power_of_2(block)
    H = A.clone()
    tau = A.new_empty(B, n)
    idx = torch.arange(block, device=A.device)
    for k in range(0, n, block):
        ib = min(block, n - k)
        panel = H[:, k:, k:k + ib].contiguous()
        Hp, tp = torch.geqrf(panel)
        H[:, k:, k:k + ib] = Hp
        tau[:, k:k + ib] = tp
        hi = k + ib
        if hi < n:
            V = Hp.tril(-1)
            V[:, idx[:ib], idx[:ib]] = 1.0
            S = V.transpose(-1, -2) @ V                  # (B, ib, ib)
            T = A.new_empty(B, OBW, OBW)
            _larft_kernel[(B,)](S, tp, T,
                                S.stride(0), S.stride(1), S.stride(2),
                                tp.stride(0), tp.stride(1),
                                T.stride(0), T.stride(1), T.stride(2), OBW=OBW)
            Tt = T[:, :ib, :ib]
            C = H[:, k:, hi:]
            W = V.transpose(-1, -2) @ C
            W = Tt.transpose(-1, -2) @ W
            C.baddbmm_(V, W, beta=1, alpha=-1)
    return H, tau


def _qr2(A, NB, OB, nw):
    # Two-level blocked QR. Small inner block NB keeps the (tall) panel tile in
    # registers (a 4096xNB tile spills at NB>=16 -> the whole register file), while
    # the wide outer block OB gives an efficient (K=OB) TF32 trailing GEMM. Runs
    # both matrices concurrently (grid=B) vs cuSolver's serial-over-batch geqrf.
    B, m, n = A.shape
    BNB = triton.next_power_of_2(NB)
    H = A.clone()
    tau = A.new_zeros(B, n)
    for K in range(0, n, OB):
        OBw = min(OB, n - K)
        for k in range(K, K + OBw, NB):
            NBw = min(NB, K + OBw - k)
            BM = triton.next_power_of_2(m - k)
            Hv = H[:, k:, k:k + NBw]
            Tt = A.new_zeros(B, BNB, BNB)
            ts = A.new_zeros(B, BNB)
            Vb = A.new_zeros(B, m - k, NBw)
            _panel_kernel[(B,)](Hv, ts, Tt, Vb, m - k, NBw,
                                Hv.stride(0), Hv.stride(1), Hv.stride(2),
                                ts.stride(0), ts.stride(1),
                                Tt.stride(0), Tt.stride(1), Tt.stride(2),
                                Vb.stride(0), Vb.stride(1), Vb.stride(2),
                                BM=BM, BNB=BNB, num_warps=nw)
            tau[:, k:k + NBw] = ts[:, :NBw]
            if k + NBw < K + OBw:                       # within-OB trailing (narrow)
                V = Vb
                T = Tt[:, :NBw, :NBw]
                C = H[:, k:, k + NBw:K + OBw]
                W = V.transpose(-1, -2) @ C
                W = T.transpose(-1, -2) @ W
                C.baddbmm_(V, W, beta=1, alpha=-1)
        if K + OBw < n:                                 # combined far trailing
            mK = m - K
            blk = H[:, K:, K:K + OBw]
            ri = torch.arange(mK, device=A.device)[:, None]
            ci = torch.arange(OBw, device=A.device)[None, :]
            V_OB = (ri == ci).to(A.dtype) + (ri > ci).to(A.dtype) * blk
            S = V_OB.transpose(-1, -2) @ V_OB
            T_OB = A.new_zeros(B, OBw, OBw)
            _larft_kernel[(B,)](S, tau, T_OB,
                                S.stride(0), S.stride(1), S.stride(2),
                                tau.stride(0), tau.stride(1),
                                T_OB.stride(0), T_OB.stride(1), T_OB.stride(2),
                                OBW=OBw, num_warps=nw)
            C = H[:, K:, K + OBw:]
            W = V_OB.transpose(-1, -2) @ C
            W = T_OB.transpose(-1, -2) @ W
            C.baddbmm_(V_OB, W, beta=1, alpha=-1)
    return H, tau


def custom_kernel(data: input_t) -> output_t:
    A = data
    n = A.shape[-1]
    if n > 2048:
        # n=4096: cuSolver geqrf is panel-serial over the batch. Two-level blocked
        # QR (inner NB=8 fits the 4096-tall tile in registers, outer OB=64 gives an
        # efficient TF32 trailing) runs both matrices concurrently. TF32 is safe at
        # n=4096 (very lenient gate) for benign inputs; else fall back to cuSolver.
        A = A.contiguous()
        if not _tf32_unsafe(A):
            torch.backends.cuda.matmul.allow_tf32 = True
            try:
                return _qr2(A, 8, 64, 8)
            finally:
                torch.backends.cuda.matmul.allow_tf32 = False
        return torch.geqrf(A)
    # B200-measured best block per size: n=1024 likes block=32 (9.43->7.29ms),
    # but n=2048 is much worse at 32 (2048x32 panel tile) -> keep block=16 there.
    if n == 2048:
        block, nw = 16, 8
    elif n >= 1024:
        block, nw = 32, 8
    else:
        block, nw = 32, 4          # nw=4 fastest for all n<1024 (less cross-warp)
    # TF32 tensor cores for the trailing GEMMs only when safe (n>=512 + benign).
    use_tf32 = n >= 512 and not _tf32_unsafe(A)
    torch.backends.cuda.matmul.allow_tf32 = use_tf32
    try:
        return _qr(A.contiguous(), block, nw)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = False
scrolls · 277 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