submission 826338
arun_k · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 165 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-826338?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:ef3e13a4cecfae1c8cded583a6ef0123effe10f5962e453c919981b9e71d0dca
license declaredunknown
license concludedunknown
authorsarun_k
imported2026-08-26
Kernel source
submission.py165 lines
import torch
import triton
import triton.language as tl
# =====================================================================
# TRITON KERNEL: Tiled Panel Factorization & In-Kernel WY Matrix Generation
# =====================================================================
@triton.jit
def _panel_kernel(
P_ptr, TAU_ptr, T_ptr, VOUT_ptr,
M, IB,
stride_pb, stride_pr, stride_pc,
stride_tb, stride_ti,
stride_Tb, stride_Tr, stride_Tc,
stride_vb, stride_vr, stride_vc,
BM: tl.constexpr, # Static compile-time power-of-2 row bound
BNB: tl.constexpr # Static compile-time power-of-2 panel width bound
):
batch_idx = tl.program_id(0)
# Static 2D index configurations
r = tl.arange(0, BM)
c = tl.arange(0, BNB)
row_mask = r < M
col_mask = c < IB
# Base memory coordinates for this batch slice
p_panel = P_ptr + batch_idx * stride_pb + r[:, None] * stride_pr + c[None, :] * stride_pc
tile = tl.load(p_panel, mask=row_mask[:, None] & col_mask[None, :], other=0.0)
tau_vec = tl.zeros((BNB,), dtype=tl.float32)
# Sequential Householder sweep over the panel's internal columns
for j in range(BNB):
# Isolate column j cleanly using register masks
colj = tl.sum(tl.where(c[None, :] == j, tile, 0.0), axis=1)
# Pull out pivot alpha and the norm-squared of elements below it
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)
# Calculate Householder scalars stably without triggering branching divergence
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)
# Construct Householder vector elements
vb = colj / denom
vmask = tl.where(r == j, 1.0, tl.where(r > j, vb, 0.0))
# Local rank-1 update to the remaining column tracks inside this panel register file
w = tl.sum(tl.where(c[None, :] > j, vmask[:, None] * tile, 0.0), axis=0)
tile = tile - tau_j * vmask[:, None] * w[None, :]
# Save structural transformations back to our loop state accumulators
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)
# Reconstruct the unit-lower triangular reflection matrix V
V = tl.where(r[:, None] == c[None, :], 1.0, tl.where(r[:, None] > c[None, :], tile, 0.0))
# Store V out to global memory
VOUT_ptrs = VOUT_ptr + batch_idx * stride_vb + r[:, None] * stride_vr + c[None, :] * stride_vc
tl.store(VOUT_ptrs, V, mask=row_mask[:, None] & col_mask[None, :])
# Compute the Compact-WY Upper Triangular Block Matrix T entirely in SRAM
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)
# Vectorized internal dot product mapping
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)
# Store computed outputs back to global memory tracks
T_ptrs = T_ptr + batch_idx * stride_Tb + c[:, None] * stride_Tr + c[None, :] * stride_Tc
tl.store(T_ptrs, Tt, mask=col_mask[:, None] & col_mask[None, :])
tl.store(p_panel, tile, mask=row_mask[:, None] & col_mask[None, :])
tl.store(TAU_ptr + batch_idx * stride_tb + c * stride_ti, tau_vec, mask=col_mask)
def qr(A: torch.Tensor, block_size: int, num_warps: int = 8):
B, m, n = A.shape
bs = int(block_size)
BNB = triton.next_power_of_2(bs)
H = A.clone()
tau = A.new_zeros(B, n)
# Process the entire matrix via sequenced block column partitions
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] # In-place tracking slice view
Tt = A.new_zeros(B, BNB, BNB)
ts = A.new_zeros(B, BNB)
Vb = A.new_zeros(B, m - k, ib)
# Spin up the specialized Triton panel kernel configuration
_panel_kernel[(B,)](
Hv, ts, Tt, Vb, m - k, ib,
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=num_warps
)
tau[:, k : k + ib] = ts[:, :ib]
hi = k + ib
# Trailing Submatrix Update: A_trail = A_trail - V @ (T.T @ (V.T @ A_trail))
# Leverages cuBLAS Tensor Cores natively for maximum performance
if hi < n:
V = Vb
T = Tt[:, :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
# =====================================================================
# LEADERBOARD ENTRY CONFIGURATION
# =====================================================================
def custom_kernel(A: torch.Tensor):
n = A.shape[-1]
# Safety fallback bounds for incredibly gargantuan tensor ranges
if n > 2048:
return torch.geqrf(A.contiguous())
# Dynamically tune block configurations and warp density depending on matrix layout scale
if n >= 1024:
block, nw = 16, 8
elif n >= 256:
block, nw = 32, (4 if n == 512 else 8)
else:
block, nw = 32, 4
return qr(A.contiguous(), block_size=block, num_warps=nw)
compute = custom_kernelscrolls · 165 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