submission 804074
oldsquaw · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 91 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-804074?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:a437c23aa6844c5f2b8cad9895534098bc887749e43db7b42afc31df64735c9d
license declaredunknown
license concludedunknown
authorsoldsquaw
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 8
_panel_qr[(b,)](A, tau, n, j, sab, sar, sac, stb, stn, M=M, BW=bw, num_warps=8)Kernel source
submission.py91 lines
import torch
from task import input_t, output_t
try:
import triton
import triton.language as tl
_HAS_TRITON = True
except Exception:
_HAS_TRITON = False
_TRITON_N = {176, 352, 512, 1024}
_BW_CAP = 32
_OB = 128
if _HAS_TRITON:
@triton.jit
def _panel_qr(A, TAU, N, J,
sab, sar, sac, stb, stn,
M: tl.constexpr, BW: tl.constexpr):
pid = tl.program_id(0)
rows = tl.arange(0, M)
cols = tl.arange(0, BW)
grow = J + rows
rmask = grow < N
cmask = (J + cols) < N
pptr = A + pid * sab + grow[:, None] * sar + (J + cols)[None, :] * sac
tile = rmask[:, None] & cmask[None, :]
P = tl.load(pptr, mask=tile, other=0.0)
for i in range(BW):
coli = tl.sum(tl.where(cols[None, :] == i, P, 0.0), axis=1)
x = tl.where((rows >= i) & rmask, coli, 0.0)
alpha = tl.sum(tl.where(rows == i, x, 0.0))
norm = tl.sqrt(tl.sum(x * x))
live = norm > 0.0
beta = tl.where(alpha >= 0, -norm, norm)
den = alpha - beta
inv = tl.where(live & (den != 0), 1.0 / den, 0.0)
v = tl.where(live & (rows == i), 1.0,
tl.where(live & (rows > i) & rmask, x * inv, 0.0))
taui = tl.where(live, (beta - alpha) / tl.where(live, beta, 1.0), 0.0)
w = tl.where(cols >= i, tl.sum(v[:, None] * P, axis=0), 0.0)
P = P - taui * v[:, None] * w[None, :]
kept = tl.sum(tl.where(cols[None, :] == i, P, 0.0), axis=1)
newcol = tl.where(rows < i, kept, tl.where(rows == i, beta, v))
P = tl.where(cols[None, :] == i, newcol[:, None], P)
tl.store(TAU + pid * stb + (J + i) * stn, taui)
tl.store(pptr, P, mask=tile)
def _apply_wy(P, tp, C, di):
w = P.shape[2]
V = P.tril(-1)
V[:, di[:w], di[:w]] = (tp != 0).to(V.dtype)
G = V.transpose(1, 2) @ V
Tinv = torch.triu(G, 1) + torch.diag_embed(1.0 / torch.where(tp == 0, torch.ones_like(tp), tp))
Y = torch.linalg.solve_triangular(Tinv.transpose(1, 2), V.transpose(1, 2) @ C, upper=False)
C.baddbmm_(V, Y, beta=1, alpha=-1)
def _blocked_qr(A):
b, n, _ = A.shape
M = triton.next_power_of_2(n)
bw = min(_BW_CAP, max(8, 16384 // M))
bw = 1 << (bw.bit_length() - 1)
ob = max(bw, (_OB // bw) * bw)
tau = A.new_zeros(b, ((n + bw - 1) // bw) * bw)
di = torch.arange(ob, device=A.device)
sab, sar, sac, stb, stn = A.stride(0), A.stride(1), A.stride(2), tau.stride(0), tau.stride(1)
for jo in range(0, n, ob):
obw = min(ob, n - jo)
for j in range(jo, jo + obw, bw):
cur = min(bw, jo + obw - j)
_panel_qr[(b,)](A, tau, n, j, sab, sar, sac, stb, stn, M=M, BW=bw, num_warps=8)
if j + cur < jo + obw:
_apply_wy(A[:, j:, j:j + cur], tau[:, j:j + cur], A[:, j:, j + cur:jo + obw], di)
if jo + obw >= n:
break
_apply_wy(A[:, jo:, jo:jo + obw], tau[:, jo:jo + obw], A[:, jo:, jo + obw:], di)
return A, tau[:, :n]
def custom_kernel(data: input_t) -> output_t:
if _HAS_TRITON and data.shape[1] in _TRITON_N:
try:
return _blocked_qr(data.clone())
except Exception:
pass
return torch.geqrf(data)
scrolls · 91 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