Skip to content
KernelIndex
Search⌘K

submission 833619

codeman62 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_cublass.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-833619?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
12.6ms
#308 of 515
2026-06-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:db93b46721c25ed8eea7a21f69fc97be33b482882ae928d52a6cbf0ed0f7cde8
license declaredunknown
license concludedunknown
authorscodeman62
imported2026-08-26

Techniques

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

num-warps = 8BLOCK_M=BLOCK_M, num_warps=8,

Kernel source

submission_cublass.py230 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl

from task import input_t, output_t

# Tile (SRAM-resident) panel kernel is used up to this size; larger n falls back
# to the per-column kernel to avoid register spills.
_TILE_MAX_N = 512

# Tile kernel is probed once per (BLOCK_N, BLOCK_PB); if it ever fails to compile
# or run, we disable it globally and use the proven column kernel everywhere.
_tile_ok = True
_tile_probed: set = set()


# ---------------------------------------------------------------------------
# Panel factor — per-column (proven, used as fallback for large n).
#
# Operates only on rows [p0:n] via local row coords (global row = p0 + r),
# so later panels load far fewer rows than the full matrix height.
# ---------------------------------------------------------------------------
@triton.jit
def householder_qr(
    H_ptr, tau_ptr,
    n, p0, pb,
    stride_hb, stride_hr, stride_hc,
    stride_tb, stride_tr,
    BLOCK_M: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    H = H_ptr + pid * stride_hb
    T = tau_ptr + pid * stride_tb
    rows = tl.arange(0, BLOCK_M)          # local rows; global row = p0 + rows
    m = n - p0

    for xx in range(0, pb):
        gcol = p0 + xx                    # diagonal/pivot at local row xx
        x_mask = (rows < m) & (rows >= xx)
        x_ptrs = H + (p0 + rows) * stride_hr + gcol * stride_hc
        x = tl.load(x_ptrs, mask=x_mask, other=0.0)

        x0 = tl.sum(tl.where(rows == xx, x, 0.0), axis=0)
        below_mask = x_mask & (rows > xx)
        tail_sq = tl.sum(tl.where(below_mask, x * x, 0.0), axis=0)

        has_refl = tail_sq > 0.0
        norm_x = tl.sqrt(x0 * x0 + tail_sq)
        sign_x0 = tl.where(x0 >= 0.0, 1.0, -1.0)
        alpha = tl.where(has_refl, -sign_x0 * norm_x, x0)
        beta = x0 - alpha
        safe = has_refl & (beta != 0.0)
        inv_beta = tl.where(safe, 1.0 / beta, 0.0)

        v = x * inv_beta
        v = tl.where(rows == xx, 1.0, v)
        v = tl.where(x_mask & safe, v, 0.0)
        tau = tl.where(safe, -beta / alpha, 0.0)

        tl.store(T + gcol * stride_tr, tau)
        tl.store(x_ptrs, alpha, mask=(rows == xx))
        tl.store(x_ptrs, v, mask=below_mask)

        for yy in range(xx + 1, pb):
            a_ptrs = H + (p0 + rows) * stride_hr + (p0 + yy) * stride_hc
            a = tl.load(a_ptrs, mask=x_mask, other=0.0)
            dot = tl.sum(v * a, axis=0)
            tl.store(a_ptrs, a - v * (tau * dot), mask=below_mask | (rows == xx))


# ---------------------------------------------------------------------------
# Panel factor — SRAM-resident tile. Loads only the (BLOCK_M x BLOCK_PB) panel
# starting at row p0 (local rows; global row = p0 + r), runs all pb reflectors
# as full-width vectorized rank-1 updates in registers, writes back once.
# Shrinking BLOCK_M per panel cuts register pressure and reductions for later
# panels (they cover far fewer rows than the full matrix height).
# ---------------------------------------------------------------------------
@triton.jit
def panel_factor_tile(
    H_ptr, tau_ptr,
    n, p0, pb,
    stride_hb, stride_hr, stride_hc,
    stride_tb, stride_tr,
    BLOCK_M: tl.constexpr, BLOCK_PB: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    Hp = H_ptr + pid * stride_hb
    Tp = tau_ptr + pid * stride_tb

    rows = tl.arange(0, BLOCK_M)          # local rows; global row = p0 + rows
    cols = tl.arange(0, BLOCK_PB)
    m = n - p0
    row_in = rows < m
    col_in = cols < pb
    gcol = p0 + cols

    ptrs = Hp + (p0 + rows)[:, None] * stride_hr + gcol[None, :] * stride_hc
    tmask = row_in[:, None] & col_in[None, :]
    tile = tl.load(ptrs, mask=tmask, other=0.0)          # (BLOCK_M, BLOCK_PB)

    for jj in range(0, pb):
        is_jj = cols == jj
        colvec = tl.sum(tl.where(is_jj[None, :], tile, 0.0), axis=1)   # (BLOCK_M,)

        active = row_in & (rows >= jj)        # local pivot row == jj
        x0 = tl.sum(tl.where(rows == jj, colvec, 0.0), axis=0)
        below = active & (rows > jj)
        tail_sq = tl.sum(tl.where(below, colvec * colvec, 0.0), axis=0)

        has_refl = tail_sq > 0.0
        norm_x = tl.sqrt(x0 * x0 + tail_sq)
        sign_x0 = tl.where(x0 >= 0.0, 1.0, -1.0)
        alpha = tl.where(has_refl, -sign_x0 * norm_x, x0)
        beta = x0 - alpha
        safe = has_refl & (beta != 0.0)
        inv_beta = tl.where(safe, 1.0 / beta, 0.0)

        v = colvec * inv_beta
        v = tl.where(rows == jj, 1.0, v)
        v = tl.where(active & safe, v, 0.0)              # v = 0 for rows < jj
        tau = tl.where(safe, -beta / alpha, 0.0)
        tl.store(Tp + (p0 + jj) * stride_tr, tau)

        # Write the compact column: keep R above the diagonal, alpha on it, v below.
        newcol = tl.where(rows < jj, colvec, tl.where(rows == jj, alpha, v))
        tile = tl.where(is_jj[None, :], newcol[:, None], tile)

        # Rank-1 update of all panel columns to the right (cols > jj), vectorized.
        w = tau * tl.sum(v[:, None] * tile, axis=0)      # (BLOCK_PB,)
        upd = cols > jj
        tile = tile - tl.where(upd[None, :], v[:, None] * w[None, :], 0.0)

    tl.store(ptrs, tile, mask=tmask)


def _factor(H, tau, batch, n, block, BLOCK_PB, cidx):
    shb, shr, shc = H.stride(0), H.stride(1), H.stride(2)
    stb, str_ = tau.stride(0), tau.stride(1)
    use_tile = (n <= _TILE_MAX_N) and _tile_ok

    for p0 in range(0, n, block):
        pb = min(block, n - p0)
        BLOCK_M = triton.next_power_of_2(n - p0)

        if use_tile:
            nw = 16 if BLOCK_M >= 512 else 8
            panel_factor_tile[(batch,)](
                H, tau, n, p0, pb,
                shb, shr, shc, stb, str_,
                BLOCK_M=BLOCK_M, BLOCK_PB=BLOCK_PB, num_warps=nw,
            )
        else:
            householder_qr[(batch,)](
                H, tau, n, p0, pb,
                shb, shr, shc, stb, str_,
                BLOCK_M=BLOCK_M, num_warps=8,
            )

        c_end = p0 + pb
        if c_end >= n:
            break

        # Block-reflector trailing update: (I - V T V^T) A_tr via cuBLAS,
        # with T^-1 = diag(1/tau) + striu(V^T V) -> one lower-triangular solve.
        c = cidx[:pb]
        V = H[:, p0:, p0:c_end].clone()
        V = torch.tril(V)
        V[:, c, c] = 1.0
        taup = tau[:, p0:c_end]
        V = torch.where((taup == 0).unsqueeze(1), torch.zeros_like(V), V)

        TinvT = torch.tril(V.transpose(1, 2) @ V, -1)
        TinvT[:, c, c] = torch.where(taup != 0, 1.0 / taup, torch.ones_like(taup))

        A_tr = H[:, p0:, c_end:]
        W = torch.linalg.solve_triangular(
            TinvT, V.transpose(1, 2) @ A_tr, upper=False, left=True
        )
        H[:, p0:, c_end:] = A_tr - V @ W


def _params(n: int, device):
    block = 64 if n >= 256 else 32
    BLOCK_PB = triton.next_power_of_2(block)
    cidx = torch.arange(block, device=device)
    return block, BLOCK_PB, cidx


def _probe_tile(n, device, dtype):
    """Compile/run the tile kernel once on scratch data; disable on any failure."""
    global _tile_ok
    if not _tile_ok or n > _TILE_MAX_N:
        return
    block, BLOCK_PB, _ = _params(n, device)
    BLOCK_M = triton.next_power_of_2(n)
    if (BLOCK_M, BLOCK_PB) in _tile_probed:
        return
    _tile_probed.add((BLOCK_M, BLOCK_PB))
    try:
        Hs = torch.randn(1, n, n, device=device, dtype=dtype)
        ts = torch.zeros(1, n, device=device, dtype=dtype)
        panel_factor_tile[(1,)](
            Hs, ts, n, 0, min(block, n),
            Hs.stride(0), Hs.stride(1), Hs.stride(2), ts.stride(0), ts.stride(1),
            BLOCK_M=BLOCK_M, BLOCK_PB=BLOCK_PB, num_warps=16 if n >= 512 else 8,
        )
    except Exception:
        _tile_ok = False


def custom_kernel(data: input_t) -> output_t:
    # Work on an independent, contiguous copy: the input must be left unmodified
    # (the benchmark times by calling repeatedly on the same tensor).
    if data.is_contiguous():
        H = data.clone()
    else:
        H = data.contiguous()
    batch, n, _ = H.shape
    device, dtype = H.device, H.dtype

    _probe_tile(n, device, dtype)

    block, BLOCK_PB, cidx = _params(n, device)
    tau = torch.zeros(batch, n, device=device, dtype=dtype)
    _factor(H, tau, batch, n, block, BLOCK_PB, cidx)
    return H, tau
scrolls · 230 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