Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
51.4ms
#190 of 286
2026-07-10

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