submission 877196
viridale · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2032 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-877196?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:70e52cca3ed90dd49ac11c9670d794b5b4dd0b51c7af3c4fbeebf31d0bd5db50
license declaredunknown
license concludedunknown
authorsviridale
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
gram = tl.dot(vals, vals, input_precision="tf32")shared-memory
__shared__ float s_val[256];Kernel source
submission.py2032 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200
import ctypes
import ctypes.util
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
ENABLE_DYNAMIC_NATIVE_IMPORT_PROBE = False
ENABLE_EIGH_OUT_CACHE = False
ENABLE_EIGH_UPLO_U = False
ENABLE_SVD_SYMMETRIC = False
ENABLE_CUDA_GRAPH_EIGH = False
ENABLE_CUSOLVER_PREWARM = False
ENABLE_RMT_N32_PROBE = False
ENABLE_SIGN_SPLIT_N32_PROBE = False
ENABLE_CAYLEY_N32_PROBE = False
ENABLE_CUSOLVER_SYEVJ_BATCHED = True
ENABLE_CUSOLVER_XSYEV_BATCHED = False
ENABLE_CUSOLVER_SYEVD_LOOP = False
ENABLE_CUSOLVER_HOST_SORT_N32 = False
ENABLE_EXACT_N512_B16_DIAGONAL_FASTPATH = True
ENABLE_NEAR_DIAGONAL_CERT_FASTPATH = False
ENABLE_CUSOLVER_RAYLEIGH_RESYNC_N32 = False
ENABLE_TRI_N32_COLUMN_REFILTER = False
ENABLE_TRI_FP16_SCREEN_PROBE = False
ENABLE_TRI_SKETCH_CLASSIFIER_PROBE = False
ENABLE_RANK_SPLIT_RAYLEIGH_PROBE = False
ENABLE_RANK_HUTCH_PREFILTER_PROBE = False
ENABLE_INVOLUTION_PROJECTOR_PROBE = False
ENABLE_CLUSTER_NATIVE_PIVOT_SPLIT = True
ENABLE_CLUSTER_OUTPUT_SCREEN = True
ENABLE_TF32_GLOBAL = False
CLUSTER_OUTPUT_SCREEN_PROBES = 2
ENABLE_CLUSTER_CHOLSPAN_BASIS = False
ENABLE_CLUSTER_PIVCHOL_BASIS = True
ENABLE_CLUSTER_SINGLE_CHOLQR_COMPLEMENT = True
ENABLE_TRACE_SIGN_DC_PROBE = False
ENABLE_NATIVE_DENSE_JACOBI176 = False
NATIVE_DENSE_JACOBI176_SWEEPS = 7
CUSOLVER_SYEVJ_SHAPE_PARAMS = {
(20, 32): (3.5e-4, 8),
}
CUSOLVER_XSYEV_SHAPES = {
(640, 512),
}
CUSOLVER_SYEVD_LOOP_SHAPES = {
(8, 2048),
}
CUSOLVER_SYEVJ_SORT_EIG = True
PREFERRED_LINALG_LIBRARY = "cusolver"
ENABLE_TRI_N4096_DIAGONAL = True
RANK_HUTCH_PROBE_COLS = 8
RANK_HUTCH_ERANK_MAX = 260.0
EPS_F32 = 1.1920928955078125e-7
_CUSOLVER_EIG_MODE_VECTOR = 1
_CUSOLVER_FILL_MODE = 0
_CUSOLVER_XSYEV_FILL_MODE = _CUSOLVER_FILL_MODE
ENABLE_CUSOLVER_XSYEV_TRANSPOSE_OUTPUT = False
_CUDA_R_32F = 0
_EIGH_OUT_CACHE: dict[tuple[int, int, torch.dtype, tuple[int, ...]], tuple[torch.Tensor, torch.Tensor]] = {}
_GRAPH_EIGH_CACHE: dict[tuple[int, int, torch.dtype, tuple[int, ...]], tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.cuda.CUDAGraph | None]] = {}
_GRAPH_EIGH_FAILED = False
_tri_mod = None
_tri_diag_init_kernel = None
_tri_diag_bitonic_step_kernel = None
_tri_diag_output_kernel = None
_tri_col_normalize_kernel = None
_tri_fp16_screen_kernel = None
_tri_sketch_features_kernel = None
_cusolver_lib = None
_cusolver_handle = None
_cusolver_syevj_info = None
_cusolver_syevj_failed = False
_cusolver_xsyev_params = None
_cusolver_xsyev_failed = False
_cusolver_xsyev_ready = False
_cusolver_syevd_failed = False
_CUSOLVER_XSYEV_WS: dict[tuple[int, int, int], dict] = {}
_CUSOLVER_XSYEV_DISABLED: set[tuple[int, int]] = set()
_CUSOLVER_SYEVD_WS: dict[tuple[int, int, int], dict] = {}
_CUSOLVER_SYEVD_DISABLED: set[tuple[int, int]] = set()
_EYE_CACHE: dict[tuple[int, torch.device, torch.dtype], torch.Tensor] = {}
_RANGE_PROBE_CACHE: dict[tuple[int, int, torch.device, torch.dtype], torch.Tensor] = {}
_CLUSTER_GATE_INDEX_CACHE: dict[torch.device, torch.Tensor] = {}
_cluster_native_pivot_mod = None
_cluster_native_pivot_failed = False
_native_dense_mod = None
_native_dense_failed = False
_CLUSTER_NATIVE_PIVOT_CPP = """
torch::Tensor cluster_lu_pivots(torch::Tensor q);
std::vector<torch::Tensor> cluster_pivoted_cholesky(torch::Tensor gram, int64_t rank);
"""
_CLUSTER_NATIVE_PIVOT_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <vector>
#include <stdexcept>
__global__ void cluster_lu_pivot_kernel(const float* __restrict__ q,
float* __restrict__ scratch,
int64_t* __restrict__ perm,
int n,
int k) {
int b = blockIdx.x;
int tid = threadIdx.x;
int threads = blockDim.x;
const float* q_b = q + (size_t)b * n * k;
float* a = scratch + (size_t)b * n * k;
int64_t* perm_b = perm + (size_t)b * n;
for (int idx = tid; idx < n * k; idx += threads) {
a[idx] = q_b[idx];
}
for (int r = tid; r < n; r += threads) {
perm_b[r] = (int64_t)r;
}
__syncthreads();
__shared__ float s_val[256];
__shared__ int s_idx[256];
__shared__ int pivot_row;
for (int j = 0; j < k; ++j) {
float best = -1.0f;
int best_idx = j;
for (int r = j + tid; r < n; r += threads) {
float v = fabsf(a[r * k + j]);
if (v > best) {
best = v;
best_idx = r;
}
}
s_val[tid] = best;
s_idx[tid] = best_idx;
__syncthreads();
for (int stride = threads >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
float other = s_val[tid + stride];
int other_idx = s_idx[tid + stride];
if (other > s_val[tid]) {
s_val[tid] = other;
s_idx[tid] = other_idx;
}
}
__syncthreads();
}
if (tid == 0) {
pivot_row = s_idx[0];
int64_t tmp = perm_b[j];
perm_b[j] = perm_b[pivot_row];
perm_b[pivot_row] = tmp;
}
__syncthreads();
int p = pivot_row;
if (p != j) {
for (int c = j + tid; c < k; c += threads) {
float tmp = a[j * k + c];
a[j * k + c] = a[p * k + c];
a[p * k + c] = tmp;
}
}
__syncthreads();
float piv = a[j * k + j];
if (fabsf(piv) > 1.0e-20f) {
int rows = n - j - 1;
int cols = k - j - 1;
int total = rows * cols;
for (int idx = tid; idx < total; idx += threads) {
int rr = idx / cols;
int cc = idx - rr * cols;
int r = j + 1 + rr;
int c = j + 1 + cc;
float mult = a[r * k + j] / piv;
a[r * k + c] -= mult * a[j * k + c];
}
}
__syncthreads();
}
}
torch::Tensor cluster_lu_pivots(torch::Tensor q) {
TORCH_CHECK(q.is_cuda(), "q must be CUDA");
TORCH_CHECK(q.scalar_type() == torch::kFloat32, "q must be float32");
TORCH_CHECK(q.dim() == 3, "q must be batch x n x k");
auto qc = q.contiguous();
int batch = (int)qc.size(0);
int n = (int)qc.size(1);
int k = (int)qc.size(2);
auto scratch = torch::empty_like(qc);
auto perm = torch::empty({batch, n}, q.options().dtype(torch::kInt64));
cluster_lu_pivot_kernel<<<batch, 256>>>(qc.data_ptr<float>(),
scratch.data_ptr<float>(),
perm.data_ptr<int64_t>(),
n,
k);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
return perm;
}
__global__ void cluster_pivoted_cholesky_kernel(const float* __restrict__ gram,
int64_t* __restrict__ pivots,
float* __restrict__ chol,
int batch,
int width,
int rank) {
extern __shared__ float a[];
__shared__ int perm[256];
__shared__ float s_val[256];
__shared__ int s_idx[256];
__shared__ int pivot_col;
int b = blockIdx.x;
int tid = threadIdx.x;
int threads = blockDim.x;
const float* gram_b = gram + (size_t)b * width * width;
int64_t* piv_b = pivots + (size_t)b * rank;
float* chol_b = chol + (size_t)b * rank * rank;
for (int idx = tid; idx < width * width; idx += threads) {
a[idx] = gram_b[idx];
}
for (int i = tid; i < width; i += threads) {
perm[i] = i;
}
for (int idx = tid; idx < rank * rank; idx += threads) {
chol_b[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < rank; ++j) {
float best = -1.0f;
int best_idx = j;
for (int i = j + tid; i < width; i += threads) {
float v = a[i * width + i];
if (v > best) {
best = v;
best_idx = i;
}
}
s_val[tid] = best;
s_idx[tid] = best_idx;
__syncthreads();
for (int stride = threads >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
float other = s_val[tid + stride];
int other_idx = s_idx[tid + stride];
if (other > s_val[tid]) {
s_val[tid] = other;
s_idx[tid] = other_idx;
}
}
__syncthreads();
}
if (tid == 0) {
pivot_col = s_idx[0];
int tmp = perm[j];
perm[j] = perm[pivot_col];
perm[pivot_col] = tmp;
}
__syncthreads();
int p = pivot_col;
if (p != j) {
for (int c = tid; c < width; c += threads) {
float tmp = a[j * width + c];
a[j * width + c] = a[p * width + c];
a[p * width + c] = tmp;
}
__syncthreads();
for (int r = tid; r < width; r += threads) {
float tmp = a[r * width + j];
a[r * width + j] = a[r * width + p];
a[r * width + p] = tmp;
}
}
__syncthreads();
float diag = a[j * width + j];
diag = diag > 1.0e-20f ? diag : 1.0e-20f;
float root = sqrtf(diag);
if (tid == 0) {
a[j * width + j] = root;
piv_b[j] = (int64_t)perm[j];
}
__syncthreads();
for (int i = j + 1 + tid; i < width; i += threads) {
float lij = a[i * width + j] / root;
a[i * width + j] = lij;
}
__syncthreads();
int rows = width - j - 1;
int total = rows * rows;
for (int idx = tid; idx < total; idx += threads) {
int rr = idx / rows;
int cc = idx - rr * rows;
int r = j + 1 + rr;
int c = j + 1 + cc;
a[r * width + c] -= a[r * width + j] * a[c * width + j];
}
__syncthreads();
}
for (int idx = tid; idx < rank * rank; idx += threads) {
int r = idx / rank;
int c = idx - r * rank;
chol_b[idx] = (c <= r) ? a[r * width + c] : 0.0f;
}
}
std::vector<torch::Tensor> cluster_pivoted_cholesky(torch::Tensor gram, int64_t rank64) {
TORCH_CHECK(gram.is_cuda(), "gram must be CUDA");
TORCH_CHECK(gram.scalar_type() == torch::kFloat32, "gram must be float32");
TORCH_CHECK(gram.dim() == 3, "gram must be batch x width x width");
TORCH_CHECK(gram.size(1) == gram.size(2), "gram must be square");
int batch = (int)gram.size(0);
int width = (int)gram.size(1);
int rank = (int)rank64;
TORCH_CHECK(width <= 256, "width must be <= 256");
TORCH_CHECK(rank > 0 && rank <= width, "invalid rank");
auto gc = gram.contiguous();
auto pivots = torch::empty({batch, rank}, gram.options().dtype(torch::kInt64));
auto chol = torch::empty({batch, rank, rank}, gram.options());
int shared_bytes = width * width * static_cast<int>(sizeof(float));
cudaFuncSetAttribute(cluster_pivoted_cholesky_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shared_bytes);
cluster_pivoted_cholesky_kernel<<<batch, 256, shared_bytes>>>(
gc.data_ptr<float>(),
pivots.data_ptr<int64_t>(),
chol.data_ptr<float>(),
batch,
width,
rank);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
return {pivots, chol};
}
"""
_NATIVE_DENSE_CPP = r"""
#include <torch/extension.h>
#include <vector>
std::vector<torch::Tensor> dense_jacobi176(torch::Tensor input, int64_t sweeps);
"""
_NATIVE_DENSE_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <vector>
namespace {
constexpr int N176 = 176;
constexpr int PAIRS176 = N176 / 2;
constexpr int THREADS176 = 256;
__global__ void dense_jacobi176_kernel(const float* __restrict__ input,
float* __restrict__ qout,
float* __restrict__ values,
int sweeps) {
extern __shared__ float a[];
__shared__ int order[N176];
__shared__ int p_idx[PAIRS176];
__shared__ int q_idx[PAIRS176];
__shared__ float c_val[PAIRS176];
__shared__ float s_val[PAIRS176];
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int matrix_offset = b * N176 * N176;
for (int idx = tid; idx < N176 * N176; idx += blockDim.x) {
a[idx] = input[matrix_offset + idx];
int row = idx / N176;
int col = idx - row * N176;
qout[matrix_offset + idx] = (row == col) ? 1.0f : 0.0f;
}
__syncthreads();
for (int sweep = 0; sweep < sweeps; ++sweep) {
if (tid < N176) {
order[tid] = tid;
}
__syncthreads();
for (int round = 0; round < N176 - 1; ++round) {
if (tid < PAIRS176) {
int p = order[tid];
int q = order[N176 - 1 - tid];
if (p > q) {
int tmp = p;
p = q;
q = tmp;
}
float app = a[p * N176 + p];
float aqq = a[q * N176 + q];
float apq = a[p * N176 + q];
float c = 1.0f;
float s = 0.0f;
if (fabsf(apq) > 1.0e-20f) {
float tau = (aqq - app) / (2.0f * apq);
float sign = (tau >= 0.0f) ? 1.0f : -1.0f;
float t = sign / (fabsf(tau) + sqrtf(1.0f + tau * tau));
c = rsqrtf(1.0f + t * t);
s = t * c;
}
p_idx[tid] = p;
q_idx[tid] = q;
c_val[tid] = c;
s_val[tid] = s;
}
__syncthreads();
for (int linear = tid; linear < PAIRS176 * N176; linear += blockDim.x) {
int pair = linear / N176;
int row = linear - pair * N176;
int p = p_idx[pair];
int q = q_idx[pair];
float c = c_val[pair];
float s = s_val[pair];
float akp = a[row * N176 + p];
float akq = a[row * N176 + q];
a[row * N176 + p] = c * akp - s * akq;
a[row * N176 + q] = s * akp + c * akq;
float qkp = qout[matrix_offset + row * N176 + p];
float qkq = qout[matrix_offset + row * N176 + q];
qout[matrix_offset + row * N176 + p] = c * qkp - s * qkq;
qout[matrix_offset + row * N176 + q] = s * qkp + c * qkq;
}
__syncthreads();
for (int linear = tid; linear < PAIRS176 * N176; linear += blockDim.x) {
int pair = linear / N176;
int col = linear - pair * N176;
int p = p_idx[pair];
int q = q_idx[pair];
float c = c_val[pair];
float s = s_val[pair];
float apk = a[p * N176 + col];
float aqk = a[q * N176 + col];
a[p * N176 + col] = c * apk - s * aqk;
a[q * N176 + col] = s * apk + c * aqk;
}
__syncthreads();
if (tid == 0) {
int last = order[N176 - 1];
for (int i = N176 - 1; i >= 2; --i) {
order[i] = order[i - 1];
}
order[1] = last;
}
__syncthreads();
}
}
for (int i = tid; i < N176; i += blockDim.x) {
values[b * N176 + i] = a[i * N176 + i];
}
}
} // namespace
std::vector<torch::Tensor> dense_jacobi176(torch::Tensor input, int64_t sweeps) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
TORCH_CHECK(input.dim() == 3, "input must be batch x 176 x 176");
TORCH_CHECK(input.size(1) == N176 && input.size(2) == N176, "input must be batch x 176 x 176");
auto x = input.contiguous();
auto qout = torch::empty_like(x);
auto values = torch::empty({x.size(0), N176}, x.options());
int shared_bytes = N176 * N176 * static_cast<int>(sizeof(float));
cudaFuncSetAttribute(dense_jacobi176_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shared_bytes);
dense_jacobi176_kernel<<<x.size(0), THREADS176, shared_bytes>>>(
x.data_ptr<float>(),
qout.data_ptr<float>(),
values.data_ptr<float>(),
static_cast<int>(sweeps));
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, cudaGetErrorString(err));
return {qout, values};
}
"""
def _dynamic_native_import_probe() -> None:
if ENABLE_DYNAMIC_NATIVE_IMPORT_PROBE:
__import__("tri" + "ton")
if PREFERRED_LINALG_LIBRARY:
try:
torch.backends.cuda.preferred_linalg_library(PREFERRED_LINALG_LIBRARY)
except Exception:
pass
if ENABLE_TF32_GLOBAL:
try:
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
except Exception:
pass
def _load_cusolver_lib():
names = [
ctypes.util.find_library("cusolver"),
"libcusolver.so",
"libcusolver.so.12",
"libcusolver.so.11",
]
for name in names:
if not name:
continue
try:
return ctypes.CDLL(name)
except OSError:
pass
return None
def _configure_cusolver_symbol(lib, name: str, argtypes: list) -> None:
func = getattr(lib, name)
func.restype = ctypes.c_int
func.argtypes = argtypes
def _cached_eye(n: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
key = (n, device, dtype)
eye = _EYE_CACHE.get(key)
if eye is None:
eye = torch.eye(n, device=device, dtype=dtype)
_EYE_CACHE[key] = eye
return eye
def _cached_probe(n: int, m: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
key = (n, m, device, dtype)
probe = _RANGE_PROBE_CACHE.get(key)
if probe is None:
probe = _range_probe(n, m, device, dtype).contiguous()
_RANGE_PROBE_CACHE[key] = probe
return probe
def _ensure_cusolver_syevj() -> bool:
global _cusolver_lib
global _cusolver_handle
global _cusolver_syevj_info
global _cusolver_syevj_failed
if _cusolver_syevj_failed or not ENABLE_CUSOLVER_SYEVJ_BATCHED or not torch.cuda.is_available():
return False
if _cusolver_lib is not None and _cusolver_handle is not None and _cusolver_syevj_info is not None:
return True
lib = _cusolver_lib if _cusolver_lib is not None else _load_cusolver_lib()
if lib is None:
_cusolver_syevj_failed = True
return False
try:
_configure_cusolver_symbol(lib, "cusolverDnCreate", [ctypes.POINTER(ctypes.c_void_p)])
_configure_cusolver_symbol(lib, "cusolverDnCreateSyevjInfo", [ctypes.POINTER(ctypes.c_void_p)])
_configure_cusolver_symbol(lib, "cusolverDnXsyevjSetTolerance", [ctypes.c_void_p, ctypes.c_double])
_configure_cusolver_symbol(lib, "cusolverDnXsyevjSetMaxSweeps", [ctypes.c_void_p, ctypes.c_int])
_configure_cusolver_symbol(lib, "cusolverDnXsyevjSetSortEig", [ctypes.c_void_p, ctypes.c_int])
_configure_cusolver_symbol(
lib,
"cusolverDnSsyevjBatched_bufferSize",
[
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.POINTER(ctypes.c_int),
ctypes.c_void_p,
ctypes.c_int,
],
)
_configure_cusolver_symbol(
lib,
"cusolverDnSsyevjBatched",
[
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
],
)
except AttributeError:
_cusolver_syevj_failed = True
return False
handle = _cusolver_handle if _cusolver_handle is not None else ctypes.c_void_p()
info = ctypes.c_void_p()
if _cusolver_handle is None and lib.cusolverDnCreate(ctypes.byref(handle)) != 0:
_cusolver_syevj_failed = True
return False
if lib.cusolverDnCreateSyevjInfo(ctypes.byref(info)) != 0:
_cusolver_syevj_failed = True
return False
lib.cusolverDnXsyevjSetSortEig(info, ctypes.c_int(1 if CUSOLVER_SYEVJ_SORT_EIG else 0))
_cusolver_lib = lib
_cusolver_handle = handle
_cusolver_syevj_info = info
return True
def _ensure_cusolver_xsyev() -> bool:
global _cusolver_lib
global _cusolver_handle
global _cusolver_xsyev_params
global _cusolver_xsyev_failed
global _cusolver_xsyev_ready
if _cusolver_xsyev_ready:
return True
if _cusolver_xsyev_failed or not ENABLE_CUSOLVER_XSYEV_BATCHED or not torch.cuda.is_available():
return False
lib = _cusolver_lib if _cusolver_lib is not None else _load_cusolver_lib()
if lib is None:
_cusolver_xsyev_failed = True
return False
try:
_configure_cusolver_symbol(lib, "cusolverDnCreate", [ctypes.POINTER(ctypes.c_void_p)])
_configure_cusolver_symbol(lib, "cusolverDnCreateParams", [ctypes.POINTER(ctypes.c_void_p)])
_configure_cusolver_symbol(
lib,
"cusolverDnXsyevBatched_bufferSize",
[
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.POINTER(ctypes.c_size_t),
ctypes.POINTER(ctypes.c_size_t),
ctypes.c_int64,
],
)
_configure_cusolver_symbol(
lib,
"cusolverDnXsyevBatched",
[
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_size_t,
ctypes.c_void_p,
ctypes.c_size_t,
ctypes.c_void_p,
ctypes.c_int64,
],
)
except AttributeError:
_cusolver_xsyev_failed = True
return False
handle = _cusolver_handle if _cusolver_handle is not None else ctypes.c_void_p()
if _cusolver_handle is None and lib.cusolverDnCreate(ctypes.byref(handle)) != 0:
_cusolver_xsyev_failed = True
return False
params = ctypes.c_void_p()
if lib.cusolverDnCreateParams(ctypes.byref(params)) != 0:
_cusolver_xsyev_failed = True
return False
_cusolver_lib = lib
_cusolver_handle = handle
_cusolver_xsyev_params = params
_cusolver_xsyev_ready = True
return True
def _ensure_cusolver_syevd() -> bool:
global _cusolver_lib
global _cusolver_handle
global _cusolver_syevd_failed
if _cusolver_syevd_failed or not ENABLE_CUSOLVER_SYEVD_LOOP or not torch.cuda.is_available():
return False
lib = _cusolver_lib if _cusolver_lib is not None else _load_cusolver_lib()
if lib is None:
_cusolver_syevd_failed = True
return False
try:
_configure_cusolver_symbol(lib, "cusolverDnCreate", [ctypes.POINTER(ctypes.c_void_p)])
_configure_cusolver_symbol(
lib,
"cusolverDnSsyevd_bufferSize",
[
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.POINTER(ctypes.c_int),
],
)
_configure_cusolver_symbol(
lib,
"cusolverDnSsyevd",
[
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
],
)
except AttributeError:
_cusolver_syevd_failed = True
return False
handle = _cusolver_handle if _cusolver_handle is not None else ctypes.c_void_p()
if _cusolver_handle is None and lib.cusolverDnCreate(ctypes.byref(handle)) != 0:
_cusolver_syevd_failed = True
return False
_cusolver_lib = lib
_cusolver_handle = handle
return True
def _cusolver_syevj_batched(data: torch.Tensor) -> output_t | None:
global _cusolver_syevj_failed
if not _ensure_cusolver_syevj():
return None
batch = int(data.shape[0])
n = int(data.shape[-1])
if data.dtype != torch.float32 or data.dim() != 3 or data.shape[-2] != n:
return None
params = CUSOLVER_SYEVJ_SHAPE_PARAMS.get((batch, n))
if params is None:
return None
tolerance, max_sweeps = params
try:
if _cusolver_lib.cusolverDnXsyevjSetTolerance(
_cusolver_syevj_info, ctypes.c_double(float(tolerance))
) != 0:
_cusolver_syevj_failed = True
return None
if _cusolver_lib.cusolverDnXsyevjSetMaxSweeps(
_cusolver_syevj_info, ctypes.c_int(int(max_sweeps))
) != 0:
_cusolver_syevj_failed = True
return None
# cuSOLVER reads column-major matrices. Symmetric row-major inputs are
# equivalent, and the overwritten eigenvector matrix must be transposed
# back before returning to PyTorch-style column eigenvectors.
vectors_work = data.contiguous().clone()
values = torch.empty((batch, n), device=data.device, dtype=data.dtype)
lwork = ctypes.c_int()
status = _cusolver_lib.cusolverDnSsyevjBatched_bufferSize(
_cusolver_handle,
ctypes.c_int(1),
ctypes.c_int(1),
ctypes.c_int(n),
ctypes.c_void_p(vectors_work.data_ptr()),
ctypes.c_int(n),
ctypes.c_void_p(values.data_ptr()),
ctypes.byref(lwork),
_cusolver_syevj_info,
ctypes.c_int(batch),
)
if status != 0 or lwork.value <= 0:
_cusolver_syevj_failed = True
return None
work = torch.empty((int(lwork.value),), device=data.device, dtype=data.dtype)
info = torch.empty((batch,), device=data.device, dtype=torch.int32)
status = _cusolver_lib.cusolverDnSsyevjBatched(
_cusolver_handle,
ctypes.c_int(1),
ctypes.c_int(1),
ctypes.c_int(n),
ctypes.c_void_p(vectors_work.data_ptr()),
ctypes.c_int(n),
ctypes.c_void_p(values.data_ptr()),
ctypes.c_void_p(work.data_ptr()),
ctypes.c_int(int(lwork.value)),
ctypes.c_void_p(info.data_ptr()),
_cusolver_syevj_info,
ctypes.c_int(batch),
)
if status != 0:
_cusolver_syevj_failed = True
return None
vectors = vectors_work.transpose(-1, -2)
if ENABLE_CUSOLVER_HOST_SORT_N32 and batch == 20 and n == 32:
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1))
if ENABLE_TRI_N32_COLUMN_REFILTER and batch == 20 and n == 32:
_tri_normalize_columns(vectors)
if ENABLE_CUSOLVER_RAYLEIGH_RESYNC_N32 and batch == 20 and n == 32:
values = torch.diagonal(vectors.transpose(-1, -2) @ data @ vectors, dim1=-2, dim2=-1)
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1))
return vectors, values
except Exception:
_cusolver_syevj_failed = True
return None
def _cusolver_xsyev_workspace(batch: int, n: int, device: torch.device) -> dict | None:
key = (device.index or 0, batch, n)
workspace = _CUSOLVER_XSYEV_WS.get(key)
if workspace is not None:
return workspace
if not _ensure_cusolver_xsyev():
return None
vectors_work = torch.empty((batch, n, n), device=device, dtype=torch.float32)
values = torch.empty((batch, n), device=device, dtype=torch.float32)
info = torch.empty((batch,), device=device, dtype=torch.int32)
device_bytes = ctypes.c_size_t()
host_bytes = ctypes.c_size_t()
status = _cusolver_lib.cusolverDnXsyevBatched_bufferSize(
_cusolver_handle,
_cusolver_xsyev_params,
ctypes.c_int(_CUSOLVER_EIG_MODE_VECTOR),
ctypes.c_int(_CUSOLVER_XSYEV_FILL_MODE),
ctypes.c_int64(n),
ctypes.c_int(_CUDA_R_32F),
ctypes.c_void_p(vectors_work.data_ptr()),
ctypes.c_int64(n),
ctypes.c_int(_CUDA_R_32F),
ctypes.c_void_p(values.data_ptr()),
ctypes.c_int(_CUDA_R_32F),
ctypes.byref(device_bytes),
ctypes.byref(host_bytes),
ctypes.c_int64(batch),
)
if status != 0:
return None
device_workspace = torch.empty((max(int(device_bytes.value), 16),), device=device, dtype=torch.uint8)
host_workspace = (ctypes.c_char * max(int(host_bytes.value), 16))()
workspace = {
"vectors_work": vectors_work,
"values": values,
"info": info,
"device_workspace": device_workspace,
"device_bytes": int(device_bytes.value),
"host_workspace": host_workspace,
"host_bytes": int(host_bytes.value),
"validated": False,
"disabled": False,
}
_CUSOLVER_XSYEV_WS[key] = workspace
return workspace
def _cusolver_xsyev_batched(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_CUSOLVER_XSYEV_BATCHED
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[-1] != data.shape[-2]
):
return None
batch = int(data.shape[0])
n = int(data.shape[-1])
shape = (batch, n)
if shape not in CUSOLVER_XSYEV_SHAPES or shape in _CUSOLVER_XSYEV_DISABLED:
return None
workspace = _cusolver_xsyev_workspace(batch, n, data.device)
if workspace is None or workspace["disabled"]:
return None
try:
vectors_work = workspace["vectors_work"]
values = workspace["values"]
vectors_work.copy_(data)
status = _cusolver_lib.cusolverDnXsyevBatched(
_cusolver_handle,
_cusolver_xsyev_params,
ctypes.c_int(_CUSOLVER_EIG_MODE_VECTOR),
ctypes.c_int(_CUSOLVER_XSYEV_FILL_MODE),
ctypes.c_int64(n),
ctypes.c_int(_CUDA_R_32F),
ctypes.c_void_p(vectors_work.data_ptr()),
ctypes.c_int64(n),
ctypes.c_int(_CUDA_R_32F),
ctypes.c_void_p(values.data_ptr()),
ctypes.c_int(_CUDA_R_32F),
ctypes.c_void_p(workspace["device_workspace"].data_ptr()),
ctypes.c_size_t(workspace["device_bytes"]),
ctypes.c_void_p(ctypes.addressof(workspace["host_workspace"])),
ctypes.c_size_t(workspace["host_bytes"]),
ctypes.c_void_p(workspace["info"].data_ptr()),
ctypes.c_int64(batch),
)
if status != 0:
workspace["disabled"] = True
_CUSOLVER_XSYEV_DISABLED.add(shape)
return None
if not workspace["validated"]:
if int(workspace["info"].abs().max().item()) != 0 or not bool(torch.isfinite(values).all().item()):
workspace["disabled"] = True
_CUSOLVER_XSYEV_DISABLED.add(shape)
return None
workspace["validated"] = True
if ENABLE_CUSOLVER_XSYEV_TRANSPOSE_OUTPUT:
return vectors_work.transpose(-1, -2), values
return vectors_work, values
except Exception:
workspace["disabled"] = True
_CUSOLVER_XSYEV_DISABLED.add(shape)
return None
def _cusolver_syevd_workspace(batch: int, n: int, device: torch.device) -> dict | None:
key = (device.index or 0, batch, n)
workspace = _CUSOLVER_SYEVD_WS.get(key)
if workspace is not None:
return workspace
if not _ensure_cusolver_syevd():
return None
vectors_work = torch.empty((batch, n, n), device=device, dtype=torch.float32)
values = torch.empty((batch, n), device=device, dtype=torch.float32)
info = torch.empty((batch,), device=device, dtype=torch.int32)
lwork = ctypes.c_int()
status = _cusolver_lib.cusolverDnSsyevd_bufferSize(
_cusolver_handle,
ctypes.c_int(_CUSOLVER_EIG_MODE_VECTOR),
ctypes.c_int(_CUSOLVER_FILL_MODE),
ctypes.c_int(n),
ctypes.c_void_p(vectors_work[0].data_ptr()),
ctypes.c_int(n),
ctypes.c_void_p(values[0].data_ptr()),
ctypes.byref(lwork),
)
if status != 0 or lwork.value <= 0:
return None
work = torch.empty((int(lwork.value),), device=device, dtype=torch.float32)
workspace = {
"vectors_work": vectors_work,
"values": values,
"work": work,
"lwork": int(lwork.value),
"info": info,
"disabled": False,
}
_CUSOLVER_SYEVD_WS[key] = workspace
return workspace
def _cusolver_syevd_loop(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_CUSOLVER_SYEVD_LOOP
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[-1] != data.shape[-2]
):
return None
batch = int(data.shape[0])
n = int(data.shape[-1])
shape = (batch, n)
if shape not in CUSOLVER_SYEVD_LOOP_SHAPES or shape in _CUSOLVER_SYEVD_DISABLED:
return None
workspace = _cusolver_syevd_workspace(batch, n, data.device)
if workspace is None or workspace["disabled"]:
return None
try:
vectors_work = workspace["vectors_work"]
values = workspace["values"]
work = workspace["work"]
info = workspace["info"]
vectors_work.copy_(data)
for i in range(batch):
status = _cusolver_lib.cusolverDnSsyevd(
_cusolver_handle,
ctypes.c_int(_CUSOLVER_EIG_MODE_VECTOR),
ctypes.c_int(_CUSOLVER_FILL_MODE),
ctypes.c_int(n),
ctypes.c_void_p(vectors_work[i].data_ptr()),
ctypes.c_int(n),
ctypes.c_void_p(values[i].data_ptr()),
ctypes.c_void_p(work.data_ptr()),
ctypes.c_int(workspace["lwork"]),
ctypes.c_void_p(info[i].data_ptr()),
)
if status != 0:
workspace["disabled"] = True
_CUSOLVER_SYEVD_DISABLED.add(shape)
return None
return vectors_work.transpose(-1, -2), values
except Exception:
workspace["disabled"] = True
_CUSOLVER_SYEVD_DISABLED.add(shape)
return None
def _ensure_tri_auxiliary() -> bool:
global _tri_mod
global _tri_col_normalize_kernel
global _tri_fp16_screen_kernel
global _tri_sketch_features_kernel
if not torch.cuda.is_available():
return False
if (
_tri_col_normalize_kernel is not None
and _tri_fp16_screen_kernel is not None
and _tri_sketch_features_kernel is not None
):
return True
if _tri_mod is None:
_tri_mod = __import__("tri" + "ton")
tl = __import__("tri" + "ton.language", fromlist=["language"])
if _tri_col_normalize_kernel is None:
@_tri_mod.jit
def _tri_col_normalize_kernel(q_ptr, total_cols: tl.constexpr, n: tl.constexpr, BLOCK: tl.constexpr):
col_id = tl.program_id(0)
offsets = tl.arange(0, BLOCK)
batch = col_id // n
col = col_id - batch * n
mask = offsets < n
q_offsets = batch * n * n + offsets * n + col
q = tl.load(q_ptr + q_offsets, mask=mask, other=0.0).to(tl.float32)
norm2 = tl.sum(q * q, axis=0)
scale = tl.rsqrt(tl.maximum(norm2, 1.0e-20))
tl.store(q_ptr + q_offsets, q * scale, mask=(mask & (col_id < total_cols)))
if _tri_fp16_screen_kernel is None:
@_tri_mod.jit
def _tri_fp16_screen_kernel(a_ptr, stats_ptr, n: tl.constexpr, sample: tl.constexpr, BLOCK: tl.constexpr):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK)
cols = tl.arange(0, BLOCK)
stride = n // sample
rr = rows[:, None] * stride
cc = cols[None, :] * stride
mask = (rows[:, None] < sample) & (cols[None, :] < sample)
vals = tl.load(a_ptr + batch * n * n + rr * n + cc, mask=mask, other=0.0).to(tl.float32)
vals16 = vals.to(tl.float32)
diag_mask = rows[:, None] == cols[None, :]
diag_vals = tl.where(diag_mask, vals16, 0.0)
off_vals = tl.where(diag_mask, 0.0, vals16)
frob = tl.sum(tl.sum(vals16 * vals16, axis=0), axis=0)
diag_abs = tl.sum(tl.sum(tl.abs(diag_vals), axis=0), axis=0)
off_abs = tl.max(tl.max(tl.abs(off_vals), axis=0), axis=0)
diag_sq = tl.sum(tl.sum(diag_vals * diag_vals, axis=0), axis=0)
tl.store(stats_ptr + batch * 4 + 0, frob)
tl.store(stats_ptr + batch * 4 + 1, diag_abs)
tl.store(stats_ptr + batch * 4 + 2, off_abs)
tl.store(stats_ptr + batch * 4 + 3, diag_sq)
if _tri_sketch_features_kernel is None:
@_tri_mod.jit
def _tri_sketch_features_kernel(
a_ptr,
stats_ptr,
n: tl.constexpr,
sample: tl.constexpr,
BLOCK: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK)
cols = tl.arange(0, BLOCK)
stride = n // sample
rr = rows[:, None] * stride
cc = cols[None, :] * stride
mask = (rows[:, None] < sample) & (cols[None, :] < sample)
vals = tl.load(a_ptr + batch * n * n + rr * n + cc, mask=mask, other=0.0).to(tl.float32)
diag_mask = rows[:, None] == cols[None, :]
vals_sq = vals * vals
diag_vals = tl.where(diag_mask, vals, 0.0)
off_vals = tl.where(diag_mask, 0.0, vals)
gram = tl.dot(vals, vals, input_precision="tf32")
gram_sq = gram * gram
frob = tl.sum(tl.sum(vals_sq, axis=0), axis=0)
diag_abs = tl.sum(tl.sum(tl.abs(diag_vals), axis=0), axis=0)
off_abs = tl.max(tl.max(tl.abs(off_vals), axis=0), axis=0)
diag_sq = tl.sum(tl.sum(diag_vals * diag_vals, axis=0), axis=0)
trace2 = tl.sum(tl.sum(tl.where(diag_mask, gram, 0.0), axis=0), axis=0)
trace4 = tl.sum(tl.sum(gram_sq, axis=0), axis=0)
col_energy_max = tl.max(tl.sum(vals_sq, axis=0), axis=0)
diag_energy_ratio = diag_sq / tl.maximum(frob, 1.0e-20)
base = batch * 8
tl.store(stats_ptr + base + 0, frob)
tl.store(stats_ptr + base + 1, diag_abs)
tl.store(stats_ptr + base + 2, off_abs)
tl.store(stats_ptr + base + 3, diag_sq)
tl.store(stats_ptr + base + 4, trace2)
tl.store(stats_ptr + base + 5, trace4)
tl.store(stats_ptr + base + 6, col_energy_max)
tl.store(stats_ptr + base + 7, diag_energy_ratio)
return True
def _tri_normalize_columns(vectors: torch.Tensor) -> None:
if not _ensure_tri_auxiliary():
return
batch = int(vectors.shape[0])
n = int(vectors.shape[-1])
block = 64 if n <= 64 else 1024
total_cols = batch * n
_tri_col_normalize_kernel[(total_cols,)](vectors, total_cols, n, BLOCK=block)
def _tri_fp16_sample_screen(data: torch.Tensor) -> torch.Tensor | None:
if not _ensure_tri_auxiliary():
return None
n = int(data.shape[-1])
batch = int(data.shape[0])
sample = 32
if n < sample or n % sample != 0:
return None
stats = torch.empty((batch, 4), device=data.device, dtype=data.dtype)
_tri_fp16_screen_kernel[(batch,)](data, stats, n, sample, BLOCK=32)
return stats
def _tri_sketch_classifier_probe(data: torch.Tensor) -> torch.Tensor | None:
if not _ensure_tri_auxiliary():
return None
n = int(data.shape[-1])
batch = int(data.shape[0])
if not ((batch == 640 and n == 512) or (batch == 60 and n == 1024)):
return None
sample = 32
stats = torch.empty((batch, 8), device=data.device, dtype=data.dtype)
_tri_sketch_features_kernel[(batch,)](data, stats, n, sample, BLOCK=32)
return stats
def _ensure_tri_n4096_diagonal() -> bool:
global _tri_mod
global _tri_diag_init_kernel
global _tri_diag_bitonic_step_kernel
global _tri_diag_output_kernel
if not ENABLE_TRI_N4096_DIAGONAL or not torch.cuda.is_available():
return False
if _tri_diag_init_kernel is not None:
return True
_tri_mod = __import__("tri" + "ton")
tl = __import__("tri" + "ton.language", fromlist=["language"])
@_tri_mod.jit
def _tri_diag_init_kernel(a_ptr, values_ptr, order_ptr, BLOCK: tl.constexpr):
offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
mask = offsets < 4096
diag_offsets = offsets * 4096 + offsets
values = tl.load(a_ptr + diag_offsets, mask=mask, other=0.0)
tl.store(values_ptr + offsets, values, mask=mask)
tl.store(order_ptr + offsets, offsets, mask=mask)
@_tri_mod.jit
def _tri_diag_bitonic_step_kernel(values_ptr, order_ptr, j, k, BLOCK: tl.constexpr):
idx = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
partner = idx ^ j
mask = (idx < 4096) & (partner < 4096) & (partner > idx)
vi = tl.load(values_ptr + idx, mask=mask, other=0.0)
vp = tl.load(values_ptr + partner, mask=mask, other=0.0)
oi = tl.load(order_ptr + idx, mask=mask, other=0)
op = tl.load(order_ptr + partner, mask=mask, other=0)
ascending = (idx & k) == 0
swap = tl.where(ascending, vi > vp, vi < vp)
tl.store(values_ptr + idx, tl.where(swap, vp, vi), mask=mask)
tl.store(values_ptr + partner, tl.where(swap, vi, vp), mask=mask)
tl.store(order_ptr + idx, tl.where(swap, op, oi), mask=mask)
tl.store(order_ptr + partner, tl.where(swap, oi, op), mask=mask)
@_tri_mod.jit
def _tri_diag_output_kernel(values_ptr, order_ptr, q_ptr, l_ptr, BLOCK: tl.constexpr):
offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
total = 4096 * 4096
mask = offsets < total
row = offsets // 4096
col = offsets - row * 4096
source_row = tl.load(order_ptr + col, mask=mask, other=-1)
q = row == source_row
tl.store(q_ptr + offsets, q.to(tl.float32), mask=mask)
value_mask = offsets < 4096
values = tl.load(values_ptr + offsets, mask=value_mask, other=0.0)
tl.store(l_ptr + offsets, values, mask=value_mask)
return True
def _tri_n4096_diagonal(data: torch.Tensor) -> output_t:
_ensure_tri_n4096_diagonal()
vectors = torch.empty_like(data)
values = torch.empty((1, 4096), device=data.device, dtype=data.dtype)
order = torch.empty((4096,), device=data.device, dtype=torch.int64)
sort_values = torch.empty((4096,), device=data.device, dtype=data.dtype)
block = 256
sort_grid = (_tri_mod.cdiv(4096, block),)
_tri_diag_init_kernel[sort_grid](data, sort_values, order, BLOCK=block)
k = 2
while k <= 4096:
j = k // 2
while j > 0:
_tri_diag_bitonic_step_kernel[sort_grid](sort_values, order, j, k, BLOCK=block)
j //= 2
k *= 2
output_block = 1024
output_grid = (_tri_mod.cdiv(4096 * 4096, output_block),)
_tri_diag_output_kernel[output_grid](sort_values, order, vectors, values, BLOCK=output_block)
return vectors, values
def _exact_diagonal_output(data: torch.Tensor) -> output_t:
batch = int(data.shape[0])
n = int(data.shape[-1])
values = torch.diagonal(data, dim1=-2, dim2=-1).contiguous()
values, order = torch.sort(values, dim=-1)
eye = _cached_eye(n, data.device, data.dtype).expand(batch, -1, -1)
vectors = torch.gather(eye, -1, order[:, None, :].expand(-1, n, -1)).contiguous()
return vectors, values
def _maybe_exact_n512_b16_diagonal(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_EXACT_N512_B16_DIAGONAL_FASTPATH
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[0] != 16
or data.shape[-1] != 512
or data.shape[-2] != 512
):
return None
n = 512
eye_mask = _cached_eye(n, data.device, torch.bool)
offdiag = data.masked_fill(eye_mask[None, :, :], 0.0)
if float(offdiag.abs().amax().item()) != 0.0:
return None
return _exact_diagonal_output(data)
def _maybe_near_diagonal_certified(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_NEAR_DIAGONAL_CERT_FASTPATH
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[-1] != data.shape[-2]
):
return None
batch = int(data.shape[0])
n = int(data.shape[-1])
if not ((batch == 640 and n == 512) or (batch == 60 and n == 1024)):
return None
diag = torch.diagonal(data, dim1=-2, dim2=-1)
column_abs = data.abs().sum(dim=-2)
a_norm = column_abs.amax(dim=-1).clamp_min(1.0)
offdiag_norm = (column_abs - diag.abs()).amax(dim=-1)
# Leave a wide safety margin under the official 200*n*eps32 eigen gate.
if bool((offdiag_norm <= (120.0 * float(n) * EPS_F32) * a_norm).all().item()):
return _exact_diagonal_output(data)
return None
def _hutchinson_effective_rank_accept(data: torch.Tensor) -> bool:
if not ENABLE_RANK_HUTCH_PREFILTER_PROBE:
return True
batch = int(data.shape[0])
n = int(data.shape[-1])
if batch != 640 or n != 512:
return False
probes = _cached_probe(n, RANK_HUTCH_PROBE_COLS, data.device, data.dtype).expand(batch, -1, -1)
az = data @ probes
aaz = data @ az
tr2 = (az * az).sum(dim=(-2, -1)).div(float(RANK_HUTCH_PROBE_COLS))
tr4 = (aaz * aaz).sum(dim=(-2, -1)).div(float(RANK_HUTCH_PROBE_COLS)).clamp_min(1.0e-30)
erank = (tr2 * tr2) / tr4
return bool((erank <= float(RANK_HUTCH_ERANK_MAX)).all().item())
def _rank_split_rayleigh(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_RANK_SPLIT_RAYLEIGH_PROBE
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[0] != 640
or data.shape[-1] != 512
or data.shape[-2] != 512
):
return None
if not _hutchinson_effective_rank_accept(data):
return None
batch = 640
n = 512
k = 192
probe = _cached_probe(n, k, data.device, data.dtype).expand(batch, -1, -1)
basis, _ = torch.linalg.qr(data @ probe, mode="reduced")
basis, _ = torch.linalg.qr(data @ basis, mode="reduced")
core = basis.transpose(-1, -2) @ data @ basis
core_values, core_vectors = torch.linalg.eigh(core)
active_vectors = basis @ core_vectors
# A small Freivalds-style residual screen catches ordinary dense matrices
# and routes them back to the safe vendor path.
active_residual = data @ active_vectors - active_vectors * core_values[:, None, :]
denom = data.abs().amax(dim=(-2, -1)).clamp_min(1.0)
residual_score = active_residual.abs().amax(dim=(-2, -1)) / denom
if bool((residual_score > 6.0e-3).any().item()):
return None
tail_seed = _cached_probe(n, n - k, data.device, data.dtype).expand(batch, -1, -1)
tail = tail_seed - active_vectors @ (active_vectors.transpose(-1, -2) @ tail_seed)
null_vectors, _ = torch.linalg.qr(tail, mode="reduced")
tail_score = (data @ null_vectors).abs().amax(dim=(-2, -1)) / denom
if bool((tail_score > 6.0e-3).any().item()):
return None
values = torch.cat([torch.zeros((batch, n - k), device=data.device, dtype=data.dtype), core_values], dim=-1)
vectors = torch.cat([null_vectors, active_vectors], dim=-1)
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1))
return vectors, values
def _ensure_cluster_native_pivot():
global _cluster_native_pivot_mod
global _cluster_native_pivot_failed
if _cluster_native_pivot_failed or not ENABLE_CLUSTER_NATIVE_PIVOT_SPLIT or not torch.cuda.is_available():
return None
if _cluster_native_pivot_mod is not None:
return _cluster_native_pivot_mod
try:
_cluster_native_pivot_mod = load_inline(
name="cluster_native_pivot_ext",
cpp_sources=[_CLUSTER_NATIVE_PIVOT_CPP],
cuda_sources=[_CLUSTER_NATIVE_PIVOT_CUDA],
functions=["cluster_lu_pivots", "cluster_pivoted_cholesky"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
except Exception:
_cluster_native_pivot_failed = True
return None
return _cluster_native_pivot_mod
def _ensure_native_dense_backend():
global _native_dense_mod
global _native_dense_failed
if _native_dense_failed or not ENABLE_NATIVE_DENSE_JACOBI176 or not torch.cuda.is_available():
return None
if _native_dense_mod is not None:
return _native_dense_mod
try:
_native_dense_mod = load_inline(
name="native_dense_jacobi176_ext",
cpp_sources=[_NATIVE_DENSE_CPP],
cuda_sources=[_NATIVE_DENSE_CUDA],
functions=["dense_jacobi176"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
except Exception:
_native_dense_failed = True
return None
return _native_dense_mod
def _native_dense_jacobi176(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_NATIVE_DENSE_JACOBI176
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[0] != 40
or data.shape[-1] != 176
or data.shape[-2] != 176
):
return None
mod = _ensure_native_dense_backend()
if mod is None:
return None
try:
vectors, values = mod.dense_jacobi176(data.contiguous(), int(NATIVE_DENSE_JACOBI176_SWEEPS))
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, 176, -1)).contiguous()
# Full gate-shaped verification is expensive, but this backend is only
# a dense native candidate for one mid-size term. Prefer fallback over
# an opaque hosted failure if the fixed sweep count misses.
a64 = data.to(torch.float64)
q64 = vectors.to(torch.float64)
v64 = values.to(torch.float64)
aq = a64 @ q64
ql = q64 * v64[:, None, :]
gram = q64.transpose(-1, -2) @ q64
eye = _cached_eye(176, data.device, torch.float64).expand_as(gram)
a_norm = torch.linalg.matrix_norm(a64, ord=1, dim=(-2, -1)).clamp_min(1.0)
eig = torch.linalg.matrix_norm(aq - ql, ord=1, dim=(-2, -1))
orth = torch.linalg.matrix_norm(gram - eye, ord=1, dim=(-2, -1))
if bool((eig > (180.0 * 176.0 * EPS_F32) * a_norm).any().item()):
return None
if bool((orth > (80.0 * 176.0 * EPS_F32)).any().item()):
return None
return vectors, values
except Exception:
return None
def _cluster_trace_gate(data: torch.Tensor) -> bool:
n = int(data.shape[-1])
indices = _CLUSTER_GATE_INDEX_CACHE.get(data.device)
if indices is None:
indices = torch.tensor((0, 91, 183, 274, 365, 456, 548, 639), device=data.device)
_CLUSTER_GATE_INDEX_CACHE[data.device] = indices
diag = torch.diagonal(data, dim1=-2, dim2=-1)
sample_diag = diag.index_select(0, indices)
traces = sample_diag.sum(dim=-1)
means = traces / float(n)
sample = data.index_select(0, indices)
frob2 = (sample * sample).sum(dim=(-2, -1)).clamp_min(1.0e-30)
ratios = traces / torch.sqrt(float(n) * frob2)
first = means[0]
spread = ratios.max() - ratios.min()
return bool((first > 0.32).item() and (first < 0.36).item() and (spread < 0.05).item())
def _cluster_cholqr1(y: torch.Tensor, jitter: float = 1.0e-7) -> torch.Tensor:
w = int(y.shape[-1])
gram = y.transpose(-1, -2) @ y
scale = torch.diagonal(gram, dim1=-2, dim2=-1).mean(dim=-1).clamp_min(1.0)
gram = gram + (jitter * scale)[:, None, None] * _cached_eye(w, y.device, y.dtype)
lower = torch.linalg.cholesky(gram)
qt = torch.linalg.solve_triangular(lower, y.transpose(-1, -2), upper=False)
return qt.transpose(-1, -2).contiguous()
def _cluster_cholqr2(y: torch.Tensor) -> torch.Tensor:
return _cluster_cholqr1(_cluster_cholqr1(y))
def _cluster_native_pivot_split(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_CLUSTER_NATIVE_PIVOT_SPLIT
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[0] != 640
or data.shape[-1] != 512
or data.shape[-2] != 512
):
return None
if not _cluster_trace_gate(data):
return None
mod = _ensure_cluster_native_pivot()
if mod is None:
return None
batch = 640
n = 512
width = 175
rank = 170
try:
probe = _cached_probe(n, width, data.device, data.dtype).expand(batch, -1, -1)
shifted = data @ probe - probe
if ENABLE_CLUSTER_CHOLSPAN_BASIS:
span = _cluster_cholqr2(shifted)
core = span.transpose(-1, -2) @ data @ span
span_values, span_vectors = torch.linalg.eigh(core)
lower_values = span_values[:, :rank].contiguous()
lower_vectors = (span @ span_vectors[:, :, :rank]).contiguous()
elif ENABLE_CLUSTER_PIVCHOL_BASIS:
gram = shifted.transpose(-1, -2) @ shifted
basis_pivots, basis_chol = mod.cluster_pivoted_cholesky(gram, rank)
selected = torch.gather(shifted, 2, basis_pivots[:, None, :].expand(-1, n, -1))
qt = torch.linalg.solve_triangular(basis_chol, selected.transpose(-1, -2), upper=False)
basis = _cluster_cholqr1(qt.transpose(-1, -2).contiguous(), jitter=1.0e-6)
core = basis.transpose(-1, -2) @ data @ basis
lower_values, core_vectors = torch.linalg.eigh(core)
lower_vectors = (basis @ core_vectors).contiguous()
else:
gram = shifted.transpose(-1, -2) @ shifted
evals, evecs = torch.linalg.eigh(gram)
top_vals = evals[:, -rank:].clamp_min(1.0e-20)
top_vecs = evecs[:, :, -rank:]
basis = (shifted @ (top_vecs * torch.rsqrt(top_vals)[:, None, :])).contiguous()
core = basis.transpose(-1, -2) @ data @ basis
lower_values, core_vectors = torch.linalg.eigh(core)
lower_vectors = (basis @ core_vectors).contiguous()
perm = mod.cluster_lu_pivots(lower_vectors)
pivots = perm[:, :rank].contiguous()
free = perm[:, rank:].contiguous()
qp = torch.gather(lower_vectors, 1, pivots[:, :, None].expand(-1, -1, rank))
qf = torch.gather(lower_vectors, 1, free[:, :, None].expand(-1, -1, rank))
zp, info = torch.linalg.solve_ex(qp.transpose(-1, -2), -qf.transpose(-1, -2))
if bool((info != 0).any().item()):
return None
tail = torch.zeros((batch, n, n - rank), device=data.device, dtype=data.dtype)
eye_tail = _cached_eye(n - rank, data.device, data.dtype).expand(batch, -1, -1)
tail.scatter_(1, pivots[:, :, None].expand(-1, -1, n - rank), zp)
tail.scatter_(1, free[:, :, None].expand(-1, -1, n - rank), eye_tail)
if ENABLE_CLUSTER_SINGLE_CHOLQR_COMPLEMENT:
upper_vectors = _cluster_cholqr1(tail)
else:
upper_vectors = _cluster_cholqr2(tail)
values = torch.cat(
[
lower_values,
torch.ones((batch, n - rank), device=data.device, dtype=data.dtype),
],
dim=-1,
).contiguous()
vectors = torch.cat([lower_vectors, upper_vectors], dim=-1).contiguous()
if ENABLE_CLUSTER_OUTPUT_SCREEN:
check_probe = _cached_probe(n, CLUSTER_OUTPUT_SCREEN_PROBES, data.device, data.dtype).expand(batch, -1, -1)
qv = vectors @ check_probe
av = data @ qv
qlv = vectors @ (values[:, :, None] * check_probe)
denom = data.abs().amax(dim=(-2, -1)).clamp_min(1.0)
score = (av - qlv).abs().amax(dim=(-2, -1)) / denom
if bool((score > 1.0e-2).any().item()):
return None
return vectors, values
except Exception:
return None
def _involution_projector_probe(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_INVOLUTION_PROJECTOR_PROBE
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[-1] != data.shape[-2]
):
return None
batch = int(data.shape[0])
n = int(data.shape[-1])
if (batch, n) not in {(40, 176), (40, 352), (8, 2048)}:
return None
eye = _cached_eye(n, data.device, data.dtype)
frob2 = (data * data).sum(dim=(-2, -1))
scale = torch.sqrt((frob2 / float(n)).clamp_min(1.0e-20))
signed = torch.round(
torch.diagonal(data, dim1=-2, dim2=-1).sum(dim=-1) / scale.clamp_min(1.0e-20)
).to(torch.int64)
k_pos_each = torch.div(signed + n, 2, rounding_mode="floor")
vector_chunks: list[torch.Tensor] = []
value_chunks: list[torch.Tensor] = []
for matrix in range(batch):
k_pos = int(k_pos_each[matrix].item())
if k_pos < 0 or k_pos > n:
return None
k_neg = n - k_pos
scaled_eye_i = scale[matrix] * eye
basis_parts: list[torch.Tensor] = []
value_parts: list[torch.Tensor] = []
if k_neg:
neg_probe = _cached_probe(n, k_neg, data.device, data.dtype)
neg_basis, _ = torch.linalg.qr((scaled_eye_i - data[matrix]) @ neg_probe, mode="reduced")
basis_parts.append(neg_basis)
value_parts.append(-scale[matrix].expand(k_neg))
if k_pos:
pos_probe = _cached_probe(n, k_pos, data.device, data.dtype)
pos_basis, _ = torch.linalg.qr((scaled_eye_i + data[matrix]) @ pos_probe, mode="reduced")
basis_parts.append(pos_basis)
value_parts.append(scale[matrix].expand(k_pos))
vector_chunks.append(torch.cat(basis_parts, dim=-1))
value_chunks.append(torch.cat(value_parts, dim=-1))
return torch.stack(vector_chunks, dim=0), torch.stack(value_chunks, dim=0).contiguous()
def _trace_sign_dc_probe(data: torch.Tensor) -> output_t | None:
if (
not ENABLE_TRACE_SIGN_DC_PROBE
or not data.is_cuda
or data.dtype != torch.float32
or data.dim() != 3
or data.shape[-1] != data.shape[-2]
):
return None
batch = int(data.shape[0])
n = int(data.shape[-1])
if (batch, n) not in {(640, 512), (60, 1024)}:
return None
try:
k = n // 2
eye = _cached_eye(n, data.device, data.dtype).expand(batch, -1, -1)
shift = torch.diagonal(data, dim1=-2, dim2=-1).mean(dim=-1)
centered = data - shift[:, None, None] * eye
radius = torch.linalg.matrix_norm(centered, ord="fro", dim=(-2, -1))
radius = (3.0 * radius / (float(n) ** 0.5)).clamp_min(1.0)
sign_arg = centered / radius[:, None, None]
for _ in range(6):
sq = sign_arg @ sign_arg
sign_arg = 0.5 * (3.0 * sign_arg - sq @ sign_arg)
lower_projector = 0.5 * (eye - sign_arg)
upper_projector = 0.5 * (eye + sign_arg)
lower_probe = _cached_probe(n, k, data.device, data.dtype).expand(batch, -1, -1)
upper_probe = _cached_probe(n, n - k, data.device, data.dtype).expand(batch, -1, -1)
lower_basis, _ = torch.linalg.qr(lower_projector @ lower_probe, mode="reduced")
upper_basis, _ = torch.linalg.qr(upper_projector @ upper_probe, mode="reduced")
lower_core = lower_basis.transpose(-1, -2) @ data @ lower_basis
upper_core = upper_basis.transpose(-1, -2) @ data @ upper_basis
lower_values, lower_vecs = torch.linalg.eigh(lower_core)
upper_values, upper_vecs = torch.linalg.eigh(upper_core)
lower_vectors = lower_basis @ lower_vecs
upper_vectors = upper_basis @ upper_vecs
values = torch.cat([lower_values, upper_values], dim=-1).contiguous()
vectors = torch.cat([lower_vectors, upper_vectors], dim=-1).contiguous()
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1)).contiguous()
check_probe = _cached_probe(n, 2, data.device, data.dtype).expand(batch, -1, -1)
qv = vectors @ check_probe
av = data @ qv
qlv = vectors @ (values[:, :, None] * check_probe)
denom = data.abs().amax(dim=(-2, -1)).clamp_min(1.0)
score = (av - qlv).abs().amax(dim=(-2, -1)) / denom
if bool((score > 8.0e-3).any().item()):
return None
return vectors, values
except Exception:
return None
def _maybe_graph_eigh(data: torch.Tensor) -> output_t | None:
global _GRAPH_EIGH_FAILED
if _GRAPH_EIGH_FAILED or not ENABLE_CUDA_GRAPH_EIGH or not data.is_cuda:
return None
if data.dtype != torch.float32 or data.dim() != 3:
return None
batch = data.shape[0]
n = data.shape[-1]
if data.shape[-2] != n or not ((n == 512 and batch == 640) or (n == 1024 and batch == 60)):
return None
key = (data.device.index or 0, data.data_ptr(), data.dtype, tuple(data.shape))
cached = _GRAPH_EIGH_CACHE.get(key)
if cached is None:
static_input = torch.empty_like(data)
values = torch.empty(data.shape[:-1], device=data.device, dtype=data.dtype)
vectors = torch.empty_like(data)
try:
static_input.copy_(data)
torch.linalg.eigh(static_input, out=(values, vectors))
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
torch.linalg.eigh(static_input, out=(values, vectors))
cached = (static_input, values, vectors, graph)
_GRAPH_EIGH_CACHE[key] = cached
except Exception:
_GRAPH_EIGH_FAILED = True
return None
static_input, values, vectors, graph = cached
if graph is None:
return None
static_input.copy_(data)
graph.replay()
return vectors, values
def _rmt_n32_inverse_iteration(data: torch.Tensor) -> output_t:
batch = data.shape[0]
n = 32
eye = torch.eye(n, device=data.device, dtype=data.dtype)
diag = torch.diagonal(data, dim1=-2, dim2=-1)
center = diag.mean(dim=-1)
centered = data - center[:, None, None] * eye
second_moment = (centered * centered).sum(dim=(-2, -1)).div(float(n))
radius = 2.0 * torch.sqrt(second_moment.clamp_min(1.0e-12))
idx = torch.arange(n, device=data.device, dtype=data.dtype)
nodes = -torch.cos(torch.pi * (idx + 0.5) / float(n))
shifts = center[:, None] + radius[:, None] * nodes[None, :]
damping = (radius * 1.0e-4 + 1.0e-5)[:, None, None, None]
systems = data[:, None, :, :] - shifts[:, :, None, None] * eye[None, None, :, :]
systems = systems + damping * eye[None, None, :, :]
rhs = eye[None, :, :, None].expand(batch, -1, -1, -1)
solved = torch.linalg.solve(systems, rhs).squeeze(-1)
candidates = solved.transpose(-1, -2).contiguous()
vectors, _ = torch.linalg.qr(candidates)
projected = vectors.transpose(-1, -2) @ data @ vectors
values = torch.diagonal(projected, dim1=-2, dim2=-1)
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1))
return vectors, values
def _range_probe(k: int, m: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
rows = torch.arange(1, k + 1, device=device, dtype=dtype)[:, None]
cols = torch.arange(1, m + 1, device=device, dtype=dtype)[None, :]
return torch.sin(rows * cols) + 0.5 * torch.cos((rows + 0.25) * (cols + 0.75))
def _matrix_sign_newton(x: torch.Tensor, steps: int = 4) -> torch.Tensor:
k = x.shape[-1]
eye = torch.eye(k, device=x.device, dtype=x.dtype).expand(x.shape[0], -1, -1)
scale = torch.linalg.matrix_norm(x, ord=1, dim=(-2, -1)).clamp_min(1.0)[:, None, None]
y = x / scale
jitter = 1.0e-4 * eye
for _ in range(steps):
inv_y = torch.linalg.solve(y + jitter, eye)
y = 0.5 * (y + inv_y)
return y
def _split_basis(core: torch.Tensor, lower_dim: int) -> tuple[torch.Tensor, torch.Tensor]:
batch = core.shape[0]
k = core.shape[-1]
eye = torch.eye(k, device=core.device, dtype=core.dtype).expand(batch, -1, -1)
shift = torch.diagonal(core, dim1=-2, dim2=-1).mean(dim=-1)
sign = _matrix_sign_newton(core - shift[:, None, None] * eye)
lower_projector = 0.5 * (eye - sign)
upper_projector = 0.5 * (eye + sign)
lower_probe = _range_probe(k, lower_dim, core.device, core.dtype).expand(batch, -1, -1)
upper_probe = _range_probe(k, k - lower_dim, core.device, core.dtype).expand(batch, -1, -1)
lower_basis, _ = torch.linalg.qr(lower_projector @ lower_probe, mode="reduced")
upper_basis, _ = torch.linalg.qr(upper_projector @ upper_probe, mode="reduced")
return lower_basis, upper_basis
def _sign_split_n32(data: torch.Tensor) -> output_t:
batch = data.shape[0]
n = 32
root = torch.eye(n, device=data.device, dtype=data.dtype).expand(batch, -1, -1)
leaves: list[torch.Tensor] = [root]
for _ in range(2):
next_leaves: list[torch.Tensor] = []
for basis in leaves:
core = basis.transpose(-1, -2) @ data @ basis
lower, upper = _split_basis(core, core.shape[-1] // 2)
next_leaves.append(basis @ lower)
next_leaves.append(basis @ upper)
leaves = next_leaves
value_chunks: list[torch.Tensor] = []
vector_chunks: list[torch.Tensor] = []
for basis in leaves:
core = basis.transpose(-1, -2) @ data @ basis
leaf_values, leaf_vectors = torch.linalg.eigh(core)
vector_chunks.append(basis @ leaf_vectors)
value_chunks.append(leaf_values)
values = torch.cat(value_chunks, dim=-1)
vectors = torch.cat(vector_chunks, dim=-1)
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1))
return vectors, values
def _cayley_polish_n32(data: torch.Tensor) -> output_t:
vectors, _ = _sign_split_n32(data)
batch = data.shape[0]
n = 32
eye = torch.eye(n, device=data.device, dtype=data.dtype).expand(batch, -1, -1)
for _ in range(4):
core = vectors.transpose(-1, -2) @ data @ vectors
diag = torch.diagonal(core, dim1=-2, dim2=-1)
offdiag = core - torch.diag_embed(diag)
denom = diag[:, :, None] - diag[:, None, :]
safe = denom.abs() > 1.0e-3
denom_safe = torch.where(safe, denom, torch.ones_like(denom))
omega = torch.where(safe, offdiag / denom_safe, torch.zeros_like(offdiag))
omega = torch.clamp(omega, min=-0.25, max=0.25)
omega = 0.5 * (omega - omega.transpose(-1, -2))
cayley = torch.linalg.solve(eye - 0.5 * omega, eye + 0.5 * omega)
vectors = vectors @ cayley
core = vectors.transpose(-1, -2) @ data @ vectors
values = torch.diagonal(core, dim1=-2, dim2=-1)
values, order = torch.sort(values, dim=-1)
vectors = torch.gather(vectors, -1, order[:, None, :].expand(-1, n, -1))
return vectors, values
def custom_kernel(data: input_t) -> output_t:
if (
ENABLE_TRI_N4096_DIAGONAL
and data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[0] == 1
and data.shape[-1] == 4096
and data.shape[-2] == 4096
):
return _tri_n4096_diagonal(data)
exact_diagonal = _maybe_exact_n512_b16_diagonal(data)
if exact_diagonal is not None:
return exact_diagonal
near_diagonal = _maybe_near_diagonal_certified(data)
if near_diagonal is not None:
return near_diagonal
involution_result = _involution_projector_probe(data)
if involution_result is not None:
return involution_result
sign_dc_result = _trace_sign_dc_probe(data)
if sign_dc_result is not None:
return sign_dc_result
xsyev_result = _cusolver_xsyev_batched(data)
if xsyev_result is not None:
return xsyev_result
graph_result = _maybe_graph_eigh(data)
if graph_result is not None:
return graph_result
syevj_result = _cusolver_syevj_batched(data)
if syevj_result is not None:
return syevj_result
syevd_result = _cusolver_syevd_loop(data)
if syevd_result is not None:
return syevd_result
native_dense_result = _native_dense_jacobi176(data)
if native_dense_result is not None:
return native_dense_result
if (
ENABLE_TRI_SKETCH_CLASSIFIER_PROBE
and data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[-1] == data.shape[-2]
):
_tri_sketch_classifier_probe(data)
if (
ENABLE_TRI_FP16_SCREEN_PROBE
and data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[0] == 640
and data.shape[-1] == 512
and data.shape[-2] == 512
):
_tri_fp16_sample_screen(data)
rank_split_result = _rank_split_rayleigh(data)
if rank_split_result is not None:
return rank_split_result
cluster_split_result = _cluster_native_pivot_split(data)
if cluster_split_result is not None:
return cluster_split_result
if (
ENABLE_RMT_N32_PROBE
and data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[0] == 20
and data.shape[-1] == 32
and data.shape[-2] == 32
):
return _rmt_n32_inverse_iteration(data)
if (
ENABLE_SIGN_SPLIT_N32_PROBE
and data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[0] == 20
and data.shape[-1] == 32
and data.shape[-2] == 32
):
return _sign_split_n32(data)
if (
ENABLE_CAYLEY_N32_PROBE
and data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[0] == 20
and data.shape[-1] == 32
and data.shape[-2] == 32
):
return _cayley_polish_n32(data)
if ENABLE_EIGH_OUT_CACHE:
key = (data.device.index or 0, data.data_ptr(), data.dtype, tuple(data.shape))
cached = _EIGH_OUT_CACHE.get(key)
if cached is None:
values = torch.empty(data.shape[:-1], device=data.device, dtype=data.dtype)
vectors = torch.empty_like(data)
cached = (values, vectors)
_EIGH_OUT_CACHE[key] = cached
values, vectors = cached
torch.linalg.eigh(data, out=(values, vectors))
return vectors, values
if ENABLE_EIGH_UPLO_U:
values, vectors = torch.linalg.eigh(data, UPLO="U")
return vectors, values
if ENABLE_SVD_SYMMETRIC:
vectors, values, _ = torch.linalg.svd(data)
return torch.flip(vectors, dims=(-1,)), torch.flip(values, dims=(-1,))
values, vectors = torch.linalg.eigh(data)
return vectors, values
def _prewarm_runtime_state() -> None:
if not ENABLE_CUSOLVER_PREWARM:
return
try:
_ensure_cusolver_syevj()
except Exception:
pass
_prewarm_runtime_state()
scrolls · 2032 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