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
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