submission 825389
umbrella___ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 277 lines, June 9 Researcher Reciprocity License v1.0.
submission_2level.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-825389?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:96696595111f46de12b8147f9f56951a92bffa85a65bc5d0165108c6f763c339
license declaredunknown
license concludedunknown
authorsumbrella___
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 8
def _qr(A, block, num_warps=8):Kernel source
submission_2level.py277 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ===========================================================================
# Batched compact-Householder QR (Triton blocked compact-WY), FP32.
#
# A Triton panel kernel does the unblocked geqr2 + in-kernel compact-WY T per
# panel; the FLOP-heavy trailing update C -= V (T^T (V^T C)) is batched cuBLAS
# FP32 GEMM (bmm/baddbmm). n>2048 falls back to cuSolver (torch.geqrf) -- a
# custom batched panel there is far slower than cuSolver.
#
# Precision: the trailing GEMMs run on TF32 tensor cores when it is SAFE, else
# FP32. Plain TF32 fails the band/rowscale gate and is too tight for n<512
# (rankdef/nearrank fail at n=384), so TF32 is gated on n>=512 AND a structural
# detector that flags rowscale (row-norm spread) and band/diagonal (both
# off-corner blocks ~0). The panel (reflectors) is always FP32. Validated by a
# 10-bit-mantissa TF32 simulation: every TF32-failing case is caught -> FP32.
# ===========================================================================
torch.backends.cuda.matmul.allow_tf32 = False
def _tf32_unsafe(A):
# True -> the trailing update must stay FP32 (TF32 would risk the gate).
# Cheap (~20us even on n=512 b=640): sub-sample columns for the row-norm
# spread, touch only the n/4 corners for band/diagonal. Avoids a full
# A.abs() materialization (which cost ~560us on the big batches).
n = A.shape[-1]
step = max(1, n // 32)
s = A[:, :, ::step]
rn2 = (s * s).sum(dim=2) # (B, n) approx squared row norms
mx = rn2.amax(dim=1)
mn = rn2.amin(dim=1).clamp_min(1e-37)
if (mx / mn).amax() > 1e6: # rowscale: norm spread (1e4)^2
return True
q = max(1, n // 4)
scale = A[:, ::step, ::step].abs().amax().clamp_min(1e-30)
tr = A[:, :q, n - q:].abs().amax() # top-right corner
bl = A[:, n - q:, :q].abs().amax() # bottom-left corner
if (tr < 1e-6 * scale) and (bl < 1e-6 * scale): # band / diagonal
return True
return False
@triton.jit
def _panel_kernel(P, TAU, T, VOUT, M, IB,
spb, spr, spc, stb, sti, sTb, sTr, sTc, svb, svr, svc,
BM: tl.constexpr, BNB: tl.constexpr):
b = tl.program_id(0)
r = tl.arange(0, BM)
c = tl.arange(0, BNB)
rm = r < M
cm = c < IB
p = P + b * spb + r[:, None] * spr + c[None, :] * spc
tile = tl.load(p, mask=rm[:, None] & cm[None, :], other=0.0)
tau_vec = tl.zeros((BNB,), dtype=tl.float32)
for j in range(BNB):
colj = tl.sum(tl.where(c[None, :] == j, tile, 0.0), axis=1)
alpha = tl.sum(tl.where(r == j, colj, 0.0))
xn2 = tl.sum(tl.where(r > j, colj * colj, 0.0))
reflect = xn2 > 0.0
sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(reflect, -sgn * tl.sqrt(alpha * alpha + xn2), alpha)
tau_j = tl.where(reflect, (beta - alpha) / tl.where(reflect, beta, 1.0), 0.0)
denom = tl.where(reflect, alpha - beta, 1.0)
vb = colj / denom
v = tl.where(r == j, 1.0, tl.where(r > j, vb, 0.0))
vmask = tl.where(r >= j, v, 0.0)
w = tl.sum(tl.where(c[None, :] > j, vmask[:, None] * tile, 0.0), axis=0)
tile = tile - tau_j * vmask[:, None] * w[None, :]
newcol = tl.where(r < j, colj, tl.where(r == j, beta, vb))
tile = tl.where(c[None, :] == j, newcol[:, None], tile)
tau_vec = tl.where(c == j, tau_j, tau_vec)
V = tl.where(r[:, None] == c[None, :], 1.0, tl.where(r[:, None] > c[None, :], tile, 0.0))
tl.store(VOUT + b * svb + r[:, None] * svr + c[None, :] * svc, V, mask=rm[:, None] & cm[None, :])
Tt = tl.zeros((BNB, BNB), dtype=tl.float32)
tau0 = tl.sum(tl.where(c == 0, tau_vec, 0.0))
Tt = tl.where((c[:, None] == 0) & (c[None, :] == 0), tau0, Tt)
for i in range(1, BNB):
tau_i = tl.sum(tl.where(c == i, tau_vec, 0.0))
Vi = tl.sum(tl.where(c[None, :] == i, V, 0.0), axis=1)
dots = tl.sum(V * Vi[:, None], axis=0)
z = tl.where(c < i, -tau_i * dots, 0.0)
Tz = tl.sum(tl.where(c[None, :] < i, Tt * z[None, :], 0.0), axis=1)
newTcol = tl.where(c < i, Tz, tl.where(c == i, tau_i, 0.0))
Tt = tl.where(c[None, :] == i, newTcol[:, None], Tt)
tl.store(T + b * sTb + c[:, None] * sTr + c[None, :] * sTc, Tt, mask=cm[:, None] & cm[None, :])
tl.store(P + b * spb + r[:, None] * spr + c[None, :] * spc, tile, mask=rm[:, None] & cm[None, :])
tl.store(TAU + b * stb + c * sti, tau_vec, mask=cm)
_WS = {}
def _ws(B, m, block, dev):
# reused (un-zeroed) workspaces: the panel fully overwrites the used region.
key = (B, m, block, str(dev))
ws = _WS.get(key)
if ws is None:
BNB = triton.next_power_of_2(block)
Vbuf = torch.empty((B, m, block), device=dev, dtype=torch.float32)
Tbuf = torch.empty((B, BNB, BNB), device=dev, dtype=torch.float32)
_WS[key] = ws = (Vbuf, Tbuf)
return ws
def _qr(A, block, num_warps=8):
B, m, n = A.shape
bs = int(block)
BNB = triton.next_power_of_2(bs)
H = A.clone()
tau = A.new_empty(B, n) # panel writes every column
Vbuf, Tbuf = _ws(B, m, bs, A.device)
for k in range(0, n, bs):
ib = min(bs, n - k)
BM = triton.next_power_of_2(m - k)
Hv = H[:, k:, k:k + ib]
Vb = Vbuf[:, :m - k, :ib]
tv = tau[:, k:] # panel writes tau in place
_panel_kernel[(B,)](Hv, tv, Tbuf, Vb, m - k, ib,
Hv.stride(0), Hv.stride(1), Hv.stride(2),
tv.stride(0), tv.stride(1),
Tbuf.stride(0), Tbuf.stride(1), Tbuf.stride(2),
Vb.stride(0), Vb.stride(1), Vb.stride(2),
BM=BM, BNB=BNB, num_warps=num_warps)
hi = k + ib
if hi < n:
V = Vb
T = Tbuf[:, :ib, :ib]
C = H[:, k:, hi:]
W = V.transpose(-1, -2) @ C
W = T.transpose(-1, -2) @ W
C.baddbmm_(V, W, beta=1, alpha=-1)
return H, tau
@triton.jit
def _larft_kernel(S, TAU, T, sSb, sSr, sSc, stb, sti, sTb, sTr, sTc,
OBW: tl.constexpr):
# build the compact-WY T (OBW x OBW upper-tri) from S = V^T V and tau.
b = tl.program_id(0)
rr = tl.arange(0, OBW)
cc = tl.arange(0, OBW)
Sm = tl.load(S + b * sSb + rr[:, None] * sSr + cc[None, :] * sSc)
tauv = tl.load(TAU + b * stb + cc * sti)
Tt = tl.zeros((OBW, OBW), dtype=tl.float32)
tau0 = tl.sum(tl.where(cc == 0, tauv, 0.0))
Tt = tl.where((rr[:, None] == 0) & (cc[None, :] == 0), tau0, Tt)
for i in range(1, OBW):
tau_i = tl.sum(tl.where(cc == i, tauv, 0.0))
Sci = tl.sum(tl.where(cc[None, :] == i, Sm, 0.0), axis=1)
z = tl.where(rr < i, -tau_i * Sci, 0.0)
Tz = tl.sum(tl.where(cc[None, :] < i, Tt * z[None, :], 0.0), axis=1)
newcol = tl.where(rr < i, Tz, tl.where(rr == i, tau_i, 0.0))
Tt = tl.where(cc[None, :] == i, newcol[:, None], Tt)
tl.store(T + b * sTb + rr[:, None] * sTr + cc[None, :] * sTc, Tt)
def _qr_blk(A, block):
# Blocked QR for large n with tiny batch: cuSolver factors each tall panel
# (FP32, well-utilized on tall-skinny), the O(n^3) trailing update runs on
# TF32 tensor cores. Beats a single big FP32 cuSolver geqrf when TF32-safe.
B, m, n = A.shape
OBW = triton.next_power_of_2(block)
H = A.clone()
tau = A.new_empty(B, n)
idx = torch.arange(block, device=A.device)
for k in range(0, n, block):
ib = min(block, n - k)
panel = H[:, k:, k:k + ib].contiguous()
Hp, tp = torch.geqrf(panel)
H[:, k:, k:k + ib] = Hp
tau[:, k:k + ib] = tp
hi = k + ib
if hi < n:
V = Hp.tril(-1)
V[:, idx[:ib], idx[:ib]] = 1.0
S = V.transpose(-1, -2) @ V # (B, ib, ib)
T = A.new_empty(B, OBW, OBW)
_larft_kernel[(B,)](S, tp, T,
S.stride(0), S.stride(1), S.stride(2),
tp.stride(0), tp.stride(1),
T.stride(0), T.stride(1), T.stride(2), OBW=OBW)
Tt = T[:, :ib, :ib]
C = H[:, k:, hi:]
W = V.transpose(-1, -2) @ C
W = Tt.transpose(-1, -2) @ W
C.baddbmm_(V, W, beta=1, alpha=-1)
return H, tau
def _qr2(A, NB, OB, nw):
# Two-level blocked QR. Small inner block NB keeps the (tall) panel tile in
# registers (a 4096xNB tile spills at NB>=16 -> the whole register file), while
# the wide outer block OB gives an efficient (K=OB) TF32 trailing GEMM. Runs
# both matrices concurrently (grid=B) vs cuSolver's serial-over-batch geqrf.
B, m, n = A.shape
BNB = triton.next_power_of_2(NB)
H = A.clone()
tau = A.new_zeros(B, n)
for K in range(0, n, OB):
OBw = min(OB, n - K)
for k in range(K, K + OBw, NB):
NBw = min(NB, K + OBw - k)
BM = triton.next_power_of_2(m - k)
Hv = H[:, k:, k:k + NBw]
Tt = A.new_zeros(B, BNB, BNB)
ts = A.new_zeros(B, BNB)
Vb = A.new_zeros(B, m - k, NBw)
_panel_kernel[(B,)](Hv, ts, Tt, Vb, m - k, NBw,
Hv.stride(0), Hv.stride(1), Hv.stride(2),
ts.stride(0), ts.stride(1),
Tt.stride(0), Tt.stride(1), Tt.stride(2),
Vb.stride(0), Vb.stride(1), Vb.stride(2),
BM=BM, BNB=BNB, num_warps=nw)
tau[:, k:k + NBw] = ts[:, :NBw]
if k + NBw < K + OBw: # within-OB trailing (narrow)
V = Vb
T = Tt[:, :NBw, :NBw]
C = H[:, k:, k + NBw:K + OBw]
W = V.transpose(-1, -2) @ C
W = T.transpose(-1, -2) @ W
C.baddbmm_(V, W, beta=1, alpha=-1)
if K + OBw < n: # combined far trailing
mK = m - K
blk = H[:, K:, K:K + OBw]
ri = torch.arange(mK, device=A.device)[:, None]
ci = torch.arange(OBw, device=A.device)[None, :]
V_OB = (ri == ci).to(A.dtype) + (ri > ci).to(A.dtype) * blk
S = V_OB.transpose(-1, -2) @ V_OB
T_OB = A.new_zeros(B, OBw, OBw)
_larft_kernel[(B,)](S, tau, T_OB,
S.stride(0), S.stride(1), S.stride(2),
tau.stride(0), tau.stride(1),
T_OB.stride(0), T_OB.stride(1), T_OB.stride(2),
OBW=OBw, num_warps=nw)
C = H[:, K:, K + OBw:]
W = V_OB.transpose(-1, -2) @ C
W = T_OB.transpose(-1, -2) @ W
C.baddbmm_(V_OB, W, beta=1, alpha=-1)
return H, tau
def custom_kernel(data: input_t) -> output_t:
A = data
n = A.shape[-1]
if n > 2048:
# n=4096: cuSolver geqrf is panel-serial over the batch. Two-level blocked
# QR (inner NB=8 fits the 4096-tall tile in registers, outer OB=64 gives an
# efficient TF32 trailing) runs both matrices concurrently. TF32 is safe at
# n=4096 (very lenient gate) for benign inputs; else fall back to cuSolver.
A = A.contiguous()
if not _tf32_unsafe(A):
torch.backends.cuda.matmul.allow_tf32 = True
try:
return _qr2(A, 8, 64, 8)
finally:
torch.backends.cuda.matmul.allow_tf32 = False
return torch.geqrf(A)
# B200-measured best block per size: n=1024 likes block=32 (9.43->7.29ms),
# but n=2048 is much worse at 32 (2048x32 panel tile) -> keep block=16 there.
if n == 2048:
block, nw = 16, 8
elif n >= 1024:
block, nw = 32, 8
else:
block, nw = 32, 4 # nw=4 fastest for all n<1024 (less cross-warp)
# TF32 tensor cores for the trailing GEMMs only when safe (n>=512 + benign).
use_tf32 = n >= 512 and not _tf32_unsafe(A)
torch.backends.cuda.matmul.allow_tf32 = use_tf32
try:
return _qr(A.contiguous(), block, nw)
finally:
torch.backends.cuda.matmul.allow_tf32 = False
scrolls · 277 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