submission 837609
arsrivish26691 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 182 lines, June 9 Researcher Reciprocity License v1.0.
qrv9.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-837609?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:d369e161965eb4ea44ae2f3792055e488921f27d8d80c8f40c48c5be4f82f08b
license declaredunknown
license concludedunknown
authorsarsrivish26691
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
G += tl.dot(tl.trans(V), V, input_precision="ieee")num-warps = 8
_copy_kernel[(triton.cdiv(total, COPY_BLOCK),)](data, H, TOTAL=total, BLOCK=COPY_BLOCK, num_warps=8)stages = 1
NB=IB, RB=RBLOCK, CB=CBLOCK, PR=PREC, num_warps=4, num_stages=1)Kernel source
qrv9.py182 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from task import input_t, output_t
IB = 16 # inner panel width: factored sequentially, K>=16, no register spill
OB = 64 # outer block width: composed block reflector -> fat trailing GEMM (K=64)
RBLOCK = 128
# Trailing-GEMM precision (wide update contracts over K=OB=64, so this bites here).
# Sweep "ieee" -> "tf32x3" -> "tf32"; read orthogonality in --mode test.
PREC = "tf32x3"
CBLOCK = 64
COPY_BLOCK = 1024
@triton.jit
def _copy_kernel(A, H, TOTAL: tl.constexpr, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < TOTAL
tl.store(H + offs, tl.load(A + offs, mask=mask, other=0.0), mask=mask)
@triton.jit
def _panel_factor_kernel(H, TAU, TBUF,
N: tl.constexpr, P, ROWS: tl.constexpr, NB: tl.constexpr):
b = tl.program_id(0)
ro = tl.arange(0, ROWS)
rows = P + ro
q = tl.arange(0, NB)
cols = P + q
X = tl.load(H + b * N * N + rows[:, None] * N + cols[None, :],
mask=(rows[:, None] < N) & (cols[None, :] < N), other=0.0)
tauv = tl.zeros((NB,), dtype=tl.float32)
for jj in tl.static_range(0, NB):
k = P + jj
active = k < N
x = tl.sum(tl.where(q[None, :] == jj, X, 0.0), axis=1)
x = tl.where((rows >= k) & active, x, 0.0)
alpha = tl.sum(tl.where(rows == k, x, 0.0), axis=0)
nrm = tl.sqrt(tl.sum(x * x, axis=0))
sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = -sgn * nrm
safe = nrm > 0.0
beta = tl.where(safe, beta, alpha)
denom = tl.where(safe, alpha - beta, 1.0)
beta_safe = tl.where(tl.abs(beta) > 0.0, beta, 1.0)
tau = tl.where(safe, (beta - alpha) / beta_safe, 0.0)
tau = tl.where(active, tau, 0.0)
v = tl.where(rows == k, 1.0, x / denom)
v = tl.where((rows >= k) & active, v, 0.0)
dots = tl.sum(v[:, None] * X, axis=0)
X = tl.where((q[None, :] > jj) & (rows[:, None] >= k) & active,
X - tau * v[:, None] * dots[None, :], X)
X = tl.where((q[None, :] == jj) & (rows[:, None] == k) & active, beta, X)
X = tl.where((q[None, :] == jj) & (rows[:, None] > k) & active, v[:, None], X)
tauv = tl.where(q == jj, tau, tauv)
tl.store(TAU + b * N + k, tau, mask=active)
tl.store(H + b * N * N + rows[:, None] * N + cols[None, :], X,
mask=(rows[:, None] < N) & (cols[None, :] < N))
V = tl.where(rows[:, None] == cols[None, :], 1.0, X)
V = tl.where((rows[:, None] >= cols[None, :]) & (cols[None, :] < N), V, 0.0)
ti = tl.arange(0, NB)
tj = tl.arange(0, NB)
T = tl.zeros((NB, NB), dtype=tl.float32)
for jj in tl.static_range(0, NB):
active = (P + jj) < N
vj = tl.sum(tl.where(q[None, :] == jj, V, 0.0), axis=1)
s = tl.where(q < jj, tl.sum(V * vj[:, None], axis=0), 0.0)
tau_j = tl.sum(tl.where(q == jj, tauv, 0.0), axis=0)
y = -tau_j * tl.sum(T * s[None, :], axis=1)
T = tl.where((tj[None, :] == jj) & (ti[:, None] < jj) & active, y[:, None], T)
T = tl.where((tj[None, :] == jj) & (ti[:, None] == jj) & active, tau_j, T)
tl.store(TBUF + b * NB * NB + ti[:, None] * NB + tj[None, :], T)
@triton.jit
def _form_T_wide_kernel(H, TAU, TBUFO,
N: tl.constexpr, P, NROWT, OBW: tl.constexpr, RB: tl.constexpr, LOG2: tl.constexpr):
b = tl.program_id(0)
qi = tl.arange(0, OBW)
rr = tl.arange(0, RB)
cols = P + qi
base = b * N * N
G = tl.zeros((OBW, OBW), dtype=tl.float32) # Gram V^T V, row-tiled
for rt in range(NROWT):
rows = P + rt * RB + rr
rm = rows < N
Xv = tl.load(H + base + rows[:, None] * N + cols[None, :],
mask=rm[:, None] & (cols[None, :] < N), other=0.0)
V = tl.where(rows[:, None] == cols[None, :], 1.0, Xv)
V = tl.where((rows[:, None] >= cols[None, :]) & (cols[None, :] < N), V, 0.0)
G += tl.dot(tl.trans(V), V, input_precision="ieee")
tauv = tl.load(TAU + b * N + cols, mask=cols < N, other=0.0)
ti = tl.arange(0, OBW)
tj = tl.arange(0, OBW)
U = tl.where(tj[None, :] > ti[:, None], G, 0.0) # strictly upper of Gram
M = -tauv[:, None] * U # -diag(tau) . striu(V^T V), nilpotent
R = tl.where(ti[:, None] == tj[None, :], 1.0, 0.0) # identity
cur = M
for _ in tl.static_range(0, LOG2): # (I-M)^-1 = prod (I + M^{2^l})
R = R + tl.dot(R, cur, input_precision="ieee") # R @ (I + cur)
cur = tl.dot(cur, cur, input_precision="ieee") # square M
T = R * tauv[None, :] # R @ diag(tau)
tl.store(TBUFO + b * OBW * OBW + ti[:, None] * OBW + tj[None, :], T)
@triton.jit
def _trailing_kernel(H, TBUF,
N: tl.constexpr, P, NROWT, COL_END,
NB: tl.constexpr, RB: tl.constexpr, CB: tl.constexpr, PR: tl.constexpr):
b = tl.program_id(0)
cbid = tl.program_id(1)
qi = tl.arange(0, NB)
rr = tl.arange(0, RB)
cj = tl.arange(0, CB)
vcols = P + qi
ccols = P + NB + cbid * CB + cj
base = b * N * N
W = tl.zeros((NB, CB), dtype=tl.float32)
for rt in range(NROWT):
rows = P + rt * RB + rr
rm = rows < N
Xv = tl.load(H + base + rows[:, None] * N + vcols[None, :],
mask=rm[:, None] & (vcols[None, :] < N), other=0.0)
V = tl.where(rows[:, None] == vcols[None, :], 1.0, Xv)
V = tl.where((rows[:, None] >= vcols[None, :]) & (vcols[None, :] < N), V, 0.0)
C = tl.load(H + base + rows[:, None] * N + ccols[None, :],
mask=rm[:, None] & (ccols[None, :] < COL_END), other=0.0)
W += tl.dot(tl.trans(V), C, input_precision=PR)
ti = tl.arange(0, NB)
tj = tl.arange(0, NB)
T = tl.load(TBUF + b * NB * NB + ti[:, None] * NB + tj[None, :])
W2 = tl.dot(tl.trans(T), W, input_precision=PR)
for rt in range(NROWT):
rows = P + rt * RB + rr
rm = rows < N
cmask = rm[:, None] & (ccols[None, :] < COL_END)
Xv = tl.load(H + base + rows[:, None] * N + vcols[None, :],
mask=rm[:, None] & (vcols[None, :] < N), other=0.0)
V = tl.where(rows[:, None] == vcols[None, :], 1.0, Xv)
V = tl.where((rows[:, None] >= vcols[None, :]) & (vcols[None, :] < N), V, 0.0)
C = tl.load(H + base + rows[:, None] * N + ccols[None, :], mask=cmask, other=0.0)
C = C - tl.dot(V, W2, input_precision=PR)
tl.store(H + base + rows[:, None] * N + ccols[None, :], C, mask=cmask)
def custom_kernel(data: input_t) -> output_t:
B = data.shape[0]
N = data.shape[1]
H = torch.empty_like(data)
tau = torch.empty((B, N), device=data.device, dtype=data.dtype)
total = B * N * N
_copy_kernel[(triton.cdiv(total, COPY_BLOCK),)](data, H, TOTAL=total, BLOCK=COPY_BLOCK, num_warps=8)
rows_pow2 = triton.next_power_of_2(N)
pf_warps = 32 if N >= 2048 else (16 if N >= 1024 else 8) # n=512 didn't spill; keep it at 8
Tin = torch.empty((B, IB, IB), device=data.device, dtype=data.dtype)
Tout = torch.empty((B, OB, OB), device=data.device, dtype=data.dtype)
for p in range(0, N, OB):
col_end = min(p + OB, N)
for ip in range(0, OB, IB):
c0 = p + ip
if c0 >= N:
break
_panel_factor_kernel[(B,)](H, tau, Tin, N=N, P=c0, ROWS=rows_pow2, NB=IB, num_warps=pf_warps)
if c0 + IB < col_end:
nrowt = triton.cdiv(N - c0, RBLOCK)
inct = triton.cdiv(col_end - (c0 + IB), CBLOCK)
_trailing_kernel[(B, inct)](H, Tin, N=N, P=c0, NROWT=nrowt, COL_END=col_end,
NB=IB, RB=RBLOCK, CB=CBLOCK, PR=PREC, num_warps=4, num_stages=1)
if p + OB < N:
nrowt = triton.cdiv(N - p, RBLOCK)
_form_T_wide_kernel[(B,)](H, tau, Tout, N=N, P=p, NROWT=nrowt, OBW=OB, RB=RBLOCK, LOG2=OB.bit_length() - 1, num_warps=pf_warps)
nct = triton.cdiv(N - (p + OB), CBLOCK)
_trailing_kernel[(B, nct)](H, Tout, N=N, P=p, NROWT=nrowt, COL_END=N,
NB=OB, RB=RBLOCK, CB=CBLOCK, PR=PREC, num_warps=8, num_stages=1)
return H, tauscrolls · 182 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