Skip to content
KernelIndex
Search⌘K

submission 822508

ozaka_8787 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-822508?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
7.67ms
#247 of 515
2026-06-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8e89a881f3e7b8c25223125cdfa4b3146a890ddc71f806481dc6b6c5b829fe97
license declaredunknown
license concludedunknown
authorsozaka_8787
imported2026-08-26

Techniques

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

num-warps = 8BLOCK_M=block_m, BLOCK_N=block_n, num_warps=8,
tile-m = 512_BLOCK_M = 512 # max panel height handled by the Triton kernel (512*64*4 = 128KB smem)

Kernel source

submission.py326 lines
import os
import time
import torch
from task import input_t, output_t

# Stage 2: batched blocked Householder QR with a Triton panel kernel + CUDA graphs.
#
# Stage 1.5 (CUDA graphs) removed CPU launch overhead but left the GPU-side
# bottleneck: the panel is factored column-by-column, so n=512 fires ~512 tiny
# sequential kernels just for the panels. This stage collapses each 64-column
# panel into ONE Triton kernel (one program per matrix, panel resident in shared
# memory, the 64 Householder reflections done sequentially inside the kernel).
# The trailing-matrix update stays as batched torch.bmm. Everything is wrapped in
# a per-shape CUDA graph.
#
# Correctness is identical Householder math to the verified Stage 1: real
# reflectors (geqrf (H,tau) convention), zero-pivot guard (tau=0), no pivoting,
# no input-pattern probing/routing, input never mutated, deterministic. Safety:
# if Triton is unavailable, a panel doesn't fit shared memory, or the kernel
# errors, we fall back to the proven pure-PyTorch panel -- so we never regress
# below the working Stage 1.5.

torch.backends.cuda.matmul.allow_tf32 = False
try:
    torch.backends.cudnn.allow_tf32 = False
    torch.set_float32_matmul_precision("highest")
except Exception:
    pass

_NB = 64          # panel width (all benchmark n are multiples of 64)
_BLOCK_M = 512    # max panel height handled by the Triton kernel (512*64*4 = 128KB smem)

try:
    import triton
    import triton.language as tl
    _HAVE_TRITON = True
except Exception:
    _HAVE_TRITON = False

_USE_TRITON = _HAVE_TRITON and os.environ.get("QR_NO_TRITON") != "1"


if _HAVE_TRITON:

    @triton.jit
    def _panel_kernel(H_ptr, tau_ptr, n, p,
                      sb, sr, sc, stb, stc,
                      BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr):
        # One program per matrix. Factor the panel rows [p, n) x cols [p, p+BLOCK_N)
        # in-place: writes R on/above the diagonal, reflector tails below it, tau.
        pid = tl.program_id(0)
        M = n - p
        rows = tl.arange(0, BLOCK_M)          # tile row i -> global row p+i
        cols = tl.arange(0, BLOCK_N)          # tile col k -> global col p+k (panel diag at i==k)
        row_mask = rows < M

        base = pid * sb + p * sr + p * sc
        offs = base + rows[:, None] * sr + cols[None, :] * sc
        mask = row_mask[:, None]
        T = tl.load(H_ptr + offs, mask=mask, other=0.0).to(tl.float32)

        for j in range(0, BLOCK_N):
            cj = tl.sum(tl.where(cols[None, :] == j, T, 0.0), axis=1)   # column j (BLOCK_M,)
            alpha = tl.sum(tl.where(rows == j, cj, 0.0))                # scalar diag
            below = (rows > j) & row_mask
            xnorm_sq = tl.sum(tl.where(below, cj * cj, 0.0))
            zero = xnorm_sq == 0.0

            fullnorm = tl.sqrt(alpha * alpha + xnorm_sq)
            sgn = tl.where(alpha < 0.0, -1.0, 1.0)
            beta = -sgn * fullnorm
            beta_safe = tl.where(zero, 1.0, beta)
            tau_j = tl.where(zero, 0.0, (beta_safe - alpha) / beta_safe)
            denom = tl.where(zero, 1.0, alpha - beta_safe)

            vtail = tl.where(below, cj / denom, 0.0)                    # 0 when zero (cj below all 0)
            v = tl.where(rows == j, 1.0, vtail)                        # reflector, unit at diag
            dval = tl.where(zero, alpha, beta)                         # R diagonal
            newcj = tl.where(rows == j, dval, tl.where(below, vtail, cj))
            T = tl.where(cols[None, :] == j, newcj[:, None], T)

            tl.store(tau_ptr + pid * stb + (p + j) * stc, tau_j)

            # apply (I - tau v v^T) to panel columns k > j
            w = tl.sum(v[:, None] * T, axis=0)                         # (BLOCK_N,)
            T = tl.where(cols[None, :] > j, T - tau_j * v[:, None] * w[None, :], T)

        tl.store(H_ptr + offs, T, mask=mask)


def _panel_triton(H, tau, p, n, block_m, block_n):
    B = H.shape[0]
    _panel_kernel[(B,)](
        H, tau, n, p,
        H.stride(0), H.stride(1), H.stride(2),
        tau.stride(0), tau.stride(1),
        BLOCK_M=block_m, BLOCK_N=block_n, num_warps=8,
    )


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


_SMEM_FLOATS = 131072 // 4  # ~128KB panel-tile budget (BLOCK_M * BLOCK_N floats)


def _panel_plan(M, cap):
    # Pick (BLOCK_M, panel_width) so the M-row panel tile fits shared memory.
    # `cap` (shape-dependent) bounds the width: small matrices like a narrow panel
    # (tiny tile, no spills, panel-dominated) while n>=512 wants a wider panel
    # (fewer/larger trailing GEMMs, and -- with TF32 trailing -- fewer panels so
    # less accumulated TF32 error).
    bm = _next_pow2(M)
    nb = _SMEM_FLOATS // bm
    if nb < 8:
        return None, None          # even nb=8 doesn't fit -> caller uses PyTorch
    nb = min(cap, nb)
    p2 = 1
    while p2 * 2 <= nb:
        p2 *= 2
    return bm, p2


def _panel_pytorch(H, tau, p, hi, n, b, dtype, ones_b):
    # Proven pure-PyTorch column-by-column panel (fallback path).
    for j in range(p, hi):
        m = n - j
        col = H[:, j:, j].double()
        alpha = col[:, 0]
        x2 = col[:, 1:]
        xnorm_sq = (x2 * x2).sum(dim=1)
        zero = xnorm_sq == 0
        fullnorm = torch.sqrt(alpha * alpha + xnorm_sq)
        sgn = torch.where(alpha < 0, -torch.ones_like(alpha), torch.ones_like(alpha))
        beta = -sgn * fullnorm
        beta_safe = torch.where(zero, torch.ones_like(beta), beta)
        tau_j = torch.where(zero, torch.zeros_like(beta), (beta_safe - alpha) / beta_safe)
        denom = torch.where(zero, torch.ones_like(beta), alpha - beta_safe)
        v2 = torch.where(zero.unsqueeze(1), torch.zeros_like(x2), x2 / denom.unsqueeze(1))
        H[:, j, j] = torch.where(zero, alpha, beta).to(dtype)
        if m > 1:
            H[:, j + 1:, j] = v2.to(dtype)
        tau[:, j] = tau_j.to(dtype)
        if j + 1 < hi:
            sub = H[:, j:, j + 1:hi]
            v = torch.cat([ones_b, H[:, j + 1:, j]], dim=1)
            tj = tau[:, j]
            w = torch.einsum('bi,bic->bc', v, sub)
            H[:, j:, j + 1:hi] = sub - tj.view(b, 1, 1) * v.unsqueeze(2) * w.unsqueeze(1)


# Single TF32 tensor-core matmul for the trailing GEMMs at big n, where the gate
# tolerance (rtol = 20*n*eps32) is loose enough to absorb TF32's ~1e-3 error.
# (At n=512 the gate is tighter and plain TF32 / 3xTF32 both lost; FP32 there.)
_TC_MIN_N = int(os.environ.get("QR_TC_MIN_N", "512"))


def _mm(a, b, tf32):
    if not tf32:
        return torch.matmul(a, b)
    prev = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    out = torch.matmul(a, b)
    torch.backends.cuda.matmul.allow_tf32 = prev   # restore -> solve stays FP32
    return out


def _build_T(V, tau_panel):
    b, _, pb = V.shape
    T = torch.zeros(b, pb, pb, device=V.device, dtype=V.dtype)
    for i in range(pb):
        T[:, i, i] = tau_panel[:, i]
        if i > 0:
            z = torch.bmm(V[:, :, :i].transpose(1, 2), V[:, :, i:i + 1])
            z = -tau_panel[:, i].view(b, 1, 1) * z
            T[:, :i, i:i + 1] = torch.bmm(T[:, :i, :i], z)
    return T


def _blocked_qr(A, prof=None):
    b, n, _ = A.shape
    device, dtype = A.device, A.dtype
    H = A.clone()
    tau = torch.zeros(b, n, device=device, dtype=dtype)
    ones_b = torch.ones(b, 1, device=device, dtype=dtype)

    triton_ok = _USE_TRITON and A.is_cuda and dtype == torch.float32
    nb_cap = 16 if n <= 352 else 64   # small n -> narrow panel; big n -> wide panel

    p = 0
    while p < n:
        M = n - p
        block_m, nb = _panel_plan(M, nb_cap) if triton_ok else (None, None)
        use_triton = block_m is not None and nb <= M
        pb = nb if use_triton else min(_NB, M)
        hi = p + pb

        if prof is not None:
            torch.cuda.synchronize(); _t = time.perf_counter()

        # panel: Triton with a shared-memory-fitting width, else PyTorch fallback
        if use_triton:
            _panel_triton(H, tau, p, n, block_m, pb)
        else:
            _panel_pytorch(H, tau, p, hi, n, b, dtype, ones_b)

        if prof is not None:
            torch.cuda.synchronize(); prof["panel"] += time.perf_counter() - _t; _t = time.perf_counter()

        # trailing update on columns [hi, n): C <- (I - V T^T V^T) C
        if hi < n:
            if prof is not None:
                torch.cuda.synchronize(); _t2 = time.perf_counter()
            V = H[:, p:, p:hi].clone()
            rr = torch.arange(M, device=device).view(1, -1, 1)
            cc = torch.arange(pb, device=device).view(1, 1, -1)
            V = torch.where(rr == cc, torch.ones_like(V), V)
            V = torch.where(rr < cc, torch.zeros_like(V), V)
            taup = tau[:, p:hi]                                  # (b, pb)
            # deflate identity reflectors (tau==0, rank-deficient cols): zero the
            # WHOLE column of V so it contributes nothing to the block reflector.
            V = torch.where((taup == 0).view(b, 1, pb), torch.zeros_like(V), V)
            if prof is not None:
                torch.cuda.synchronize(); prof["vbuild"] += time.perf_counter() - _t2; _t2 = time.perf_counter()
            # UT transform (Joffrain et al.): never form T. Minv = T^{-1} =
            # striu(V^T V) + diag(1/tau); apply Q^T C = C - V (M^{-T} (V^T C)).
            G = torch.matmul(V.transpose(1, 2), V)               # one batched GEMM (was ~128 tiny bmms)
            d = 1.0 / torch.where(taup == 0, torch.ones_like(taup), taup)
            Minv = torch.triu(G, 1) + torch.diag_embed(d)
            if prof is not None:
                torch.cuda.synchronize(); prof["buildT"] += time.perf_counter() - _t2; _t2 = time.perf_counter()
            C = H[:, p:, hi:]
            tc = n >= _TC_MIN_N                       # TF32 tensor cores for big-n
            W = _mm(V.transpose(1, 2), C, tc)
            X = torch.linalg.solve_triangular(Minv.transpose(1, 2), W, upper=False)
            H[:, p:, hi:] = C - _mm(V, X, tc)
            if prof is not None:
                torch.cuda.synchronize(); prof["bmm"] += time.perf_counter() - _t2

        if prof is not None:
            torch.cuda.synchronize(); prof["trail"] += time.perf_counter() - _t

        p = hi

    return H, tau


_triton_checked = False


def _ensure_triton():
    # One-time probe: run the Triton panel path on a small input and sanity-check
    # against torch.geqrf. If it crashes / NaNs / is grossly wrong, permanently
    # disable Triton and fall back to the proven pure-PyTorch panel.
    global _USE_TRITON, _triton_checked
    if _triton_checked:
        return
    _triton_checked = True
    if not _USE_TRITON:
        return
    try:
        A = torch.randn(2, 128, 128, device="cuda", dtype=torch.float32)
        H, tau = _blocked_qr(A)
        Q = torch.linalg.householder_product(H, tau)
        R = torch.triu(H)
        resid = (R - Q.transpose(-1, -2) @ A).abs().amax()
        scale = A.abs().amax()
        if not torch.isfinite(resid).item() or resid.item() > 1e-2 * scale.item():
            _USE_TRITON = False
    except Exception:
        _USE_TRITON = False


# Per-shape CUDA graph cache.
_graphs: dict = {}


def _try_capture(data):
    # Warm up first (JIT Triton, init cuBLAS workspaces) so nothing compiles or
    # allocates during capture; torch.cuda.graph handles capture internally.
    try:
        static_in = data.clone()
        for _ in range(3):
            _blocked_qr(static_in)
        torch.cuda.synchronize()
        g = torch.cuda.CUDAGraph()
        with torch.cuda.graph(g):
            out_H, out_tau = _blocked_qr(static_in)
        return (static_in, g, out_H, out_tau)
    except Exception:
        return None


def _use_geqrf(n, batch):
    # Size-based dispatch only (allowed; never data/position-based). torch.geqrf
    # (cuSOLVER/cuBLAS) is the reference -> always correct, and measured fastest
    # for (a) tiny matrices (batched cuBLAS path) and (b) huge matrices with tiny
    # batch (one large factorization fills all 148 SMs, vs our 1-CTA-per-matrix).
    return n <= 64 or (n >= 3072 and batch <= 4)


def custom_kernel(data: input_t) -> output_t:
    if not data.is_cuda:
        return _blocked_qr(data)

    if _use_geqrf(int(data.shape[1]), int(data.shape[0])):
        return torch.geqrf(data)

    _ensure_triton()
    key = (int(data.shape[0]), int(data.shape[1]))
    if key not in _graphs:
        _graphs[key] = _try_capture(data)

    entry = _graphs[key]
    if entry is None:                       # capture unsupported -> eager
        return _blocked_qr(data)

    static_in, g, out_H, out_tau = entry
    static_in.copy_(data)
    g.replay()
    return out_H.clone(), out_tau.clone()
scrolls · 326 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