Skip to content
KernelIndex
Search⌘K

submission 851792

bellz199 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1074 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-851792?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
46.1ms
#107 of 286
2026-07-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b38264efd419db1f486bdabf719ec1270113d5fac85795cbf9b150157efee65d
license declaredunknown
license concludedunknown
authorsbellz199
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fused-epilogue"""X <- 1.5 X - 0.5 X^3 with a fused epilogue (GEMMs stay cuBLAS)."""
shared-memory__shared__ float S_[W32][n * LD];
tile-k = 64_k_rayleigh[(bs, triton.cdiv(r, 64))](AU, U, out, n, r, BI=32, BK=64)

Kernel source

submission.py1074 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200

# ==========================================================
#  AUTO-GENERATED by tools/submission_generator.py -- DO NOT EDIT.
#  Edit src/batched_jacobi.cu and src/submission_template.py, then
#  run: python tools/submission_generator.py
# ==========================================================

# Tiered batched symmetric eigendecomposition with custom Triton kernels for
# every non-GEMM hot operation (cuBLAS keeps the large GEMMs).
#   Tier 0: exactly-diagonal inputs -> Triton permutation writer.
#   Tier 1: two-point spectra (`clustered`) -> one-shot spectral projector
#           split (Triton projector formation + Rayleigh + verify kernels).
#   Tier 2: truncated Rayleigh-Ritz for concentrated (graded) spectra at
#           n~512 (fp16 filter chain, keep-converged Ritz pairs).
#   Tier 3: cuSOLVER fallback. Every fast-path result is verified per matrix
#           against a conservative fraction of the real gates; failures fall
#           back, so correctness holds by construction.

import torch
import triton
import triton.language as tl
from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = False  # accuracy of fp32 GEMM paths

_EPS = torch.finfo(torch.float32).eps
_MODEL_TOL = 1e-2         # two-point minimal-poly fit acceptance
_GATE_FRACTION = 0.5      # accept fast-path result only if residual < 50% of gate

# custom CUDA eigensolver for n == 32. jacobi32x4 (4 warps/matrix): 5.9x
# faster than the 1-warp kernel locally, 2.9x vs cuSOLVER, racecheck-clean.
# PARKED 2026-07-03: the runner no longer completes CUDA load_inline within
# the benchmark-mode budget (v1 AND x4 both fail public benchmark; R12 passed
# 2 days ago). Re-enable if the runner env improves, or port to Triton.
_USE_WARP32 = False
_mod32 = None
_CUDA_SRC = r"""
// ============================================================================
//  batched_jacobi.cu - warp-level batched eigensolver for n == 32
// ----------------------------------------------------------------------------
//  One WARP solves one 32x32 symmetric matrix entirely in shared memory:
//  n == warpSize, so lane <-> matrix row/column maps 1:1. A single kernel
//  launch does: load -> parallel-order Jacobi sweeps -> in-warp rank sort ->
//  write Q (sorted eigenvector columns) and L (ascending eigenvalues).
//  cuSOLVER needs a multi-kernel sequence (~140us for batch=20); this is one
//  launch. Warp-synchronous: no block barriers, only __syncwarp().
// ============================================================================

#include <torch/extension.h>
#include <vector>

#define W32 4    // warps (matrices) per block
#define LD 33    // row stride padding: kills 32-way bank conflicts
__global__ void jacobi32_kernel(const float* __restrict__ Ag,
                                float* __restrict__ Qg,
                                float* __restrict__ Lg,
                                const int* __restrict__ sched,   // (31,16,2)
                                int nsweeps, int batch)
{
    constexpr int n = 32;
    __shared__ float S_[W32][n * LD];
    __shared__ float V_[W32][n * LD];
    __shared__ float dg_[W32][n];
    __shared__ int   rk_[W32][n];
    const int warp = threadIdx.x / 32;
    const int lane = threadIdx.x % 32;
    const int mat  = blockIdx.x * W32 + warp;
    if (mat >= batch) return;
    float* S = S_[warp];  float* V = V_[warp];
    float* dg = dg_[warp]; int* rk = rk_[warp];
    const size_t base = (size_t)mat * n * n;

    for (int i = lane; i < n * n; i += 32) S[(i / n) * LD + (i % n)] = Ag[base + i];
    for (int i = lane; i < n * LD; i += 32) V[i] = 0.0f;
    __syncwarp();
    V[lane * LD + lane] = 1.0f;
    __syncwarp();

    for (int sw = 0; sw < nsweeps; ++sw) {
        for (int r = 0; r < n - 1; ++r) {
            const int* sp = sched + r * 32;
            // every lane computes ALL 16 rotations in registers (no hazard:
            // c,s only read rows p,q which only their own pair writes)
            float cr[16], sr[16];
            int pr[16], qr[16];
            #pragma unroll
            for (int k = 0; k < 16; ++k) {
                const int p = sp[2 * k], q = sp[2 * k + 1];
                pr[k] = p; qr[k] = q;
                const float app = S[p * LD + p];
                const float aqq = S[q * LD + q];
                const float apq = S[p * LD + q];
                float c = 1.0f, sv = 0.0f;
                if (fabsf(apq) > 1e-20f) {
                    const float tau = (aqq - app) / (2.0f * apq);
                    const float t = copysignf(1.0f, tau) /
                                    (fabsf(tau) + sqrtf(1.0f + tau * tau));
                    c = rsqrtf(1.0f + t * t);
                    sv = t * c;
                }
                cr[k] = c; sr[k] = sv;
            }
            // rotate ROWS: lane owns column `lane`; pairs are disjoint rows
            #pragma unroll
            for (int k = 0; k < 16; ++k) {
                const float c = cr[k], s = sr[k];
                const float a = S[pr[k] * LD + lane], b = S[qr[k] * LD + lane];
                S[pr[k] * LD + lane] = c * a - s * b;
                S[qr[k] * LD + lane] = s * a + c * b;
            }
            __syncwarp();
            // rotate COLS of S and V: lane owns row `lane`; disjoint columns
            #pragma unroll
            for (int k = 0; k < 16; ++k) {
                const float c = cr[k], s = sr[k];
                const float a = S[lane * LD + pr[k]], b = S[lane * LD + qr[k]];
                S[lane * LD + pr[k]] = c * a - s * b;
                S[lane * LD + qr[k]] = s * a + c * b;
                const float u = V[lane * LD + pr[k]], w = V[lane * LD + qr[k]];
                V[lane * LD + pr[k]] = c * u - s * w;
                V[lane * LD + qr[k]] = s * u + c * w;
            }
            __syncwarp();
        }
    }

    // in-warp rank sort of the diagonal
    dg[lane] = S[lane * LD + lane];
    __syncwarp();
    const float di = dg[lane];
    int rank = 0;
    #pragma unroll
    for (int j = 0; j < n; ++j) {
        const float dj = dg[j];
        rank += (dj < di) || (dj == di && j < lane);
    }
    Lg[(size_t)mat * n + rank] = di;
    rk[lane] = rank;
    __syncwarp();
    // Q column rk[i] = V column i ; lane writes its row across all columns
    for (int i = 0; i < n; ++i)
        Qg[base + lane * n + rk[i]] = V[lane * LD + i];
}

std::vector<torch::Tensor> jacobi32_eigh(torch::Tensor A, int64_t nsweeps)
{
    TORCH_CHECK(A.dim() == 3 && A.size(1) == 32 && A.size(2) == 32,
                "jacobi32 expects (batch, 32, 32)");
    TORCH_CHECK(A.scalar_type() == torch::kFloat32 && A.is_cuda(), "fp32 CUDA only");
    A = A.contiguous();
    const int batch = (int)A.size(0);
    auto Q = torch::empty_like(A);
    auto L = torch::empty({batch, 32}, A.options());

    static torch::Tensor sched_gpu;                // round-robin, built once
    if (!sched_gpu.defined() || sched_gpu.device() != A.device()) {
        const int n = 32, halfN = 16;
        std::vector<int> arr(n);
        for (int i = 0; i < n; ++i) arr[i] = i;
        std::vector<int> sched;
        sched.reserve((n - 1) * n);
        for (int r = 0; r < n - 1; ++r) {
            for (int k = 0; k < halfN; ++k) {
                sched.push_back(arr[k]);
                sched.push_back(arr[n - 1 - k]);
            }
            const int last = arr[n - 1];
            for (int i = n - 1; i >= 2; --i) arr[i] = arr[i - 1];
            arr[1] = last;
        }
        sched_gpu = torch::from_blob(sched.data(), {(long)sched.size()},
                                     torch::TensorOptions().dtype(torch::kInt32))
                        .clone().to(A.device());
    }

    const int blocks = (batch + W32 - 1) / W32;
    jacobi32_kernel<<<blocks, W32 * 32>>>(A.data_ptr<float>(), Q.data_ptr<float>(),
                                     L.data_ptr<float>(), sched_gpu.data_ptr<int>(),
                                     (int)nsweeps, batch);
    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "jacobi32 launch failed: ", cudaGetErrorString(err));
    return {Q, L};
}

// ============================================================================
//  jacobi32x4 - 4 warps (128 threads) cooperate on ONE 32x32 matrix.
//  Same math as jacobi32_kernel but each phase's work is split 4 ways:
//  ~4x less serial latency per round, 8.8KB static smem per block (the
//  W32=8 variant's 69KB static allocation exceeds the 48KB static limit).
// ============================================================================
__global__ void jacobi32x4_kernel(const float* __restrict__ Ag,
                                  float* __restrict__ Qg,
                                  float* __restrict__ Lg,
                                  const int* __restrict__ sched,   // (31,16,2)
                                  int nsweeps, int batch)
{
    constexpr int n = 32;
    __shared__ float S[n * LD];
    __shared__ float V[n * LD];
    __shared__ float cs[16], sn[16];
    __shared__ float dg[n];
    __shared__ int   rk[n];
    const int tid = threadIdx.x;          // 0..127
    const int mat = blockIdx.x;
    if (mat >= batch) return;
    const size_t base = (size_t)mat * n * n;

    for (int i = tid; i < n * n; i += 128) {
        S[(i / n) * LD + (i % n)] = Ag[base + i];
        V[(i / n) * LD + (i % n)] = (i / n == i % n) ? 1.0f : 0.0f;
    }
    __syncthreads();

    const int c  = tid & 31;              // owned column / row
    const int k0 = (tid >> 5) * 4;        // 4 pairs per thread

    for (int sw = 0; sw < nsweeps; ++sw) {
        for (int r = 0; r < n - 1; ++r) {
            const int* sp = sched + r * n;
            if (tid < 16) {
                const int p = sp[2 * tid], q = sp[2 * tid + 1];
                const float app = S[p * LD + p];
                const float aqq = S[q * LD + q];
                const float apq = S[p * LD + q];
                float cc = 1.0f, sv = 0.0f;
                if (fabsf(apq) > 1e-20f) {
                    const float tau = (aqq - app) / (2.0f * apq);
                    const float t = copysignf(1.0f, tau) /
                                    (fabsf(tau) + sqrtf(1.0f + tau * tau));
                    cc = rsqrtf(1.0f + t * t);
                    sv = t * cc;
                }
                cs[tid] = cc; sn[tid] = sv;
            }
            __syncthreads();
            #pragma unroll
            for (int kk = 0; kk < 4; ++kk) {
                const int k = k0 + kk;
                const int p = sp[2 * k], q = sp[2 * k + 1];
                const float cc = cs[k], s = sn[k];
                const float a = S[p * LD + c], b = S[q * LD + c];
                S[p * LD + c] = cc * a - s * b;
                S[q * LD + c] = s * a + cc * b;
            }
            __syncthreads();
            #pragma unroll
            for (int kk = 0; kk < 4; ++kk) {
                const int k = k0 + kk;
                const int p = sp[2 * k], q = sp[2 * k + 1];
                const float cc = cs[k], s = sn[k];
                const int r0 = c * LD;
                float a = S[r0 + p], b = S[r0 + q];
                S[r0 + p] = cc * a - s * b;
                S[r0 + q] = s * a + cc * b;
                a = V[r0 + p]; b = V[r0 + q];
                V[r0 + p] = cc * a - s * b;
                V[r0 + q] = s * a + cc * b;
            }
            __syncthreads();
        }
    }

    if (tid < n) dg[tid] = S[tid * LD + tid];
    __syncthreads();
    if (tid < n) {
        const float di = dg[tid];
        int rank = 0;
        #pragma unroll
        for (int j = 0; j < n; ++j) {
            const float dj = dg[j];
            rank += (dj < di) || (dj == di && j < tid);
        }
        Lg[(size_t)mat * n + rank] = di;
        rk[tid] = rank;
    }
    __syncthreads();
    for (int i = tid; i < n * n; i += 128) {
        const int row = i / n, col = i % n;
        Qg[base + row * n + rk[col]] = V[row * LD + col];
    }
}

std::vector<torch::Tensor> jacobi32x4_eigh(torch::Tensor A, int64_t nsweeps)
{
    TORCH_CHECK(A.dim() == 3 && A.size(1) == 32 && A.size(2) == 32,
                "jacobi32x4 expects (batch, 32, 32)");
    TORCH_CHECK(A.scalar_type() == torch::kFloat32 && A.is_cuda(), "fp32 CUDA only");
    A = A.contiguous();
    const int batch = (int)A.size(0);
    auto Q = torch::empty_like(A);
    auto L = torch::empty({batch, 32}, A.options());
    static torch::Tensor sched_gpu2;
    if (!sched_gpu2.defined() || sched_gpu2.device() != A.device()) {
        const int n = 32, halfN = 16;
        std::vector<int> arr(n);
        for (int i = 0; i < n; ++i) arr[i] = i;
        std::vector<int> sched;
        for (int r = 0; r < n - 1; ++r) {
            for (int k = 0; k < halfN; ++k) {
                sched.push_back(arr[k]);
                sched.push_back(arr[n - 1 - k]);
            }
            const int last = arr[n - 1];
            for (int i = n - 1; i >= 2; --i) arr[i] = arr[i - 1];
            arr[1] = last;
        }
        sched_gpu2 = torch::from_blob(sched.data(), {(long)sched.size()},
                                      torch::TensorOptions().dtype(torch::kInt32))
                         .clone().to(A.device());
    }
    jacobi32x4_kernel<<<batch, 128>>>(
        A.data_ptr<float>(), Q.data_ptr<float>(), L.data_ptr<float>(),
        sched_gpu2.data_ptr<int>(), (int)nsweeps, batch);
    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "jacobi32x4 launch failed: ",
                cudaGetErrorString(err));
    return {Q, L};
}
"""
if _USE_WARP32:
    from torch.utils.cpp_extension import load_inline
    _CPP_SRC = ("std::vector<torch::Tensor> jacobi32_eigh(torch::Tensor A, int64_t nsweeps);\n"
                "std::vector<torch::Tensor> jacobi32x4_eigh(torch::Tensor A, int64_t nsweeps);")
    try:
        _mod32 = load_inline(name="jacobi32", cpp_sources=[_CPP_SRC],
                             cuda_sources=[_CUDA_SRC],
                             functions=["jacobi32_eigh", "jacobi32x4_eigh"],
                             extra_cuda_cflags=["-O3"], verbose=False)
    except Exception:
        _mod32 = None


# ===========================================================================
# Triton kernels (the custom-kernel layer)
# ===========================================================================
@triton.jit
def _k_shift_scale(A_ptr, P_ptr, a_ptr, inv_ptr, total, n, BLOCK: tl.constexpr):
    """P = (A - a*I) * inv   (per-matrix scalars a, inv), fused single pass."""
    offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    m = offs < total
    nn = n * n
    b = offs // nn
    rem = offs - b * nn
    i = rem // n
    j = rem - i * n
    a = tl.load(a_ptr + b, mask=m, other=0.0)
    inv = tl.load(inv_ptr + b, mask=m, other=0.0)
    x = tl.load(A_ptr + offs, mask=m, other=0.0)
    x = tl.where(i == j, x - a, x)
    tl.store(P_ptr + offs, x * inv, mask=m)


@triton.jit
def _k_rayleigh(AU_ptr, U_ptr, out_ptr, n, r, BI: tl.constexpr, BK: tl.constexpr):
    """out[b, k] = sum_i AU[b, i, k] * U[b, i, k]  (diag of U^T A U)."""
    pb = tl.program_id(0)
    pk = tl.program_id(1)
    k = pk * BK + tl.arange(0, BK)
    mk = k < r
    acc = tl.zeros((BK,), dtype=tl.float32)
    for i0 in range(0, n, BI):
        i = i0 + tl.arange(0, BI)
        mi = i < n
        ptr = pb * n * r + i[:, None] * r + k[None, :]
        msk = mi[:, None] & mk[None, :]
        x = tl.load(AU_ptr + ptr, mask=msk, other=0.0)
        y = tl.load(U_ptr + ptr, mask=msk, other=0.0)
        acc += tl.sum(x * y, 0)
    tl.store(out_ptr + pb * r + k, acc, mask=mk)


@triton.jit
def _k_qres_cols(AQ_ptr, Q_ptr, L_ptr, out_ptr, n, BI: tl.constexpr, BK: tl.constexpr):
    """out[b, j] = sum_i |AQ[b,i,j] - Q[b,i,j]*L[b,j]|  (eigen residual colsums)."""
    pb = tl.program_id(0)
    pk = tl.program_id(1)
    j = pk * BK + tl.arange(0, BK)
    mj = j < n
    lam = tl.load(L_ptr + pb * n + j, mask=mj, other=0.0)
    acc = tl.zeros((BK,), dtype=tl.float32)
    for i0 in range(0, n, BI):
        i = i0 + tl.arange(0, BI)
        mi = i < n
        ptr = pb * n * n + i[:, None] * n + j[None, :]
        msk = mi[:, None] & mj[None, :]
        x = tl.load(AQ_ptr + ptr, mask=msk, other=0.0)
        q = tl.load(Q_ptr + ptr, mask=msk, other=0.0)
        acc += tl.sum(tl.abs(x - q * lam[None, :]), 0)
    tl.store(out_ptr + pb * n + j, acc, mask=mj)


@triton.jit
def _k_absl1_cols(X_ptr, out_ptr, n, BI: tl.constexpr, BK: tl.constexpr):
    """out[b, j] = sum_i |X[b,i,j]|  (L1 column sums)."""
    pb = tl.program_id(0)
    pk = tl.program_id(1)
    j = pk * BK + tl.arange(0, BK)
    mj = j < n
    acc = tl.zeros((BK,), dtype=tl.float32)
    for i0 in range(0, n, BI):
        i = i0 + tl.arange(0, BI)
        mi = i < n
        ptr = pb * n * n + i[:, None] * n + j[None, :]
        x = tl.load(X_ptr + ptr, mask=mi[:, None] & mj[None, :], other=0.0)
        acc += tl.sum(tl.abs(x), 0)
    tl.store(out_ptr + pb * n + j, acc, mask=mj)


@triton.jit
def _k_offident_cols(G_ptr, out_ptr, n, BI: tl.constexpr, BK: tl.constexpr):
    """out[b, j] = sum_i |G[b,i,j] - (i==j)|  (orthogonality residual colsums)."""
    pb = tl.program_id(0)
    pk = tl.program_id(1)
    j = pk * BK + tl.arange(0, BK)
    mj = j < n
    acc = tl.zeros((BK,), dtype=tl.float32)
    for i0 in range(0, n, BI):
        i = i0 + tl.arange(0, BI)
        mi = i < n
        ptr = pb * n * n + i[:, None] * n + j[None, :]
        x = tl.load(G_ptr + ptr, mask=mi[:, None] & mj[None, :], other=0.0)
        x = tl.where(i[:, None] == j[None, :], x - 1.0, x)
        acc += tl.sum(tl.abs(x), 0)
    tl.store(out_ptr + pb * n + j, acc, mask=mj)


@triton.jit
def _k_perm_eye(out_ptr, perm_ptr, total, n, BLOCK: tl.constexpr):
    """out[b, i, k] = (i == perm[b, k])  -- permutation-matrix writer."""
    offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    m = offs < total
    nn = n * n
    b = offs // nn
    rem = offs - b * nn
    i = rem // n
    k = rem - i * n
    src = tl.load(perm_ptr + b * n + k, mask=m, other=-1)
    tl.store(out_ptr + offs, (i == src).to(tl.float32), mask=m)


_B_ELEM = 1024
_GEN = None


def _randn(bs, n, r, device):
    global _GEN
    if _GEN is None or _GEN.device != torch.device(device):
        _GEN = torch.Generator(device=device)
        _GEN.manual_seed(0x5EED)
    return torch.randn(bs, n, r, device=device, generator=_GEN)


def _shift_scale(A, a, inv):
    P = torch.empty_like(A)
    total = A.numel()
    n = A.size(-1)
    _k_shift_scale[(triton.cdiv(total, _B_ELEM),)](A, P, a.contiguous(), inv.contiguous(),
                                                   total, n, BLOCK=_B_ELEM)
    return P


def _rayleigh(AU, U):
    bs, n, r = AU.shape
    out = torch.empty(bs, r, device=AU.device, dtype=torch.float32)
    _k_rayleigh[(bs, triton.cdiv(r, 64))](AU, U, out, n, r, BI=32, BK=64)
    return out


def _perm_eye(perm, n):
    bs = perm.size(0)
    Q = torch.empty(bs, n, n, device=perm.device, dtype=torch.float32)
    total = Q.numel()
    _k_perm_eye[(triton.cdiv(total, _B_ELEM),)](Q, perm.contiguous(), total, n, BLOCK=_B_ELEM)
    return Q


def _eigen_residual_ok(A, Q, L, n, frac=_GATE_FRACTION):
    """Per-matrix eigen + orthogonality residuals vs conservative gate fractions.
    Custom Triton reductions; the two GEMMs stay on cuBLAS."""
    bs = A.size(0)
    AQ = A @ Q
    res_c = torch.empty(bs, n, device=A.device, dtype=torch.float32)
    _k_qres_cols[(bs, triton.cdiv(n, 64))](AQ, Q, L.contiguous(), res_c, n, BI=32, BK=64)
    scale_c = torch.empty(bs, n, device=A.device, dtype=torch.float32)
    _k_absl1_cols[(bs, triton.cdiv(n, 64))](A, scale_c, n, BI=32, BK=64)
    G = Q.transpose(-1, -2) @ Q
    orth_c = torch.empty(bs, n, device=A.device, dtype=torch.float32)
    _k_offident_cols[(bs, triton.cdiv(n, 64))](G, orth_c, n, BI=32, BK=64)
    res = res_c.amax(-1)
    scale = scale_c.amax(-1)
    orth = orth_c.amax(-1)
    gate = 200.0 * n * _EPS * scale
    ortho_gate = 100.0 * n * _EPS
    return (res < frac * gate) & (orth < frac * ortho_gate)


# ===========================================================================
# building blocks
# ===========================================================================
def _diagonal_path(A):
    n = A.size(-1)
    L, idx = torch.sort(torch.diagonal(A, dim1=-2, dim2=-1), dim=-1)
    return _perm_eye(idx, n), L.contiguous()


def _cholqr(X, jit):
    """CholeskyQR with right-side triangular solve (no transposes/copies)."""
    G = X.transpose(-1, -2) @ X
    d = G.diagonal(dim1=-2, dim2=-1).amax(-1, keepdim=True).clamp_min(1e-30)
    G.diagonal(dim1=-2, dim2=-1).add_(jit * d)          # in-place jitter
    R = torch.linalg.cholesky(G)                        # lower: G = R R^T
    return torch.linalg.solve_triangular(R.transpose(-1, -2), X, upper=True, left=False)


def _two_point_fit(A):
    """Per-matrix minimal-polynomial fit A^2 ~ s*A - p*I; returns fit, a, b, w."""
    n = A.size(-1)
    M2 = A @ A
    m1 = torch.diagonal(A, dim1=-2, dim2=-1).sum(-1) / n
    m2 = torch.diagonal(M2, dim1=-2, dim2=-1).sum(-1) / n
    m3 = (M2 * A).sum((-2, -1)) / n
    den = m1 * m1 - m2
    safe = den.abs() > 1e-12
    s = torch.where(safe, (m2 * m1 - m3) / den, torch.zeros_like(den))
    p = torch.where(safe, (m2 * m2 - m1 * m3) / den, torch.zeros_like(den))
    I = torch.eye(n, device=A.device)
    R = M2 - s[:, None, None] * A + p[:, None, None] * I
    fit = R.norm(dim=(-2, -1)) / M2.norm(dim=(-2, -1)).clamp_min(1e-30)
    fit = torch.where(safe, fit, torch.ones_like(fit))
    disc = (s * s / 4 - p).clamp_min(0)
    a_val = s / 2 - disc.sqrt()
    b_val = s / 2 + disc.sqrt()
    w = ((m1 - a_val) / (b_val - a_val).clamp_min(1e-30)).clamp(0, 1)
    return fit, a_val, b_val, w


def _two_point_sample_hit(sample):
    n = sample.size(-1)
    fit, _, _, w = _two_point_fit(sample)
    cand = (fit < _MODEL_TOL) & (w * n > 0.5) & (w * n < n - 0.5)
    return float(cand.float().mean().item()) >= 0.25


def _two_point_solve(A, a, b, kb):
    """Spectral split for two-cluster spectra {a, b}; kb = dim of b-cluster."""
    bs, n, _ = A.shape
    ka = n - kb
    inv = 1.0 / (b - a)
    Pb = _shift_scale(A, a, inv)                        # (A - a I)/(b - a), fused

    Ph = Pb.half()

    def subspace(P, r, flip):
        # flip: use (I - P) x = x - P x without materializing I - P
        # fp16 filter applies (tensor cores); final round fp32 for clean margins
        X = _randn(bs, n, r, A.device)
        Y = (Ph @ X.half()).float()
        U = _cholqr(X - Y if flip else Y, 1e-4)
        Y = (Ph @ U.half()).float()
        U = _cholqr(U - Y if flip else Y, 1e-6)
        Y = P @ U
        U = _cholqr(U - Y if flip else Y, 1e-8)
        return U

    Ub = subspace(Pb, kb, False)
    Ua = subspace(Pb, ka, True)
    U = torch.cat([Ua, Ub], dim=-1)
    L = _rayleigh(A @ U, U)                             # diag(U^T A U), fused
    Ls, idx = torch.sort(L, dim=-1)
    Q = torch.gather(U, -1, idx.unsqueeze(-2).expand(-1, n, -1))
    return Q, Ls


def _trr_mass_probe(sample, scale_s):
    """Top-96 Frobenius mass of the (normalized) sample."""
    bs, n, _ = sample.shape
    S = sample / scale_s[:, None, None]
    Sh = S.half()
    W = _cholqr((Sh @ _randn(bs, n, 96, S.device).half()).float(), 1e-4)
    W = _cholqr((Sh @ W.half()).float(), 1e-6)
    W = _cholqr(W, 1e-10)
    B = W.transpose(-1, -2) @ (S @ W)
    B = 0.5 * (B + B.transpose(-1, -2))
    th = torch.linalg.eigvalsh(B)
    return (th * th).sum(-1) / (S * S).sum((-2, -1)).clamp_min(1e-30)


def _trr_solve(A, r):
    """Truncated Rayleigh-Ritz: fp16 filter chain, keep converged Ritz pairs,
    complete with unkept Ritz vectors + random complement. A must be O(1)-scaled."""
    bs, n, _ = A.shape
    Ah = A.half()
    W = _cholqr((Ah @ _randn(bs, n, r, A.device).half()).float(), 1e-4)
    W = _cholqr((Ah @ W.half()).float(), 1e-6)
    W = _cholqr(W, 1e-10)                                # ortho polish
    AW = A @ W
    B = W.transpose(-1, -2) @ AW
    B = 0.5 * (B + B.transpose(-1, -2))
    theta, V = torch.linalg.eigh(B)
    U = W @ V
    G = AW - W @ B
    GG = G.transpose(-1, -2) @ G
    rn = ((GG @ V) * V).sum(-2).clamp_min(0).sqrt()      # per-pair residuals
    scale = A.abs().sum(-2).amax(-1)
    gate = 200.0 * n * _EPS * scale
    conv = rn < (0.5 * gate / (n ** 0.5)).unsqueeze(-1)
    kcnt = conv.sum(-1)
    k = int(kcnt.median().item())
    if k < r // 4:
        raise RuntimeError("insufficient convergence")
    score = torch.where(conv, theta.abs(), torch.full_like(theta, -1.0))
    top = score.topk(k, dim=-1).indices
    Uk = torch.gather(U, -1, top.unsqueeze(-2).expand(-1, n, -1))
    Tk = torch.gather(theta, -1, top)
    mask = torch.ones_like(theta, dtype=torch.bool)
    mask.scatter_(-1, top, False)
    rest = mask.nonzero(as_tuple=True)[1].view(bs, r - k)
    Ur = torch.gather(U, -1, rest.unsqueeze(-2).expand(-1, n, -1))
    Tr = torch.gather(theta, -1, rest)
    Om2 = _randn(bs, n, n - r, A.device)
    P = _cholqr(Om2 - U @ (U.transpose(-1, -2) @ Om2), 1e-6)
    P = _cholqr(P - U @ (U.transpose(-1, -2) @ P), 1e-8)
    mu = _rayleigh(A @ P, P)
    Q = torch.cat([Uk, Ur, P], dim=-1)
    L = torch.cat([Tk, Tr, mu], dim=-1)
    Ls, idx = torch.sort(L, dim=-1)
    Q = torch.gather(Q, -1, idx.unsqueeze(-2).expand(-1, n, -1))
    return Q, Ls


# ===========================================================================
# sign-based divide-and-conquer solver (handwritten Triton; runner-portable)
# ===========================================================================
_SD_GEN = None
_QUINTIC = (3.4445, -4.7750, 2.0315)


def _sd_randn(*shape, device):
    global _SD_GEN
    if _SD_GEN is None or _SD_GEN.device != torch.device(device):
        _SD_GEN = torch.Generator(device=device)
        _SD_GEN.manual_seed(0xC0FFEE)
    return torch.randn(*shape, device=device, generator=_SD_GEN)


def _sd_cholqr(X, jit):
    torch.backends.cuda.matmul.allow_tf32 = True   # Gram on tensor cores
    try:
        G = X.transpose(-1, -2) @ X
    finally:
        torch.backends.cuda.matmul.allow_tf32 = False
    d = G.diagonal(dim1=-2, dim2=-1).amax(-1, keepdim=True).clamp_min(1e-30)
    G.diagonal(dim1=-2, dim2=-1).add_(jit * d)
    R = torch.linalg.cholesky(G)
    return torch.linalg.solve_triangular(R.transpose(-1, -2), X, upper=True, left=False)


def _sd_specnorm(M, iters=8):
    b, n, _ = M.shape
    v = _sd_randn(b, n, 1, device=M.device)
    v = v / v.norm(dim=-2, keepdim=True)
    for _ in range(iters):
        v = M @ v
        v = v / v.norm(dim=-2, keepdim=True).clamp_min(1e-30)
    return (M @ v).norm(dim=-2, keepdim=True).clamp_min(1e-30)


_TB = 2048


@triton.jit
def _kt_poly5(x_ptr, x3_ptr, x5_ptr, o_ptr, numel,
              A: tl.constexpr, B: tl.constexpr, C: tl.constexpr, BLOCK: tl.constexpr):
    off = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    m = off < numel
    x = tl.load(x_ptr + off, mask=m, other=0.0).to(tl.float32)
    x3 = tl.load(x3_ptr + off, mask=m, other=0.0).to(tl.float32)
    x5 = tl.load(x5_ptr + off, mask=m, other=0.0).to(tl.float32)
    tl.store(o_ptr + off, (A * x + B * x3 + C * x5).to(tl.float16), mask=m)


@triton.jit
def _kt_ns(x_ptr, y_ptr, o_ptr, numel, F16: tl.constexpr, BLOCK: tl.constexpr):
    off = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    m = off < numel
    x = tl.load(x_ptr + off, mask=m, other=0.0).to(tl.float32)
    y = tl.load(y_ptr + off, mask=m, other=0.0).to(tl.float32)
    r = 1.5 * x - 0.5 * y
    if F16:
        tl.store(o_ptr + off, r.to(tl.float16), mask=m)
    else:
        tl.store(o_ptr + off, r, mask=m)


@triton.jit
def _kt_sym(x_ptr, o_ptr, n, BLOCK: tl.constexpr):
    pb = tl.program_id(0).to(tl.int64)
    pi = tl.program_id(1)
    pj = tl.program_id(2)
    i = pi * BLOCK + tl.arange(0, BLOCK)
    j = pj * BLOCK + tl.arange(0, BLOCK)
    mi = i < n
    mj = j < n
    base = pb * n * n
    p_ij = base + i[:, None] * n + j[None, :]
    p_ji = base + j[:, None] * n + i[None, :]
    a = tl.load(x_ptr + p_ij, mask=mi[:, None] & mj[None, :], other=0.0).to(tl.float32)
    b = tl.load(x_ptr + p_ji, mask=mj[:, None] & mi[None, :], other=0.0).to(tl.float32)
    tl.store(o_ptr + p_ij, 0.5 * (a + tl.trans(b)), mask=mi[:, None] & mj[None, :])


@triton.jit
def _kt_blend(s_ptr, sgn_ptr, p_ptr, ph_ptr, numel, n, BLOCK: tl.constexpr):
    off = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    m = off < numel
    nn = n * n
    b = off // nn
    rem = off - b * nn
    i = rem // n
    j = rem - i * n
    sg = tl.load(sgn_ptr + b, mask=m, other=0.0)
    sv = tl.load(s_ptr + off, mask=m, other=0.0)
    pv = 0.5 * sg * sv + tl.where(i == j, 0.5, 0.0)
    tl.store(p_ptr + off, pv, mask=m)
    tl.store(ph_ptr + off, pv.to(tl.float16), mask=m)


@triton.jit
def _kt_scale_sgn(x_ptr, sgn_ptr, o_ptr, numel, nn, BLOCK: tl.constexpr):
    off = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    m = off < numel
    b = off // nn
    sg = tl.load(sgn_ptr + b, mask=m, other=0.0)
    x = tl.load(x_ptr + off, mask=m, other=0.0).to(tl.float32)
    tl.store(o_ptr + off, (sg * x).to(tl.float16), mask=m)


def _poly5(X, X3, X5):
    o = torch.empty_like(X)
    numel = X.numel()
    a, b, c = _QUINTIC
    _kt_poly5[(triton.cdiv(numel, _TB),)](X, X3, X5, o, numel, a, b, c, BLOCK=_TB)
    return o


def _ns_step(X):
    """X <- 1.5 X - 0.5 X^3 with a fused epilogue (GEMMs stay cuBLAS)."""
    Y = (X @ X) @ X
    o = torch.empty_like(X)
    numel = X.numel()
    _kt_ns[(triton.cdiv(numel, _TB),)](X, Y, o, numel, F16=(X.dtype == torch.float16), BLOCK=_TB)
    return o


def _sym(X):
    o = torch.empty_like(X)
    n = X.size(-1)
    _kt_sym[(X.size(0), triton.cdiv(n, 32), triton.cdiv(n, 32))](X, o, n, BLOCK=32)
    return o


def msign(A, iters=5, polish=3):
    X = (A / (_sd_specnorm(A, 8) * 1.10)).half()
    for _ in range(2):                            # safe cubic pre-contraction
        X = _ns_step(X)
    for _ in range(iters):
        X2 = X @ X
        X3 = X2 @ X
        X5 = X2 @ X3
        X = _poly5(X, X3, X5)
    X = _sym(X.float())
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        for _ in range(polish):                   # TF32 coarse walk
            X = _ns_step(X)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = False
    for _ in range(2):                            # fp32 landing
        X = _ns_step(X)
    return X


def count_above(A, sigma, n, I):
    M = A - sigma * I
    X = (M / (_sd_specnorm(M, 4) * 1.15)).half()
    for _ in range(2):
        X = _ns_step(X)
    for _ in range(4):
        X2 = X @ X
        X3 = X2 @ X
        X5 = X2 @ X3
        X = _poly5(X, X3, X5)
    X = _ns_step(X)
    return 0.5 * (n + X.float().diagonal(dim1=-2, dim2=-1).sum(-1)), X


def _sd_cholqr_UNUSED(X, jit):
    G = X.transpose(-1, -2) @ X
    d = G.diagonal(dim1=-2, dim2=-1).amax(-1, keepdim=True).clamp_min(1e-30)
    G.diagonal(dim1=-2, dim2=-1).add_(jit * d)
    R = torch.linalg.cholesky(G)
    return torch.linalg.solve_triangular(R.transpose(-1, -2), X, upper=True, left=False)


def _level_core(A, S, sigma, Ashn_h, Om1, Om2, deep: bool, rank_complement: bool):
    """All the level math after the sign is known. Eager + Triton kernels."""
    bs, n, _ = A.shape
    h = n // 2
    k = 0.5 * (n + S.diagonal(dim1=-2, dim2=-1).sum(-1))
    sgn_v = torch.where(k < h, -torch.ones_like(k), torch.ones_like(k))
    numel = S.numel()
    P = torch.empty_like(S)
    Ph = torch.empty(S.shape, device=S.device, dtype=torch.float16)
    _kt_blend[(triton.cdiv(numel, _TB),)](S, sgn_v, P, Ph, numel, n, BLOCK=_TB)
    Anh = torch.empty_like(Ashn_h)
    _kt_scale_sgn[(triton.cdiv(numel, _TB),)](Ashn_h, sgn_v, Anh, numel, n * n, BLOCK=_TB)

    def cholqr2(X, jit):
        return _sd_cholqr(_sd_cholqr(X, jit), 1e-8)

    U = _sd_cholqr((Ph @ Om1.half()).float(), 1e-4)
    # rank by |lambda-sigma|^2: spill must be exactly the smallest members
    U = cholqr2((Ph @ (Anh @ (Anh @ U.half()))).float(), 1e-4)
    if deep:
        U = cholqr2((Ph @ (Anh @ (Anh @ U.half()))).float(), 1e-4)
    U1 = cholqr2(P @ U, 1e-6)

    U2 = _sd_cholqr(Om2 - U1 @ (U1.transpose(-1, -2) @ Om2), 1e-4)
    if rank_complement:                           # ordered diag for recursion
        for jit in (1e-4, 1e-5, 1e-6):
            Z = 0.05 * U2 - (Anh @ U2.half()).float()
            Z = Z - U1 @ (U1.transpose(-1, -2) @ Z)
            U2 = _sd_cholqr(Z, jit)
            U2 = _sd_cholqr(U2 - U1 @ (U1.transpose(-1, -2) @ U2), 1e-8)
    else:
        U2 = _sd_cholqr(U2 - U1 @ (U1.transpose(-1, -2) @ U2), 1e-6)
        U2 = _sd_cholqr(U2 - U1 @ (U1.transpose(-1, -2) @ U2), 1e-8)

    Uc = torch.cat([U1, U2], dim=-1)
    torch.backends.cuda.matmul.allow_tf32 = True   # transforms on tensor cores
    try:
        B = _sym(Uc.transpose(-1, -2) @ (A @ Uc))
    finally:
        torch.backends.cuda.matmul.allow_tf32 = False
    return Uc, B[:, :h, :h].contiguous(), B[:, h:, h:].contiguous()


# ---------------------------------------------------------------------------
# python driver: all branching / randomness here
# ---------------------------------------------------------------------------
def _signdc_level(A, rank_complement=True):
    bs, n, _ = A.shape
    h = n // 2
    I = torch.eye(n, device=A.device)
    sigma = A.diagonal(dim1=-2, dim2=-1).median(-1).values[:, None, None]
    S = msign(A - sigma * I)                      # the full sign IS the count
    k = 0.5 * (n + S.diagonal(dim1=-2, dim2=-1).sum(-1))
    deep = bool(((k - h).abs() > 6).any().item())
    if deep:                                      # asymmetric spectra: refine
        lmax = _sd_specnorm(A, 6).squeeze(-1).squeeze(-1)
        k0, s0 = k, sigma.squeeze(-1).squeeze(-1)
        need = (k0 - h).abs() > 6
        s1 = s0 + torch.where(need, 0.15 * lmax * torch.sign(k0 - h), torch.zeros_like(s0))
        k1, Xc = count_above(A, s1[:, None, None], n, I)
        for _ in range(2):                        # masked secant steps
            den = k1 - k0
            safe = den.abs() > 0.5
            step = (s1 - s0) * (h - k1) / torch.where(safe, den, torch.ones_like(den))
            s2 = s1 + torch.where(need & safe, step, torch.zeros_like(step))
            s0, s1, k0 = s1, s2, k1
            k1, Xc = count_above(A, s1[:, None, None], n, I)
        sigma = torch.where(need[:, None, None], s1[:, None, None], sigma)
        S = msign(A - sigma * I)                  # full-quality re-sign
        # (reusing the shallow counting-sign lift broke lapack_even: 226/640)
    Ash = A - sigma * I
    Ashn_h = (Ash / _sd_specnorm(Ash, 8)).half()
    Om1 = _sd_randn(bs, n, h, device=A.device)
    Om2 = _sd_randn(bs, n, h, device=A.device)
    return _level_core(A, S, sigma, Ashn_h, Om1, Om2, deep, rank_complement)


def _signdc_solve_impl(A, levels=1):
    bs, n, _ = A.shape
    if levels == 0:
        L, Q = torch.linalg.eigh(A)
        return Q, L
    U, B1, B2 = _signdc_level(A, rank_complement=(levels > 1))
    h = n // 2
    Bs = torch.cat([B1, B2], dim=0)
    Qs, Ls = _signdc_solve_impl(Bs, levels - 1)
    Qsub = torch.zeros(bs, n, n, device=A.device)
    Qsub[:, :h, :h] = Qs[:bs]
    Qsub[:, h:, h:] = Qs[bs:]
    Q = U @ Qsub
    L = torch.cat([Ls[:bs], Ls[bs:]], dim=-1)
    Ls_, idx = torch.sort(L, dim=-1)
    Q = torch.gather(Q, -1, idx.unsqueeze(-2).expand(-1, n, -1))
    return Q, Ls_




def _signdc_solve(A):
    return _signdc_solve_impl(A, 2)


# ===========================================================================
# dispatcher
# ===========================================================================
def custom_kernel(data: input_t) -> output_t:
    A = data
    bs, n, _ = A.shape

    # ---- warp kernel: n == 32 solved in ONE launch (one warp per matrix) ----
    if n == 32 and _mod32 is not None:
        Q, L = _mod32.jacobi32_eigh(A, 6)
        return Q, L

    # small sizes: no tier engages below n=500 - skip all probing/dispatch
    # python (these cases are latency-bound; every microsecond is geomean)
    if n < 500:
        L, Q = torch.linalg.eigh(A)
        return Q, L

    # sampled pre-probe: look at a few matrices before any full-batch scan
    sample = A[: min(bs, 8)]

    # ---- tier 0: exactly diagonal ----
    if n >= 500:
        s_off = torch.count_nonzero(sample) - torch.count_nonzero(
            torch.diagonal(sample, dim1=-2, dim2=-1))
        if int(s_off.item()) == 0:
            off_nnz = torch.count_nonzero(A) - torch.count_nonzero(
                torch.diagonal(A, dim1=-2, dim2=-1))
            if int(off_nnz.item()) == 0:
                return _diagonal_path(A)

    # ---- tier 1: two-point spectra (clustered) ----
    if n >= 500 and bs >= 2 and _two_point_sample_hit(sample):
        try:
            fit, a_val, b_val, w = _two_point_fit(A)
            cand = (fit < _MODEL_TOL) & (w * n > 0.5) & (w * n < n - 0.5)
            if float(cand.float().mean().item()) >= 0.25:
                idxs = cand.nonzero(as_tuple=True)[0]
                kbs = (w[idxs] * n).round().long()
                Qout = None
                Lout = None
                done = torch.zeros(bs, dtype=torch.bool, device=A.device)
                for kb in kbs.unique().tolist():
                    grp = idxs[kbs == kb]
                    Ag = A[grp]
                    Qg, Lg = _two_point_solve(Ag, a_val[grp], b_val[grp], int(kb))
                    ok = _eigen_residual_ok(Ag, Qg, Lg, n)
                    if bool(ok.any()):
                        if Qout is None:
                            Qout = torch.empty_like(A)
                            Lout = torch.empty(bs, n, device=A.device, dtype=A.dtype)
                        sel = grp[ok]
                        Qout[sel] = Qg[ok]
                        Lout[sel] = Lg[ok]
                        done[sel] = True
                if Qout is not None:
                    rest = ~done
                    if bool(rest.any()):
                        Lr, Qr = torch.linalg.eigh(A[rest])
                        Qout[rest] = Qr
                        Lout[rest] = Lr
                    return Qout, Lout
        except Exception:
            pass

    # ---- sign-D&C solver: measured on B200 (dense 180 vs 141, rankdef 239
    # vs 170, even 267 vs 206) - loses to deployed paths; disabled until the
    # handwritten-Triton fusion arc closes the ~30% gap ----
    _SOLVER_ENABLED = False
    _use_solver = False
    if _SOLVER_ENABLED and n == 512 and bs >= 100:
        try:
            _sc = sample.abs().amax((-2, -1)).clamp_min(1e-30)
            _sm = _trr_mass_probe(sample, _sc)
            _use_solver = bool((_sm.max() - _sm.min()).item() < 0.25)  # not mixed
        except Exception:
            _use_solver = False
    if _use_solver:
        try:
            Qs, Lsv = _signdc_solve(A)
            ok = _eigen_residual_ok(A, Qs, Lsv, n, frac=0.9)
            if bool(ok.all()):
                return Qs, Lsv
            rest = ~ok
            Lr, Qr = torch.linalg.eigh(A[rest])
            Qs[rest] = Qr
            Lsv[rest] = Lr
            return Qs, Lsv
        except Exception:
            pass

    # ---- tier 2: truncated Rayleigh-Ritz for concentrated spectra ----
    smass = None
    if 500 <= n <= 1100 and bs >= 4:
        try:
            scale_s = sample.abs().amax((-2, -1)).clamp_min(1e-30)
            mass = _trr_mass_probe(sample, scale_s)
            smass = mass
            mmin = float(mass.min().item())
            mmed = float(mass.median().item())
            engage = (mmin > 0.70) if n <= 640 else (mmin > 0.35)
            if engage:
                # steep spectra (very high probe mass) have shallow numerical
                # rank: a deep sketch hits rank-deficient Grams (junk columns).
                if n > 640 and mmed >= 0.75:
                    r = int(0.39 * n)          # e.g. lapack_geometric
                elif n <= 640:
                    r = (21 * n) // 32         # 336 @ n=512 (validated margin)
                else:
                    r = (83 * n) // 128        # 664 @ n=1024 (validated margin)
                r = min(max((r // 16) * 16, 128), n - 64)
                sc = A.abs().amax((-2, -1)).clamp_min(1e-30)
                An = A / sc[:, None, None]
                Qt, Lt = _trr_solve(An, r)
                ok = _eigen_residual_ok(An, Qt, Lt, n, frac=0.8)
                Lt = Lt * sc[:, None]
                if bool(ok.all()):
                    return Qt, Lt
                if bool(ok.any()):
                    rest = ~ok
                    Lr, Qr = torch.linalg.eigh(A[rest])
                    Qt[rest] = Qr
                    Lt[rest] = Lr
                    return Qt, Lt
        except Exception:
            pass

    # ---- tier 2b: heterogeneous (mixed) batches - per-matrix routing ----
    # sample shows a SPREAD of masses (some graded, some not) -> probe every
    # matrix, truncate the graded majority, cuSOLVER the rest, scatter back.
    _MIXED_SPLIT = False   # B200-measured: split loses (190 vs 167);
    # cuSOLVER's batch amortization beats sub-size savings when ~50% remains
    if (_MIXED_SPLIT and 500 <= n <= 640 and bs >= 64 and smass is not None
            and float(smass.max().item()) > 0.70
            and float(smass.min().item()) <= 0.70):
        try:
            sc = A.abs().amax((-2, -1)).clamp_min(1e-30)
            fmass = _trr_mass_probe(A, sc)
            grad = fmass > 0.85
            frac = float(grad.float().mean().item())
            if 0.25 <= frac <= 0.97:
                gidx = grad.nonzero(as_tuple=True)[0]
                r = min(max(((3 * n // 4) // 16) * 16, 128), n - 64)
                An = A[gidx] / sc[gidx][:, None, None]
                Qt, Lt = _trr_solve(An, r)
                ok = _eigen_residual_ok(An, Qt, Lt, n, frac=0.8)
                Qout = torch.empty_like(A)
                Lout = torch.empty(bs, n, device=A.device, dtype=A.dtype)
                sel = gidx[ok]
                Qout[sel] = Qt[ok]
                Lout[sel] = Lt[ok] * sc[sel][:, None]
                done = torch.zeros(bs, dtype=torch.bool, device=A.device)
                done[sel] = True
                rest = ~done
                if bool(rest.any()):
                    Lr, Qr = torch.linalg.eigh(A[rest])
                    Qout[rest] = Qr
                    Lout[rest] = Lr
                return Qout, Lout
        except Exception:
            pass

    # ---- tier 3: cuSOLVER ----
    L, Q = torch.linalg.eigh(A)
    return Q, L
scrolls · 1074 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