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
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 = 8
BLOCK_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