Skip to content
KernelIndex
Search⌘K

submission 837609

arsrivish26691 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

qrv9.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-837609?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
6.03ms
#209 of 515
2026-06-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d369e161965eb4ea44ae2f3792055e488921f27d8d80c8f40c48c5be4f82f08b
license declaredunknown
license concludedunknown
authorsarsrivish26691
imported2026-08-26

Techniques

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

mmaG += tl.dot(tl.trans(V), V, input_precision="ieee")
num-warps = 8_copy_kernel[(triton.cdiv(total, COPY_BLOCK),)](data, H, TOTAL=total, BLOCK=COPY_BLOCK, num_warps=8)
stages = 1NB=IB, RB=RBLOCK, CB=CBLOCK, PR=PREC, num_warps=4, num_stages=1)

Kernel source

qrv9.py182 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

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

IB = 16          # inner panel width: factored sequentially, K>=16, no register spill
OB = 64          # outer block width: composed block reflector -> fat trailing GEMM (K=64)
RBLOCK = 128
# Trailing-GEMM precision (wide update contracts over K=OB=64, so this bites here).
# Sweep "ieee" -> "tf32x3" -> "tf32"; read orthogonality in --mode test.
PREC = "tf32x3"
CBLOCK = 64
COPY_BLOCK = 1024


@triton.jit
def _copy_kernel(A, H, TOTAL: tl.constexpr, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < TOTAL
    tl.store(H + offs, tl.load(A + offs, mask=mask, other=0.0), mask=mask)


@triton.jit
def _panel_factor_kernel(H, TAU, TBUF,
                         N: tl.constexpr, P, ROWS: tl.constexpr, NB: tl.constexpr):
    b = tl.program_id(0)
    ro = tl.arange(0, ROWS)
    rows = P + ro
    q = tl.arange(0, NB)
    cols = P + q
    X = tl.load(H + b * N * N + rows[:, None] * N + cols[None, :],
                mask=(rows[:, None] < N) & (cols[None, :] < N), other=0.0)
    tauv = tl.zeros((NB,), dtype=tl.float32)
    for jj in tl.static_range(0, NB):
        k = P + jj
        active = k < N
        x = tl.sum(tl.where(q[None, :] == jj, X, 0.0), axis=1)
        x = tl.where((rows >= k) & active, x, 0.0)
        alpha = tl.sum(tl.where(rows == k, x, 0.0), axis=0)
        nrm = tl.sqrt(tl.sum(x * x, axis=0))
        sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = -sgn * nrm
        safe = nrm > 0.0
        beta = tl.where(safe, beta, alpha)
        denom = tl.where(safe, alpha - beta, 1.0)
        beta_safe = tl.where(tl.abs(beta) > 0.0, beta, 1.0)
        tau = tl.where(safe, (beta - alpha) / beta_safe, 0.0)
        tau = tl.where(active, tau, 0.0)
        v = tl.where(rows == k, 1.0, x / denom)
        v = tl.where((rows >= k) & active, v, 0.0)
        dots = tl.sum(v[:, None] * X, axis=0)
        X = tl.where((q[None, :] > jj) & (rows[:, None] >= k) & active,
                     X - tau * v[:, None] * dots[None, :], X)
        X = tl.where((q[None, :] == jj) & (rows[:, None] == k) & active, beta, X)
        X = tl.where((q[None, :] == jj) & (rows[:, None] > k) & active, v[:, None], X)
        tauv = tl.where(q == jj, tau, tauv)
        tl.store(TAU + b * N + k, tau, mask=active)
    tl.store(H + b * N * N + rows[:, None] * N + cols[None, :], X,
             mask=(rows[:, None] < N) & (cols[None, :] < N))
    V = tl.where(rows[:, None] == cols[None, :], 1.0, X)
    V = tl.where((rows[:, None] >= cols[None, :]) & (cols[None, :] < N), V, 0.0)
    ti = tl.arange(0, NB)
    tj = tl.arange(0, NB)
    T = tl.zeros((NB, NB), dtype=tl.float32)
    for jj in tl.static_range(0, NB):
        active = (P + jj) < N
        vj = tl.sum(tl.where(q[None, :] == jj, V, 0.0), axis=1)
        s = tl.where(q < jj, tl.sum(V * vj[:, None], axis=0), 0.0)
        tau_j = tl.sum(tl.where(q == jj, tauv, 0.0), axis=0)
        y = -tau_j * tl.sum(T * s[None, :], axis=1)
        T = tl.where((tj[None, :] == jj) & (ti[:, None] < jj) & active, y[:, None], T)
        T = tl.where((tj[None, :] == jj) & (ti[:, None] == jj) & active, tau_j, T)
    tl.store(TBUF + b * NB * NB + ti[:, None] * NB + tj[None, :], T)


@triton.jit
def _form_T_wide_kernel(H, TAU, TBUFO,
                        N: tl.constexpr, P, NROWT, OBW: tl.constexpr, RB: tl.constexpr, LOG2: tl.constexpr):
    b = tl.program_id(0)
    qi = tl.arange(0, OBW)
    rr = tl.arange(0, RB)
    cols = P + qi
    base = b * N * N
    G = tl.zeros((OBW, OBW), dtype=tl.float32)             # Gram V^T V, row-tiled
    for rt in range(NROWT):
        rows = P + rt * RB + rr
        rm = rows < N
        Xv = tl.load(H + base + rows[:, None] * N + cols[None, :],
                     mask=rm[:, None] & (cols[None, :] < N), other=0.0)
        V = tl.where(rows[:, None] == cols[None, :], 1.0, Xv)
        V = tl.where((rows[:, None] >= cols[None, :]) & (cols[None, :] < N), V, 0.0)
        G += tl.dot(tl.trans(V), V, input_precision="ieee")
    tauv = tl.load(TAU + b * N + cols, mask=cols < N, other=0.0)
    ti = tl.arange(0, OBW)
    tj = tl.arange(0, OBW)
    U = tl.where(tj[None, :] > ti[:, None], G, 0.0)        # strictly upper of Gram
    M = -tauv[:, None] * U                                 # -diag(tau) . striu(V^T V), nilpotent
    R = tl.where(ti[:, None] == tj[None, :], 1.0, 0.0)     # identity
    cur = M
    for _ in tl.static_range(0, LOG2):                     # (I-M)^-1 = prod (I + M^{2^l})
        R = R + tl.dot(R, cur, input_precision="ieee")     # R @ (I + cur)
        cur = tl.dot(cur, cur, input_precision="ieee")     # square M
    T = R * tauv[None, :]                                  # R @ diag(tau)
    tl.store(TBUFO + b * OBW * OBW + ti[:, None] * OBW + tj[None, :], T)


@triton.jit
def _trailing_kernel(H, TBUF,
                     N: tl.constexpr, P, NROWT, COL_END,
                     NB: tl.constexpr, RB: tl.constexpr, CB: tl.constexpr, PR: tl.constexpr):
    b = tl.program_id(0)
    cbid = tl.program_id(1)
    qi = tl.arange(0, NB)
    rr = tl.arange(0, RB)
    cj = tl.arange(0, CB)
    vcols = P + qi
    ccols = P + NB + cbid * CB + cj
    base = b * N * N
    W = tl.zeros((NB, CB), dtype=tl.float32)
    for rt in range(NROWT):
        rows = P + rt * RB + rr
        rm = rows < N
        Xv = tl.load(H + base + rows[:, None] * N + vcols[None, :],
                     mask=rm[:, None] & (vcols[None, :] < N), other=0.0)
        V = tl.where(rows[:, None] == vcols[None, :], 1.0, Xv)
        V = tl.where((rows[:, None] >= vcols[None, :]) & (vcols[None, :] < N), V, 0.0)
        C = tl.load(H + base + rows[:, None] * N + ccols[None, :],
                    mask=rm[:, None] & (ccols[None, :] < COL_END), other=0.0)
        W += tl.dot(tl.trans(V), C, input_precision=PR)
    ti = tl.arange(0, NB)
    tj = tl.arange(0, NB)
    T = tl.load(TBUF + b * NB * NB + ti[:, None] * NB + tj[None, :])
    W2 = tl.dot(tl.trans(T), W, input_precision=PR)
    for rt in range(NROWT):
        rows = P + rt * RB + rr
        rm = rows < N
        cmask = rm[:, None] & (ccols[None, :] < COL_END)
        Xv = tl.load(H + base + rows[:, None] * N + vcols[None, :],
                     mask=rm[:, None] & (vcols[None, :] < N), other=0.0)
        V = tl.where(rows[:, None] == vcols[None, :], 1.0, Xv)
        V = tl.where((rows[:, None] >= vcols[None, :]) & (vcols[None, :] < N), V, 0.0)
        C = tl.load(H + base + rows[:, None] * N + ccols[None, :], mask=cmask, other=0.0)
        C = C - tl.dot(V, W2, input_precision=PR)
        tl.store(H + base + rows[:, None] * N + ccols[None, :], C, mask=cmask)


def custom_kernel(data: input_t) -> output_t:
    B = data.shape[0]
    N = data.shape[1]
    H = torch.empty_like(data)
    tau = torch.empty((B, N), device=data.device, dtype=data.dtype)
    total = B * N * N
    _copy_kernel[(triton.cdiv(total, COPY_BLOCK),)](data, H, TOTAL=total, BLOCK=COPY_BLOCK, num_warps=8)

    rows_pow2 = triton.next_power_of_2(N)
    pf_warps = 32 if N >= 2048 else (16 if N >= 1024 else 8)   # n=512 didn't spill; keep it at 8
    Tin = torch.empty((B, IB, IB), device=data.device, dtype=data.dtype)
    Tout = torch.empty((B, OB, OB), device=data.device, dtype=data.dtype)

    for p in range(0, N, OB):
        col_end = min(p + OB, N)
        for ip in range(0, OB, IB):
            c0 = p + ip
            if c0 >= N:
                break
            _panel_factor_kernel[(B,)](H, tau, Tin, N=N, P=c0, ROWS=rows_pow2, NB=IB, num_warps=pf_warps)
            if c0 + IB < col_end:
                nrowt = triton.cdiv(N - c0, RBLOCK)
                inct = triton.cdiv(col_end - (c0 + IB), CBLOCK)
                _trailing_kernel[(B, inct)](H, Tin, N=N, P=c0, NROWT=nrowt, COL_END=col_end,
                                            NB=IB, RB=RBLOCK, CB=CBLOCK, PR=PREC, num_warps=4, num_stages=1)
        if p + OB < N:
            nrowt = triton.cdiv(N - p, RBLOCK)
            _form_T_wide_kernel[(B,)](H, tau, Tout, N=N, P=p, NROWT=nrowt, OBW=OB, RB=RBLOCK, LOG2=OB.bit_length() - 1, num_warps=pf_warps)
            nct = triton.cdiv(N - (p + OB), CBLOCK)
            _trailing_kernel[(B, nct)](H, Tout, N=N, P=p, NROWT=nrowt, COL_END=N,
                                       NB=OB, RB=RBLOCK, CB=CBLOCK, PR=PREC, num_warps=8, num_stages=1)
    return H, tau
scrolls · 182 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