Skip to content
KernelIndex
Search⌘K

submission 811076

jeeva2812 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v2_optimized.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-811076?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
85.9ms
#414 of 515
2026-06-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8fe99e4505d617301e0c963b68aa1f06aaf917d1f44a8299bea6548defe1ff72
license declaredunknown
license concludedunknown
authorsjeeva2812
imported2026-08-26

Techniques

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

num-warps = 4num_warps=4 if block_n <= 64 else 8,
tile-n = 2562. **Triton fused QR only viable for n ≤ 256** (BLOCK_N = 256 sits at the

Kernel source

submission_v2_optimized.py296 lines
"""
submission_v2_optimized.py — best-per-shape dispatch in a single file.

Synthesizes what the per-shape data on the new (updated) leaderboard shows:

    shape                    best path           measured µs (separate runs)
    ------------------------ ------------------- ---------------------------
    n = 32   b = 20          Triton fused QR      61 µs   (was 369 µs geqrf)
    n = 176  b = 40          Triton fused QR      2,160 µs (was 21,600 µs geqrf)
    n = 352  b = 40          blocked panel geqrf  28,100 µs (was 50,000 µs geqrf)
    n = 512  b = 640         blocked panel geqrf  611,000 µs (was 1,071,000 µs geqrf)
    n = 1024 b = 60          CholeskyQR3+Yamamoto 206,000 µs (vs 224,000 µs blocked)
    n = 2048 b = 8           CholeskyQR3+Yamamoto 137,000 µs (vs 139,000 µs blocked)
    n = 4096 b = 2           CholeskyQR3+Yamamoto 185,000 µs (vs 186,000 µs blocked)

Two key learnings driving the dispatch:

1. **Panel-geqrf beats full-matrix-geqrf at n ∈ [352, 768]**. cuSOLVER's
   `geqrfBatched` is serialized at large n; calling it on small (rows × 32)
   panel blocks exposes more parallelism. v1's "slow" blocked path is
   actually the fastest cuBLAS/cuSOLVER variant for medium n.

2. **Triton fused QR only viable for n ≤ 256** (BLOCK_N = 256 sits at the
   B200 SMEM/CTA ceiling at 256 KB FP32). For n = 352 we'd need BLOCK_N =
   512 = 1 MB, which spills. So panel-geqrf carries the middle.

Routing
-------
    n ≤ 256          → Triton fused QR (one CTA per matrix)
    256 < n ≤ 768    → blocked panel geqrf + compact-WY trailing
    n > 768          → CholeskyQR3 (FP32) + Yamamoto LU (FP64) + per-matrix
                       geqrf fallback for gate failures

Upper-triangular fast-path applies to all branches.

Fallback / safety
-----------------
- Triton kernel guarded by env flag (QR_DISABLE_TRITON=1) and try/except.
- Yamamoto LU gates (chol info, lu info, finite, tau range, L growth)
  catch instability and route those indices to geqrf.
- All output FP32, all per-matrix shapes guaranteed.
"""

import os
import torch
from task import input_t, output_t


# =========================================================================
# Triton fused QR for small n (n ≤ 256)
# =========================================================================

_HAVE_TRITON = False
if os.environ.get("QR_DISABLE_TRITON") != "1":
    try:
        import triton
        import triton.language as tl
        _HAVE_TRITON = True
    except Exception:
        pass


if _HAVE_TRITON:

    @triton.jit
    def _qr_fused_kernel(
        A_ptr, Tau_ptr,
        stride_ab, stride_an, stride_am,
        N: tl.constexpr,
        BLOCK_N: tl.constexpr,
    ):
        bid = tl.program_id(0)

        row_off = tl.arange(0, BLOCK_N)
        col_off = tl.arange(0, BLOCK_N)
        row_mask = row_off < N
        col_mask = col_off < N
        valid = row_mask[:, None] & col_mask[None, :]

        A_ptrs = (
            A_ptr
            + bid * stride_ab
            + row_off[:, None] * stride_an
            + col_off[None, :] * stride_am
        )
        A_blk = tl.load(A_ptrs, mask=valid, other=0.0)
        tau_vec = tl.zeros([BLOCK_N], dtype=tl.float32)

        for j in range(0, N):
            rows_below_eq = (row_off >= j) & row_mask
            rows_strict_below = (row_off > j) & row_mask
            col_eq_j = (col_off == j)

            col_j = tl.sum(A_blk * col_eq_j[None, :].to(tl.float32), axis=1)
            alpha = tl.sum(col_j * (row_off == j).to(tl.float32))
            col_strict = tl.where(rows_strict_below, col_j, 0.0)
            sigma = tl.sum(col_strict * col_strict)

            norm_sq = sigma + alpha * alpha
            need = sigma > 0.0
            norm = tl.sqrt(norm_sq)
            beta_active = tl.where(alpha >= 0.0, -norm, norm)
            beta = tl.where(need, beta_active, alpha)

            denom_raw = alpha - beta
            denom = tl.where(denom_raw == 0.0, 1.0, denom_raw)
            tau_j = tl.where(need, (beta - alpha) / beta, 0.0)

            v_below = tl.where(rows_strict_below, col_strict / denom, 0.0)
            v = tl.where(row_off == j, 1.0, v_below)
            v = tl.where(need, v, 0.0)

            cols_below_eq = (col_off >= j) & col_mask
            w_full = tl.sum(v[:, None] * A_blk, axis=0)
            w = tl.where(cols_below_eq, w_full, 0.0)

            A_blk = A_blk - tau_j * v[:, None] * w[None, :]
            store_mask = col_eq_j[None, :] & rows_strict_below[:, None]
            A_blk = tl.where(store_mask, v[:, None], A_blk)
            tau_vec = tl.where(col_off == j, tau_j, tau_vec)

        tl.store(A_ptrs, A_blk, mask=valid)
        tau_ptrs = Tau_ptr + bid * N + col_off
        tl.store(tau_ptrs, tau_vec, mask=col_mask)


    def _qr_triton(A: torch.Tensor):
        b, n, _ = A.shape
        block_n = 1
        while block_n < n:
            block_n *= 2
        H = A.clone().contiguous()
        tau = A.new_zeros(b, n)
        _qr_fused_kernel[(b,)](
            H, tau,
            H.stride(0), H.stride(1), H.stride(2),
            N=n, BLOCK_N=block_n,
            num_warps=4 if block_n <= 64 else 8,
        )
        return H, tau


# =========================================================================
# Blocked panel-geqrf + compact-WY trailing (the cuSOLVER-aware mid-n path)
# =========================================================================

_PANEL = 32


def _compact_wy(A: torch.Tensor, col: int, panel: int, tau: torch.Tensor):
    b, n, _ = A.shape
    rows = n - col

    Y = A.new_zeros(b, rows, panel)
    raw = A[:, col:col + rows, col:col + panel]
    lower = torch.ones(rows, panel, dtype=torch.bool, device=A.device).tril(-1)
    Y[:, lower] = raw[:, lower]
    diag_idx = torch.arange(panel, device=A.device)
    Y[:, diag_idx, diag_idx] = 1.0

    S = torch.bmm(Y.transpose(1, 2).contiguous(), Y)

    T = A.new_zeros(b, panel, panel)
    for j in range(panel):
        tau_j = tau[:, col + j]
        T[:, j, j] = tau_j
        if j > 0:
            Tz = torch.bmm(T[:, :j, :j], S[:, :j, j:j+1]).squeeze(-1)
            T[:, :j, j] = -tau_j.unsqueeze(-1) * Tz
    return Y, T


def _qr_blocked_panel(A: torch.Tensor):
    b, n, _ = A.shape
    tau = A.new_zeros(b, n)
    A = A.clone().contiguous()

    for col in range(0, n, _PANEL):
        panel = min(_PANEL, n - col)
        blk = A[:, col:, col:col + panel].contiguous()
        H_p, tau_p = torch.geqrf(blk)
        A[:, col:, col:col + panel] = H_p
        tau[:, col:col + panel] = tau_p

        if col + panel >= n:
            break

        Y, T = _compact_wy(A, col, panel, tau)
        trail = A[:, col:, col + panel:].contiguous()
        W = torch.bmm(Y.transpose(1, 2).contiguous(), trail)
        W = torch.bmm(T.transpose(1, 2).contiguous(), W)
        A[:, col:, col + panel:] = trail - torch.bmm(Y, W)

    return A, tau


# =========================================================================
# CholeskyQR3 + Yamamoto LU reconstruction (the fast large-n path)
# =========================================================================

def _qr_choleskyqr3_yamamoto(A: torch.Tensor):
    b, n, _ = A.shape

    G1 = torch.bmm(A.transpose(-2, -1), A)
    R1, inf1 = torch.linalg.cholesky_ex(G1, upper=True)
    if not (inf1 == 0).all():
        bad = inf1 != 0
        R1 = R1.clone()
        R1[bad] = torch.eye(n, dtype=A.dtype, device=A.device)
    Q1 = torch.linalg.solve_triangular(R1, A, upper=True, left=False)

    G2 = torch.bmm(Q1.transpose(-2, -1), Q1)
    R2, inf2 = torch.linalg.cholesky_ex(G2, upper=True)
    if not (inf2 == 0).all():
        bad2 = inf2 != 0
        R2 = R2.clone()
        R2[bad2] = torch.eye(n, dtype=A.dtype, device=A.device)
    Q2 = torch.linalg.solve_triangular(R2, Q1, upper=True, left=False)

    G3 = torch.bmm(Q2.transpose(-2, -1), Q2)
    R3, inf3 = torch.linalg.cholesky_ex(G3, upper=True)
    if not (inf3 == 0).all():
        bad3 = inf3 != 0
        R3 = R3.clone()
        R3[bad3] = torch.eye(n, dtype=A.dtype, device=A.device)
    Q = torch.linalg.solve_triangular(R3, Q2, upper=True, left=False)
    R = torch.bmm(R3, torch.bmm(R2, R1))

    ok_chol = (inf1 == 0) & (inf2 == 0) & (inf3 == 0)

    diag_Q = torch.diagonal(Q, dim1=-2, dim2=-1)
    s = torch.where(diag_Q >= 0,
                    torch.full_like(diag_Q, -1.0),
                    torch.full_like(diag_Q,  1.0))
    Q_s = Q * s.unsqueeze(-2)
    R_s = R * s.unsqueeze(-1)

    eye = torch.eye(n, dtype=A.dtype, device=A.device).expand(b, n, n)
    M = eye - Q_s
    LU64, _, inf_lu = torch.linalg.lu_factor_ex(
        M.double(), pivot=False, check_errors=False)
    L_strict = torch.tril(LU64, -1).float()
    tau = torch.diagonal(LU64, dim1=-2, dim2=-1).float()
    H = torch.triu(R_s) + L_strict

    ok_lu = inf_lu == 0
    ok_fin = torch.isfinite(tau).all(-1) & torch.isfinite(L_strict).flatten(1).all(-1)
    ok_tau = (tau > 1e-6).all(-1) & (tau <= 2.001).all(-1)
    ok_lgrowth = L_strict.abs().amax(dim=(-2, -1)) < 10.0
    ok = ok_chol & ok_lu & ok_fin & ok_tau & ok_lgrowth

    if not ok.all():
        bad = (~ok).nonzero(as_tuple=True)[0]
        # Use blocked-panel path for the fallback at n in [some range, 768]
        # but Yamamoto only fires for n > 768, so fallback uses geqrf which is
        # fine here (large enough that cuSOLVER per-matrix cost amortizes).
        H_fb, tau_fb = torch.geqrf(A.index_select(0, bad).contiguous())
        H = H.index_copy(0, bad, H_fb)
        tau = tau.index_copy(0, bad, tau_fb)

    return H, tau


# =========================================================================
# Dispatch
# =========================================================================

_TRITON_MAX_N        = 256   # BLOCK_N=256 at SMEM ceiling on B200
_BLOCKED_PANEL_MAX_N = 768   # above this, CholeskyQR3 wins


def _dispatch(A: torch.Tensor):
    b, n, _ = A.shape

    # Upper-triangular fast-path
    if A.tril(diagonal=-1).abs().max() < 1e-6:
        return A.clone(), A.new_zeros(b, n)

    if _HAVE_TRITON and n <= _TRITON_MAX_N:
        try:
            return _qr_triton(A)
        except Exception:
            pass

    if n <= _BLOCKED_PANEL_MAX_N:
        return _qr_blocked_panel(A)

    return _qr_choleskyqr3_yamamoto(A)


def custom_kernel(data: input_t) -> output_t:
    if data.dim() == 2:
        H, tau = _dispatch(data.unsqueeze(0))
        return H.squeeze(0), tau.squeeze(0)
    return _dispatch(data)
scrolls · 296 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