submission 875600
shraderdm · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 387 lines, June 9 Researcher Reciprocity License v1.0.
submission_shredder.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-875600?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:71eece6c81ba7008f9e6e66d8e71477712e7c395164b3a9a9d8f7571fc017f57
license declaredunknown
license concludedunknown
authorsshraderdm
imported2026-08-26
Kernel source
submission_shredder.py387 lines
"""Batched symmetric eigh with a conservative two-cluster (near-involution) fast path.
Dispatcher:
1. DETECT: cheap batched probe (4 GEMV pairs) fits the quadratic annihilating
polynomial A^2 - s*A + p*I ~ 0 per matrix by least squares. A matrix is
accepted only if the fit residual is tiny (rel < 3e-4), the two roots
c-, c+ are real and well separated, and the implied multiplicity
k- = (n*c+ - tr A) / (c+ - c-) is within 0.05 of an integer in [1, n-1].
This is msuiche's "clustered spectra: use the projector already present
in A" structure, generalized to arbitrary cluster centers via the shift
B = (A - mid*I)/half with mid = (c+ + c-)/2, half = (c+ - c-)/2, B^2 ~ I.
2. FAST PATH (only if >= 90% of the batch accepts): randomized range finder
on the small-cluster projector P = (I -/+ B)/2 (CholQR2 + one projector
refinement pass), Householder-free orthogonal complement via randomized
completion + block Gram-Schmidt, eigenvalues by Rayleigh quotient on the
ORIGINAL fp32 A, full sort with matching column permutation.
3. SELF-VERIFY: the actual reference.py gates (eigen / orthogonality /
reconstruction, L1 matrix norms) per matrix at gate/4 margin in fp32.
Failures retry once with fresh randomness, then fall back to cuSOLVER
and scatter back.
4. Everything else (dense, mixed, rankdef, geometric, band, lapack_*, ...)
goes to cusolverDnXsyevBatched, copied verbatim from the working
submission.py binding.
Hard rules honored: no memoization/answer caching,
input is never mutated (fallback clones; fast path only reads), pure torch
(no Triton dependency), custom_kernel(data) -> (Q, L) with Q columns =
orthonormal eigenvectors and L ascending.
"""
import ctypes
import math
import torch
from task import input_t, output_t
# The fast path's Rayleigh quotients and self-verification GEMMs must be true
# fp32 (the checker recomputes residuals in fp64).
torch.backends.cuda.matmul.allow_tf32 = False
# ----------------------------------------------------------------------------
# cuSOLVER XsyevBatched fallback (verbatim binding from submission.py)
# ----------------------------------------------------------------------------
CUSOLVER_EIG_MODE_VECTOR = 1
CUBLAS_FILL_MODE_LOWER = 0
CUDA_R_32F = 0
def _load_cusolver():
"""Prefer the newest cuSOLVER (CUDA 13, libcusolver.so.12 - a measured ~5-10%
faster XsyevBatched on the large shapes) if it is present AND loadable with its
cu13 deps; otherwise fall back to the known-good cu12.8 (.so.11). Any failure in
the cu13 path silently drops to the fallback, so a runner without cu13 keeps the
exact current behavior - zero downside."""
import glob as _glob
_cu13 = sorted(
_glob.glob("/root/cu13redist/libcusolver*/lib/libcusolver.so.12.*") +
_glob.glob("/usr/local/lib/python3*/dist-packages/nvidia/cusolver/lib/libcusolver.so.12*") +
_glob.glob("/usr/local/cuda-13*/targets/x86_64-linux/lib/libcusolver.so.12*") +
_glob.glob("/usr/local/cuda*/lib64/libcusolver.so.12*"))
if _cu13:
try:
for _dep in (
_glob.glob("/root/cu13redist/libnvjitlink*/lib/libnvJitLink.so.13*") +
_glob.glob("/usr/local/lib/python3*/dist-packages/nvidia/nvjitlink/lib/libnvJitLink.so.13*") +
_glob.glob("/root/cu13redist/libcublas*/lib/libcublasLt.so.13*") +
_glob.glob("/usr/local/lib/python3*/dist-packages/nvidia/cublas/lib/libcublasLt.so.13*") +
_glob.glob("/root/cu13redist/libcublas*/lib/libcublas.so.13*") +
_glob.glob("/usr/local/lib/python3*/dist-packages/nvidia/cublas/lib/libcublas.so.13*")):
try:
ctypes.CDLL(_dep, mode=ctypes.RTLD_GLOBAL)
except OSError:
pass
return ctypes.CDLL(_cu13[-1], mode=ctypes.RTLD_GLOBAL)
except OSError:
pass
for _p in ("libcusolver.so.11",
"/usr/local/lib/python3.11/dist-packages/nvidia/cusolver/lib/libcusolver.so.11",
"/usr/local/cuda/lib64/libcusolver.so",
"/usr/local/cuda-12.8/targets/x86_64-linux/lib/libcusolver.so"):
try:
return ctypes.CDLL(_p)
except OSError:
continue
return None
_lib = _load_cusolver()
assert _lib is not None, "libcusolver not found"
_lib.cusolverDnCreate.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
_lib.cusolverDnCreateParams.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
_bs = _lib.cusolverDnXsyevBatched_bufferSize
_bs.argtypes = [ctypes.c_void_p, ctypes.c_void_p, ctypes.c_int, ctypes.c_int, ctypes.c_int64,
ctypes.c_int, ctypes.c_void_p, ctypes.c_int64, ctypes.c_int, ctypes.c_void_p,
ctypes.c_int, ctypes.POINTER(ctypes.c_size_t), ctypes.POINTER(ctypes.c_size_t), ctypes.c_int64]
_ev = _lib.cusolverDnXsyevBatched
_ev.argtypes = [ctypes.c_void_p, ctypes.c_void_p, ctypes.c_int, ctypes.c_int, ctypes.c_int64,
ctypes.c_int, ctypes.c_void_p, ctypes.c_int64, ctypes.c_int, ctypes.c_void_p,
ctypes.c_int, ctypes.c_void_p, ctypes.c_size_t, ctypes.c_void_p, ctypes.c_size_t,
ctypes.c_void_p, ctypes.c_int64]
_handle = ctypes.c_void_p(); _lib.cusolverDnCreate(ctypes.byref(_handle))
_params = ctypes.c_void_p(); _lib.cusolverDnCreateParams(ctypes.byref(_params))
def _cusolver_eigh(A: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
batch, n, _ = A.shape
Aw = A.contiguous().clone() # cuSOLVER overwrites A with eigenvectors; also protects the reused input tensor
W = torch.empty((batch, n), dtype=torch.float32, device=A.device)
info = torch.empty((batch,), dtype=torch.int32, device=A.device)
db = ctypes.c_size_t(0); hb = ctypes.c_size_t(0)
_bs(_handle, _params, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_LOWER, ctypes.c_int64(n),
CUDA_R_32F, ctypes.c_void_p(Aw.data_ptr()), ctypes.c_int64(n), CUDA_R_32F,
ctypes.c_void_p(W.data_ptr()), CUDA_R_32F, ctypes.byref(db), ctypes.byref(hb), ctypes.c_int64(batch))
dev = torch.empty(max(db.value, 1), dtype=torch.uint8, device=A.device)
host = (ctypes.c_char * max(hb.value, 1))()
_ev(_handle, _params, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_LOWER, ctypes.c_int64(n),
CUDA_R_32F, ctypes.c_void_p(Aw.data_ptr()), ctypes.c_int64(n), CUDA_R_32F,
ctypes.c_void_p(W.data_ptr()), CUDA_R_32F, ctypes.c_void_p(dev.data_ptr()), db,
ctypes.cast(host, ctypes.c_void_p), hb, ctypes.c_void_p(info.data_ptr()), ctypes.c_int64(batch))
return Aw.transpose(-1, -2).contiguous(), W
# ----------------------------------------------------------------------------
# Two-cluster detection + fast path
# ----------------------------------------------------------------------------
_FAST_MIN_N = 256 # below this cuSOLVER is already fast; skip probing
_DETECT_PROBES = 4
_DETECT_RES_TOL = 3.0e-4 # rel residual of the quadratic fit (clustered ~2e-5)
_MIN_GAP_REL = 1.0e-3 # clusters must be separated by > 1e-3 * scale
_DET_REL_TOL = 1.0e-6 # LS system must be non-degenerate (rejects A ~ c*I)
_K_INT_TOL = 0.05 # implied multiplicity must be near an integer
_ACCEPT_FRAC = 0.9 # engage fast path only if >= 90% of batch accepts
_VERIFY_MARGIN = 4.0 # self-verify at gate/4
_PROBE_SEED = 0x5EED0001
_BASIS_SEED = 0x5EED0002
_EIGEN_RTOL_FACTOR = 200.0
_RECON_RTOL_FACTOR = 400.0
_ORTH_RTOL_FACTOR = 100.0
def _detect_two_cluster(A: torch.Tensor):
"""Per-matrix probe: does A^2 - s*A + p*I annihilate random vectors?
Returns (accept, cminus, cplus, kminus_float) as device tensors.
"""
batch, n, _ = A.shape
gen = torch.Generator(device=A.device)
gen.manual_seed(_PROBE_SEED)
V = torch.randn((n, _DETECT_PROBES), device=A.device, dtype=torch.float32, generator=gen)
V = V / V.norm(dim=0, keepdim=True).clamp_min(1e-30)
W1 = A @ V # (batch, n, probes)
W2 = A @ W1
# Least squares for W2 ~ s*W1 + t*V (t = -p), stacked over all probes.
a11 = (W1 * W1).sum(dim=(-2, -1))
a12 = (W1 * V).sum(dim=(-2, -1))
a22 = float(_DETECT_PROBES) # V columns are unit norm
b1 = (W1 * W2).sum(dim=(-2, -1))
b2 = (V * W2).sum(dim=(-2, -1))
det = a11 * a22 - a12 * a12
det_ok = det > _DET_REL_TOL * (a11 * a22).clamp_min(1e-30)
det_safe = torch.where(det.abs() > 1e-30, det, torch.full_like(det, 1e-30))
s = (b1 * a22 - b2 * a12) / det_safe
t = (a11 * b2 - a12 * b1) / det_safe
R = W2 - s[:, None, None] * W1 - t[:, None, None] * V
w2n = W2.reshape(batch, -1).norm(dim=-1)
res = R.reshape(batch, -1).norm(dim=-1) / w2n.clamp_min(1e-30)
disc = s * s + 4.0 * t
gap = torch.sqrt(disc.clamp_min(0.0))
cplus = 0.5 * (s + gap)
cminus = 0.5 * (s - gap)
scale = torch.maximum(cplus.abs(), cminus.abs())
tr = A.diagonal(dim1=-2, dim2=-1).sum(dim=-1)
kf = (float(n) * cplus - tr) / gap.clamp_min(1e-30)
kr = torch.round(kf)
accept = (
det_ok
& torch.isfinite(s) & torch.isfinite(t)
& torch.isfinite(res) & (res < _DETECT_RES_TOL)
& (disc > 0) & (scale > 0) & (w2n > 0)
& (gap > _MIN_GAP_REL * scale)
& torch.isfinite(kf) & ((kf - kr).abs() < _K_INT_TOL)
& (kr >= 1.0) & (kr <= float(n - 1))
)
return accept, cminus, cplus, kr
def _cholqr(Y: torch.Tensor, shift: float | None = None):
"""Batched Cholesky-QR. Returns (Q, ok) where ok flags a clean Cholesky.
A shift (shifted CholQR, Fukaya et al.) regularizes the Gram matrix so
the round survives the heavy condition-number tail of square Gaussian
coefficient matrices; each properly-shifted round reduces cond(Y) by
~sqrt(s / (s + ||Y||_2^2)) and follow-up plain rounds restore
eps-level orthogonality. Shifts only right-multiply Y, so the SPAN is
untouched (the projector refinement pass owns span quality).
"""
kk = Y.shape[-1]
Gm = Y.transpose(-1, -2) @ Y
if shift is not None:
Gm = Gm + shift * torch.eye(kk, device=Y.device, dtype=Y.dtype)
Lc, info = torch.linalg.cholesky_ex(Gm)
ok = info == 0
Q = torch.linalg.solve_triangular(Lc.transpose(-1, -2), Y, upper=True, left=False)
return Q, ok
def _shift_value(nn: int, kk: int, norm2sq: float) -> float:
"""Aggressive shifted-CholQR shift: 30 * k * u * ||Y||_2^2, u = eps/2.
Smaller than the Fukaya et al. worst-case shift (11*(n*k + k*(k+1))*u*.),
so each shifted round reduces cond(Y) by ~sqrt(30*k*u) (~40x) instead of
~2.5x, while the shifted Gram's condition number stays ~1/(30*k*u) (~2e3),
far below the fp32 Cholesky breakdown point (~1/(k*u)). Breakdown, if the
gamble ever loses, is flagged by cholesky_ex and handled by retry/rescue.
"""
u = 0.5 * torch.finfo(torch.float32).eps
return 30.0 * kk * u * norm2sq
def _orth_chain(Y: torch.Tensor, norm2sq: float, rounds: tuple[bool, ...]):
"""Orthonormalization chain: True = shifted round, False = plain round.
norm2sq is an a-priori bound on ||Y||_2^2 for the first round; after any
round all singular values are <= ~1, so later shifted rounds use 2.0.
"""
nn, kk = Y.shape[-2], Y.shape[-1]
ok = None
bound = norm2sq
for shifted in rounds:
s = _shift_value(nn, kk, bound) if shifted else None
Y, oki = _cholqr(Y, s)
ok = oki if ok is None else ok & oki
bound = 2.0
return Y, ok
def _fast_two_cluster(Asub: torch.Tensor, cminus: torch.Tensor, cplus: torch.Tensor,
kminus: int, seed_offset: int = 0):
"""Projector-based eigendecomposition for near-two-cluster matrices.
Asub: (m, n, n) fp32, all sharing the same negative-cluster multiplicity.
Returns (Q, L, AQ, ok) with L ascending; ok is a per-matrix health flag
(self-verification still runs on top of it).
"""
m, n, _ = Asub.shape
kplus = n - kminus
ks = min(kminus, kplus)
kb = n - ks
small_is_minus = kminus <= kplus
mid = (0.5 * (cplus + cminus))[:, None, None]
half = (0.5 * (cplus - cminus))[:, None, None]
gen = torch.Generator(device=Asub.device)
gen.manual_seed(_BASIS_SEED + 1315423911 * seed_offset + 131 * n + kminus)
Gall = torch.randn((n, n), device=Asub.device, dtype=torch.float32, generator=gen)
G = Gall[:, :ks]
G2 = Gall[:, ks:]
# Small-cluster basis: Y = P_small @ G, P-/+ = (I -/+ B)/2,
# B = (A - mid)/half. ||Y||_2 <= ||G||_2 ~ (sqrt(n) + sqrt(ks)).
# The square Gaussian coefficient matrix U^T G has a heavy (linear)
# sigma_min tail, so first rounds are aggressively shifted (they cannot
# break down and each cuts cond by ~40x); plain rounds finish. Two
# shifted rounds here and no plain: the refinement chain re-orthonormalizes.
AG = Asub @ G
BG = (AG - mid * G) / half
Ys = 0.5 * (G - BG) if small_is_minus else 0.5 * (G + BG)
g_norm2sq = 1.1 * (math.sqrt(n) + math.sqrt(ks)) ** 2
Qs, ok = _orth_chain(Ys, g_norm2sq, rounds=(True, True))
# One projector refinement pass: squashes cross-cluster contamination
# (amplified by the random-mixing coefficient matrix) down to the
# jitter/rounding floor. Shifted first round again in case a bad draw
# collapsed a column; any orthonormal basis of the eigenspace is equally
# valid, so shifts are harmless.
AQs = Asub @ Qs
BQs = (AQs - mid * Qs) / half
Ys2 = 0.5 * (Qs - BQs) if small_is_minus else 0.5 * (Qs + BQs)
Qs, ok2 = _orth_chain(Ys2, 2.0, rounds=(True, False, False))
ok = ok & ok2
# Orthogonal complement (spans the big cluster's eigenspace exactly, up to
# the tilt of Qs): randomized completion + block Gram-Schmidt. Same heavy
# coefficient tail: two aggressive shifted rounds, one plain; the polish
# round below doubles as the second plain round.
Y2 = G2 - Qs @ (Qs.transpose(-1, -2) @ G2)
Y2 = Y2 - Qs @ (Qs.transpose(-1, -2) @ Y2)
b_norm2sq = 1.1 * (math.sqrt(n) + math.sqrt(kb)) ** 2
Qb, ok3 = _orth_chain(Y2, b_norm2sq, rounds=(True, True, False))
# Final polish: re-project rounding leakage now that Qb is well conditioned.
Qb = Qb - Qs @ (Qs.transpose(-1, -2) @ Qb)
Qb, ok4 = _cholqr(Qb)
ok = ok & ok3 & ok4
Q = torch.cat([Qs, Qb], dim=-1)
# Rayleigh quotients on the ORIGINAL fp32 A, then a full ascending sort
# with the matching column permutation (handles cluster order and the
# repeated-eigenvalue rotation freedom the checker allows).
AQ = Asub @ Q
lvals = (Q * AQ).sum(dim=-2)
lvals, idx = torch.sort(lvals, dim=-1)
gather_idx = idx.unsqueeze(-2).expand(m, n, n)
Q = torch.gather(Q, -1, gather_idx).contiguous()
AQ = torch.gather(AQ, -1, gather_idx)
return Q, lvals.contiguous(), AQ, ok
def _verify_gates(Asub: torch.Tensor, Q: torch.Tensor, L: torch.Tensor,
AQ: torch.Tensor) -> torch.Tensor:
"""Per-matrix reference.py gates (eigen/orth/recon, L1 norms) at gate/4.
fp32 evaluation: the fp32-vs-fp64 GEMM discrepancy (~1e-5 relative) is far
inside the 4x margin against gates that sit at ~1e-2 relative for n=512.
"""
m, n, _ = Asub.shape
eps = torch.finfo(torch.float32).eps
l1A = torch.linalg.matrix_norm(Asub, ord=1, dim=(-2, -1))
QL = Q * L.unsqueeze(-2)
eig_res = torch.linalg.matrix_norm(AQ - QL, ord=1, dim=(-2, -1))
ok = eig_res <= (_EIGEN_RTOL_FACTOR * n * eps / _VERIFY_MARGIN) * l1A
eye = torch.eye(n, device=Q.device, dtype=Q.dtype)
orth_res = torch.linalg.matrix_norm(Q.transpose(-1, -2) @ Q - eye, ord=1, dim=(-2, -1))
ok &= orth_res <= (_ORTH_RTOL_FACTOR * n * eps / _VERIFY_MARGIN)
recon_res = torch.linalg.matrix_norm(QL @ Q.transpose(-1, -2) - Asub, ord=1, dim=(-2, -1))
ok &= recon_res <= (_RECON_RTOL_FACTOR * n * eps / _VERIFY_MARGIN) * l1A
ok &= torch.isfinite(Q).all(-1).all(-1) & torch.isfinite(L).all(-1)
return ok
# ----------------------------------------------------------------------------
# Dispatcher
# ----------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
A = data
batch, n, _ = A.shape
if n < _FAST_MIN_N:
return _cusolver_eigh(A)
try:
accept, cminus, cplus, kr = _detect_two_cluster(A)
n_accept = int(accept.sum().item())
if n_accept < max(int(math.ceil(_ACCEPT_FRAC * batch)), 1):
return _cusolver_eigh(A)
Qout = torch.empty((batch, n, n), dtype=torch.float32, device=A.device)
Lout = torch.empty((batch, n), dtype=torch.float32, device=A.device)
rej_idx = torch.nonzero(~accept, as_tuple=False).flatten()
if rej_idx.numel() > 0:
Qr, Lr = _cusolver_eigh(A.index_select(0, rej_idx))
Qout[rej_idx] = Qr
Lout[rej_idx] = Lr
acc_idx = torch.nonzero(accept, as_tuple=False).flatten()
k_acc = kr.index_select(0, acc_idx).to(torch.int64)
for kv in sorted(set(k_acc.tolist())):
g_idx = acc_idx[k_acc == kv]
sub = A.index_select(0, g_idx)
cm_g = cminus.index_select(0, g_idx)
cp_g = cplus.index_select(0, g_idx)
Qg, Lg, AQg, okg = _fast_two_cluster(sub, cm_g, cp_g, int(kv), seed_offset=0)
okg &= _verify_gates(sub, Qg, Lg, AQg)
bad = torch.nonzero(~okg, as_tuple=False).flatten()
if bad.numel() > 0:
# Retry with fresh randomness first (failures are usually a
# property of the random draw, not the matrix)...
sub2 = sub.index_select(0, bad)
Q2, L2, AQ2, ok2 = _fast_two_cluster(
sub2, cm_g.index_select(0, bad), cp_g.index_select(0, bad),
int(kv), seed_offset=1)
ok2 &= _verify_gates(sub2, Q2, L2, AQ2)
bad2 = torch.nonzero(~ok2, as_tuple=False).flatten()
if bad2.numel() > 0:
# ...then hand the survivors to cuSOLVER.
Q3, L3 = _cusolver_eigh(sub2.index_select(0, bad2))
Q2[bad2] = Q3
L2[bad2] = L3
Qg[bad] = Q2
Lg[bad] = L2
Qout[g_idx] = Qg
Lout[g_idx] = Lg
return Qout, Lout
except Exception:
# Any structural surprise: recompute everything with cuSOLVER.
return _cusolver_eigh(A)
# ----------------------------------------------------------------------------
scrolls · 387 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