Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
43.3ms
#89 of 286
2026-07-14

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