submission 836219
whao89 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 140 lines, June 9 Researcher Reciprocity License v1.0.
submission_geo80ms.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-836219?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:86bdfcdf7f9b7e30eb5a1759a393a9f1364afc24fa0bafe8beda2b2d54f746af
license declaredunknown
license concludedunknown
authorswhao89
imported2026-08-26
Kernel source
submission_geo80ms.py140 lines
from __future__ import annotations
import torch
import triton
import triton.language as tl
@triton.jit
def _panel_kernel(H_ptr, tau_ptr, k, m, n, B: tl.constexpr, BM: tl.constexpr, BB: tl.constexpr):
pid = tl.program_id(0)
rows = tl.arange(0, BM)
cols = tl.arange(0, BB)
rmask = rows < m
cmask = cols < B
base = pid * n * n + k * n + k
offs = rows[:, None] * n + cols[None, :]
pmask = rmask[:, None] & cmask[None, :]
panel = tl.load(H_ptr + base + offs, mask=pmask, other=0.0)
for jj in tl.static_range(B):
sel = cols == jj
colj = tl.sum(tl.where(sel[None, :], panel, 0.0), axis=1)
on_diag = rows == jj
below = rows > jj
x = tl.where(below & rmask, colj, 0.0)
alpha = tl.sum(tl.where(on_diag, colj, 0.0))
xfull = tl.where((rows >= jj) & rmask, colj, 0.0)
xnorm = tl.sqrt(tl.sum(xfull * xfull))
beta = -tl.where(alpha >= 0, xnorm, -xnorm)
degen = xnorm == 0.0
beta_safe = tl.where(degen, 1.0, beta)
tau_j = tl.where(degen, 0.0, (beta_safe - alpha) / beta_safe)
denom = alpha - beta_safe
v = tl.where(below, x / denom, 0.0)
v = tl.where(on_diag, 1.0, v)
v = tl.where(degen, 0.0, v)
tl.store(tau_ptr + pid * n + (k + jj), tau_j)
diagval = tl.where(degen, 0.0, beta)
newcol = tl.where(on_diag, diagval, tl.where(below, v, colj))
panel = tl.where(sel[None, :], newcol[:, None], panel)
w = tl.sum(v[:, None] * panel, axis=0)
upd = tau_j * v[:, None] * w[None, :]
panel = tl.where((cols > jj)[None, :], panel - upd, panel)
tl.store(H_ptr + base + offs, panel, mask=pmask)
def _next_pow2(x: int) -> int:
return 1 << (x - 1).bit_length()
def _unit_trapezoid(block, b):
ar = torch.arange(b, device=block.device)
lower = (ar[:, None] > ar[None, :]).to(block.dtype)
eye = torch.eye(b, device=block.device, dtype=block.dtype)
top = block[:, :b, :] * lower + eye
return torch.cat([top, block[:, b:, :]], dim=1)
def _build_T(V, tau_blk, b, batch, device, dtype):
S = torch.bmm(V.transpose(1, 2), V)
ar = torch.arange(b, device=device)
su = ar[:, None] < ar[None, :]
eye = torch.eye(b, device=device, dtype=dtype)
M = torch.where(su, S * tau_blk[:, None, :], torch.zeros_like(S)) + eye
rhs = eye.unsqueeze(0).expand(batch, b, b).contiguous()
invM = torch.linalg.solve_triangular(M, rhs, upper=True, unitriangular=True)
return tau_blk[:, :, None] * invM
def fused_panel_qr(A: torch.Tensor, nb: int = 64):
batch, n, _ = A.shape
device, dtype = A.device, A.dtype
H = A.clone()
tau = torch.zeros(batch, n, device=device, dtype=dtype)
BM = _next_pow2(n)
nwarps = 8 if BM <= 256 else 16
for k in range(0, n, nb):
b = min(nb, n - k)
m = n - k
_panel_kernel[(batch,)](H, tau, k, m, n, B=b, BM=BM, BB=_next_pow2(b), num_warps=nwarps)
V = _unit_trapezoid(H[:, k:, k:k + b], b)
T = _build_T(V, tau[:, k:k + b], b, batch, device, dtype)
if k + b < n:
C = H[:, k:, k + b:]
W = torch.bmm(V.transpose(1, 2), C)
C -= torch.bmm(V, torch.bmm(T.transpose(1, 2), W))
return H, tau
_graph_cache: dict = {}
def _capture(A, nb):
static_in = A.clone()
for _ in range(3):
fused_panel_qr(static_in, nb=nb)
torch.cuda.synchronize()
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
H, tau = fused_panel_qr(static_in, nb=nb)
return static_in, g, H, tau
def graphed_fused_qr(A: torch.Tensor, nb: int = 64):
key = (tuple(A.shape), A.dtype, A.device.index, nb)
entry = _graph_cache.get(key)
if entry is None:
entry = _capture(A, nb)
_graph_cache[key] = entry
static_in, g, H, tau = entry
static_in.copy_(A)
g.replay()
return H.clone(), tau.clone()
def _use_fused(batch: int, n: int) -> bool:
return n <= 1280 and batch * n >= 4096
def custom_kernel(data):
A = data
batch, n, _ = A.shape
if _use_fused(batch, n):
return graphed_fused_qr(A, nb=32)
H, tau = torch.geqrf(A)
return H.contiguous(), tau.contiguous()
scrolls · 140 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