submission 809781
Justin Arndt · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 278 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-809781?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:b0e739d212aba0a91e901d096c6702d55c9e4f016d05baabd83a5b532e14158b
license declaredunknown
license concludedunknown
authorsJustin Arndt
imported2026-08-26
Kernel source
submission.py278 lines
"""
Multi-Strategy Batched Householder QR Factorization
====================================================
Matches torch.geqrf convention exactly:
H = upper triangle R + lower triangle Householder vectors (v[0]=1 implicit)
tau = reflector coefficients
Reflector: Q_k = I - tau_k * v_k * v_k^T
Strategies:
A. Small n (<=128) or tiny batch (<=2): Direct cuSOLVER via torch.geqrf
B. Medium n (176-512), large batch: Blocked Householder QR, panel_width=32
C. Large n (1024+), moderate batch: Blocked Householder QR, panel_width=64
The blocked algorithm:
1. Panel factorization: column-by-column with branchless Householder (DLARFG)
2. WY representation: accumulate T matrix for block reflector I - V*T*V^T
3. Trailing update: 3 batched GEMMs via torch.bmm (FP32 CUDA cores)
Key design decisions:
- TF32 tensor cores DISABLED: QR error O(n*eps_TF32) exceeds tolerance 20*n*eps_FP32
- Branchless tau/v computation: torch.where avoids warp divergence on ill-conditioned columns
- V^T V pre-computed once per panel to reduce T-matrix kernel launches
- Safety fallback to torch.geqrf if custom path produces NaN/Inf
"""
import torch
from task import input_t, output_t
def custom_kernel(data: input_t) -> output_t:
"""Batched Householder QR factorization with strategy routing."""
batch, n, _ = data.shape
# ── Strategy A: cuSOLVER fast path ──────────────────────────────
# Profiling shows torch.geqrf (cuSOLVER/MAGMA) is faster for:
# - Small n, small batch, large n with small batch
# - batch=40 n=176/352, batch=60 n=1024, batch=8 n=2048
# cuSOLVER loops over batch internally; Python overhead of custom
# blocked QR only pays off when batch is very large (>=200).
# ── Strategy B: Blocked Householder QR ──────────────────────────
# For LARGE batch (>=200) with MODERATE n (256-768), custom blocked
# QR with batched GEMM trailing updates is 1.85x faster than cuSOLVER.
# The large batch makes trailing GEMMs dominate over panel loop overhead.
# Profiled: B640-N512 custom=510ms vs geqrf=942ms (1.85x speedup)
if batch >= 200 and 256 <= n <= 768:
H, tau = _blocked_householder_qr(data)
# Safety: fall back to cuSOLVER if NaN/Inf detected
if torch.isfinite(H).all() and torch.isfinite(tau).all():
return H, tau
return torch.geqrf(data)
# ── Default: cuSOLVER for everything else ───────────────────────
return torch.geqrf(data)
# ====================================================================
# Block size selection (inspired by recursive panel scaling pattern)
# ====================================================================
def _select_block_size(n: int, batch: int) -> int:
"""Select panel width based on matrix size.
Smaller panels = less panel overhead, more trailing GEMM steps.
Larger panels = more panel overhead, fewer (larger) trailing GEMMs.
For B200 with 192KB+ shared memory, larger panels are viable.
"""
if n <= 256:
return 32
elif n <= 512:
# For critical batch=640, n=512 shape: nb=32 gives good trailing GEMM sizes
return 32 if batch >= 64 else 64
elif n <= 1024:
return 64
else:
return 64
# ====================================================================
# Blocked Householder QR
# ====================================================================
def _blocked_householder_qr(A: torch.Tensor):
"""Blocked Householder QR with batched GEMM trailing updates.
Panel factorization is column-by-column (unblocked) with branchless
Householder reflector computation matching LAPACK DLARFG convention.
Trailing matrix updated via WY block reflector (I - V*T*V^T) using
three large batched GEMMs that maximize GPU utilization.
"""
batch, n, _ = A.shape
device = A.device
dtype = A.dtype
# CRITICAL: Disable TF32 for matmul precision.
# TF32 mantissa is 10 bits (eps ~ 5e-4), QR tolerance is 20*n*eps32 ~ 20*n*1.2e-7.
# For n=512: TF32 error ~ O(n * 5e-4) = 0.25 >> tolerance ~ 1.2e-3.
prev_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = False
try:
nb = _select_block_size(n, batch)
H = A.clone()
tau = torch.zeros(batch, n, device=device, dtype=dtype)
for j in range(0, n, nb):
jb = min(nb, n - j)
# 1. Panel factorization: column-by-column Householder reflectors
_panel_factorize(H, tau, j, jb, batch, n)
# 2. Trailing update via WY representation + batched GEMM
if j + jb < n:
_trailing_update(H, tau, j, jb, batch, n)
return H, tau
finally:
torch.backends.cuda.matmul.allow_tf32 = prev_tf32
# ====================================================================
# Panel Factorization (LAPACK DLARFG convention)
# ====================================================================
def _panel_factorize(H, tau, j, jb, batch, n):
"""Column-by-column Householder panel factorization.
For each column k in the panel [j, j+jb):
1. Compute Householder reflector (branchless, matching DLARFG)
2. Apply rank-1 update to remaining panel columns
3. Store reflector vector below diagonal and tau coefficient
LAPACK DLARFG convention:
beta = -sign(x[0]) * ||x|| (becomes R[k,k])
tau = (beta - x[0]) / beta (reflector coefficient, in [1,2])
v = x / (x[0] - beta), v[0] = 1 (Householder vector)
When subdiagonal ||x[1:]|| == 0: tau = 0, no reflection.
"""
for k in range(jb):
col = j + k
m = n - col # subcolumn length from diagonal down
if m <= 1:
# Scalar: no reflection possible
tau[:, col] = 0.0
continue
# ── Extract column data ──────────────────────────────────
x0 = H[:, col, col].clone() # (batch,) diagonal element
x_tail = H[:, col + 1:, col] # (batch, m-1) view below diagonal
# ── Compute norms ────────────────────────────────────────
norm_tail_sq = (x_tail * x_tail).sum(dim=-1) # (batch,)
norm_x = torch.sqrt(x0 * x0 + norm_tail_sq) # (batch,)
# ── Branchless sign: sign(0) = 1 ─────────────────────────
s = x0.sign()
s = torch.where(s == 0, torch.ones_like(s), s)
# ── beta = -sign(x0) * ||x|| ─────────────────────────────
beta = -s * norm_x
# ── Branchless tau and v computation ──────────────────────
# Only reflect when subdiagonal is nonzero (branchless via torch.where)
needs_ref = norm_tail_sq > 0 # (batch,) mask
# tau = (beta - x0) / beta, safe for beta=0
safe_beta = torch.where(beta.abs() > 0, beta, torch.ones_like(beta))
tau_k = torch.where(needs_ref,
(beta - x0) / safe_beta,
torch.zeros_like(x0))
# v_below = x_tail / (x0 - beta)
# denom = x0 - beta = x0 + sign(x0)*||x||, magnitude >= ||x||, always safe
denom = x0 - beta
safe_denom = torch.where(needs_ref, denom, torch.ones_like(denom))
v_below = x_tail / safe_denom.unsqueeze(-1)
v_below = torch.where(needs_ref.unsqueeze(-1), v_below,
torch.zeros_like(v_below))
# ── Store results in H ────────────────────────────────────
H[:, col, col] = torch.where(needs_ref, beta, x0)
H[:, col + 1:, col] = v_below
tau[:, col] = tau_k
# ── Apply reflector to remaining panel columns ────────────
# H_k = I - tau_k * v * v^T, where v = [1, v_below]^T
# rem = H[:, col:, col+1:j+jb]
# rem -= tau_k * v * (v^T @ rem)
if k + 1 < jb:
rem = H[:, col:, col + 1:j + jb] # (batch, m, jb-k-1) view
# w = v^T @ rem = rem[0,:] + v_below^T @ rem[1:,:]
# Using bmm for the v_below^T @ rem[1:,:] part
w = rem[:, 0:1, :].clone() # (batch, 1, jb-k-1)
w = w + torch.bmm(
v_below.unsqueeze(1), # (batch, 1, m-1)
rem[:, 1:, :] # (batch, m-1, jb-k-1)
)
# rank-1 update: rem -= tau_k * v * w
tw = tau_k.view(batch, 1, 1) * w # (batch, 1, jb-k-1)
rem[:, 0:1, :] -= tw # row 0: -= tau_k * 1 * w
rem[:, 1:, :] -= v_below.unsqueeze(-1) * tw # rows 1+: -= tau_k * v_below * w
# ====================================================================
# Trailing Matrix Update (WY Block Reflector)
# ====================================================================
def _trailing_update(H, tau, j, jb, batch, n):
"""Apply block reflector to trailing matrix using WY representation.
Block reflector: P = I - V * T * V^T
where V is (batch, m, jb) unit lower triangular (Householder vectors)
and T is (batch, jb, jb) upper triangular (WY coefficients)
Trailing update via 3 batched GEMMs:
W = V^T @ A_trail (batch, jb, trailing_cols)
W = T @ W (batch, jb, trailing_cols)
A_trail -= V @ W (batch, m, trailing_cols)
"""
m = n - j
trailing_cols = n - j - jb
device = H.device
dtype = H.dtype
if trailing_cols <= 0:
return
# ── Build V: unit lower triangular from stored Householder vectors ──
V = H[:, j:, j:j + jb].clone() # (batch, m, jb) contiguous copy
V.tril_(diagonal=-1) # zero upper triangle (remove R entries)
diag_idx = torch.arange(min(m, jb), device=device)
V[:, diag_idx, diag_idx] = 1.0 # set unit diagonal (v[0]=1)
# ── Build T: upper triangular WY matrix via DLARFT recurrence ───────
T = _build_T_matrix(V, tau[:, j:j + jb], jb, batch, device, dtype)
# ── Trailing update: A = Q^T @ A where Q = I - V*T*V^T ────────────
# Q^T = I - V * T^T * V^T (transpose T for the adjoint)
# These are the large GEMMs that benefit from batch parallelism
W = torch.bmm(V.transpose(1, 2), H[:, j:, j + jb:]) # (batch, jb, trailing)
W = torch.bmm(T.transpose(1, 2), W) # (batch, jb, trailing) T^T!
H[:, j:, j + jb:] -= torch.bmm(V, W) # (batch, m, trailing)
# ====================================================================
# T Matrix Construction (LAPACK DLARFT)
# ====================================================================
def _build_T_matrix(V, tau_panel, jb, batch, device, dtype):
"""Build upper triangular T matrix for WY block reflector.
Recurrence from LAPACK DLARFT:
T[k,k] = tau[k]
T[0:k, k] = -tau[k] * T[0:k, 0:k] @ (V^T V)[0:k, k]
Pre-computes V^T V as a single batched GEMM to reduce kernel launches
(1 large GEMM instead of jb small GEMMs).
"""
# Single GEMM for all inner products between Householder vectors
VTV = torch.bmm(V.transpose(1, 2), V) # (batch, jb, jb)
T = torch.zeros(batch, jb, jb, device=device, dtype=dtype)
T[:, 0, 0] = tau_panel[:, 0]
for k in range(1, jb):
tk = tau_panel[:, k] # (batch,)
z = VTV[:, :k, k:k + 1] # (batch, k, 1) pre-computed inner products
Tz = torch.bmm(T[:, :k, :k], z) # (batch, k, 1) triangular matvec
T[:, :k, k:k + 1] = -tk.view(batch, 1, 1) * Tz
T[:, k, k] = tk
return T
scrolls · 278 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