Skip to content
KernelIndex
Search⌘K

submission 836219

whao89 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_geo80ms.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-836219?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.98ms
#251 of 515
2026-06-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:86bdfcdf7f9b7e30eb5a1759a393a9f1364afc24fa0bafe8beda2b2d54f746af
license declaredunknown
license concludedunknown
authorswhao89
imported2026-08-26

Kernel source

submission_geo80ms.py140 lines
from __future__ import annotations

import torch
import triton
import triton.language as tl


@triton.jit
def _panel_kernel(H_ptr, tau_ptr, k, m, n, B: tl.constexpr, BM: tl.constexpr, BB: tl.constexpr):
    pid = tl.program_id(0)
    rows = tl.arange(0, BM)
    cols = tl.arange(0, BB)
    rmask = rows < m
    cmask = cols < B

    base = pid * n * n + k * n + k
    offs = rows[:, None] * n + cols[None, :]
    pmask = rmask[:, None] & cmask[None, :]
    panel = tl.load(H_ptr + base + offs, mask=pmask, other=0.0)

    for jj in tl.static_range(B):
        sel = cols == jj
        colj = tl.sum(tl.where(sel[None, :], panel, 0.0), axis=1)

        on_diag = rows == jj
        below = rows > jj
        x = tl.where(below & rmask, colj, 0.0)
        alpha = tl.sum(tl.where(on_diag, colj, 0.0))
        xfull = tl.where((rows >= jj) & rmask, colj, 0.0)
        xnorm = tl.sqrt(tl.sum(xfull * xfull))

        beta = -tl.where(alpha >= 0, xnorm, -xnorm)
        degen = xnorm == 0.0
        beta_safe = tl.where(degen, 1.0, beta)
        tau_j = tl.where(degen, 0.0, (beta_safe - alpha) / beta_safe)

        denom = alpha - beta_safe
        v = tl.where(below, x / denom, 0.0)
        v = tl.where(on_diag, 1.0, v)
        v = tl.where(degen, 0.0, v)

        tl.store(tau_ptr + pid * n + (k + jj), tau_j)

        diagval = tl.where(degen, 0.0, beta)
        newcol = tl.where(on_diag, diagval, tl.where(below, v, colj))
        panel = tl.where(sel[None, :], newcol[:, None], panel)

        w = tl.sum(v[:, None] * panel, axis=0)
        upd = tau_j * v[:, None] * w[None, :]
        panel = tl.where((cols > jj)[None, :], panel - upd, panel)

    tl.store(H_ptr + base + offs, panel, mask=pmask)


def _next_pow2(x: int) -> int:
    return 1 << (x - 1).bit_length()


def _unit_trapezoid(block, b):
    ar = torch.arange(b, device=block.device)
    lower = (ar[:, None] > ar[None, :]).to(block.dtype)
    eye = torch.eye(b, device=block.device, dtype=block.dtype)
    top = block[:, :b, :] * lower + eye
    return torch.cat([top, block[:, b:, :]], dim=1)


def _build_T(V, tau_blk, b, batch, device, dtype):
    S = torch.bmm(V.transpose(1, 2), V)
    ar = torch.arange(b, device=device)
    su = ar[:, None] < ar[None, :]
    eye = torch.eye(b, device=device, dtype=dtype)
    M = torch.where(su, S * tau_blk[:, None, :], torch.zeros_like(S)) + eye
    rhs = eye.unsqueeze(0).expand(batch, b, b).contiguous()
    invM = torch.linalg.solve_triangular(M, rhs, upper=True, unitriangular=True)
    return tau_blk[:, :, None] * invM


def fused_panel_qr(A: torch.Tensor, nb: int = 64):
    batch, n, _ = A.shape
    device, dtype = A.device, A.dtype
    H = A.clone()
    tau = torch.zeros(batch, n, device=device, dtype=dtype)
    BM = _next_pow2(n)
    nwarps = 8 if BM <= 256 else 16

    for k in range(0, n, nb):
        b = min(nb, n - k)
        m = n - k
        _panel_kernel[(batch,)](H, tau, k, m, n, B=b, BM=BM, BB=_next_pow2(b), num_warps=nwarps)

        V = _unit_trapezoid(H[:, k:, k:k + b], b)
        T = _build_T(V, tau[:, k:k + b], b, batch, device, dtype)

        if k + b < n:
            C = H[:, k:, k + b:]
            W = torch.bmm(V.transpose(1, 2), C)
            C -= torch.bmm(V, torch.bmm(T.transpose(1, 2), W))

    return H, tau


_graph_cache: dict = {}


def _capture(A, nb):
    static_in = A.clone()
    for _ in range(3):
        fused_panel_qr(static_in, nb=nb)
    torch.cuda.synchronize()

    g = torch.cuda.CUDAGraph()
    with torch.cuda.graph(g):
        H, tau = fused_panel_qr(static_in, nb=nb)
    return static_in, g, H, tau


def graphed_fused_qr(A: torch.Tensor, nb: int = 64):
    key = (tuple(A.shape), A.dtype, A.device.index, nb)
    entry = _graph_cache.get(key)
    if entry is None:
        entry = _capture(A, nb)
        _graph_cache[key] = entry
    static_in, g, H, tau = entry
    static_in.copy_(A)
    g.replay()
    return H.clone(), tau.clone()


def _use_fused(batch: int, n: int) -> bool:
    return n <= 1280 and batch * n >= 4096


def custom_kernel(data):
    A = data
    batch, n, _ = A.shape
    if _use_fused(batch, n):
        return graphed_fused_qr(A, nb=32)
    H, tau = torch.geqrf(A)
    return H.contiguous(), tau.contiguous()
scrolls · 140 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