Skip to content
KernelIndex
Search⌘K

submission 830796

Vedanth Chamala · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-830796?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
10.2ms
#286 of 515
2026-06-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:26deae83148d96801db5b270cd7d4b8e456c5fac7c897e18eaa641e1c24b7161
license declaredunknown
license concludedunknown
authorsVedanth Chamala
imported2026-08-26

Kernel source

submission_v2.py160 lines
"""qr_v2 submission v2 — Triton fused-panel blocked Householder QR.

Architecture:
  * n <= 1024  : custom batched blocked Householder. Panel factorization is a
    single fused Triton kernel (one program/matrix, panel resident in SRAM, all
    kb=32 reflectors per launch). T via closed form. Trailing update + T-solve
    via torch GEMM/trsm (BF16x9 FP32-emulation enabled for tensor-core speed).
  * n >= 2048  : dispatch to torch.geqrf (few efficient single-matrix cuSOLVER
    calls; tensor-core-accelerated via the emulation env var). Dispatch is on
    SHAPE only, never on matrix contents.

Returns LAPACK-geqrf-compatible (H, tau). Does NOT mutate input.
The Triton kernel mirrors local/panel_mirror.py, which is validated on CPU
against the real checker.
"""
import os

os.environ.setdefault("CUBLAS_EMULATE_SINGLE_PRECISION", "1")

import torch
from task import input_t, output_t

try:
    import triton
    import triton.language as tl
    HAVE_TRITON = torch.cuda.is_available()
except Exception:
    HAVE_TRITON = False

NB = 32


# ----------------------------------------------------------------------------
# torch panel factorization (fallback: CPU, partial panels, non-Triton)
# ----------------------------------------------------------------------------
def _panel_factor_torch(H, j, kb, tau):
    B = H.shape[0]
    for i in range(kb):
        col = j + i
        alpha = H[:, col, col]
        tail = H[:, col + 1:, col]
        xnorm = torch.linalg.vector_norm(tail, dim=1)
        normfull = torch.sqrt(alpha * alpha + xnorm * xnorm)
        sign = torch.where(alpha >= 0, 1.0, -1.0)
        beta = -sign * normfull
        active = xnorm > 0
        denom_safe = torch.where(active, alpha - beta, torch.ones_like(alpha))
        taui = torch.where(active, (beta - alpha) / torch.where(active, beta, torch.ones_like(beta)),
                           torch.zeros_like(beta))
        H[:, col, col] = torch.where(active, beta, alpha)
        vtail = torch.where(active.unsqueeze(1), tail / denom_safe.unsqueeze(1), torch.zeros_like(tail))
        H[:, col + 1:, col] = vtail
        tau[:, col] = taui
        nc = (j + kb) - (col + 1)
        if nc > 0:
            P = H[:, col:, col + 1:j + kb]
            ones = torch.ones(B, 1, device=H.device, dtype=H.dtype)
            v = torch.cat([ones, vtail], dim=1)
            w = torch.bmm(v.unsqueeze(1), P).squeeze(1)
            P.sub_(taui.view(B, 1, 1) * v.unsqueeze(2) * w.unsqueeze(1))


# ----------------------------------------------------------------------------
# Triton fused panel kernel (one program per matrix; full panel kb==KB)
# ----------------------------------------------------------------------------
if HAVE_TRITON:
    @triton.jit
    def _panel_kernel(H_ptr, TAU_ptr, j, m,
                      s_b, s_r, s_c, s_tb,
                      BLOCK_M: tl.constexpr, KB: tl.constexpr):
        b = tl.program_id(0)
        row = tl.arange(0, BLOCK_M)
        colk = tl.arange(0, KB)
        rmask = row < m
        base = b * s_b + j * s_r + j * s_c
        offs = base + row[:, None] * s_r + colk[None, :] * s_c
        load_mask = rmask[:, None]
        P = tl.load(H_ptr + offs, mask=load_mask, other=0.0)
        tau_vec = tl.zeros([KB], dtype=tl.float32)
        for i in tl.static_range(KB):
            is_i = row == i
            below = (row > i) & rmask
            col_i = tl.sum(tl.where(colk[None, :] == i, P, 0.0), axis=1)
            alpha = tl.sum(tl.where(is_i, col_i, 0.0))
            xnorm2 = tl.sum(tl.where(below, col_i * col_i, 0.0))
            normfull = tl.sqrt(alpha * alpha + xnorm2)
            active = xnorm2 > 0.0
            sign = tl.where(alpha >= 0.0, 1.0, -1.0)
            beta = -sign * normfull
            safe_denom = tl.where(active, alpha - beta, 1.0)
            safe_beta = tl.where(active, beta, 1.0)
            inv_denom = tl.where(active, 1.0 / safe_denom, 0.0)
            tau_i = tl.where(active, (beta - alpha) / safe_beta, 0.0)
            vtail = tl.where(below & active, col_i * inv_denom, 0.0)
            v = tl.where(is_i, 1.0, vtail)
            new_diag = tl.where(active, beta, alpha)
            new_col_i = tl.where(is_i, new_diag, tl.where(below, vtail, col_i))
            P = tl.where(colk[None, :] == i, new_col_i[:, None], P)
            tau_vec = tl.where(colk == i, tau_i, tau_vec)
            w = tl.sum(v[:, None] * P, axis=0)
            cmask = colk > i
            P = P - tl.where(cmask[None, :], tau_i * v[:, None] * w[None, :], 0.0)
        tl.store(H_ptr + offs, P, mask=load_mask)
        tl.store(TAU_ptr + b * s_tb + (j + colk), tau_vec)

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

    def _panel_factor_triton(H, j, kb, tau):
        B, n, _ = H.shape
        m = n - j
        BLOCK_M = _next_pow2(n)
        nw = 8 if BLOCK_M <= 512 else 16
        _panel_kernel[(B,)](H, tau, j, m,
                            H.stride(0), H.stride(1), H.stride(2), tau.stride(0),
                            BLOCK_M=BLOCK_M, KB=kb, num_warps=nw)


# ----------------------------------------------------------------------------
# closed-form compact-WY T  (one GEMM + one batched triangular solve)
# ----------------------------------------------------------------------------
def _build_T_fast(V, taus):
    kb = V.shape[2]
    G = V.transpose(1, 2) @ V
    N = taus.unsqueeze(2) * torch.triu(G, 1)
    M = torch.eye(kb, dtype=V.dtype, device=V.device).expand(V.shape[0], kb, kb) + N
    return torch.linalg.solve_triangular(M, torch.diag_embed(taus), upper=True,
                                         left=True, unitriangular=True)


def _qr_blocked(A, use_triton):
    B, n, _ = A.shape
    H = A.clone()
    tau = H.new_zeros(B, n)
    didx = torch.arange(NB, device=H.device)
    for j in range(0, n, NB):
        kb = min(NB, n - j)
        if use_triton and kb == NB:
            _panel_factor_triton(H, j, kb, tau)
        else:
            _panel_factor_torch(H, j, kb, tau)
        panel = H[:, j:, j:j + kb]
        V = torch.tril(panel, -1).contiguous()
        V[:, didx[:kb], didx[:kb]] = 1.0
        if j + kb < n:
            T = _build_T_fast(V, tau[:, j:j + kb])
            C = H[:, j:, j + kb:]
            W = torch.bmm(V.transpose(1, 2), C)
            W = torch.bmm(T.transpose(1, 2), W)
            C.sub_(torch.bmm(V, W))
    return H, tau


def custom_kernel(data: input_t) -> output_t:
    A = data
    n = A.shape[-1]
    if (not HAVE_TRITON) or n >= 2048:
        return torch.geqrf(A)
    return _qr_blocked(A, use_triton=True)
scrolls · 160 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