Skip to content
KernelIndex
Search⌘K

submission 804074

oldsquaw · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-804074?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
8.13ms
#254 of 515
2026-06-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a437c23aa6844c5f2b8cad9895534098bc887749e43db7b42afc31df64735c9d
license declaredunknown
license concludedunknown
authorsoldsquaw
imported2026-08-26

Techniques

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

num-warps = 8_panel_qr[(b,)](A, tau, n, j, sab, sar, sac, stb, stn, M=M, BW=bw, num_warps=8)

Kernel source

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

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

_TRITON_N = {176, 352, 512, 1024}
_BW_CAP = 32
_OB = 128


if _HAS_TRITON:

    @triton.jit
    def _panel_qr(A, TAU, N, J,
                  sab, sar, sac, stb, stn,
                  M: tl.constexpr, BW: tl.constexpr):
        pid = tl.program_id(0)
        rows = tl.arange(0, M)
        cols = tl.arange(0, BW)
        grow = J + rows
        rmask = grow < N
        cmask = (J + cols) < N
        pptr = A + pid * sab + grow[:, None] * sar + (J + cols)[None, :] * sac
        tile = rmask[:, None] & cmask[None, :]
        P = tl.load(pptr, mask=tile, other=0.0)
        for i in range(BW):
            coli = tl.sum(tl.where(cols[None, :] == i, P, 0.0), axis=1)
            x = tl.where((rows >= i) & rmask, coli, 0.0)
            alpha = tl.sum(tl.where(rows == i, x, 0.0))
            norm = tl.sqrt(tl.sum(x * x))
            live = norm > 0.0
            beta = tl.where(alpha >= 0, -norm, norm)
            den = alpha - beta
            inv = tl.where(live & (den != 0), 1.0 / den, 0.0)
            v = tl.where(live & (rows == i), 1.0,
                         tl.where(live & (rows > i) & rmask, x * inv, 0.0))
            taui = tl.where(live, (beta - alpha) / tl.where(live, beta, 1.0), 0.0)
            w = tl.where(cols >= i, tl.sum(v[:, None] * P, axis=0), 0.0)
            P = P - taui * v[:, None] * w[None, :]
            kept = tl.sum(tl.where(cols[None, :] == i, P, 0.0), axis=1)
            newcol = tl.where(rows < i, kept, tl.where(rows == i, beta, v))
            P = tl.where(cols[None, :] == i, newcol[:, None], P)
            tl.store(TAU + pid * stb + (J + i) * stn, taui)
        tl.store(pptr, P, mask=tile)


def _apply_wy(P, tp, C, di):
    w = P.shape[2]
    V = P.tril(-1)
    V[:, di[:w], di[:w]] = (tp != 0).to(V.dtype)
    G = V.transpose(1, 2) @ V
    Tinv = torch.triu(G, 1) + torch.diag_embed(1.0 / torch.where(tp == 0, torch.ones_like(tp), tp))
    Y = torch.linalg.solve_triangular(Tinv.transpose(1, 2), V.transpose(1, 2) @ C, upper=False)
    C.baddbmm_(V, Y, beta=1, alpha=-1)


def _blocked_qr(A):
    b, n, _ = A.shape
    M = triton.next_power_of_2(n)
    bw = min(_BW_CAP, max(8, 16384 // M))
    bw = 1 << (bw.bit_length() - 1)
    ob = max(bw, (_OB // bw) * bw)
    tau = A.new_zeros(b, ((n + bw - 1) // bw) * bw)
    di = torch.arange(ob, device=A.device)
    sab, sar, sac, stb, stn = A.stride(0), A.stride(1), A.stride(2), tau.stride(0), tau.stride(1)
    for jo in range(0, n, ob):
        obw = min(ob, n - jo)
        for j in range(jo, jo + obw, bw):
            cur = min(bw, jo + obw - j)
            _panel_qr[(b,)](A, tau, n, j, sab, sar, sac, stb, stn, M=M, BW=bw, num_warps=8)
            if j + cur < jo + obw:
                _apply_wy(A[:, j:, j:j + cur], tau[:, j:j + cur], A[:, j:, j + cur:jo + obw], di)
        if jo + obw >= n:
            break
        _apply_wy(A[:, jo:, jo:jo + obw], tau[:, jo:jo + obw], A[:, jo:, jo + obw:], di)
    return A, tau[:, :n]


def custom_kernel(data: input_t) -> output_t:
    if _HAS_TRITON and data.shape[1] in _TRITON_N:
        try:
            return _blocked_qr(data.clone())
        except Exception:
            pass
    return torch.geqrf(data)
scrolls · 91 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