submission 830796
Vedanth Chamala · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 160 lines, June 9 Researcher Reciprocity License v1.0.
submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-830796?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:26deae83148d96801db5b270cd7d4b8e456c5fac7c897e18eaa641e1c24b7161
license declaredunknown
license concludedunknown
authorsVedanth Chamala
imported2026-08-26
Kernel source
submission_v2.py160 lines
"""qr_v2 submission v2 — Triton fused-panel blocked Householder QR.
Architecture:
* n <= 1024 : custom batched blocked Householder. Panel factorization is a
single fused Triton kernel (one program/matrix, panel resident in SRAM, all
kb=32 reflectors per launch). T via closed form. Trailing update + T-solve
via torch GEMM/trsm (BF16x9 FP32-emulation enabled for tensor-core speed).
* n >= 2048 : dispatch to torch.geqrf (few efficient single-matrix cuSOLVER
calls; tensor-core-accelerated via the emulation env var). Dispatch is on
SHAPE only, never on matrix contents.
Returns LAPACK-geqrf-compatible (H, tau). Does NOT mutate input.
The Triton kernel mirrors local/panel_mirror.py, which is validated on CPU
against the real checker.
"""
import os
os.environ.setdefault("CUBLAS_EMULATE_SINGLE_PRECISION", "1")
import torch
from task import input_t, output_t
try:
import triton
import triton.language as tl
HAVE_TRITON = torch.cuda.is_available()
except Exception:
HAVE_TRITON = False
NB = 32
# ----------------------------------------------------------------------------
# torch panel factorization (fallback: CPU, partial panels, non-Triton)
# ----------------------------------------------------------------------------
def _panel_factor_torch(H, j, kb, tau):
B = H.shape[0]
for i in range(kb):
col = j + i
alpha = H[:, col, col]
tail = H[:, col + 1:, col]
xnorm = torch.linalg.vector_norm(tail, dim=1)
normfull = torch.sqrt(alpha * alpha + xnorm * xnorm)
sign = torch.where(alpha >= 0, 1.0, -1.0)
beta = -sign * normfull
active = xnorm > 0
denom_safe = torch.where(active, alpha - beta, torch.ones_like(alpha))
taui = torch.where(active, (beta - alpha) / torch.where(active, beta, torch.ones_like(beta)),
torch.zeros_like(beta))
H[:, col, col] = torch.where(active, beta, alpha)
vtail = torch.where(active.unsqueeze(1), tail / denom_safe.unsqueeze(1), torch.zeros_like(tail))
H[:, col + 1:, col] = vtail
tau[:, col] = taui
nc = (j + kb) - (col + 1)
if nc > 0:
P = H[:, col:, col + 1:j + kb]
ones = torch.ones(B, 1, device=H.device, dtype=H.dtype)
v = torch.cat([ones, vtail], dim=1)
w = torch.bmm(v.unsqueeze(1), P).squeeze(1)
P.sub_(taui.view(B, 1, 1) * v.unsqueeze(2) * w.unsqueeze(1))
# ----------------------------------------------------------------------------
# Triton fused panel kernel (one program per matrix; full panel kb==KB)
# ----------------------------------------------------------------------------
if HAVE_TRITON:
@triton.jit
def _panel_kernel(H_ptr, TAU_ptr, j, m,
s_b, s_r, s_c, s_tb,
BLOCK_M: tl.constexpr, KB: tl.constexpr):
b = tl.program_id(0)
row = tl.arange(0, BLOCK_M)
colk = tl.arange(0, KB)
rmask = row < m
base = b * s_b + j * s_r + j * s_c
offs = base + row[:, None] * s_r + colk[None, :] * s_c
load_mask = rmask[:, None]
P = tl.load(H_ptr + offs, mask=load_mask, other=0.0)
tau_vec = tl.zeros([KB], dtype=tl.float32)
for i in tl.static_range(KB):
is_i = row == i
below = (row > i) & rmask
col_i = tl.sum(tl.where(colk[None, :] == i, P, 0.0), axis=1)
alpha = tl.sum(tl.where(is_i, col_i, 0.0))
xnorm2 = tl.sum(tl.where(below, col_i * col_i, 0.0))
normfull = tl.sqrt(alpha * alpha + xnorm2)
active = xnorm2 > 0.0
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = -sign * normfull
safe_denom = tl.where(active, alpha - beta, 1.0)
safe_beta = tl.where(active, beta, 1.0)
inv_denom = tl.where(active, 1.0 / safe_denom, 0.0)
tau_i = tl.where(active, (beta - alpha) / safe_beta, 0.0)
vtail = tl.where(below & active, col_i * inv_denom, 0.0)
v = tl.where(is_i, 1.0, vtail)
new_diag = tl.where(active, beta, alpha)
new_col_i = tl.where(is_i, new_diag, tl.where(below, vtail, col_i))
P = tl.where(colk[None, :] == i, new_col_i[:, None], P)
tau_vec = tl.where(colk == i, tau_i, tau_vec)
w = tl.sum(v[:, None] * P, axis=0)
cmask = colk > i
P = P - tl.where(cmask[None, :], tau_i * v[:, None] * w[None, :], 0.0)
tl.store(H_ptr + offs, P, mask=load_mask)
tl.store(TAU_ptr + b * s_tb + (j + colk), tau_vec)
def _next_pow2(x):
return 1 << (x - 1).bit_length()
def _panel_factor_triton(H, j, kb, tau):
B, n, _ = H.shape
m = n - j
BLOCK_M = _next_pow2(n)
nw = 8 if BLOCK_M <= 512 else 16
_panel_kernel[(B,)](H, tau, j, m,
H.stride(0), H.stride(1), H.stride(2), tau.stride(0),
BLOCK_M=BLOCK_M, KB=kb, num_warps=nw)
# ----------------------------------------------------------------------------
# closed-form compact-WY T (one GEMM + one batched triangular solve)
# ----------------------------------------------------------------------------
def _build_T_fast(V, taus):
kb = V.shape[2]
G = V.transpose(1, 2) @ V
N = taus.unsqueeze(2) * torch.triu(G, 1)
M = torch.eye(kb, dtype=V.dtype, device=V.device).expand(V.shape[0], kb, kb) + N
return torch.linalg.solve_triangular(M, torch.diag_embed(taus), upper=True,
left=True, unitriangular=True)
def _qr_blocked(A, use_triton):
B, n, _ = A.shape
H = A.clone()
tau = H.new_zeros(B, n)
didx = torch.arange(NB, device=H.device)
for j in range(0, n, NB):
kb = min(NB, n - j)
if use_triton and kb == NB:
_panel_factor_triton(H, j, kb, tau)
else:
_panel_factor_torch(H, j, kb, tau)
panel = H[:, j:, j:j + kb]
V = torch.tril(panel, -1).contiguous()
V[:, didx[:kb], didx[:kb]] = 1.0
if j + kb < n:
T = _build_T_fast(V, tau[:, j:j + kb])
C = H[:, j:, j + kb:]
W = torch.bmm(V.transpose(1, 2), C)
W = torch.bmm(T.transpose(1, 2), W)
C.sub_(torch.bmm(V, W))
return H, tau
def custom_kernel(data: input_t) -> output_t:
A = data
n = A.shape[-1]
if (not HAVE_TRITON) or n >= 2048:
return torch.geqrf(A)
return _qr_blocked(A, use_triton=True)
scrolls · 160 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