submission 865080
floatingswitch_50642 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 430 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-865080?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:bee5ec9ae779a85a7efcebd70240824c6f043e1ea545af20cf28ceae6b830ea4
license declaredunknown
license concludedunknown
authorsfloatingswitch_50642
imported2026-08-26
Kernel source
submission.py430 lines
from typing import Tuple
import torch
import triton
import triton.language as tl
try:
from task import input_t, output_t
except ImportError:
input_t = torch.Tensor
output_t = Tuple[torch.Tensor, torch.Tensor]
# ---------------------------------------------------------------------------
# Custom batched symmetric eigensolver.
#
# torch.linalg.eigh on CUDA loops over the batch dimension one matrix at a
# time (each iteration syncing device<->host), which dominates runtime for
# large batches. This module implements the classical pipeline --
# 1) blocked Householder tridiagonalization (WY / compact-WY updates,
# touching the full trailing matrix once per panel instead of once per
# column -- avoids the memory-bandwidth blowup of a naive per-column
# reduction)
# 2) bisection eigenvalues on the tridiagonal form, via a Triton kernel
# that fuses the whole O(n) Sturm-count recurrence into one kernel
# launch per bisection iteration (pure-PyTorch/eager sequential ops
# here cost ~50x more due to per-launch overhead across ~55*n steps)
# 3) inverse iteration for eigenvectors, using a Triton-fused Thomas solve
# for the same reason, with per-iteration QR re-orthogonalization
# (needed for repeated/clustered eigenvalues) and every-other-iteration
# QR frequency (torch.linalg.qr itself loops per-matrix for large
# batches, so calling it every iteration would reintroduce the exact
# bottleneck this is trying to avoid)
# 4) blocked Householder back-transformation to recover eigenvectors of
# the original (dense) matrix
# -- entirely with batched PyTorch/Triton ops, so the whole batch is
# processed together instead of looped.
#
# This wins decisively for large batches at moderate n (the n=512, batch=640
# shapes this task is centered on) but is NOT worth it for small batches or
# very large n, where per-call/per-panel overhead isn't amortized and
# torch.linalg.eigh's vendor-library path is already competitive or better.
# custom_kernel therefore dispatches based on shape, with a safe fallback to
# torch.linalg.eigh on any error.
# ---------------------------------------------------------------------------
# ===================== Triton: fused Sturm-count (bisection) =====================
@triton.jit
def _sturm_count_kernel(
diag_ptr, offdiag_ptr, mid_ptr, count_ptr,
n, n_eig,
stride_diag_b, stride_offdiag_b, stride_mid_b, stride_mid_e, stride_count_b, stride_count_e,
BLOCK_E: tl.constexpr,
):
b = tl.program_id(0)
e_block = tl.program_id(1)
e_offs = e_block * BLOCK_E + tl.arange(0, BLOCK_E)
e_mask = e_offs < n_eig
mid = tl.load(mid_ptr + b * stride_mid_b + e_offs * stride_mid_e, mask=e_mask, other=0.0)
tiny = 1e-30
d0 = tl.load(diag_ptr + b * stride_diag_b + 0)
q = d0 - mid
q = tl.where(tl.abs(q) < tiny, -tiny, q)
count = tl.where(q < 0, 1.0, 0.0)
for i in range(1, n):
di = tl.load(diag_ptr + b * stride_diag_b + i)
ei_1 = tl.load(offdiag_ptr + b * stride_offdiag_b + (i - 1))
q = (di - mid) - (ei_1 * ei_1) / q
q = tl.where(tl.abs(q) < tiny, -tiny, q)
count = count + tl.where(q < 0, 1.0, 0.0)
tl.store(count_ptr + b * stride_count_b + e_offs * stride_count_e, count, mask=e_mask)
def _sturm_count_triton(diag, offdiag, mid, BLOCK_E: int = 128):
batch, n = diag.shape
n_eig = mid.shape[-1]
count = torch.empty(batch, n_eig, device=diag.device, dtype=diag.dtype)
diag = diag.contiguous()
offdiag = offdiag.contiguous()
mid = mid.contiguous()
grid = (batch, triton.cdiv(n_eig, BLOCK_E))
_sturm_count_kernel[grid](
diag, offdiag, mid, count, n, n_eig,
diag.stride(0), offdiag.stride(0), mid.stride(0), mid.stride(1),
count.stride(0), count.stride(1), BLOCK_E=BLOCK_E,
)
return count
def _gershgorin_bounds(diag, offdiag):
batch, n = diag.shape
od_pad = torch.nn.functional.pad(offdiag.abs(), (1, 1))
radius = od_pad[:, :-1] + od_pad[:, 1:]
lo = (diag - radius).min(dim=-1).values
hi = (diag + radius).max(dim=-1).values
return lo, hi
def _bisection_eigvals(diag, offdiag, iters: int = 55):
batch, n = diag.shape
device, dtype = diag.device, diag.dtype
lo0, hi0 = _gershgorin_bounds(diag, offdiag)
pad = (hi0 - lo0).clamp_min(1.0) * 1e-3
lo0 = lo0 - pad
hi0 = hi0 + pad
lo = lo0.unsqueeze(-1).expand(batch, n).contiguous()
hi = hi0.unsqueeze(-1).expand(batch, n).contiguous()
ranks = torch.arange(n, device=device, dtype=dtype).unsqueeze(0)
for _ in range(iters):
mid = 0.5 * (lo + hi)
cnt = _sturm_count_triton(diag, offdiag, mid)
go_hi = cnt > ranks
hi = torch.where(go_hi, mid, hi)
lo = torch.where(go_hi, lo, mid)
return 0.5 * (lo + hi)
# ===================== Triton: fused Thomas solve (inverse iteration) =====================
@triton.jit
def _thomas_kernel(
diag_shift_ptr, offdiag_ptr, rhs_ptr, out_ptr,
cprime_ptr, dprime_ptr,
n, n_eig,
stride_ds_b, stride_ds_n, stride_ds_e,
stride_od_b,
stride_rhs_b, stride_rhs_n, stride_rhs_e,
stride_out_b, stride_out_n, stride_out_e,
stride_scr_b, stride_scr_n, stride_scr_e,
BLOCK_E: tl.constexpr,
):
b = tl.program_id(0)
e_block = tl.program_id(1)
e_offs = e_block * BLOCK_E + tl.arange(0, BLOCK_E)
e_mask = e_offs < n_eig
tiny = 1e-30
ds_base = diag_shift_ptr + b * stride_ds_b + e_offs * stride_ds_e
rhs_base = rhs_ptr + b * stride_rhs_b + e_offs * stride_rhs_e
out_base = out_ptr + b * stride_out_b + e_offs * stride_out_e
cs_base = cprime_ptr + b * stride_scr_b + e_offs * stride_scr_e
ds2_base = dprime_ptr + b * stride_scr_b + e_offs * stride_scr_e
od_base = offdiag_ptr + b * stride_od_b
d0 = tl.load(ds_base + 0 * stride_ds_n, mask=e_mask, other=1.0)
d0 = tl.where(tl.abs(d0) < tiny, tiny, d0)
r0 = tl.load(rhs_base + 0 * stride_rhs_n, mask=e_mask, other=0.0)
dprime_prev = r0 / d0
e0 = tl.load(od_base + 0)
cprime_prev = e0 / d0
tl.store(cs_base + 0 * stride_scr_n, cprime_prev, mask=e_mask)
tl.store(ds2_base + 0 * stride_scr_n, dprime_prev, mask=e_mask)
for i in range(1, n):
e_im1 = tl.load(od_base + (i - 1))
di = tl.load(ds_base + i * stride_ds_n, mask=e_mask, other=1.0)
ri = tl.load(rhs_base + i * stride_rhs_n, mask=e_mask, other=0.0)
denom = di - e_im1 * cprime_prev
denom = tl.where(tl.abs(denom) < tiny, tiny, denom)
dprime_prev = (ri - e_im1 * dprime_prev) / denom
if i < n - 1:
e_i = tl.load(od_base + i)
cprime_prev = e_i / denom
tl.store(cs_base + i * stride_scr_n, cprime_prev, mask=e_mask)
tl.store(ds2_base + i * stride_scr_n, dprime_prev, mask=e_mask)
y_next = dprime_prev
tl.store(out_base + (n - 1) * stride_out_n, y_next, mask=e_mask)
for ii in range(1, n):
i = n - 1 - ii
c_i = tl.load(cs_base + i * stride_scr_n, mask=e_mask, other=0.0)
d_i = tl.load(ds2_base + i * stride_scr_n, mask=e_mask, other=0.0)
y_i = d_i - c_i * y_next
tl.store(out_base + i * stride_out_n, y_i, mask=e_mask)
y_next = y_i
def _thomas_solve_triton(diag_shift, offdiag, rhs, BLOCK_E: int = 64):
batch, n, n_eig = rhs.shape
out = torch.empty_like(rhs)
diag_shift = diag_shift.contiguous()
offdiag = offdiag.contiguous()
rhs = rhs.contiguous()
cprime = torch.empty(batch, n, n_eig, device=rhs.device, dtype=rhs.dtype)
dprime = torch.empty(batch, n, n_eig, device=rhs.device, dtype=rhs.dtype)
grid = (batch, triton.cdiv(n_eig, BLOCK_E))
_thomas_kernel[grid](
diag_shift, offdiag, rhs, out, cprime, dprime, n, n_eig,
diag_shift.stride(0), diag_shift.stride(1), diag_shift.stride(2),
offdiag.stride(0),
rhs.stride(0), rhs.stride(1), rhs.stride(2),
out.stride(0), out.stride(1), out.stride(2),
cprime.stride(0), cprime.stride(1), cprime.stride(2),
BLOCK_E=BLOCK_E,
)
return out
def _inverse_iteration(diag, offdiag, eigvals, iters: int = 6):
batch, n = diag.shape
n_eig = eigvals.shape[-1]
device, dtype = diag.device, diag.dtype
mat_scale = diag.abs().amax(dim=-1, keepdim=True).clamp_min(1.0)
big = mat_scale.expand(batch, n_eig) * 1e6
gap_below = torch.cat([big[:, :1], eigvals[:, 1:] - eigvals[:, :-1]], dim=-1)
gap_above = torch.cat([eigvals[:, 1:] - eigvals[:, :-1], big[:, :1]], dim=-1)
local_gap = torch.minimum(gap_below, gap_above)
floor = torch.maximum(eigvals.abs(), mat_scale * 1e-10) * 1e-7
perturb = torch.maximum(local_gap * 0.1, floor)
lam = eigvals - perturb
diag_shift = diag.unsqueeze(-1) - lam.unsqueeze(1)
gen = torch.Generator(device=device)
gen.manual_seed(0)
gen_mat = torch.randn(n, n_eig, device=device, dtype=dtype, generator=gen)
y = gen_mat.unsqueeze(0).expand(batch, n, n_eig).contiguous()
for it in range(iters):
y = _thomas_solve_triton(diag_shift, offdiag, y)
norm = torch.linalg.vector_norm(y, dim=1, keepdim=True).clamp_min(1e-30)
y = y / norm
if (it + 1) % 2 == 0 or it == iters - 1:
y, _ = torch.linalg.qr(y)
return y
# ===================== Blocked Householder tridiagonalization =====================
def _tridiagonalize_blocked(A: torch.Tensor, nb: int = 32):
batch, n, _ = A.shape
device, dtype = A.device, A.dtype
nsteps = max(0, n - 2)
V_all = torch.zeros(batch, nsteps, n, device=device, dtype=dtype)
beta_all = torch.zeros(batch, nsteps, device=device, dtype=dtype)
diag_list = []
offdiag_list = []
eps = torch.finfo(dtype).tiny
Ta = A.clone()
k0 = 0
while k0 < nsteps:
m = n - k0
cur_nb = min(nb, nsteps - k0)
row_idx = torch.arange(m, device=device)
Vp = torch.zeros(batch, m, cur_nb, device=device, dtype=dtype)
Wp = torch.zeros(batch, m, cur_nb, device=device, dtype=dtype)
for j in range(cur_nb):
col = Ta[:, :, j].clone()
if j > 0:
vj_row = Vp[:, j, :j]
wj_row = Wp[:, j, :j]
corr = torch.einsum('bmj,bj->bm', Vp[:, :, :j], wj_row) \
+ torch.einsum('bmj,bj->bm', Wp[:, :, :j], vj_row)
col = col - corr
diag_list.append(col[:, j].clone())
mask = (row_idx > j).to(dtype)
x = col * mask
norm_x = torch.linalg.vector_norm(x, dim=-1)
x0 = col[:, j + 1]
sign_val = torch.where(x0 >= 0, 1.0, -1.0).to(dtype)
alpha = -sign_val * norm_x
offdiag_list.append(alpha.clone())
v = x.clone()
v[:, j + 1] = x0 - alpha
vtv = (v * v).sum(-1)
beta = 2.0 / vtv.clamp_min(eps)
Tav = torch.einsum('bij,bj->bi', Ta, v)
p_raw = Tav
if j > 0:
VtV = torch.einsum('bmj,bm->bj', Vp[:, :, :j], v)
WtV = torch.einsum('bmj,bm->bj', Wp[:, :, :j], v)
corr_p = torch.einsum('bmj,bj->bm', Vp[:, :, :j], WtV) \
+ torch.einsum('bmj,bj->bm', Wp[:, :, :j], VtV)
p_raw = p_raw - corr_p
p = beta.unsqueeze(-1) * p_raw
vp = (v * p).sum(-1, keepdim=True)
w = p - 0.5 * beta.unsqueeze(-1) * vp * v
Vp[:, :, j] = v
Wp[:, :, j] = w
k = k0 + j
V_all[:, k, k0:] = v
beta_all[:, k] = beta
Ta = Ta - torch.bmm(Vp, Wp.transpose(-1, -2)) - torch.bmm(Wp, Vp.transpose(-1, -2))
Ta = Ta[:, cur_nb:, cur_nb:].clone()
k0 += cur_nb
rem = Ta.shape[-1]
for i in range(rem):
diag_list.append(Ta[:, i, i].clone())
if rem > 1:
offdiag_list.append(Ta[:, 0, 1].clone())
diag = torch.stack(diag_list, dim=-1)
offdiag = torch.stack(offdiag_list, dim=-1) if offdiag_list else torch.zeros(batch, 0, device=device, dtype=dtype)
return diag, offdiag, V_all, beta_all
# ===================== Blocked Householder back-transformation =====================
def _apply_householders_blocked(V, beta, Y, nb: int = 32):
batch, nsteps, n = V.shape
device, dtype = V.device, V.dtype
k = nsteps - 1
while k >= 0:
cur_nb = min(nb, k + 1)
lo = k - cur_nb + 1
idx = torch.arange(lo, k + 1, device=device)
Vp = V[:, idx, :].transpose(-1, -2)
bp = beta[:, idx]
T = torch.zeros(batch, cur_nb, cur_nb, device=device, dtype=dtype)
T[:, 0, 0] = bp[:, 0]
for j in range(1, cur_nb):
vj = Vp[:, :, j]
Vprev = Vp[:, :, :j]
w = torch.einsum('bnj,bn->bj', Vprev, vj)
Tprev = T[:, :j, :j]
col = -bp[:, j:j + 1] * torch.einsum('bij,bj->bi', Tprev, w)
T[:, :j, j] = col
T[:, j, j] = bp[:, j]
VtY = torch.einsum('bnj,bnm->bjm', Vp, Y)
TVtY = torch.einsum('bij,bjm->bim', T, VtY)
Y = Y - torch.einsum('bnj,bjm->bnm', Vp, TVtY)
k = lo - 1
return Y
# ===================== Full pipeline + dispatch =====================
def _eigh_custom(A: torch.Tensor, nb: int = 32, bisect_iters: int = 55, ii_iters: int = 6):
diag, offdiag, V, beta = _tridiagonalize_blocked(A, nb=nb)
L = _bisection_eigvals(diag, offdiag, iters=bisect_iters)
Y = _inverse_iteration(diag, offdiag, L, iters=ii_iters)
Q = _apply_householders_blocked(V, beta, Y, nb=nb)
return Q, L
def _verify_and_patch(A: torch.Tensor, Q: torch.Tensor, L: torch.Tensor, safety: float = 0.5):
# The heuristic pipeline (random-start inverse iteration + periodic QR)
# occasionally fails to resolve a handful of near-degenerate eigenvalue
# clusters (gaps below fp32 precision) -- rare (~1 in several hundred
# matrices) but a real correctness risk since it depends on the specific
# random orthogonal transform of each matrix, not just its case type.
# Rather than chase individual hyperparameters (which just moves the
# failure to a different case, as observed empirically), verify every
# row's residual in fp64 and recompute only the rows that fail via
# torch.linalg.eigh -- guaranteed correct, and cheap since failures are
# rare so the per-matrix-loop fallback only ever touches a few rows.
batch, n, _ = A.shape
device = A.device
tol = (5e-5 * (n ** 0.5) + 1e-6) * safety
budget_bytes = 24 * 1024 * 1024
chunk = max(1, min(batch, budget_bytes // (n * n * 8) + 1))
I = torch.eye(n, device=device, dtype=torch.float64)
bad_chunks = []
for s in range(0, batch, chunk):
e = min(batch, s + chunk)
A64 = A[s:e].to(torch.float64)
Q64 = Q[s:e].to(torch.float64)
L64 = L[s:e].to(torch.float64)
A_l1 = A64.abs().sum(dim=(-2, -1)).clamp_min(1e-30)
AQ = A64 @ Q64
QL = Q64 * L64.unsqueeze(-2)
eig_rel = (AQ - QL).abs().sum(dim=(-2, -1)) / A_l1
recon = Q64 @ (L64.unsqueeze(-1) * Q64.transpose(-1, -2))
recon_rel = (recon - A64).abs().sum(dim=(-2, -1)) / A_l1
orth_rel = (Q64.transpose(-1, -2) @ Q64 - I).abs().sum(dim=(-2, -1)) / n
bad_chunks.append((eig_rel > tol) | (recon_rel > tol) | (orth_rel > tol))
bad_mask = torch.cat(bad_chunks)
if bool(bad_mask.any()):
idx = bad_mask.nonzero().flatten()
vals, vecs = torch.linalg.eigh(A[idx])
Q = Q.clone()
L = L.clone()
Q[idx] = vecs
L[idx] = vals
return Q, L
def custom_kernel(data: input_t) -> output_t:
A = data
# The custom batched pipeline (tridiagonalize -> bisect -> inverse
# iterate -> back-transform) was designed to beat torch.linalg.eigh's
# per-matrix loop for large batches at moderate n. On the actual grading
# GPU it loses badly instead (~4s vs ~0.17s at n=512,batch=640): the
# panel-construction stages of the tridiagonalization/back-transform are
# inherently O(n) sequential Python-level tensor ops (~500+ small kernel
# launches per stage, independent of block size nb), which is dominated
# by host-dispatch/launch overhead on fast hardware where the GPU
# compute itself is nearly instant. That overhead barely mattered on a
# slower local GPU (compute-bound there) but dominates on faster
# hardware (launch-bound there), causing a net regression instead of a
# win. Disabled pending a rewrite that fuses the per-column loop into a
# single internally-looping Triton kernel; always defer to eigh for now.
values, vectors = torch.linalg.eigh(data)
return vectors, values
scrolls · 430 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