submission 872271
dbuddha · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2473 lines, June 9 Researcher Reciprocity License v1.0.
sub_r54_fp16_eigen_screen.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-872271?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:1fc2978279b3d969db536ee9c259087f51097eec0753d2ab745f69f4e4ab1362
license declaredunknown
license concludedunknown
authorsdbuddha
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float smem[];Kernel source
sub_r54_fp16_eigen_screen.py2473 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200
"""R22 deferred factor-tree S2 on the incumbent R21 rational pipeline.
The custom n=512 path is guarded per matrix by the official eigen-equation
and orthogonality metrics. Other sizes use torch directly.
"""
"""NVRTC compile+launch kit (zhong mechanism, hosted-probe-verified stage 4).
Usage:
from nvrtc_kit import NvrtcKernel
k = NvrtcKernel(SRC, "my_kernel") # compiles for sm_100a
k.launch(grid=(640,1,1), block=(256,1,1), args=[tensor_a, np.int64(n)],
shared=0)
Args may be torch tensors (data_ptr taken), python ints (c_longlong), or
floats (c_float). Modules cached per (source-hash, name) in _KMOD.
"""
import atexit
import ctypes
import ctypes.util
import hashlib
import os
import torch
_LIBS: dict = {}
_KMOD: dict = {}
def _lib(kind: str):
if kind in _LIBS:
return _LIBS[kind]
if kind == "driver":
names = [ctypes.util.find_library("cuda"), "libcuda.so.1", "libcuda.so"]
else:
names = [ctypes.util.find_library("nvrtc"), "libnvrtc.so",
"libnvrtc.so.13", "libnvrtc.so.12"]
package_root = os.path.abspath(
os.path.join(os.path.dirname(torch.__file__), "..", "nvidia"))
if os.path.isdir(package_root):
for child in os.listdir(package_root):
lib_dir = os.path.join(package_root, child, "lib")
for soname in ("libnvrtc.so", "libnvrtc.so.13",
"libnvrtc.so.12"):
names.append(os.path.join(lib_dir, soname))
last = None
for name in names:
if not name:
continue
try:
lib = ctypes.CDLL(name)
_LIBS[kind] = lib
return lib
except OSError as exc:
last = exc
raise RuntimeError(f"no CUDA {kind} library: {last}")
def _nvrtc():
lib = _lib("nvrtc")
if getattr(lib, "_sigs_done", False):
return lib
lib.nvrtcCreateProgram.restype = ctypes.c_int
lib.nvrtcCreateProgram.argtypes = [
ctypes.POINTER(ctypes.c_void_p), ctypes.c_char_p, ctypes.c_char_p,
ctypes.c_int, ctypes.POINTER(ctypes.c_char_p),
ctypes.POINTER(ctypes.c_char_p)]
lib.nvrtcCompileProgram.restype = ctypes.c_int
lib.nvrtcCompileProgram.argtypes = [
ctypes.c_void_p, ctypes.c_int, ctypes.POINTER(ctypes.c_char_p)]
for fn in ("nvrtcGetProgramLogSize", "nvrtcGetCUBINSize"):
getattr(lib, fn).restype = ctypes.c_int
getattr(lib, fn).argtypes = [ctypes.c_void_p,
ctypes.POINTER(ctypes.c_size_t)]
for fn in ("nvrtcGetProgramLog", "nvrtcGetCUBIN"):
getattr(lib, fn).restype = ctypes.c_int
getattr(lib, fn).argtypes = [ctypes.c_void_p, ctypes.c_void_p]
lib.nvrtcDestroyProgram.restype = ctypes.c_int
lib.nvrtcDestroyProgram.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
lib._sigs_done = True
return lib
def _driver():
lib = _lib("driver")
if getattr(lib, "_sigs_done", False):
return lib
lib.cuModuleLoadData.restype = ctypes.c_int
lib.cuModuleLoadData.argtypes = [ctypes.POINTER(ctypes.c_void_p),
ctypes.c_void_p]
lib.cuModuleGetFunction.restype = ctypes.c_int
lib.cuModuleGetFunction.argtypes = [ctypes.POINTER(ctypes.c_void_p),
ctypes.c_void_p, ctypes.c_char_p]
lib.cuFuncSetAttribute.restype = ctypes.c_int
lib.cuFuncSetAttribute.argtypes = [ctypes.c_void_p, ctypes.c_int,
ctypes.c_int]
lib.cuLaunchKernel.restype = ctypes.c_int
lib.cuLaunchKernel.argtypes = [
ctypes.c_void_p, ctypes.c_uint, ctypes.c_uint, ctypes.c_uint,
ctypes.c_uint, ctypes.c_uint, ctypes.c_uint, ctypes.c_uint,
ctypes.c_void_p, ctypes.POINTER(ctypes.c_void_p), ctypes.c_void_p]
lib._sigs_done = True
return lib
def _include_dirs():
out = []
seen = set()
package_root = os.path.abspath(
os.path.join(os.path.dirname(torch.__file__), "..", "nvidia"))
cands = []
if os.path.isdir(package_root):
for child in sorted(os.listdir(package_root)):
cands.append(os.path.join(package_root, child, "include"))
for d in cands:
if not d or d in seen or not os.path.isdir(d):
continue
seen.add(d)
out.append(d)
cccl = os.path.join(d, "cccl")
if os.path.exists(os.path.join(cccl, "cuda", "std")) and cccl not in seen:
seen.add(cccl)
out.append(cccl)
return out
def compile_cubin(source: str, name: str = "kernel.cu") -> bytes:
torch.cuda.init()
torch.empty(0, device="cuda")
nvrtc = _nvrtc()
major, minor = torch.cuda.get_device_capability()
sm = major * 10 + minor
arch = f"sm_{sm}a" if sm >= 90 else f"sm_{sm}"
opts = [f"--gpu-architecture={arch}", "-std=c++17", "-default-device",
"--use_fast_math"]
for d in _include_dirs():
opts.append(f"-I{d}")
prog = ctypes.c_void_p()
err = nvrtc.nvrtcCreateProgram(ctypes.byref(prog), source.encode(),
name.encode(), 0, None, None)
if err != 0:
raise RuntimeError(f"nvrtcCreateProgram={err}")
try:
enc = [o.encode() for o in opts]
arr = (ctypes.c_char_p * len(enc))(*enc)
cerr = nvrtc.nvrtcCompileProgram(prog, len(enc), arr)
if cerr != 0:
size = ctypes.c_size_t()
log = ""
if (nvrtc.nvrtcGetProgramLogSize(prog, ctypes.byref(size)) == 0
and size.value):
buf = ctypes.create_string_buffer(size.value)
nvrtc.nvrtcGetProgramLog(prog, buf)
log = buf.value.decode(errors="replace")
raise RuntimeError(f"NVRTC compile failed ({name}):\n{log}")
size = ctypes.c_size_t()
nvrtc.nvrtcGetCUBINSize(prog, ctypes.byref(size))
image = ctypes.create_string_buffer(size.value)
nvrtc.nvrtcGetCUBIN(prog, image)
return bytes(image.raw)
finally:
nvrtc.nvrtcDestroyProgram(ctypes.byref(prog))
class NvrtcKernel:
def __init__(self, source: str, func_name: str, shared_carveout: int = 0):
key = (hashlib.sha1(source.encode()).hexdigest(), func_name)
if key in _KMOD:
self._func = _KMOD[key]
return
cubin = compile_cubin(source, func_name + ".cu")
drv = _driver()
module = ctypes.c_void_p()
buf = ctypes.create_string_buffer(cubin, len(cubin))
err = drv.cuModuleLoadData(ctypes.byref(module), buf)
if err != 0:
raise RuntimeError(f"cuModuleLoadData={err}")
func = ctypes.c_void_p()
err = drv.cuModuleGetFunction(ctypes.byref(func), module,
func_name.encode())
if err != 0:
raise RuntimeError(f"cuModuleGetFunction={err} ({func_name})")
if shared_carveout > 0:
# CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES = 8
e = drv.cuFuncSetAttribute(func, 8, shared_carveout)
if e != 0:
raise RuntimeError(f"cuFuncSetAttribute smem={shared_carveout} "
f"err={e} (exceeds device max?)")
# keep module alive via the buffer reference
_KMOD[key] = func
_KMOD[(key, "keepalive")] = (module, buf)
self._func = func
def launch(self, grid, block, args, shared: int = 0):
drv = _driver()
holders = []
ptrs = (ctypes.c_void_p * len(args))()
for i, a in enumerate(args):
if isinstance(a, torch.Tensor):
h = ctypes.c_void_p(a.data_ptr())
elif isinstance(a, bool):
h = ctypes.c_int(int(a))
elif isinstance(a, int):
h = ctypes.c_longlong(a)
elif isinstance(a, float):
h = ctypes.c_float(a)
elif isinstance(a, ctypes._SimpleCData):
h = a
else:
raise TypeError(f"unsupported arg type {type(a)}")
holders.append(h)
ptrs[i] = ctypes.cast(ctypes.byref(h), ctypes.c_void_p)
err = drv.cuLaunchKernel(
self._func, int(grid[0]), int(grid[1]), int(grid[2]),
int(block[0]), int(block[1]), int(block[2]), int(shared),
None, ptrs, None)
if err != 0:
raise RuntimeError(
f"cuLaunchKernel={err} grid={grid} block={block} "
f"shared={shared} nargs={len(args)}")
"""R34 keeps R10's column-major V/W storage and replaces the serialized
prior-column dot/correction loop with warp-per-column ownership."""
import torch
NB = 48
_SRC = r"""
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#define NB 48
extern "C" __global__ void latrd_panel_r34_half_matrix(
const __half* __restrict__ Ain,
float* __restrict__ Vout,
float* __restrict__ Wout,
float* __restrict__ TauOut,
float* __restrict__ Dout,
float* __restrict__ Eout,
long long rows_ll,
long long ld_ll)
{
const int rows = (int)rows_ll;
const int ld = (int)ld_ll;
const int mat = blockIdx.x;
const int tid = threadIdx.x;
const int T = blockDim.x;
extern __shared__ float smem[];
float* V = smem;
float* W = V + rows * NB;
float* red = W + rows * NB;
float* bcastV = red + T;
float* bcastW = bcastV + NB;
float* crossW = bcastW + NB;
float* crossV = crossW + NB;
float* colbuf = crossV + NB;
const __half* Ag = Ain + (long)mat * ld * ld;
// V/W are column-major in shared memory and row-major in output.
for (int idx = tid; idx < rows * NB; idx += T) { V[idx] = 0.0f; W[idx] = 0.0f; }
__syncthreads();
for (int k = 0; k < NB; ++k) {
if (tid < k) { bcastV[tid] = V[tid * rows + k]; bcastW[tid] = W[tid * rows + k]; }
__syncthreads();
for (int i = k + tid; i < rows; i += T) {
float c = Ag[(long)i * ld + k];
for (int j = 0; j < k; ++j) c -= V[j * rows + i] * bcastW[j] + W[j * rows + i] * bcastV[j];
colbuf[i] = c;
}
__syncthreads();
float d_entry = colbuf[k];
float local = 0.0f;
for (int i = k + 2 + tid; i < rows; i += T) { float v = colbuf[i]; local += v * v; }
#pragma unroll
for (int off = 16; off > 0; off >>= 1) local += __shfl_down_sync(0xffffffffu, local, off);
{
const int lane0 = tid & 31, warp_id0 = tid >> 5, nwarps0 = T >> 5;
if (lane0 == 0) red[warp_id0] = local;
__syncthreads();
if (tid == 0) { float s = 0.0f; for (int w = 0; w < nwarps0; ++w) s += red[w]; red[0] = s; }
__syncthreads();
}
float sumsq = red[0]; __syncthreads();
float alpha = (k + 1 < rows) ? colbuf[k + 1] : 0.0f;
float normx = sqrtf(alpha * alpha + sumsq);
float beta = (alpha >= 0.0f) ? -normx : normx;
float denom = alpha - beta;
bool good = normx > 1e-20f && fabsf(denom) > 1e-20f;
float invdenom = good ? 1.0f / denom : 0.0f;
if (tid == 0 && k + 1 < rows) colbuf[k + 1] = 1.0f;
for (int i = k + 2 + tid; i < rows; i += T) colbuf[i] = good ? (colbuf[i] * invdenom) : 0.0f;
__syncthreads();
float vv = 1.0f + sumsq * invdenom * invdenom;
float tau_k = good ? (2.0f / vv) : 0.0f;
float sub = good ? beta : alpha;
if (tid == 0) {
Dout[(long)mat * NB + k] = d_entry;
Eout[(long)mat * NB + k] = sub;
TauOut[(long)mat * NB + k] = tau_k;
}
for (int i = tid; i <= k; i += T) V[k * rows + i] = 0.0f;
for (int i = k + 1 + tid; i < rows; i += T) V[k * rows + i] = good ? colbuf[i] : 0.0f;
__syncthreads();
{
const int lane = tid & 31;
const int warp_id = tid >> 5;
const int nwarps = T >> 5;
for (int i = k + 1 + warp_id; i < rows; i += nwarps) {
double p = 0.0;
for (int c = k + 1 + lane; c < rows; c += 32)
p += (double)__half2float(Ag[(long)i * ld + c]) * (double)V[k * rows + c];
#pragma unroll
for (int off = 16; off > 0; off >>= 1)
p += __shfl_down_sync(0xffffffffu, p, off);
if (lane == 0) colbuf[i] = (float)p;
}
}
__syncthreads();
{
const int lane = tid & 31;
const int warp_id = tid >> 5;
const int nwarps = T >> 5;
#pragma unroll
for (int round = 0; round < 2; ++round) {
const int j = round * nwarps + warp_id;
if (j < k) {
float local_wtv = 0.0f;
float local_vtv = 0.0f;
for (int i = k + 1 + lane; i < rows; i += 32) {
const float v_i = V[k * rows + i];
local_wtv += W[j * rows + i] * v_i;
local_vtv += V[j * rows + i] * v_i;
}
#pragma unroll
for (int off = 16; off > 0; off >>= 1) {
local_wtv += __shfl_down_sync(
0xffffffffu, local_wtv, off);
local_vtv += __shfl_down_sync(
0xffffffffu, local_vtv, off);
}
if (lane == 0) {
crossW[j] = local_wtv;
crossV[j] = local_vtv;
}
}
}
__syncthreads();
for (int i = k + 1 + tid; i < rows; i += T) {
float corrected = colbuf[i];
for (int j = 0; j < k; ++j)
corrected -= V[j * rows + i] * crossW[j]
+ W[j * rows + i] * crossV[j];
colbuf[i] = corrected;
}
__syncthreads();
}
float local_pv = 0.0f;
for (int i = k + 1 + tid; i < rows; i += T) local_pv += colbuf[i] * V[k * rows + i];
#pragma unroll
for (int off = 16; off > 0; off >>= 1) local_pv += __shfl_down_sync(0xffffffffu, local_pv, off);
{
const int lane1 = tid & 31, warp_id1 = tid >> 5, nwarps1 = T >> 5;
if (lane1 == 0) red[warp_id1] = local_pv;
__syncthreads();
if (tid == 0) { float s = 0.0f; for (int w = 0; w < nwarps1; ++w) s += red[w]; red[0] = s; }
__syncthreads();
}
float pv = red[0]; __syncthreads();
float corr = 0.5f * tau_k * tau_k * pv;
for (int i = tid; i <= k; i += T) W[k * rows + i] = 0.0f;
for (int i = k + 1 + tid; i < rows; i += T) {
W[k * rows + i] = tau_k * colbuf[i] - corr * V[k * rows + i];
}
__syncthreads();
}
// Vout/Wout are (batch, rows, NB) ROW-MAJOR for the later torch code
// (_build_t, trailing GEMM) -- transpose back out of column-major smem.
for (int idx = tid; idx < rows * NB; idx += T) {
int row = idx / NB, col = idx % NB;
Vout[(long)mat * rows * NB + idx] = V[col * rows + row];
Wout[(long)mat * rows * NB + idx] = W[col * rows + row];
}
}
"""
_COMPILE_OK = False
try:
_k = NvrtcKernel(
_SRC, "latrd_panel_r34_half_matrix", shared_carveout=205 * 1024)
_COMPILE_OK = True
except Exception:
_COMPILE_OK = False
_COMPACT_SOURCE = r"""
#include <cuda_runtime.h>
extern "C" __global__ void r34_compact_factor(
const float* __restrict__ gram,
const float* __restrict__ tau,
float* __restrict__ factor,
long long width_ll)
{
const int width = (int)width_ll;
const int matrix = blockIdx.x;
const int tid = threadIdx.x;
const long long offset = (long long)matrix * width * width;
const float* matrix_gram = gram + offset;
float* matrix_factor = factor + offset;
const float* matrix_tau = tau + (long long)matrix * width;
for (int index = tid; index < width * width; index += blockDim.x)
matrix_factor[index] = 0.0f;
__syncthreads();
for (int column = 0; column < width; ++column) {
if (tid < column) {
float value = 0.0f;
for (int inner = 0; inner < column; ++inner)
value = fmaf(
matrix_factor[tid * width + inner],
matrix_gram[inner * width + column],
value);
matrix_factor[tid * width + column] =
-matrix_tau[column] * value;
}
if (tid == 0)
matrix_factor[column * width + column] = matrix_tau[column];
__syncthreads();
}
}
"""
_compact_k = NvrtcKernel(_COMPACT_SOURCE, "r34_compact_factor")
_TAIL_SOURCE = (
_SRC.replace("#define NB 48", "#define NB 32", 1)
.replace("latrd_panel_r34_half_matrix", "latrd_panel_r34_tail32")
)
_tail_k = NvrtcKernel(
_TAIL_SOURCE, "latrd_panel_r34_tail32", shared_carveout=32 * 1024)
def _smem_bytes(rows, block=256):
return (rows * NB * 2 + block + NB * 4 + rows) * 4
def _launch_panel(A_view, batch, rows, ld, block=1024):
"""A_view: a (batch, rows, rows) VIEW into a live (batch, n, n) tensor
(e.g. work[:, j:, j:], NOT .contiguous()'d) -- ld is the ORIGINAL full
matrix width (constant across panels), used by the kernel to address
directly into the live tensor with no per-panel copy. Zero-copy: avoids
the previous ~600MB+ .contiguous() allocation for the first panel alone
at batch=640."""
Vout = torch.zeros(batch, rows, NB, device="cuda", dtype=torch.float32)
Wout = torch.zeros(batch, rows, NB, device="cuda", dtype=torch.float32)
TauOut = torch.zeros(batch, NB, device="cuda", dtype=torch.float32)
Dout = torch.zeros(batch, NB, device="cuda", dtype=torch.float32)
Eout = torch.zeros(batch, NB, device="cuda", dtype=torch.float32)
_k.launch((batch, 1, 1), (block, 1, 1),
[A_view, Vout.reshape(-1), Wout.reshape(-1), TauOut.reshape(-1),
Dout.reshape(-1), Eout.reshape(-1), rows, ld],
shared=_smem_bytes(rows, block))
return Vout, Wout, TauOut, Dout, Eout
def _build_t(V, tau):
batch, _, width = V.shape
gram = V.mT @ V
factor = torch.empty(
batch, width, width, dtype=torch.float32, device=V.device)
_compact_k.launch(
(batch, 1, 1), (64, 1, 1),
[gram, tau, factor, width], shared=0)
return factor
def _tail_reduce(A):
batch, rows, _ = A.shape
if rows != 32:
raise RuntimeError("unexpected tail width")
matrix = A.to(torch.float16)
vectors = torch.empty(
batch, rows, 32, device=A.device, dtype=torch.float32)
work = torch.empty_like(vectors)
tau = torch.empty(batch, 32, device=A.device, dtype=torch.float32)
diagonal = torch.empty_like(tau)
subdiagonal = torch.empty_like(tau)
shared = (rows * 32 * 2 + 1024 + 32 * 4 + rows) * 4
_tail_k.launch(
(batch, 1, 1), (1024, 1, 1),
[matrix, vectors, work, tau, diagonal, subdiagonal, rows, rows],
shared=shared)
compact = _build_t(vectors[:, 1:, :].contiguous(), tau)
return (
diagonal,
subdiagonal[:, :rows - 1].contiguous(),
[(1, vectors[:, 1:, :].contiguous(), compact)],
)
def back_transform(Q, factors):
for r0, V, T in reversed(factors):
sub = Q[:, r0:, :]
Q[:, r0:, :] = sub - V @ (T @ (V.transpose(-1, -2) @ sub))
return Q
def s1_kernel_reduce(A, tail=48):
bsz, n, _ = A.shape
work = A.to(torch.float16)
factors = []
d_parts, e_parts = [], []
j = 0
while (n - j) > tail:
rows = n - j
view = work[:, j:, j:] # zero-copy view, NOT .contiguous()
Vout, Wout, TauOut, Dout, Eout = _launch_panel(view, bsz, rows, n)
vt = Vout[:, NB:, :].to(torch.float16); wt = Wout[:, NB:, :].to(torch.float16)
upd = vt @ wt.mT
work[:, j + NB:, j + NB:] -= (upd + upd.mT) # in-place on the live tensor
factors.append((j + 1, Vout[:, 1:, :].contiguous(), _build_t(Vout[:, 1:, :], TauOut)))
d_parts.append(Dout); e_parts.append(Eout)
j += NB
tblock = work[:, j:, j:].contiguous()
d_t, e_t, factors_t = _tail_reduce(tblock)
for (r0, v_mat, t_mat) in factors_t:
factors.append((r0 + j, v_mat, t_mat))
d_parts.append(d_t); e_parts.append(e_t)
d_full = torch.cat(d_parts, dim=-1)
e_full = torch.cat(e_parts, dim=-1)
return d_full, e_full, factors
LEAF = 32
_SRC = r"""
#include <cuda_runtime.h>
extern "C" __global__ void rank1_merge(
const double* __restrict__ ds_all, // (nsub, m) ascending
const double* __restrict__ zs_all, // (nsub, m)
const double* __restrict__ rho_all, // (nsub,) > 0
double* __restrict__ mu_all, // (nsub, m) out
double* __restrict__ U_all, // (nsub, m, m) out: U[i*m+k]=u_k[i]
long long m)
{
const int sub = blockIdx.x;
const int tid = threadIdx.x;
const int nt = blockDim.x;
const double* ds = ds_all + (long)sub * m;
const double* zs = zs_all + (long)sub * m;
const double rho = rho_all[sub];
double* mu = mu_all + (long)sub * m;
double* U = U_all + (long)sub * (long)m * m;
extern __shared__ double sh[];
double* sd = sh; // m
double* sz = sh + m; // m
double* smu = sh + 2 * m; // m
double* sact = sh + 3 * m; // m (1.0 active / 0.0 deflated)
double* red = sh + 4 * m; // nt (reduction scratch)
for (int i = tid; i < m; i += nt) { sd[i] = ds[i]; sz[i] = zs[i]; }
__syncthreads();
// --- scale = max|d| + rho * sum z^2 (block reductions)
double amax = 0.0, sz2 = 0.0;
for (int i = tid; i < m; i += nt) {
double a = fabs(sd[i]); if (a > amax) amax = a;
sz2 += sz[i] * sz[i];
}
red[tid] = amax; __syncthreads();
for (int s = nt / 2; s > 0; s >>= 1) { if (tid < s) red[tid] = fmax(red[tid], red[tid + s]); __syncthreads(); }
double smax = red[0]; __syncthreads();
red[tid] = sz2; __syncthreads();
for (int s = nt / 2; s > 0; s >>= 1) { if (tid < s) red[tid] += red[tid + s]; __syncthreads(); }
double sumz2 = red[0]; __syncthreads();
double scale = smax + rho * sumz2;
double wtol = 1e-13 * scale, gtol = 1e-10 * scale;
// --- deflation: active[i] = weight big AND not near-equal to prev
for (int i = tid; i < m; i += nt) {
bool w = rho * sz[i] * sz[i] > wtol;
bool nt_tie = (i == 0) || ((sd[i] - sd[i - 1]) > gtol);
sact[i] = (w && nt_tie) ? 1.0 : 0.0;
}
__syncthreads();
// --- secular roots: SAFEGUARDED NEWTON over ACTIVE set (F3/e081: ~40 iters
// needed for full fp64 precision on hard cases; vectorization/parallelism
// is the real lever, not iteration count -- kept per-thread-per-pole like
// the original, just Newton instead of bisection).
for (int k = tid; k < m; k += nt) {
if (sact[k] == 0.0) { smu[k] = sd[k]; continue; }
double lo = sd[k];
double hi = sd[m - 1] + rho * sumz2;
for (int j = k + 1; j < m; ++j) { if (sact[j] != 0.0) { hi = sd[j]; break; } }
double x = 0.5 * (lo + hi);
for (int it = 0; it < 18; ++it) {
double f = 1.0, fp = 0.0;
for (int i = 0; i < m; ++i) {
if (sact[i] == 0.0) continue;
double di = sd[i] - x;
if (fabs(di) < 1e-300) di = (di < 0 ? -1e-300 : 1e-300);
double term = rho * sz[i] * sz[i] / di;
f += term;
fp += term / di;
}
double xn = (fabs(fp) > 1e-300) ? (x - f / fp) : (0.5 * (lo + hi));
if (f < 0.0) lo = x; else hi = x;
if (xn <= lo || xn >= hi || !isfinite(xn)) xn = 0.5 * (lo + hi);
x = xn;
}
smu[k] = x;
}
__syncthreads();
// --- Gu-Eisenstat zhat: RATIO-PRODUCT (F3/e081: validated more accurate
// AND avoids log/exp SFU calls -- O(1)-magnitude terms, no overflow).
for (int i = tid; i < m; i += nt) {
double zhat = 0.0;
if (sact[i] != 0.0) {
// prod = (mu_i-d_i) * PRODUCT_{k active, k!=i} [(mu_k-d_i)/(d_k-d_i)]
// -- simple running product from 1.0, no special-case ordering bug.
double prod = 1.0;
for (int k = 0; k < m; ++k) {
if (sact[k] == 0.0) continue;
double num = smu[k] - sd[i];
if (k == i) {
prod *= num;
} else {
double den = sd[k] - sd[i];
if (fabs(den) < 1e-300) den = (den < 0 ? -1e-300 : 1e-300);
prod *= num / den;
}
}
double s = (sz[i] < 0.0) ? -1.0 : 1.0;
zhat = s * sqrt(fabs(prod));
}
// fill row i of U for all columns k (unnormalized)
for (int k = 0; k < m; ++k) {
double v;
if (sact[k] != 0.0) {
if (sact[i] != 0.0) {
double di = sd[i] - smu[k];
if (fabs(di) < 1e-300) di = (di < 0 ? -1e-300 : 1e-300);
v = zhat / di;
} else v = 0.0;
} else {
v = (i == k) ? 1.0 : 0.0; // deflated eigenvector = e_k
}
U[(long)i * m + k] = v;
}
}
__syncthreads();
// --- normalize each active column k (deflated cols already unit)
for (int k = tid; k < m; k += nt) {
mu[k] = smu[k];
if (sact[k] == 0.0) continue;
double nrm = 0.0;
for (int i = 0; i < m; ++i) { double v = U[(long)i * m + k]; nrm += v * v; }
nrm = sqrt(nrm); if (nrm < 1e-300) nrm = 1e-300;
for (int i = 0; i < m; ++i) U[(long)i * m + k] /= nrm;
}
}
"""
_KERN = {}
try:
_KERN["rank1_merge"] = NvrtcKernel(
_SRC, "rank1_merge", shared_carveout=100 * 1024)
except Exception:
_COMPILE_OK = False
def _merge_cuda(ds, zs, rho_abs):
"""ds,zs (nsub,m) fp64 ascending; rho_abs (nsub,) >0. -> mu (nsub,m), U (nsub,m,m)."""
nsub, m = ds.shape
dev = ds.device
mu = torch.empty(nsub, m, device=dev, dtype=torch.float64)
U = torch.empty(nsub, m, m, device=dev, dtype=torch.float64)
k = _KERN["rank1_merge"]
block = 256
smem = (4 * m + block) * 8
k.launch((nsub, 1, 1), (block, 1, 1),
[ds.reshape(-1), zs.reshape(-1), rho_abs, mu.reshape(-1),
U.reshape(-1), m], shared=smem)
return mu, U
def rank1_eigh_cuda(d, rho, z):
"""Signed rho, unsorted d; returns lam ascending, U (B,m,m) cols=vecs,
rows in original d order. Torch does sort/sign; CUDA does the O(m^2) core."""
bsz, m = d.shape
dev = d.device
d = d.double(); z = z.double()
if not torch.is_tensor(rho):
rho = torch.full((bsz,), float(rho), device=dev, dtype=torch.float64)
else:
rho = rho.double()
s = torch.where(rho >= 0, 1.0, -1.0) # (bsz,)
dd = s.view(bsz, 1) * d # broadcast over m
ds, perm = torch.sort(dd, dim=-1)
zs = torch.gather(z, 1, perm)
mu, U = _merge_cuda(ds.contiguous(), zs.contiguous(), rho.abs().contiguous())
inv = torch.argsort(perm, dim=-1)
U = torch.gather(U, 1, inv.unsqueeze(-1).expand(bsz, m, m))
lam = s.view(bsz, 1) * mu
lam, kperm = torch.sort(lam, dim=-1)
U = torch.gather(U, 2, kperm.unsqueeze(1).expand(bsz, m, m))
return lam, U
def dc_tridiag_cuda(d, e):
"""Iterative batched D&C with the CUDA merge kernel. d(B,n),e(B,n-1)."""
bsz, n = d.shape
dev = d.device
L = LEAF
assert n % L == 0, n
nb = n // L
dmod = d.clone()
bpos = torch.arange(L, n, L, device=dev)
beta_b = e[:, bpos - 1]
dmod[:, bpos - 1] -= beta_b
dmod[:, bpos] -= beta_b
efull = torch.zeros(bsz, nb, L - 1, device=dev)
for j in range(nb):
efull[:, j, :] = e[:, j * L:j * L + L - 1]
Tl = (torch.diag_embed(dmod.reshape(bsz * nb, L))
+ torch.diag_embed(efull.reshape(bsz * nb, L - 1), 1)
+ torch.diag_embed(efull.reshape(bsz * nb, L - 1), -1))
lam_l, Q_l = torch.linalg.eigh(Tl)
cur_lam = lam_l.reshape(bsz, nb, L)
cur_Q = Q_l.reshape(bsz, nb, L, L)
s = L
nblk = nb
while nblk > 1:
pairs = nblk // 2
bp = ((2 * torch.arange(pairs, device=dev) + 1) * s) - 1
beta = e[:, bp]
lamL = cur_lam[:, 0::2, :]; lamR = cur_lam[:, 1::2, :]
QL = cur_Q[:, 0::2, :, :]; QR = cur_Q[:, 1::2, :, :]
D = torch.cat([lamL, lamR], dim=-1)
z = torch.cat([QL[:, :, -1, :], QR[:, :, 0, :]], dim=-1)
BP = bsz * pairs
Lam2, U2 = rank1_eigh_cuda(D.reshape(BP, 2 * s), beta.reshape(BP),
z.reshape(BP, 2 * s))
U2 = U2.to(cur_Q.dtype).reshape(bsz, pairs, 2 * s, 2 * s)
newQ = torch.empty(bsz, pairs, 2 * s, 2 * s, device=dev, dtype=cur_Q.dtype)
newQ[:, :, :s, :] = QL @ U2[:, :, :s, :]
newQ[:, :, s:, :] = QR @ U2[:, :, s:, :]
cur_lam = Lam2.to(cur_Q.dtype).reshape(bsz, pairs, 2 * s)
cur_Q = newQ
s *= 2; nblk = pairs
return cur_lam[:, 0, :], cur_Q[:, 0, :, :]
_SRC_F32 = r"""
#include <cuda_runtime.h>
#define R20_FLT_MIN 1.1754943508222875e-38f
#define R20_FLT_EPSILON 1.1920928955078125e-7f
__device__ __forceinline__ float safe_div(float numerator, float denominator)
{
if (fabsf(denominator) < R20_FLT_MIN)
denominator = copysignf(
R20_FLT_MIN, denominator == 0.0f ? 1.0f : denominator);
return numerator / denominator;
}
__device__ __forceinline__ float strict_midpoint(float lo, float hi)
{
const float lo_in = nextafterf(lo, hi);
const float hi_in = nextafterf(hi, lo);
const float midpoint = lo + (hi - lo) * 0.5f;
if (midpoint <= lo)
return lo_in;
if (midpoint >= hi)
return hi_in;
return midpoint;
}
__device__ __forceinline__ bool strict_candidate(
float candidate, float lo, float hi, float* result)
{
if (!isfinite(candidate))
return false;
const float lo_in = nextafterf(lo, hi);
const float hi_in = nextafterf(hi, lo);
if (lo_in > hi_in)
return false;
if (candidate <= lo)
candidate = lo_in;
else if (candidate >= hi)
candidate = hi_in;
if (!(candidate > lo && candidate < hi))
return false;
*result = candidate;
return true;
}
extern "C" __global__ void rank1_merge_f32_givens(
const float* __restrict__ ds_all,
const float* __restrict__ zs_all,
const float* __restrict__ rho_all,
float* __restrict__ mu_all,
float* __restrict__ U_all,
int* __restrict__ nactive_all,
int* __restrict__ iterations_all,
int* __restrict__ status_all,
long long collect_stats,
long long m)
{
const int sub = blockIdx.x;
const int tid = threadIdx.x;
const int nt = blockDim.x;
const float* ds = ds_all + (long)sub * m;
const float* zs = zs_all + (long)sub * m;
const float rho = rho_all[sub];
float* mu = mu_all + (long)sub * m;
float* U = U_all + (long)sub * (long)m * m;
extern __shared__ float sh[];
float* sd = sh;
float* sz = sd + m;
float* smu = sz + m;
float* shat = smu + m;
float* rotc = shat + m;
float* rots = rotc + m;
float* red = rots + m;
int* apos = reinterpret_cast<int*>(red + nt);
int* aidx = apos + m;
int* rotp = aidx + m;
int* rotq = rotp + m;
int* counts = rotq + m;
for (int i = tid; i < m; i += nt) {
sd[i] = ds[i];
sz[i] = zs[i];
smu[i] = ds[i];
shat[i] = 0.0f;
apos[i] = -1;
}
__syncthreads();
float amax = 0.0f;
float sumz2_local = 0.0f;
for (int i = tid; i < m; i += nt) {
amax = fmaxf(amax, fabsf(sd[i]));
sumz2_local += sz[i] * sz[i];
}
red[tid] = amax;
__syncthreads();
for (int stride = nt / 2; stride > 0; stride >>= 1) {
if (tid < stride)
red[tid] = fmaxf(red[tid], red[tid + stride]);
__syncthreads();
}
const float smax = red[0];
red[tid] = sumz2_local;
__syncthreads();
for (int stride = nt / 2; stride > 0; stride >>= 1) {
if (tid < stride)
red[tid] += red[tid + stride];
__syncthreads();
}
const float initial_sumz2 = red[0];
const float scale = smax + rho * initial_sumz2;
const float wtol = 1.0e-6f * scale;
const float gtol = 1.0e-5f * scale;
if (tid == 0) {
int nr = 0;
int previous = -1;
for (int i = 0; i < m; ++i) {
const bool weighted = rho * sz[i] * sz[i] > wtol;
apos[i] = weighted ? 0 : -1;
if (!weighted)
sz[i] = 0.0f;
}
for (int i = 0; i < m; ++i) {
if (apos[i] < 0)
continue;
if (previous >= 0 && (sd[i] - sd[previous]) <= gtol) {
const float zi = sz[previous];
const float zj = sz[i];
const float radius = hypotf(zi, zj);
if (radius > 1.0e-30f) {
rotp[nr] = previous;
rotq[nr] = i;
rotc[nr] = zi / radius;
rots[nr] = zj / radius;
++nr;
sz[previous] = radius;
}
sz[i] = 0.0f;
apos[i] = -1;
} else {
previous = i;
}
}
int na = 0;
float active_sumz2 = 0.0f;
for (int i = 0; i < m; ++i) {
if (apos[i] >= 0) {
apos[i] = na;
aidx[na] = i;
++na;
active_sumz2 += sz[i] * sz[i];
}
}
counts[0] = nr;
counts[1] = na;
red[0] = active_sumz2;
if (collect_stats)
nactive_all[sub] = na;
}
__syncthreads();
const int nr = counts[0];
const int na = counts[1];
const float active_sumz2 = red[0];
const float inv_active_sumz2 =
active_sumz2 > 0.0f ? 1.0f / active_sumz2 : 0.0f;
const float rho_normalized = rho * active_sumz2;
const float rho_inverse = safe_div(1.0f, rho_normalized);
for (int active_root = tid; active_root < na; active_root += nt) {
const int pole_slot = aidx[active_root];
const bool largest = active_root == na - 1;
const float global_lo = sd[pole_slot];
const float global_hi = largest
? sd[pole_slot] + rho_normalized
: sd[aidx[active_root + 1]];
float lo;
float hi;
float x;
bool origin_at_i = false;
if (!largest) {
const int next_slot = aidx[active_root + 1];
const float di = sd[pole_slot];
const float dip1 = sd[next_slot];
const float gap = dip1 - di;
const float midpoint = di + gap * 0.5f;
float psi = 0.0f;
float phi = 0.0f;
for (int b = 0; b < active_root; ++b) {
const int i = aidx[b];
const float z2 = sz[i] * sz[i] * inv_active_sumz2;
psi += safe_div(z2, sd[i] - midpoint);
}
for (int b = active_root + 2; b < na; ++b) {
const int i = aidx[b];
const float z2 = sz[i] * sz[i] * inv_active_sumz2;
phi += safe_div(z2, sd[i] - midpoint);
}
const float zi2 =
sz[pole_slot] * sz[pole_slot] * inv_active_sumz2;
const float zip12 =
sz[next_slot] * sz[next_slot] * inv_active_sumz2;
const float c = rho_inverse + psi + phi;
const float delta_i = di - midpoint;
const float delta_ip1 = dip1 - midpoint;
const float w = c + safe_div(zi2, delta_i)
+ safe_div(zip12, delta_ip1);
float tau;
if (w > 0.0f) {
origin_at_i = true;
const float a = c * gap + zi2 + zip12;
const float b = zi2 * gap;
const float discriminant = fabsf(a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (fabsf(c) <= R20_FLT_MIN)
tau = gap * 0.25f;
else if (a > 0.0f)
tau = safe_div(2.0f * b, a + root);
else
tau = safe_div(a - root, 2.0f * c);
lo = di;
hi = midpoint;
x = di + tau;
} else {
origin_at_i = false;
const float a = c * gap - zi2 - zip12;
const float b = zip12 * gap;
const float discriminant = fabsf(a * a + 4.0f * b * c);
const float root = sqrtf(discriminant);
if (fabsf(c) <= R20_FLT_MIN)
tau = -gap * 0.25f;
else if (a < 0.0f)
tau = safe_div(2.0f * b, a - root);
else
tau = safe_div(-(a + root), 2.0f * c);
lo = midpoint;
hi = dip1;
x = dip1 + tau;
}
float strict_x;
x = strict_candidate(x, lo, hi, &strict_x)
? strict_x : strict_midpoint(lo, hi);
} else {
const int previous_slot = na > 1 ? aidx[na - 2] : pole_slot;
const float dn = sd[pole_slot];
const float gap = dn - sd[previous_slot];
const float upper = dn + rho_normalized;
const float midpoint_tau = rho_normalized * 0.5f;
const float midpoint = dn + midpoint_tau;
float psi = 0.0f;
for (int b = 0; b < na - 2; ++b) {
const int i = aidx[b];
const float z2 = sz[i] * sz[i] * inv_active_sumz2;
psi += safe_div(z2, sd[i] - midpoint);
}
const float znm12 =
sz[previous_slot] * sz[previous_slot] * inv_active_sumz2;
const float zn2 =
sz[pole_slot] * sz[pole_slot] * inv_active_sumz2;
const float c = rho_inverse + psi;
const float w = c
+ safe_div(znm12, sd[previous_slot] - midpoint)
+ safe_div(zn2, dn - midpoint);
const float a = -c * gap + znm12 + zn2;
const float b = zn2 * gap;
const float discriminant = fabsf(a * a + 4.0f * b * c);
const float root = sqrtf(discriminant);
float tau_model;
if (fabsf(c) <= R20_FLT_MIN)
tau_model = midpoint_tau;
else if (a < 0.0f)
tau_model = safe_div(2.0f * b, root - a);
else
tau_model = safe_div(a + root, 2.0f * c);
float tau;
if (w <= 0.0f) {
const float temp =
safe_div(znm12, gap + rho_normalized)
+ safe_div(zn2, rho_normalized);
tau = c <= temp ? rho_normalized : tau_model;
lo = midpoint;
hi = upper;
} else {
tau = tau_model;
lo = dn;
hi = midpoint;
}
x = dn + tau;
float strict_x;
x = strict_candidate(x, lo, hi, &strict_x)
? strict_x : strict_midpoint(lo, hi);
}
bool converged = false;
bool used_safeguard = false;
int iterations = 0;
const int origin_active = largest
? na - 1
: (origin_at_i ? active_root : active_root + 1);
const int origin_slot = aidx[origin_active];
#pragma unroll
for (int iteration = 1; iteration <= 16; ++iteration) {
if (converged)
continue;
iterations = iteration;
float psi = 0.0f;
float phi = 0.0f;
float dpsi = 0.0f;
float dphi = 0.0f;
float abs_sum = 0.0f;
for (int b = 0; b < origin_active; ++b) {
const int i = aidx[b];
const float delta = sd[i] - x;
const float z2 = sz[i] * sz[i] * inv_active_sumz2;
const float term = safe_div(z2, delta);
psi += term;
dpsi += safe_div(z2, delta * delta);
abs_sum += fabsf(term);
}
for (int b = origin_active + 1; b < na; ++b) {
const int i = aidx[b];
const float delta = sd[i] - x;
const float z2 = sz[i] * sz[i] * inv_active_sumz2;
const float term = safe_div(z2, delta);
phi += term;
dphi += safe_div(z2, delta * delta);
abs_sum += fabsf(term);
}
const float origin_delta = sd[origin_slot] - x;
const float origin_z2 =
sz[origin_slot] * sz[origin_slot] * inv_active_sumz2;
const float central = safe_div(origin_z2, origin_delta);
const float dcentral =
safe_div(origin_z2, origin_delta * origin_delta);
const float derivative = dpsi + dphi + dcentral;
const float w = rho_inverse + psi + phi + central;
const float error_scale =
8.0f * (fabsf(psi) + fabsf(phi))
+ abs_sum + 2.0f * fabsf(rho_inverse)
+ 3.0f * fabsf(central) + fabsf(x) * derivative;
if (
fabsf(w)
<= R20_FLT_EPSILON * fmaxf(error_scale, 1.0f)
) {
converged = true;
continue;
}
if (w < 0.0f)
lo = fmaxf(lo, x);
else
hi = fminf(hi, x);
if (nextafterf(lo, hi) >= hi) {
x = strict_midpoint(global_lo, global_hi);
converged = true;
continue;
}
float rational = 0.0f;
bool have_rational = iteration <= 12;
if (have_rational && !largest) {
const int next_slot = aidx[active_root + 1];
const float delta_i = sd[pole_slot] - x;
const float delta_ip1 = sd[next_slot] - x;
const float gap = sd[next_slot] - sd[pole_slot];
float c;
if (origin_at_i) {
const float zi2 =
sz[pole_slot] * sz[pole_slot] * inv_active_sumz2;
const float local = safe_div(zi2, delta_i);
c = w - delta_ip1 * derivative + gap * local * local;
} else {
const float zip12 =
sz[next_slot] * sz[next_slot] * inv_active_sumz2;
const float local = safe_div(zip12, delta_ip1);
c = w - delta_i * derivative - gap * local * local;
}
const float a =
(delta_i + delta_ip1) * w
- delta_i * delta_ip1 * derivative;
const float b = delta_i * delta_ip1 * w;
float eta;
if (fabsf(c) <= R20_FLT_MIN) {
if (fabsf(a) <= R20_FLT_MIN) {
have_rational = false;
eta = 0.0f;
} else {
eta = safe_div(b, a);
}
} else {
const float discriminant =
fabsf(a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
eta = a <= 0.0f
? safe_div(a - root, 2.0f * c)
: safe_div(2.0f * b, a + root);
}
rational = x + eta;
if (!isfinite(rational))
have_rational = false;
} else if (have_rational) {
const int previous_slot =
na > 1 ? aidx[na - 2] : pole_slot;
const float delta_nm1 = sd[previous_slot] - x;
const float delta_n = sd[pole_slot] - x;
const float last_z2 =
sz[pole_slot] * sz[pole_slot] * inv_active_sumz2;
const float last_derivative =
safe_div(last_z2, delta_n * delta_n);
const float c = fabsf(
w - delta_nm1 * (derivative - last_derivative)
- delta_n * last_derivative);
if (fabsf(c) <= R20_FLT_MIN) {
have_rational = false;
} else {
const float a =
(delta_nm1 + delta_n) * w
- delta_nm1 * delta_n * derivative;
const float b = delta_nm1 * delta_n * w;
const float discriminant =
fabsf(a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
const float eta = a >= 0.0f
? safe_div(a + root, 2.0f * c)
: safe_div(2.0f * b, a - root);
rational = x + eta;
if (!isfinite(rational))
have_rational = false;
}
}
float candidate;
bool strict_ok = have_rational
&& strict_candidate(rational, lo, hi, &candidate);
const bool correct_direction =
have_rational && w * (rational - x) < 0.0f;
if (!strict_ok || !correct_direction) {
used_safeguard = true;
const float newton = x - safe_div(w, derivative);
if (!strict_candidate(newton, lo, hi, &candidate))
candidate = strict_midpoint(lo, hi);
}
x = candidate;
}
smu[pole_slot] = x;
if (collect_stats) {
iterations_all[(long)sub * m + pole_slot] = iterations;
int status = 0;
if (!(x > global_lo && x < global_hi))
status |= 1;
if (!converged)
status |= 2;
if (used_safeguard)
status |= 4;
status_all[(long)sub * m + pole_slot] = status;
}
}
__syncthreads();
for (int active_row = tid; active_row < na; active_row += nt) {
const int i = aidx[active_row];
float product = 1.0f;
for (int active_root = 0; active_root < na; ++active_root) {
const int k = aidx[active_root];
const float numerator = smu[k] - sd[i];
if (k == i) {
product *= numerator;
} else {
float denominator = sd[k] - sd[i];
if (fabsf(denominator) < 1.0e-30f)
denominator = copysignf(
1.0e-30f, denominator == 0.0f ? 1.0f : denominator);
product *= numerator / denominator;
}
}
shat[i] = copysignf(sqrtf(fabsf(product)), sz[i]);
}
__syncthreads();
const bool sparse = 4 * na <= 3 * m;
if (sparse) {
for (long index = tid; index < (long)m * m; index += nt)
U[index] = 0.0f;
__syncthreads();
for (long index = tid; index < (long)na * na; index += nt) {
const int active_row = (int)(index / na);
const int active_column =
(int)(index - (long)active_row * na);
const int i = aidx[active_row];
const int k = aidx[active_column];
float denominator = sd[i] - smu[k];
if (fabsf(denominator) < 1.0e-30f)
denominator = copysignf(
1.0e-30f, denominator == 0.0f ? 1.0f : denominator);
U[(long)i * m + k] = shat[i] / denominator;
}
for (int i = tid; i < m; i += nt) {
if (apos[i] < 0)
U[(long)i * m + i] = 1.0f;
}
} else {
for (int i = tid; i < m; i += nt) {
for (int k = 0; k < m; ++k) {
float value;
if (apos[k] >= 0) {
if (apos[i] >= 0) {
float denominator = sd[i] - smu[k];
if (fabsf(denominator) < 1.0e-30f)
denominator = copysignf(
1.0e-30f,
denominator == 0.0f ? 1.0f : denominator);
value = shat[i] / denominator;
} else {
value = 0.0f;
}
} else {
value = i == k ? 1.0f : 0.0f;
}
U[(long)i * m + k] = value;
}
}
}
__syncthreads();
for (int active_root = tid; active_root < na; active_root += nt) {
const int k = aidx[active_root];
float norm2 = 0.0f;
for (int active_row = 0; active_row < na; ++active_row) {
const int i = aidx[active_row];
const float value = U[(long)i * m + k];
norm2 += value * value;
}
const float inverse_norm = rsqrtf(fmaxf(norm2, 1.0e-30f));
for (int active_row = 0; active_row < na; ++active_row) {
const int i = aidx[active_row];
U[(long)i * m + k] *= inverse_norm;
}
}
for (int i = tid; i < m; i += nt)
mu[i] = smu[i];
__syncthreads();
for (int rotation = nr - 1; rotation >= 0; --rotation) {
const int p = rotp[rotation];
const int q = rotq[rotation];
const float cosine = rotc[rotation];
const float sine = rots[rotation];
for (int k = tid; k < m; k += nt) {
const float up = U[(long)p * m + k];
const float uq = U[(long)q * m + k];
U[(long)p * m + k] = cosine * up - sine * uq;
U[(long)q * m + k] = sine * up + cosine * uq;
}
__syncthreads();
}
}
"""
_KERN_F32 = {}
try:
_KERN_F32["rank1_merge_f32_givens"] = NvrtcKernel(
_SRC_F32, "rank1_merge_f32_givens",
shared_carveout=100 * 1024)
except Exception:
_COMPILE_OK = False
def _merge_cuda_f32(ds, zs, rho_abs):
"""Solve sorted positive-rho compact rational merges in FP32."""
nsub, m = ds.shape
dev = ds.device
mu = torch.empty(nsub, m, device=dev, dtype=torch.float32)
U = torch.empty(nsub, m, m, device=dev, dtype=torch.float32)
dummy = torch.empty(1, device=dev, dtype=torch.int32)
k = _KERN_F32["rank1_merge_f32_givens"]
block = 256
smem = (10 * m + block) * 4 + 8
k.launch((nsub, 1, 1), (block, 1, 1),
[ds.reshape(-1), zs.reshape(-1), rho_abs, mu.reshape(-1),
U.reshape(-1), dummy, dummy, dummy, 0, m], shared=smem)
return mu, U
def rank1_eigh_cuda_f32(d, rho, z):
"""Signed rho, unsorted d; returns lam ascending, U (B,m,m) cols=vecs,
rows in original d order. The merge applies proper near-tie rotations."""
bsz, m = d.shape
dev = d.device
d = d.float(); z = z.float()
if not torch.is_tensor(rho):
rho = torch.full((bsz,), float(rho), device=dev, dtype=torch.float32)
else:
rho = rho.float()
s = torch.where(rho >= 0, 1.0, -1.0)
dd = s.view(bsz, 1) * d
ds, perm = torch.sort(dd, dim=-1)
zs = torch.gather(z, 1, perm)
rho_abs = rho.abs().contiguous()
mu, U = _merge_cuda_f32(
ds.contiguous(), zs.contiguous(), rho_abs)
inv = torch.argsort(perm, dim=-1)
U = torch.gather(U, 1, inv.unsqueeze(-1).expand(bsz, m, m))
lam = s.view(bsz, 1) * mu
lam, kperm = torch.sort(lam, dim=-1)
U = torch.gather(U, 2, kperm.unsqueeze(1).expand(bsz, m, m))
return lam, U
_DEFERRED_SRC = r"""
#include <cuda_runtime.h>
#include <math_constants.h>
extern "C" __global__ void r22_pack_sorted_children(
const float* __restrict__ child_values,
const float* __restrict__ child_first,
const float* __restrict__ child_last,
const float* __restrict__ off_diagonal,
float* __restrict__ sorted_poles,
float* __restrict__ sorted_weights,
float* __restrict__ rho_abs,
int* __restrict__ original_rows,
signed char* __restrict__ signs,
long long batch,
long long child_count,
long long width,
long long n)
{
const int parent = blockIdx.x;
const int tid = threadIdx.x;
const int pairs = (int)(child_count / 2);
const int matrix = parent / pairs;
const int pair = parent - matrix * pairs;
const int merged_width = 2 * (int)width;
const long long left_child =
((long long)matrix * child_count + 2 * pair) * width;
const long long right_child = left_child + width;
const long long output = (long long)parent * merged_width;
const long long boundary =
(long long)matrix * (n - 1)
+ (2LL * pair + 1LL) * width - 1LL;
const float beta = off_diagonal[boundary];
const int sign = beta >= 0.0f ? 1 : -1;
if (tid == 0) {
rho_abs[parent] = fabsf(beta);
signs[parent] = (signed char)sign;
int left_position = sign > 0 ? 0 : (int)width - 1;
int right_position = sign > 0 ? 0 : (int)width - 1;
for (int destination = 0; destination < merged_width; ++destination) {
const bool left_valid = sign > 0
? left_position < width
: left_position >= 0;
const bool right_valid = sign > 0
? right_position < width
: right_position >= 0;
const float left_value = left_valid
? sign * child_values[left_child + left_position]
: CUDART_INF_F;
const float right_value = right_valid
? sign * child_values[right_child + right_position]
: CUDART_INF_F;
const bool take_left =
left_valid && (!right_valid || left_value <= right_value);
const int child_position =
take_left ? left_position : right_position;
const int original =
take_left ? child_position : (int)width + child_position;
sorted_poles[output + destination] =
take_left ? left_value : right_value;
original_rows[output + destination] = original;
if (take_left)
left_position += sign;
else
right_position += sign;
}
}
__syncthreads();
for (int destination = tid;
destination < merged_width;
destination += blockDim.x) {
const int original = original_rows[output + destination];
sorted_weights[output + destination] =
original < width
? child_last[left_child + original]
: child_first[right_child + original - width];
}
}
extern "C" __global__ void r22_restore_and_propagate(
const float* __restrict__ sorted_values,
const float* __restrict__ sorted_vectors,
const int* __restrict__ original_rows,
const signed char* __restrict__ signs,
const float* __restrict__ child_first,
const float* __restrict__ child_last,
float* __restrict__ values,
float* __restrict__ factors,
float* __restrict__ parent_first,
float* __restrict__ parent_last,
long long child_count,
long long width,
long long propagate_boundaries)
{
const int parent = blockIdx.x;
const int tid = threadIdx.x;
const int pairs = (int)(child_count / 2);
const int matrix = parent / pairs;
const int pair = parent - matrix * pairs;
const int merged_width = 2 * (int)width;
const long long matrix_offset =
(long long)parent * merged_width * merged_width;
const long long vector_offset = (long long)parent * merged_width;
const long long left_child =
((long long)matrix * child_count + 2 * pair) * width;
const long long right_child = left_child + width;
const int sign = (int)signs[parent];
extern __shared__ unsigned char dynamic_shared[];
float* eigen_keys = reinterpret_cast<float*>(dynamic_shared);
int* eigen_order = reinterpret_cast<int*>(
eigen_keys + merged_width);
for (int index = tid; index < merged_width; index += blockDim.x) {
eigen_keys[index] =
sign * sorted_values[vector_offset + index];
eigen_order[index] = index;
}
__syncthreads();
for (int span = 2; span <= merged_width; span <<= 1) {
for (int stride = span >> 1; stride > 0; stride >>= 1) {
for (int index = tid;
index < merged_width;
index += blockDim.x) {
const int partner = index ^ stride;
if (partner > index) {
const bool ascending = (index & span) == 0;
const float left_key = eigen_keys[index];
const float right_key = eigen_keys[partner];
const bool exchange = ascending
? left_key > right_key
: left_key < right_key;
if (exchange) {
eigen_keys[index] = right_key;
eigen_keys[partner] = left_key;
const int left_order = eigen_order[index];
eigen_order[index] = eigen_order[partner];
eigen_order[partner] = left_order;
}
}
}
__syncthreads();
}
}
for (long long index = tid;
index < (long long)merged_width * merged_width;
index += blockDim.x) {
const int sorted_row = (int)(index / merged_width);
const int output_column =
(int)(index - (long long)sorted_row * merged_width);
const int original_row =
original_rows[vector_offset + sorted_row];
const int sorted_column = eigen_order[output_column];
factors[
matrix_offset
+ (long long)original_row * merged_width
+ output_column
] = sorted_vectors[
matrix_offset
+ (long long)sorted_row * merged_width
+ sorted_column
];
}
for (int sorted_column = tid;
sorted_column < merged_width;
sorted_column += blockDim.x) {
values[vector_offset + sorted_column] =
eigen_keys[sorted_column];
}
__syncthreads();
if (propagate_boundaries) {
for (int column = tid; column < merged_width; column += blockDim.x) {
float first_sum = 0.0f;
float last_sum = 0.0f;
for (int index = 0; index < width; ++index) {
first_sum +=
child_first[left_child + index]
* factors[
matrix_offset
+ (long long)index * merged_width
+ column
];
last_sum +=
child_last[right_child + index]
* factors[
matrix_offset
+ (long long)(width + index) * merged_width
+ column
];
}
parent_first[vector_offset + column] = first_sum;
parent_last[vector_offset + column] = last_sum;
}
}
}
"""
_DEFERRED_PACK = None
_DEFERRED_RESTORE = None
try:
_DEFERRED_PACK = NvrtcKernel(
_DEFERRED_SRC, "r22_pack_sorted_children")
_DEFERRED_RESTORE = NvrtcKernel(
_DEFERRED_SRC, "r22_restore_and_propagate")
except Exception:
_COMPILE_OK = False
def _deferred_merge_level(
child_values, child_first, child_last, off_diagonal, factors
):
batch, child_count, width = child_values.shape
pairs = child_count // 2
merged_width = 2 * width
parent_count = batch * pairs
device = child_values.device
sorted_poles = torch.empty(
parent_count, merged_width, device=device, dtype=torch.float32)
sorted_weights = torch.empty_like(sorted_poles)
rho_abs = torch.empty(
parent_count, device=device, dtype=torch.float32)
original_rows = torch.empty(
parent_count, merged_width, device=device, dtype=torch.int32)
signs = torch.empty(parent_count, device=device, dtype=torch.int8)
sorted_values = torch.empty_like(sorted_poles)
parent_values = torch.empty_like(sorted_poles)
sorted_vectors = torch.empty(
parent_count, merged_width, merged_width,
device=device, dtype=torch.float32)
parent_first = torch.empty_like(sorted_poles)
parent_last = torch.empty_like(sorted_poles)
output_factors = torch.empty_like(sorted_vectors)
_DEFERRED_PACK.launch(
(parent_count, 1, 1), (256, 1, 1),
[
child_values.reshape(-1),
child_first.reshape(-1),
child_last.reshape(-1),
off_diagonal.reshape(-1),
sorted_poles.reshape(-1),
sorted_weights.reshape(-1),
rho_abs,
original_rows.reshape(-1),
signs,
batch,
child_count,
width,
off_diagonal.shape[1] + 1,
],
)
dummy = torch.empty(1, device=device, dtype=torch.int32)
block = 256
shared = (10 * merged_width + block) * 4 + 8
_KERN_F32["rank1_merge_f32_givens"].launch(
(parent_count, 1, 1), (block, 1, 1),
[
sorted_poles.reshape(-1),
sorted_weights.reshape(-1),
rho_abs,
sorted_values.reshape(-1),
sorted_vectors.reshape(-1),
dummy,
dummy,
dummy,
0,
merged_width,
],
shared=shared,
)
_DEFERRED_RESTORE.launch(
(parent_count, 1, 1), (256, 1, 1),
[
sorted_values.reshape(-1),
sorted_vectors.reshape(-1),
original_rows.reshape(-1),
signs,
child_first.reshape(-1),
child_last.reshape(-1),
parent_values.reshape(-1),
output_factors.reshape(-1),
parent_first.reshape(-1),
parent_last.reshape(-1),
child_count,
width,
int(pairs > 1),
],
shared=2 * merged_width * 4,
)
factors.append(
output_factors.reshape(
batch, pairs, merged_width, merged_width))
return (
parent_values.reshape(batch, pairs, merged_width),
parent_first.reshape(batch, pairs, merged_width),
parent_last.reshape(batch, pairs, merged_width),
)
def _deferred_assemble_level(current, factor):
batch, child_count, width, _ = current.shape
pairs = child_count // 2
merged_width = 2 * width
children = current.reshape(
batch * child_count, width, width)
halves = factor.reshape(
batch, pairs, 2, width, merged_width
).reshape(batch * child_count, width, merged_width)
return torch.bmm(children, halves).reshape(
batch, pairs, merged_width, merged_width)
def dc_tridiag_cuda_deferred(diagonal, off_diagonal, leaf=32):
diagonal_work = diagonal.float()
off_diagonal_work = off_diagonal.float()
batch, n = diagonal_work.shape
if n % leaf:
raise ValueError(f"n={n} must be divisible by leaf={leaf}")
blocks = n // leaf
modified = diagonal_work.clone()
boundaries = torch.arange(
leaf, n, leaf, device=diagonal_work.device)
boundary_beta = off_diagonal_work[:, boundaries - 1]
modified[:, boundaries - 1] -= boundary_beta
modified[:, boundaries] -= boundary_beta
leaf_e = torch.zeros(
batch, blocks, leaf - 1,
device=diagonal_work.device, dtype=torch.float32)
for block_index in range(blocks):
leaf_e[:, block_index, :] = off_diagonal_work[
:, block_index * leaf:block_index * leaf + leaf - 1]
leaf_matrix = (
torch.diag_embed(modified.reshape(batch * blocks, leaf))
+ torch.diag_embed(
leaf_e.reshape(batch * blocks, leaf - 1), 1)
+ torch.diag_embed(
leaf_e.reshape(batch * blocks, leaf - 1), -1)
)
leaf_values, leaf_vectors = torch.linalg.eigh(leaf_matrix)
current_values = leaf_values.reshape(batch, blocks, leaf)
leaf_vectors = leaf_vectors.reshape(
batch, blocks, leaf, leaf).float()
current_first = leaf_vectors[:, :, 0, :].contiguous()
current_last = leaf_vectors[:, :, -1, :].contiguous()
e32 = off_diagonal_work.contiguous()
factors = []
while current_values.shape[1] > 1:
current_values, current_first, current_last = (
_deferred_merge_level(
current_values,
current_first,
current_last,
e32,
factors,
)
)
current_vectors = _deferred_assemble_level(
leaf_vectors, factors[0])
current_vectors = _deferred_assemble_level(
current_vectors, factors[1])
current_vectors = _deferred_assemble_level(
current_vectors, factors[2])
current_vectors = _deferred_assemble_level(
current_vectors, factors[3])
return (
current_values[:, 0, :].contiguous(),
current_vectors[:, 0, :, :].contiguous(),
)
def dc_tridiag_cuda_hybrid(d, e, leaf=32):
"""D&C with FP64 m64/m128 and compact rational FP32 m256/m512 merges."""
batch, n = d.shape
device = d.device
d64 = d.double()
e64 = e.double()
if n % leaf:
raise ValueError(f"n={n} must be divisible by leaf={leaf}")
blocks = n // leaf
modified = d64.clone()
boundaries = torch.arange(leaf, n, leaf, device=device)
boundary_beta = e64[:, boundaries - 1]
modified[:, boundaries - 1] -= boundary_beta
modified[:, boundaries] -= boundary_beta
leaf_e = torch.zeros(
batch, blocks, leaf - 1, device=device, dtype=torch.float64
)
for block_index in range(blocks):
leaf_e[:, block_index, :] = e64[
:, block_index * leaf : block_index * leaf + leaf - 1
]
leaf_matrix = (
torch.diag_embed(modified.reshape(batch * blocks, leaf))
+ torch.diag_embed(
leaf_e.reshape(batch * blocks, leaf - 1), 1
)
+ torch.diag_embed(
leaf_e.reshape(batch * blocks, leaf - 1), -1
)
)
leaf_values, leaf_vectors = torch.linalg.eigh(leaf_matrix)
current_values = leaf_values.reshape(batch, blocks, leaf)
current_vectors = leaf_vectors.reshape(
batch, blocks, leaf, leaf
)
width = leaf
block_count = blocks
while block_count > 1:
pairs = block_count // 2
positions = (
(2 * torch.arange(pairs, device=device) + 1) * width
) - 1
beta = e64[:, positions]
left_values = current_values[:, 0::2, :]
right_values = current_values[:, 1::2, :]
left_vectors = current_vectors[:, 0::2, :, :]
right_vectors = current_vectors[:, 1::2, :, :]
poles = torch.cat([left_values, right_values], dim=-1)
weights = torch.cat(
[
left_vectors[:, :, -1, :],
right_vectors[:, :, 0, :],
],
dim=-1,
)
problem_count = batch * pairs
merge_width = 2 * width
if merge_width >= 256:
merged_values, merge_vectors = (
rank1_eigh_cuda_f32(
poles.reshape(problem_count, merge_width).float(),
beta.reshape(problem_count).float(),
weights.reshape(problem_count, merge_width).float(),
)
)
merged_values = merged_values.double()
merge_vectors = merge_vectors.double()
else:
merged_values, merge_vectors = rank1_eigh_cuda(
poles.reshape(problem_count, merge_width),
beta.reshape(problem_count),
weights.reshape(problem_count, merge_width),
)
merge_vectors = merge_vectors.reshape(
batch, pairs, merge_width, merge_width
).to(current_vectors.dtype)
next_vectors = torch.empty(
batch,
pairs,
merge_width,
merge_width,
device=device,
dtype=current_vectors.dtype,
)
next_vectors[:, :, :width, :] = (
left_vectors @ merge_vectors[:, :, :width, :]
)
next_vectors[:, :, width:, :] = (
right_vectors @ merge_vectors[:, :, width:, :]
)
current_values = merged_values.reshape(
batch, pairs, merge_width
).to(current_vectors.dtype)
current_vectors = next_vectors
width = merge_width
block_count = pairs
return (
current_values[:, 0, :].float(),
current_vectors[:, 0, :, :].float(),
)
_QL32_SRC = r"""
#include <cuda_runtime.h>
#include <math_constants.h>
#define LEAF 32
#define LEAF_WARPS 8
#define LEAF_ITERS 24
#define LEAF_TOL 1.0e-5f
#define LOWER_THREADS 512
extern "C" __global__ void fused_leaf_ql32(
const float* __restrict__ diagonal,
const float* __restrict__ off_diagonal,
float* __restrict__ values,
float* __restrict__ vectors,
float* __restrict__ first_rows,
float* __restrict__ last_rows,
int* __restrict__ info,
long long batch_ll,
long long n_ll)
{
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int batch = (int)batch_ll;
const int n = (int)n_ll;
const int blocks = n / LEAF;
const int leaf_id = blockIdx.x * LEAF_WARPS + warp;
if (leaf_id >= batch * blocks)
return;
const int matrix = leaf_id / blocks;
const int leaf = leaf_id - matrix * blocks;
const int first = leaf * LEAF;
__shared__ float all_d[LEAF_WARPS][LEAF];
__shared__ float all_e[LEAF_WARPS][LEAF];
__shared__ float all_q[LEAF_WARPS][LEAF * LEAF];
__shared__ float all_rotation[LEAF_WARPS][2];
__shared__ int all_middle[LEAF_WARPS];
__shared__ int all_done[LEAF_WARPS];
__shared__ int all_insert[LEAF_WARPS];
__shared__ int all_failures[LEAF_WARPS];
float* d = all_d[warp];
float* e = all_e[warp];
float* q = all_q[warp];
float* rotation = all_rotation[warp];
int* middle_at = all_middle + warp;
int* done = all_done + warp;
int* insert_at = all_insert + warp;
int* failures = all_failures + warp;
float value = diagonal[(long)matrix * n + first + lane];
if (lane == 0 && leaf > 0)
value -= off_diagonal[(long)matrix * (n - 1) + first - 1];
if (lane == LEAF - 1 && leaf + 1 < blocks)
value -= off_diagonal[(long)matrix * (n - 1) + first + lane];
d[lane] = value;
e[lane] = lane + 1 < LEAF
? off_diagonal[(long)matrix * (n - 1) + first + lane]
: 0.0f;
for (int column = 0; column < LEAF; ++column)
q[column * LEAF + lane] = lane == column ? 1.0f : 0.0f;
if (lane == 0)
*failures = 0;
__syncwarp();
for (int left = 0; left < LEAF; ++left) {
for (int iteration = 0; iteration < LEAF_ITERS; ++iteration) {
if (lane == 0) {
int middle = left;
for (; middle < LEAF - 1; ++middle) {
const float scale = fabsf(d[middle]) + fabsf(d[middle + 1]);
if (fabsf(e[middle]) <= LEAF_TOL * (scale + 1.0e-30f))
break;
}
*middle_at = middle;
*done = middle == left;
if (*done)
e[left] = 0.0f;
}
__syncwarp();
if (*done)
break;
float g = 0.0f;
float sine = 1.0f;
float cosine = 1.0f;
float shift = 0.0f;
if (lane == 0) {
const float edge = e[left];
g = (d[left + 1] - d[left]) / (2.0f * edge);
const float radius = sqrtf(g * g + 1.0f);
float denominator = g + copysignf(radius, g);
if (fabsf(denominator) < 1.0e-30f)
denominator = copysignf(1.0e-30f, denominator);
g = d[*middle_at] - d[left] + edge / denominator;
}
for (int index = *middle_at - 1; index >= left; --index) {
if (lane == 0) {
const float f = sine * e[index];
const float b = cosine * e[index];
const float radius = sqrtf(f * f + g * g);
e[index + 1] = radius;
if (radius > 1.0e-30f) {
sine = f / radius;
cosine = g / radius;
} else {
sine = 0.0f;
cosine = 1.0f;
}
g = d[index + 1] - shift;
const float term =
(d[index] - g) * sine + 2.0f * cosine * b;
shift = sine * term;
d[index + 1] = g + shift;
g = cosine * term - b;
rotation[0] = cosine;
rotation[1] = sine;
}
__syncwarp();
const float old_left = q[index * LEAF + lane];
const float old_right = q[(index + 1) * LEAF + lane];
q[(index + 1) * LEAF + lane] =
rotation[1] * old_left + rotation[0] * old_right;
q[index * LEAF + lane] =
rotation[0] * old_left - rotation[1] * old_right;
}
if (lane == 0) {
d[left] -= shift;
e[left] = g;
e[*middle_at] = 0.0f;
if (iteration + 1 == LEAF_ITERS)
*failures += 1;
}
__syncwarp();
}
}
for (int source = 1; source < LEAF; ++source) {
const float held = q[source * LEAF + lane];
if (lane == 0) {
const float key = d[source];
int position = source;
while (position > 0 && d[position - 1] > key) {
d[position] = d[position - 1];
--position;
}
d[position] = key;
*insert_at = position;
}
__syncwarp();
for (int column = source; column > *insert_at; --column)
q[column * LEAF + lane] = q[(column - 1) * LEAF + lane];
q[*insert_at * LEAF + lane] = held;
}
values[(long)leaf_id * LEAF + lane] = d[lane];
first_rows[(long)leaf_id * LEAF + lane] = q[lane * LEAF];
last_rows[(long)leaf_id * LEAF + lane] =
q[lane * LEAF + LEAF - 1];
for (int column = 0; column < LEAF; ++column)
vectors[(long)leaf_id * LEAF * LEAF + lane * LEAF + column] =
q[column * LEAF + lane];
if (lane == 0)
info[leaf_id] = *failures;
}
"""
_QL32_KERNEL = None
def _ql32_leaves(diagonal, off_diagonal):
global _QL32_KERNEL
if _QL32_KERNEL is None:
_QL32_KERNEL = NvrtcKernel(_QL32_SRC, "fused_leaf_ql32")
batch, n = diagonal.shape
blocks = n // LEAF
device = diagonal.device
values = torch.empty(
batch, blocks, LEAF, device=device, dtype=torch.float32)
vectors = torch.empty(
batch, blocks, LEAF, LEAF, device=device, dtype=torch.float32)
first = torch.empty_like(values)
last = torch.empty_like(values)
info = torch.empty(
batch, blocks, device=device, dtype=torch.int32)
total = batch * blocks
_QL32_KERNEL.launch(
((total + 7) // 8, 1, 1),
(256, 1, 1),
[
diagonal.contiguous().reshape(-1),
off_diagonal.contiguous().reshape(-1),
values.reshape(-1),
vectors.reshape(-1),
first.reshape(-1),
last.reshape(-1),
info.reshape(-1),
batch,
n,
],
)
return values, vectors, first, last
def _deferred_state(diagonal, off_diagonal):
diagonal_work = diagonal.float().contiguous()
off_diagonal_work = off_diagonal.float().contiguous()
current_values, leaf_vectors, current_first, current_last = (
_ql32_leaves(diagonal_work, off_diagonal_work)
)
e32 = off_diagonal_work
factors = []
while current_values.shape[1] > 1:
current_values, current_first, current_last = _deferred_merge_level(
current_values, current_first, current_last, e32, factors)
return current_values[:, 0, :].contiguous(), leaf_vectors, factors
def _apply_panel(matrix, local_row, vectors, compact):
sub = matrix[:, local_row:, :]
matrix[:, local_row:, :] = sub - vectors @ (
compact @ (vectors.transpose(-1, -2) @ sub))
def _apply_eligible(
current, block_width, application_order, next_factor
):
block_start = 512 - block_width
trailing = current[:, -1, :, :]
while (
next_factor < len(application_order)
and application_order[next_factor][0] >= block_start
):
row, vectors, compact = application_order[next_factor]
_apply_panel(trailing, row - block_start, vectors, compact)
next_factor += 1
return next_factor
def _staged_pre_top(leaf_vectors, factors, householder):
application_order = list(reversed(householder))
next_factor = 0
current = _deferred_assemble_level(leaf_vectors, factors[0])
next_factor = _apply_eligible(
current, 64, application_order, next_factor)
current = _deferred_assemble_level(current, factors[1])
next_factor = _apply_eligible(
current, 128, application_order, next_factor)
current = _deferred_assemble_level(current, factors[2])
_apply_eligible(current, 256, application_order, next_factor)
return current
def _combine_top_compact(householder):
selected = [factor for factor in householder if factor[0] <= 241]
batch = selected[0][1].shape[0]
device = selected[0][1].device
total_width = sum(vectors.shape[2] for _, vectors, _ in selected)
combined_vectors = torch.zeros(
batch, 512, total_width, device=device, dtype=torch.float32)
combined_factor = torch.zeros(
batch, total_width, total_width, device=device, dtype=torch.float32)
used = 0
for row, vectors, factor in selected:
width = vectors.shape[2]
padded = torch.zeros(
batch, 512, width, device=device, dtype=torch.float32)
padded[:, row:, :] = vectors
if used:
gram = combined_vectors[:, :, :used].transpose(-1, -2) @ padded
combined_factor[:, :used, used:used + width] = -(
combined_factor[:, :used, :used] @ gram @ factor)
combined_vectors[:, :, used:used + width] = padded
combined_factor[
:, used:used + width, used:used + width] = factor
used += width
return combined_vectors, combined_factor
def _grouped_projection(leaf_vectors, factors, householder):
children = _staged_pre_top(leaf_vectors, factors, householder)
combined_vectors, combined_factor = _combine_top_compact(householder)
current = _deferred_assemble_level(children, factors[3])
output = current[:, 0]
output = output - combined_vectors @ (
combined_factor @ (combined_vectors.transpose(-1, -2) @ output))
return output.contiguous()
def scale_and_reduce_512(A):
"""S0 (pre-scale) + S1 (LATRD reduction), n=512 only."""
absmax = A.abs().amax(dim=(-2, -1)).clamp_min(1e-300)
scale_exp = torch.round(torch.log2(absmax))
scale = torch.pow(2.0, scale_exp)
work = A / scale.view(-1, 1, 1)
d, e, factors = s1_kernel_reduce(work)
return d, e, factors, scale
def full_eigh_512(A):
"""S0/S1 followed by grouped deferred D&C and Householder projection."""
d, e, householder, scale = scale_and_reduce_512(A)
lam, leaf_vectors, dc_factors = _deferred_state(d, e)
Q = _grouped_projection(leaf_vectors, dc_factors, householder)
L = lam * scale.view(-1, 1)
return Q, L
def _valved_eigh_512(data):
"""R54: FP16 eigen screen → FP32 band → FP64 hot; no orth GEMM."""
q, lam = full_eigh_512(data)
n = data.shape[-1]
eps32 = torch.finfo(torch.float32).eps
scale = torch.linalg.matrix_norm(data, ord=1, dim=(-2, -1))
eigen_allowed = 200.0 * n * eps32 * scale
# Cheap FP16 residual screen
q16 = q.to(torch.float16)
a16 = data.to(torch.float16)
lam16 = lam.to(torch.float16)
eigen16 = torch.linalg.matrix_norm(
a16 @ q16 - q16 * lam16[:, None, :],
ord=1, dim=(-2, -1)).float()
needs_fp32 = (eigen16 > 0.35 * eigen_allowed) | ~torch.isfinite(eigen16)
bad = torch.zeros(data.shape[0], dtype=torch.bool, device=data.device)
eigen_res = eigen16
if bool(needs_fp32.any()):
idx = needs_fp32.nonzero(as_tuple=True)[0]
q32 = q.index_select(0, idx)
a32 = data.index_select(0, idx)
lam32 = lam.index_select(0, idx)
eigen32 = torch.linalg.matrix_norm(
a32 @ q32 - q32 * lam32[:, None, :],
ord=1, dim=(-2, -1))
eigen_res = eigen_res.clone()
eigen_res[idx] = eigen32
needs_exact = (eigen32 > 0.5 * eigen_allowed.index_select(0, idx)) | ~torch.isfinite(eigen32)
if bool(needs_exact.any()):
exact_local = needs_exact.nonzero(as_tuple=True)[0]
exact_idx = idx.index_select(0, exact_local)
ad = data.index_select(0, exact_idx).double()
qd = q.index_select(0, exact_idx).double()
ld = lam.index_select(0, exact_idx).double()
eigen_exact = torch.linalg.matrix_norm(
ad @ qd - qd * ld[:, None, :],
ord=1, dim=(-2, -1))
exact_allowed = eigen_allowed.index_select(0, exact_idx)
exact_bad = (eigen_exact > exact_allowed) | ~torch.isfinite(eigen_exact)
bad[exact_idx] = exact_bad
# clear FP32-over marks for exact-checked rows
over32 = (eigen32 > eigen_allowed.index_select(0, idx)) & ~needs_exact
bad[idx] = bad.index_select(0, idx) | over32
else:
bad[idx] = eigen32 > eigen_allowed.index_select(0, idx)
else:
bad = eigen16 > eigen_allowed
bad = bad | ~torch.isfinite(eigen_res)
if bool(bad.any()):
idx = bad.nonzero(as_tuple=True)[0]
fallback_lam, fallback_q = torch.linalg.eigh(data[idx])
q = q.clone()
lam = lam.clone()
q[idx] = fallback_q
lam[idx] = fallback_lam
return q.contiguous(), lam.contiguous()
_CUSOLVER_EIG_MODE_VECTOR = 1
_CUBLAS_FILL_MODE_LOWER = 0
_CUDA_R_32F = 0
_XSYEV_STATE = None
_XSYEV_DISABLED = False
def _solver_status(status, operation):
if status != 0:
raise RuntimeError(f"{operation} failed with status {status}")
def _load_solver_library():
names = [
ctypes.util.find_library("cusolver"),
"libcusolver.so",
]
for root in (
os.environ.get("CUDA_HOME"),
"/usr/local/cuda",
"/usr/local/cuda-12.9",
"/usr/local/cuda-13.3",
):
if root:
names.append(os.path.join(root, "lib64", "libcusolver.so"))
names.append(
os.path.join(
root, "targets", "x86_64-linux", "lib", "libcusolver.so"))
package_root = os.path.abspath(
os.path.join(os.path.dirname(torch.__file__), "..", "nvidia"))
if os.path.isdir(package_root):
for child in os.listdir(package_root):
names.append(
os.path.join(package_root, child, "lib", "libcusolver.so"))
last = None
seen = set()
for name in names:
if not name or name in seen:
continue
seen.add(name)
try:
return ctypes.CDLL(name, mode=ctypes.RTLD_GLOBAL)
except OSError as exc:
last = exc
raise RuntimeError(f"no cuSOLVER library: {last}")
def _declare_solver_signatures(lib):
lib.cusolverDnCreate.restype = ctypes.c_int
lib.cusolverDnCreate.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
lib.cusolverDnDestroy.restype = ctypes.c_int
lib.cusolverDnDestroy.argtypes = [ctypes.c_void_p]
lib.cusolverDnCreateParams.restype = ctypes.c_int
lib.cusolverDnCreateParams.argtypes = [
ctypes.POINTER(ctypes.c_void_p)]
lib.cusolverDnDestroyParams.restype = ctypes.c_int
lib.cusolverDnDestroyParams.argtypes = [ctypes.c_void_p]
lib.cusolverDnXsyevBatched_bufferSize.restype = ctypes.c_int
lib.cusolverDnXsyevBatched_bufferSize.argtypes = [
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.POINTER(ctypes.c_size_t),
ctypes.POINTER(ctypes.c_size_t),
ctypes.c_int64,
]
lib.cusolverDnXsyevBatched.restype = ctypes.c_int
lib.cusolverDnXsyevBatched.argtypes = [
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int64,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_void_p,
ctypes.c_size_t,
ctypes.c_void_p,
ctypes.c_size_t,
ctypes.c_void_p,
ctypes.c_int64,
]
class _XsyevWorkspace:
def __init__(self, device, device_bytes, host_bytes):
self.device_bytes = device_bytes
self.host_bytes = host_bytes
self.device = (
torch.empty(
device_bytes, dtype=torch.uint8, device=device)
if device_bytes else None)
self.host = (
ctypes.create_string_buffer(host_bytes) if host_bytes else None)
@property
def device_pointer(self):
pointer = self.device.data_ptr() if self.device is not None else 0
return ctypes.c_void_p(pointer)
@property
def host_pointer(self):
pointer = ctypes.addressof(self.host) if self.host is not None else 0
return ctypes.c_void_p(pointer)
class _Xsyev176State:
def __init__(self, device):
torch.cuda.set_device(device)
self.device_index = torch.cuda.current_device()
self.lib = _load_solver_library()
_declare_solver_signatures(self.lib)
self.handle = ctypes.c_void_p()
self.params = ctypes.c_void_p()
self.workspaces = {}
try:
_solver_status(
self.lib.cusolverDnCreate(ctypes.byref(self.handle)),
"cusolver create")
_solver_status(
self.lib.cusolverDnCreateParams(ctypes.byref(self.params)),
"cusolver create params")
except Exception:
if self.params.value:
self.lib.cusolverDnDestroyParams(self.params)
self.params = ctypes.c_void_p()
if self.handle.value:
self.lib.cusolverDnDestroy(self.handle)
self.handle = ctypes.c_void_p()
raise
def _device_index(self, matrix):
index = matrix.device.index
return torch.cuda.current_device() if index is None else index
def workspace(self, matrix, values):
batch, n, _ = matrix.shape
device_index = self._device_index(matrix)
key = (device_index, batch, n)
cached = self.workspaces.get(key)
if cached is not None:
return cached
device_bytes = ctypes.c_size_t()
host_bytes = ctypes.c_size_t()
_solver_status(
self.lib.cusolverDnXsyevBatched_bufferSize(
self.handle,
self.params,
_CUSOLVER_EIG_MODE_VECTOR,
_CUBLAS_FILL_MODE_LOWER,
n,
_CUDA_R_32F,
ctypes.c_void_p(matrix.data_ptr()),
n,
_CUDA_R_32F,
ctypes.c_void_p(values.data_ptr()),
_CUDA_R_32F,
ctypes.byref(device_bytes),
ctypes.byref(host_bytes),
batch,
),
"xsyev workspace query")
cached = _XsyevWorkspace(
matrix.device, device_bytes.value, host_bytes.value)
self.workspaces[key] = cached
return cached
def solve(self, matrix):
if matrix.dtype != torch.float32 or not matrix.is_cuda:
raise ValueError("xsyev requires FP32 CUDA data")
batch, n, m = matrix.shape
if n != 176 or m != 176:
raise ValueError("xsyev route is restricted to n=176")
if self._device_index(matrix) != self.device_index:
raise ValueError("xsyev handle belongs to another CUDA device")
work = matrix.clone(memory_format=torch.contiguous_format)
values = torch.empty(
(batch, n), dtype=torch.float32, device=matrix.device)
info = torch.empty(
(batch,), dtype=torch.int32, device=matrix.device)
workspace = self.workspace(work, values)
_solver_status(
self.lib.cusolverDnXsyevBatched(
self.handle,
self.params,
_CUSOLVER_EIG_MODE_VECTOR,
_CUBLAS_FILL_MODE_LOWER,
n,
_CUDA_R_32F,
ctypes.c_void_p(work.data_ptr()),
n,
_CUDA_R_32F,
ctypes.c_void_p(values.data_ptr()),
_CUDA_R_32F,
workspace.device_pointer,
workspace.device_bytes,
workspace.host_pointer,
workspace.host_bytes,
ctypes.c_void_p(info.data_ptr()),
batch,
),
"xsyev batched")
return work.transpose(-2, -1), values
def close(self):
self.workspaces.clear()
if self.params.value:
self.lib.cusolverDnDestroyParams(self.params)
self.params = ctypes.c_void_p()
if self.handle.value:
self.lib.cusolverDnDestroy(self.handle)
self.handle = ctypes.c_void_p()
def _get_xsyev_state(device):
global _XSYEV_STATE
if _XSYEV_STATE is None:
_XSYEV_STATE = _Xsyev176State(device)
return _XSYEV_STATE
def _close_xsyev_state():
global _XSYEV_STATE
if _XSYEV_STATE is not None:
_XSYEV_STATE.close()
_XSYEV_STATE = None
atexit.register(_close_xsyev_state)
def custom_kernel(data):
global _XSYEV_DISABLED
if data.shape[-1] == 176 and not _XSYEV_DISABLED:
try:
return _get_xsyev_state(data.device).solve(data)
except Exception:
_XSYEV_DISABLED = True
_close_xsyev_state()
if data.shape[-1] == 512 and _COMPILE_OK:
try:
return _valved_eigh_512(data)
except Exception:
pass
values, vectors = torch.linalg.eigh(data)
return vectors.contiguous(), values.contiguous()
scrolls · 2473 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