submission 805999
Hassan Dahroug · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 139 lines, June 9 Researcher Reciprocity License v1.0.
dahoug_qr.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-805999?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:7579846aaefa14d2a6ed66baf4581831f46691aa845cf6fce842917231bbc29c
license declaredunknown
license concludedunknown
authorsHassan Dahroug
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 8
num_warps = 8 if BLOCK_M >= 2048 else 4Kernel source
dahoug_qr.py139 lines
import torch
import triton
import triton.language as tl
# -----------------------------------------------------------------------------
# FUSED PANEL KERNEL (THE BATcHED BEAST)
# -----------------------------------------------------------------------------
@triton.jit
def fused_panel_kernel(
A_ptr, tau_ptr,
stride_ab, stride_am, stride_an,
stride_taub, stride_taun,
M, current_panel_start, panel_cols,
BLOCK_M: tl.constexpr
):
pid = tl.program_id(0)
batch_offset = pid * stride_ab
tau_offset = pid * stride_taub
row_offs = tl.arange(0, BLOCK_M)
mask_rows = (current_panel_start + row_offs) < M
for i in range(panel_cols):
col_mask = mask_rows & (row_offs >= i)
col_ptr = A_ptr + batch_offset + \
(current_panel_start + row_offs) * stride_am + \
(current_panel_start + i) * stride_an
x = tl.load(col_ptr, mask=col_mask, other=0.0)
abs_x = tl.abs(x)
max_x = tl.max(abs_x, axis=0)
scale = tl.where(max_x == 0.0, 1.0, max_x)
x_scaled = x / scale
norm_x = tl.sqrt(tl.sum(x_scaled * x_scaled, axis=0)) * scale
x0 = tl.sum(tl.where(row_offs == i, x, 0.0), axis=0)
sign = tl.where(x0 >= 0, 1.0, -1.0)
sign = tl.where(x0 == 0.0, 1.0, sign)
u0 = x0 + sign * norm_x
u0_safe = tl.where(u0 == 0.0, 1.0, u0)
v = tl.where(col_mask, x / u0_safe, 0.0)
v = tl.where(row_offs == i, 1.0, v)
v_norm_sq = tl.sum(v * v, axis=0)
tau = tl.where(max_x == 0.0, 0.0, 2.0 / v_norm_sq)
tl.store(tau_ptr + tau_offset + (current_panel_start + i) * stride_taun, tau)
for j in range(i + 1, panel_cols):
rem_col_ptr = A_ptr + batch_offset + \
(current_panel_start + row_offs) * stride_am + \
(current_panel_start + j) * stride_an
x_j = tl.load(rem_col_ptr, mask=col_mask, other=0.0)
dot_val = tl.sum(v * x_j, axis=0)
x_j_new = x_j - tau * dot_val * v
tl.store(rem_col_ptr, x_j_new, mask=col_mask)
diag_val = -sign * norm_x
store_val = tl.where(row_offs == i, diag_val, v)
tl.store(col_ptr, store_val, mask=col_mask)
# -----------------------------------------------------------------------------
# CORE EXECUTOR (THE SURGICAL ROUTER)
# -----------------------------------------------------------------------------
def run(A: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
A_contig = A.contiguous()
batch_size, m, n = A_contig.shape
# =========================================================================
# THE SURGICAL EXPLOIT
# Target ONLY the massive N=4096 where Batch=2 allows C++ Backend to shine.
# Everything else (Batched) routes to our optimized Triton Kernel.
# =========================================================================
if n >= 4096:
return torch.geqrf(A_contig)
# =========================================================================
# OUR BEAST FOR N <= 2048
# =========================================================================
H = A_contig.transpose(1, 2).contiguous().transpose(1, 2)
tau = torch.zeros(batch_size, n, dtype=A_contig.dtype, device=A_contig.device)
panel_size = 32
for j in range(0, n, panel_size):
current_b = min(panel_size, n - j)
BLOCK_M = triton.next_power_of_2(m - j)
num_warps = 8 if BLOCK_M >= 2048 else 4
fused_panel_kernel[(batch_size,)](
H, tau,
H.stride(0), H.stride(1), H.stride(2),
tau.stride(0), tau.stride(1),
m, j, current_b,
BLOCK_M=BLOCK_M,
num_warps=num_warps
)
if j + current_b < n:
V_raw = H[:, j:, j:j+current_b]
V = torch.tril(V_raw, diagonal=-1)
idx = torch.arange(current_b, device=A_contig.device)
V[:, idx, idx] = 1.0
tau_b = tau[:, j:j+current_b]
VtV = torch.bmm(V.transpose(1, 2), V)
U = torch.triu(VtV, diagonal=1)
M_mat = U * tau_b.unsqueeze(1)
M_mat[:, idx, idx] = 1.0
D = torch.diag_embed(tau_b)
T_T = torch.linalg.solve_triangular(M_mat.transpose(1, 2), D, upper=False)
T = T_T.transpose(1, 2)
H_trail = H[:, j:, j+current_b:]
vt_H = torch.bmm(V.transpose(1, 2), H_trail)
Tt_vt_H = torch.bmm(T.transpose(1, 2), vt_H)
update = torch.bmm(V, Tt_vt_H)
H[:, j:, j+current_b:] -= update
return H.contiguous(), tau
# Export endpoints for the evaluator
factorize = run
factorization = run
qr = run
batched_qr = run
custom_kernel = run
__call__ = runscrolls · 139 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