submission 811076
jeeva2812 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 296 lines, June 9 Researcher Reciprocity License v1.0.
submission_v2_optimized.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-811076?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:8fe99e4505d617301e0c963b68aa1f06aaf917d1f44a8299bea6548defe1ff72
license declaredunknown
license concludedunknown
authorsjeeva2812
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4 if block_n <= 64 else 8,tile-n = 256
2. **Triton fused QR only viable for n ≤ 256** (BLOCK_N = 256 sits at theKernel source
submission_v2_optimized.py296 lines
"""
submission_v2_optimized.py — best-per-shape dispatch in a single file.
Synthesizes what the per-shape data on the new (updated) leaderboard shows:
shape best path measured µs (separate runs)
------------------------ ------------------- ---------------------------
n = 32 b = 20 Triton fused QR 61 µs (was 369 µs geqrf)
n = 176 b = 40 Triton fused QR 2,160 µs (was 21,600 µs geqrf)
n = 352 b = 40 blocked panel geqrf 28,100 µs (was 50,000 µs geqrf)
n = 512 b = 640 blocked panel geqrf 611,000 µs (was 1,071,000 µs geqrf)
n = 1024 b = 60 CholeskyQR3+Yamamoto 206,000 µs (vs 224,000 µs blocked)
n = 2048 b = 8 CholeskyQR3+Yamamoto 137,000 µs (vs 139,000 µs blocked)
n = 4096 b = 2 CholeskyQR3+Yamamoto 185,000 µs (vs 186,000 µs blocked)
Two key learnings driving the dispatch:
1. **Panel-geqrf beats full-matrix-geqrf at n ∈ [352, 768]**. cuSOLVER's
`geqrfBatched` is serialized at large n; calling it on small (rows × 32)
panel blocks exposes more parallelism. v1's "slow" blocked path is
actually the fastest cuBLAS/cuSOLVER variant for medium n.
2. **Triton fused QR only viable for n ≤ 256** (BLOCK_N = 256 sits at the
B200 SMEM/CTA ceiling at 256 KB FP32). For n = 352 we'd need BLOCK_N =
512 = 1 MB, which spills. So panel-geqrf carries the middle.
Routing
-------
n ≤ 256 → Triton fused QR (one CTA per matrix)
256 < n ≤ 768 → blocked panel geqrf + compact-WY trailing
n > 768 → CholeskyQR3 (FP32) + Yamamoto LU (FP64) + per-matrix
geqrf fallback for gate failures
Upper-triangular fast-path applies to all branches.
Fallback / safety
-----------------
- Triton kernel guarded by env flag (QR_DISABLE_TRITON=1) and try/except.
- Yamamoto LU gates (chol info, lu info, finite, tau range, L growth)
catch instability and route those indices to geqrf.
- All output FP32, all per-matrix shapes guaranteed.
"""
import os
import torch
from task import input_t, output_t
# =========================================================================
# Triton fused QR for small n (n ≤ 256)
# =========================================================================
_HAVE_TRITON = False
if os.environ.get("QR_DISABLE_TRITON") != "1":
try:
import triton
import triton.language as tl
_HAVE_TRITON = True
except Exception:
pass
if _HAVE_TRITON:
@triton.jit
def _qr_fused_kernel(
A_ptr, Tau_ptr,
stride_ab, stride_an, stride_am,
N: tl.constexpr,
BLOCK_N: tl.constexpr,
):
bid = tl.program_id(0)
row_off = tl.arange(0, BLOCK_N)
col_off = tl.arange(0, BLOCK_N)
row_mask = row_off < N
col_mask = col_off < N
valid = row_mask[:, None] & col_mask[None, :]
A_ptrs = (
A_ptr
+ bid * stride_ab
+ row_off[:, None] * stride_an
+ col_off[None, :] * stride_am
)
A_blk = tl.load(A_ptrs, mask=valid, other=0.0)
tau_vec = tl.zeros([BLOCK_N], dtype=tl.float32)
for j in range(0, N):
rows_below_eq = (row_off >= j) & row_mask
rows_strict_below = (row_off > j) & row_mask
col_eq_j = (col_off == j)
col_j = tl.sum(A_blk * col_eq_j[None, :].to(tl.float32), axis=1)
alpha = tl.sum(col_j * (row_off == j).to(tl.float32))
col_strict = tl.where(rows_strict_below, col_j, 0.0)
sigma = tl.sum(col_strict * col_strict)
norm_sq = sigma + alpha * alpha
need = sigma > 0.0
norm = tl.sqrt(norm_sq)
beta_active = tl.where(alpha >= 0.0, -norm, norm)
beta = tl.where(need, beta_active, alpha)
denom_raw = alpha - beta
denom = tl.where(denom_raw == 0.0, 1.0, denom_raw)
tau_j = tl.where(need, (beta - alpha) / beta, 0.0)
v_below = tl.where(rows_strict_below, col_strict / denom, 0.0)
v = tl.where(row_off == j, 1.0, v_below)
v = tl.where(need, v, 0.0)
cols_below_eq = (col_off >= j) & col_mask
w_full = tl.sum(v[:, None] * A_blk, axis=0)
w = tl.where(cols_below_eq, w_full, 0.0)
A_blk = A_blk - tau_j * v[:, None] * w[None, :]
store_mask = col_eq_j[None, :] & rows_strict_below[:, None]
A_blk = tl.where(store_mask, v[:, None], A_blk)
tau_vec = tl.where(col_off == j, tau_j, tau_vec)
tl.store(A_ptrs, A_blk, mask=valid)
tau_ptrs = Tau_ptr + bid * N + col_off
tl.store(tau_ptrs, tau_vec, mask=col_mask)
def _qr_triton(A: torch.Tensor):
b, n, _ = A.shape
block_n = 1
while block_n < n:
block_n *= 2
H = A.clone().contiguous()
tau = A.new_zeros(b, n)
_qr_fused_kernel[(b,)](
H, tau,
H.stride(0), H.stride(1), H.stride(2),
N=n, BLOCK_N=block_n,
num_warps=4 if block_n <= 64 else 8,
)
return H, tau
# =========================================================================
# Blocked panel-geqrf + compact-WY trailing (the cuSOLVER-aware mid-n path)
# =========================================================================
_PANEL = 32
def _compact_wy(A: torch.Tensor, col: int, panel: int, tau: torch.Tensor):
b, n, _ = A.shape
rows = n - col
Y = A.new_zeros(b, rows, panel)
raw = A[:, col:col + rows, col:col + panel]
lower = torch.ones(rows, panel, dtype=torch.bool, device=A.device).tril(-1)
Y[:, lower] = raw[:, lower]
diag_idx = torch.arange(panel, device=A.device)
Y[:, diag_idx, diag_idx] = 1.0
S = torch.bmm(Y.transpose(1, 2).contiguous(), Y)
T = A.new_zeros(b, panel, panel)
for j in range(panel):
tau_j = tau[:, col + j]
T[:, j, j] = tau_j
if j > 0:
Tz = torch.bmm(T[:, :j, :j], S[:, :j, j:j+1]).squeeze(-1)
T[:, :j, j] = -tau_j.unsqueeze(-1) * Tz
return Y, T
def _qr_blocked_panel(A: torch.Tensor):
b, n, _ = A.shape
tau = A.new_zeros(b, n)
A = A.clone().contiguous()
for col in range(0, n, _PANEL):
panel = min(_PANEL, n - col)
blk = A[:, col:, col:col + panel].contiguous()
H_p, tau_p = torch.geqrf(blk)
A[:, col:, col:col + panel] = H_p
tau[:, col:col + panel] = tau_p
if col + panel >= n:
break
Y, T = _compact_wy(A, col, panel, tau)
trail = A[:, col:, col + panel:].contiguous()
W = torch.bmm(Y.transpose(1, 2).contiguous(), trail)
W = torch.bmm(T.transpose(1, 2).contiguous(), W)
A[:, col:, col + panel:] = trail - torch.bmm(Y, W)
return A, tau
# =========================================================================
# CholeskyQR3 + Yamamoto LU reconstruction (the fast large-n path)
# =========================================================================
def _qr_choleskyqr3_yamamoto(A: torch.Tensor):
b, n, _ = A.shape
G1 = torch.bmm(A.transpose(-2, -1), A)
R1, inf1 = torch.linalg.cholesky_ex(G1, upper=True)
if not (inf1 == 0).all():
bad = inf1 != 0
R1 = R1.clone()
R1[bad] = torch.eye(n, dtype=A.dtype, device=A.device)
Q1 = torch.linalg.solve_triangular(R1, A, upper=True, left=False)
G2 = torch.bmm(Q1.transpose(-2, -1), Q1)
R2, inf2 = torch.linalg.cholesky_ex(G2, upper=True)
if not (inf2 == 0).all():
bad2 = inf2 != 0
R2 = R2.clone()
R2[bad2] = torch.eye(n, dtype=A.dtype, device=A.device)
Q2 = torch.linalg.solve_triangular(R2, Q1, upper=True, left=False)
G3 = torch.bmm(Q2.transpose(-2, -1), Q2)
R3, inf3 = torch.linalg.cholesky_ex(G3, upper=True)
if not (inf3 == 0).all():
bad3 = inf3 != 0
R3 = R3.clone()
R3[bad3] = torch.eye(n, dtype=A.dtype, device=A.device)
Q = torch.linalg.solve_triangular(R3, Q2, upper=True, left=False)
R = torch.bmm(R3, torch.bmm(R2, R1))
ok_chol = (inf1 == 0) & (inf2 == 0) & (inf3 == 0)
diag_Q = torch.diagonal(Q, dim1=-2, dim2=-1)
s = torch.where(diag_Q >= 0,
torch.full_like(diag_Q, -1.0),
torch.full_like(diag_Q, 1.0))
Q_s = Q * s.unsqueeze(-2)
R_s = R * s.unsqueeze(-1)
eye = torch.eye(n, dtype=A.dtype, device=A.device).expand(b, n, n)
M = eye - Q_s
LU64, _, inf_lu = torch.linalg.lu_factor_ex(
M.double(), pivot=False, check_errors=False)
L_strict = torch.tril(LU64, -1).float()
tau = torch.diagonal(LU64, dim1=-2, dim2=-1).float()
H = torch.triu(R_s) + L_strict
ok_lu = inf_lu == 0
ok_fin = torch.isfinite(tau).all(-1) & torch.isfinite(L_strict).flatten(1).all(-1)
ok_tau = (tau > 1e-6).all(-1) & (tau <= 2.001).all(-1)
ok_lgrowth = L_strict.abs().amax(dim=(-2, -1)) < 10.0
ok = ok_chol & ok_lu & ok_fin & ok_tau & ok_lgrowth
if not ok.all():
bad = (~ok).nonzero(as_tuple=True)[0]
# Use blocked-panel path for the fallback at n in [some range, 768]
# but Yamamoto only fires for n > 768, so fallback uses geqrf which is
# fine here (large enough that cuSOLVER per-matrix cost amortizes).
H_fb, tau_fb = torch.geqrf(A.index_select(0, bad).contiguous())
H = H.index_copy(0, bad, H_fb)
tau = tau.index_copy(0, bad, tau_fb)
return H, tau
# =========================================================================
# Dispatch
# =========================================================================
_TRITON_MAX_N = 256 # BLOCK_N=256 at SMEM ceiling on B200
_BLOCKED_PANEL_MAX_N = 768 # above this, CholeskyQR3 wins
def _dispatch(A: torch.Tensor):
b, n, _ = A.shape
# Upper-triangular fast-path
if A.tril(diagonal=-1).abs().max() < 1e-6:
return A.clone(), A.new_zeros(b, n)
if _HAVE_TRITON and n <= _TRITON_MAX_N:
try:
return _qr_triton(A)
except Exception:
pass
if n <= _BLOCKED_PANEL_MAX_N:
return _qr_blocked_panel(A)
return _qr_choleskyqr3_yamamoto(A)
def custom_kernel(data: input_t) -> output_t:
if data.dim() == 2:
H, tau = _dispatch(data.unsqueeze(0))
return H.squeeze(0), tau.squeeze(0)
return _dispatch(data)
scrolls · 296 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