submission 870585
Dortamac · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1140 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-870585?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:14c652303f160f5e768174aa502f46f4398108a527e3bb9346a3d45b70956081
license declaredunknown
license concludedunknown
authorsDortamac
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float b_tile[N * STRIDE];Kernel source
submission.py1140 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_SRC = r"""
#include <torch/extension.h>
#include <vector>
std::vector<torch::Tensor> onesided32_cuda(torch::Tensor input);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("onesided32", &onesided32_cuda, "Shifted one-sided Jacobi eigensolver");
}
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cfloat>
#include <cmath>
#include <vector>
namespace {
constexpr int N = 32;
constexpr int PAIRS = N / 2;
constexpr int ROUNDS = N - 1;
constexpr int SWEEPS = 10;
constexpr int STRIDE = N + 1;
constexpr unsigned FULL_MASK = 0xffffffffu;
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_down_sync(FULL_MASK, value, offset);
}
return value;
}
__device__ __forceinline__ float warp_max(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = fmaxf(value, __shfl_down_sync(FULL_MASK, value, offset));
}
return value;
}
__device__ __forceinline__ float warp_sum_all(float value) {
#pragma unroll
for (int mask = 16; mask > 0; mask >>= 1) {
value += __shfl_xor_sync(FULL_MASK, value, mask);
}
return value;
}
__device__ __forceinline__ int label_at_position(int position, int round) {
return position == 0
? 0
: 1 + (position - 1 - round + 2 * ROUNDS) % ROUNDS;
}
__global__ __launch_bounds__(512, 1)
void onesided32_kernel(
const float* __restrict__ input,
float* __restrict__ eigenvectors,
float* __restrict__ eigenvalues,
int batch) {
const int matrix = blockIdx.x;
if (matrix >= batch) return;
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const float* matrix_input = input + matrix * N * N;
float* matrix_q = eigenvectors + matrix * N * N;
__shared__ float b_tile[N * STRIDE];
__shared__ float column_norms[N];
__shared__ float warp_values[PAIRS];
__shared__ float matrix_scale;
__shared__ int warp_rotated[PAIRS];
__shared__ int continue_sweeps;
__shared__ int permutation[N];
float local_maximum = 0.0f;
#pragma unroll
for (int index = tid; index < N * N; index += blockDim.x) {
local_maximum = fmaxf(local_maximum, fabsf(matrix_input[index]));
}
local_maximum = warp_max(local_maximum);
if (lane == 0) warp_values[warp] = local_maximum;
__syncthreads();
if (warp == 0) {
float value = lane < PAIRS ? warp_values[lane] : 0.0f;
value = warp_max(value);
if (lane == 0) matrix_scale = value;
}
__syncthreads();
const float scale = matrix_scale;
const float inverse_scale = scale > 0.0f ? 1.0f / scale : 0.0f;
constexpr float alpha = 33.0f;
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int column = index - row * N;
const int lower_row = row > column ? row : column;
const int lower_column = row > column ? column : row;
const float value = matrix_input[
lower_row * N + lower_column] * inverse_scale;
b_tile[row * STRIDE + column] =
value + static_cast<float>(row == column) * alpha;
}
__syncthreads();
#pragma unroll
for (int phase = 0; phase < 2; ++phase) {
const int column = warp + phase * PAIRS;
const float value = b_tile[lane * STRIDE + column];
const float norm = warp_sum(value * value);
if (lane == 0) column_norms[column] = norm;
}
__syncthreads();
#pragma unroll 1
for (int sweep = 0; sweep < SWEEPS; ++sweep) {
bool warp_did_rotate = false;
#pragma unroll 1
for (int round = 0; round < ROUNDS; ++round) {
const int first = label_at_position(warp, round);
const int second = label_at_position(N - 1 - warp, round);
const int p = first < second ? first : second;
const int q = first < second ? second : first;
const float b_p = b_tile[lane * STRIDE + p];
const float b_q = b_tile[lane * STRIDE + q];
const float apq = warp_sum_all(b_p * b_q);
float c = 1.0f;
float s = 0.0f;
const float app = column_norms[p];
const float aqq = column_norms[q];
const float rotation_floor =
64.0f * FLT_EPSILON * fmaxf(app, aqq);
if (fabsf(apq) > rotation_floor) {
const float tau = (aqq - app) / (2.0f * apq);
const float root = sqrtf(fmaf(tau, tau, 1.0f));
const float t =
copysignf(1.0f, tau) / (fabsf(tau) + root);
c = rsqrtf(fmaf(t, t, 1.0f));
s = t * c;
if (lane == 0) {
const float cc = c * c;
const float ss = s * s;
const float two_cs_apq = 2.0f * c * s * apq;
column_norms[p] =
fmaf(cc, app, fmaf(ss, aqq, -two_cs_apq));
column_norms[q] =
fmaf(ss, app, fmaf(cc, aqq, two_cs_apq));
}
warp_did_rotate = true;
}
b_tile[lane * STRIDE + p] = fmaf(-s, b_q, c * b_p);
b_tile[lane * STRIDE + q] = fmaf( s, b_p, c * b_q);
__syncthreads();
}
if (lane == 0) warp_rotated[warp] = static_cast<int>(warp_did_rotate);
__syncthreads();
if (tid == 0) {
int any_rotation = 0;
#pragma unroll
for (int pair = 0; pair < PAIRS; ++pair) {
any_rotation |= warp_rotated[pair];
}
continue_sweeps = any_rotation;
}
__syncthreads();
if (continue_sweeps == 0) break;
}
#pragma unroll
for (int phase = 0; phase < 2; ++phase) {
const int column = warp + phase * PAIRS;
const float value = b_tile[lane * STRIDE + column];
const float norm = warp_sum(value * value);
if (lane == 0) column_norms[column] = norm;
}
__syncthreads();
if (warp == 0) {
const float norm = sqrtf(fmaxf(column_norms[lane], 0.0f));
float key = (norm - alpha) * scale;
int column_index = lane;
column_norms[lane] = norm > 0.0f ? 1.0f / norm : 0.0f;
#pragma unroll
for (int width = 2; width <= N; width <<= 1) {
#pragma unroll
for (int stride = width >> 1; stride > 0; stride >>= 1) {
const float other_key =
__shfl_xor_sync(FULL_MASK, key, stride);
const int other_index =
__shfl_xor_sync(FULL_MASK, column_index, stride);
const bool ascending = (lane & width) == 0;
const bool lower_lane = (lane & stride) == 0;
const bool keep_minimum = ascending == lower_lane;
const bool take_other = keep_minimum
? other_key < key
: other_key > key;
if (take_other) {
key = other_key;
column_index = other_index;
}
}
}
eigenvalues[matrix * N + lane] = key;
permutation[lane] = column_index;
}
__syncthreads();
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int column = index - row * N;
const int source_column = permutation[column];
const float inverse_norm = column_norms[source_column];
matrix_q[index] = inverse_norm > 0.0f
? b_tile[row * STRIDE + source_column] * inverse_norm
: static_cast<float>(row == source_column);
}
}
} // namespace
std::vector<torch::Tensor> onesided32_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch, 32, 32]");
TORCH_CHECK(input.size(1) == N && input.size(2) == N, "matrix size must be 32");
const int batch = static_cast<int>(input.size(0));
auto eigenvectors = torch::empty_like(input);
auto eigenvalues = torch::empty({batch, N}, input.options());
onesided32_kernel<<<batch, 512>>>(
input.data_ptr<float>(),
eigenvectors.data_ptr<float>(),
eigenvalues.data_ptr<float>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {eigenvectors, eigenvalues};
}
"""
_module = load_inline(
name="eigh_onesided32_shifted_v11_precomputed_norm",
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=None,
extra_cflags=["-O3", "-std=c++20"],
extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++20"],
with_cuda=True,
verbose=False,
)
TRIDIAG176_CPP_SRC = r"""
#include <torch/extension.h>
#include <vector>
std::vector<torch::Tensor> tridiag176_cuda(torch::Tensor input);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("tridiag176", &tridiag176_cuda, "Householder tridiagonal n=176 eigensolver");
}
"""
TRIDIAG176_CUDA_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cfloat>
#include <cmath>
#include <vector>
namespace {
constexpr int N = 176;
constexpr int LD = N + 1;
constexpr int THREADS = 256;
constexpr int TRI_SMEM_FLOATS = N * LD + 2 * N + THREADS;
constexpr int BACK_SMEM_FLOATS = N * LD + N;
constexpr int MAX_QL_ROTATIONS = 8 * N * N;
__device__ __forceinline__ float stable_hypot(float a, float b) {
a = fabsf(a);
b = fabsf(b);
const float high = fmaxf(a, b);
const float low = fminf(a, b);
if (high == 0.0f) return 0.0f;
const float ratio = low / high;
return high * sqrtf(fmaf(ratio, ratio, 1.0f));
}
__global__ __launch_bounds__(THREADS, 1)
void tridiagonalize176_kernel(
const float* __restrict__ input,
float* __restrict__ reflectors,
float* __restrict__ diagonal,
float* __restrict__ offdiagonal,
float* __restrict__ taus,
int batch) {
const int matrix = blockIdx.x;
const int tid = threadIdx.x;
if (matrix >= batch) return;
extern __shared__ float storage[];
float* T = storage;
float* v = T + N * LD;
float* w = v + N;
float* reduction = w + N;
const int base = matrix * N * N;
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
const float aij = input[base + row * N + col];
const float aji = input[base + col * N + row];
T[row * LD + col] = 0.5f * (aij + aji);
}
__syncthreads();
for (int k = 0; k < N - 2; ++k) {
const int m = N - k - 1;
float sum = 0.0f;
for (int i = tid; i < m; i += blockDim.x) {
const float value = T[(k + 1 + i) * LD + k];
sum = fmaf(value, value, sum);
}
reduction[tid] = sum;
__syncthreads();
for (int stride = THREADS / 2; stride > 0; stride >>= 1) {
if (tid < stride) {
reduction[tid] += reduction[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const float norm = sqrtf(fmaxf(reduction[0], 0.0f));
const float x0 = T[(k + 1) * LD + k];
const float alpha = norm > 0.0f ? -copysignf(norm, x0) : 0.0f;
const float denominator = x0 - alpha;
const float tau = alpha != 0.0f ? (alpha - x0) / alpha : 0.0f;
reduction[0] = alpha;
reduction[1] = tau;
reduction[2] = denominator;
taus[matrix * N + k] = tau;
}
__syncthreads();
const float alpha = reduction[0];
const float tau = reduction[1];
const float denominator = reduction[2];
for (int i = tid; i < m; i += blockDim.x) {
const int row = k + 1 + i;
float value = 0.0f;
if (i == 0) {
value = 1.0f;
T[row * LD + k] = alpha;
T[k * LD + row] = alpha;
} else {
value = denominator != 0.0f
? T[row * LD + k] / denominator
: 0.0f;
T[row * LD + k] = value;
}
v[i] = value;
}
__syncthreads();
for (int i = tid; i < m; i += blockDim.x) {
float value = 0.0f;
const int row = k + 1 + i;
for (int j = 0; j < m; ++j) {
value = fmaf(T[row * LD + (k + 1 + j)], v[j], value);
}
w[i] = tau * value;
}
__syncthreads();
float dot = 0.0f;
for (int i = tid; i < m; i += blockDim.x) {
dot = fmaf(v[i], w[i], dot);
}
reduction[tid] = dot;
__syncthreads();
for (int stride = THREADS / 2; stride > 0; stride >>= 1) {
if (tid < stride) {
reduction[tid] += reduction[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
reduction[0] = -0.5f * tau * reduction[0];
}
__syncthreads();
const float correction = reduction[0];
for (int i = tid; i < m; i += blockDim.x) {
w[i] = fmaf(correction, v[i], w[i]);
}
__syncthreads();
for (int index = tid; index < m * m; index += blockDim.x) {
const int i = index / m;
const int j = index - i * m;
const int row = k + 1 + i;
const int col = k + 1 + j;
T[row * LD + col] -= v[i] * w[j] + w[i] * v[j];
}
__syncthreads();
}
if (tid < N) {
diagonal[matrix * N + tid] = T[tid * LD + tid];
offdiagonal[matrix * N + tid] = tid == 0
? 0.0f
: T[tid * LD + tid - 1];
if (tid >= N - 2) taus[matrix * N + tid] = 0.0f;
}
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
reflectors[base + index] = T[row * LD + col];
}
}
__global__ void tridiagonal_ql_values176_kernel(
const float* __restrict__ diagonal_input,
const float* __restrict__ offdiagonal_input,
int* __restrict__ rotation_indices,
float* __restrict__ rotation_cosines,
float* __restrict__ rotation_sines,
int* __restrict__ rotation_counts,
int* __restrict__ permutations,
float* __restrict__ eigenvalues,
int batch) {
const int matrix = blockIdx.x;
if (matrix >= batch || threadIdx.x != 0) return;
extern __shared__ float storage[];
float* d = storage;
float* e = d + N;
for (int i = 0; i < N; ++i) {
d[i] = diagonal_input[matrix * N + i];
e[i] = i < N - 1
? offdiagonal_input[matrix * N + i + 1]
: 0.0f;
}
const int rotation_base = matrix * MAX_QL_ROTATIONS;
int rotation_count = 0;
for (int l = 0; l < N; ++l) {
int iteration = 0;
while (true) {
int m = l;
for (; m < N - 1; ++m) {
const float scale = fabsf(d[m]) + fabsf(d[m + 1]);
if (fabsf(e[m]) <= 8.0f * FLT_EPSILON * scale) break;
}
if (m == l) break;
if (iteration++ >= 64) {
e[m] = 0.0f;
break;
}
float g = (d[l + 1] - d[l]) / (2.0f * e[l]);
float r = stable_hypot(g, 1.0f);
g = d[m] - d[l] + e[l] / (g + copysignf(r, g));
float s = 1.0f;
float c = 1.0f;
float p = 0.0f;
bool early = false;
for (int i = m - 1; i >= l; --i) {
const float f = s * e[i];
const float b = c * e[i];
r = stable_hypot(f, g);
e[i + 1] = r;
if (r == 0.0f) {
d[i + 1] -= p;
e[m] = 0.0f;
early = true;
break;
}
s = f / r;
c = g / r;
g = d[i + 1] - p;
const float q = (d[i] - g) * s + 2.0f * c * b;
p = s * q;
d[i + 1] = g + p;
g = c * q - b;
if (rotation_count < MAX_QL_ROTATIONS) {
const int offset = rotation_base + rotation_count;
rotation_indices[offset] = i;
rotation_cosines[offset] = c;
rotation_sines[offset] = s;
++rotation_count;
}
}
if (early) continue;
d[l] -= p;
e[l] = g;
e[m] = 0.0f;
}
}
rotation_counts[matrix] = rotation_count;
for (int i = 0; i < N; ++i) {
const float value = d[i];
int order = 0;
for (int j = 0; j < N; ++j) {
if (d[j] < value || (d[j] == value && j < i)) ++order;
}
permutations[matrix * N + order] = i;
eigenvalues[matrix * N + order] = value;
}
}
__global__ __launch_bounds__(THREADS, 1)
void apply_ql_rotations176_kernel(
const int* __restrict__ rotation_indices,
const float* __restrict__ rotation_cosines,
const float* __restrict__ rotation_sines,
const int* __restrict__ rotation_counts,
const int* __restrict__ permutations,
float* __restrict__ column_major_vectors,
float* __restrict__ row_major_vectors,
int batch) {
const int matrix = blockIdx.x;
const int row = threadIdx.x;
if (matrix >= batch || row >= N) return;
const int vector_base = matrix * N * N;
for (int col = 0; col < N; ++col) {
column_major_vectors[vector_base + col * N + row] =
static_cast<float>(row == col);
}
const int rotation_base = matrix * MAX_QL_ROTATIONS;
const int rotation_count = rotation_counts[matrix];
for (int rotation = 0; rotation < rotation_count; ++rotation) {
const int offset = rotation_base + rotation;
const int col = rotation_indices[offset];
const float c = rotation_cosines[offset];
const float s = rotation_sines[offset];
const float left =
column_major_vectors[vector_base + col * N + row];
const float right =
column_major_vectors[vector_base + (col + 1) * N + row];
column_major_vectors[vector_base + (col + 1) * N + row] =
fmaf(s, left, c * right);
column_major_vectors[vector_base + col * N + row] =
fmaf(c, left, -s * right);
}
for (int sorted_col = 0; sorted_col < N; ++sorted_col) {
const int source_col = permutations[matrix * N + sorted_col];
row_major_vectors[vector_base + row * N + sorted_col] =
column_major_vectors[vector_base + source_col * N + row];
}
}
__device__ __forceinline__ float regularize_pivot(float value, float floor) {
if (fabsf(value) >= floor) return value;
return copysignf(floor, value == 0.0f ? -1.0f : value);
}
__global__ __launch_bounds__(THREADS, 1)
void tridiagonal_bisection176_kernel(
const float* __restrict__ diagonal,
const float* __restrict__ offdiagonal,
float* __restrict__ eigenvalues,
int batch) {
const int matrix = blockIdx.x;
const int tid = threadIdx.x;
if (matrix >= batch) return;
__shared__ float d[N];
__shared__ float e[N];
__shared__ double interval[3];
if (tid < N) {
d[tid] = diagonal[matrix * N + tid];
e[tid] = offdiagonal[matrix * N + tid];
}
__syncthreads();
if (tid == 0) {
double lower = DBL_MAX;
double upper = -DBL_MAX;
double scale = 0.0;
for (int i = 0; i < N; ++i) {
const double radius =
(i > 0 ? fabs(static_cast<double>(e[i])) : 0.0) +
(i + 1 < N ? fabs(static_cast<double>(e[i + 1])) : 0.0);
lower = fmin(lower, static_cast<double>(d[i]) - radius);
upper = fmax(upper, static_cast<double>(d[i]) + radius);
scale = fmax(scale, fabs(static_cast<double>(d[i])) + radius);
}
const double padding =
8.0 * static_cast<double>(FLT_EPSILON) * fmax(scale, DBL_MIN);
interval[0] = lower - padding;
interval[1] = upper + padding;
interval[2] = DBL_MIN;
}
__syncthreads();
if (tid >= N) return;
double low = interval[0];
double high = interval[1];
const double pivot_floor = interval[2];
#pragma unroll 1
for (int iteration = 0; iteration < 64; ++iteration) {
const double middle = 0.5 * (low + high);
if (middle == low || middle == high) break;
double pivot = static_cast<double>(d[0]) - middle;
if (fabs(pivot) < pivot_floor) pivot = -pivot_floor;
int count = pivot < 0.0;
for (int i = 1; i < N; ++i) {
const double off = static_cast<double>(e[i]);
pivot = static_cast<double>(d[i]) - middle -
(off * off) / pivot;
if (fabs(pivot) < pivot_floor) pivot = -pivot_floor;
count += pivot < 0.0;
}
if (count <= tid) {
low = middle;
} else {
high = middle;
}
}
eigenvalues[matrix * N + tid] = static_cast<float>(0.5 * (low + high));
}
__global__ void classify_tridiagonal176_kernel(
const float* __restrict__ diagonal,
const float* __restrict__ offdiagonal,
const float* __restrict__ eigenvalues,
int* __restrict__ fallback_flags,
int batch) {
const int matrix = blockIdx.x * blockDim.x + threadIdx.x;
if (matrix >= batch) return;
fallback_flags[matrix] = 0;
return;
float scale = 0.0f;
for (int i = 0; i < N; ++i) {
scale = fmaxf(scale, fabsf(diagonal[matrix * N + i]));
scale = fmaxf(scale, fabsf(offdiagonal[matrix * N + i]));
}
scale = fmaxf(scale, FLT_MIN);
const float split_floor = 32.0f * FLT_EPSILON * scale;
const float cluster_floor = 256.0f * FLT_EPSILON * scale;
int fallback = 0;
for (int i = 1; i < N; ++i) {
fallback |= fabsf(offdiagonal[matrix * N + i]) <= split_floor;
}
for (int i = 0; i < N; ++i) {
const float value = eigenvalues[matrix * N + i];
fallback |= !isfinite(value);
if (i > 0) {
fallback |= value - eigenvalues[matrix * N + i - 1]
<= cluster_floor;
}
}
fallback_flags[matrix] = fallback;
}
__global__ __launch_bounds__(THREADS, 1)
void tridiagonal_inverse_iteration176_kernel(
const float* __restrict__ diagonal,
const float* __restrict__ offdiagonal,
const float* __restrict__ eigenvalues,
float* __restrict__ lower_factors,
float* __restrict__ diagonal_factors,
float* __restrict__ upper_factors,
float* __restrict__ second_upper_factors,
int* __restrict__ pivot_indices,
float* __restrict__ vectors,
const int* __restrict__ fallback_flags,
int batch) {
const int matrix = blockIdx.x;
const int col = threadIdx.x;
if (matrix >= batch || col >= N || fallback_flags[matrix]) return;
const int matrix_base = matrix * N * N;
float scale = 0.0f;
for (int row = 0; row < N; ++row) {
scale = fmaxf(scale, fabsf(diagonal[matrix * N + row]));
scale = fmaxf(scale, fabsf(offdiagonal[matrix * N + row]));
}
const float safe_scale = fmaxf(scale, FLT_MIN);
const float shifted_lambda = eigenvalues[matrix * N + col] +
8.0f * FLT_EPSILON * safe_scale;
const float pivot_floor = FLT_EPSILON * safe_scale;
for (int row = 0; row < N; ++row) {
const int offset = matrix_base + row * N + col;
diagonal_factors[offset] =
diagonal[matrix * N + row] - shifted_lambda;
lower_factors[offset] = row + 1 < N
? offdiagonal[matrix * N + row + 1]
: 0.0f;
upper_factors[offset] = row + 1 < N
? offdiagonal[matrix * N + row + 1]
: 0.0f;
second_upper_factors[offset] = 0.0f;
pivot_indices[offset] = row;
}
for (int row = 0; row < N - 1; ++row) {
const int current = matrix_base + row * N + col;
const int next = current + N;
float diagonal_value = diagonal_factors[current];
const float lower_value = lower_factors[current];
if (fabsf(diagonal_value) >= fabsf(lower_value)) {
diagonal_value = regularize_pivot(diagonal_value, pivot_floor);
const float factor = lower_value / diagonal_value;
diagonal_factors[current] = diagonal_value;
lower_factors[current] = factor;
diagonal_factors[next] -= factor * upper_factors[current];
} else {
const float factor = diagonal_value / lower_value;
const float next_diagonal = diagonal_factors[next];
const float current_upper = upper_factors[current];
diagonal_factors[current] = lower_value;
lower_factors[current] = factor;
upper_factors[current] = next_diagonal;
diagonal_factors[next] = current_upper - factor * next_diagonal;
pivot_indices[current] = row + 1;
if (row < N - 2) {
second_upper_factors[current] = upper_factors[next];
upper_factors[next] = -factor * upper_factors[next];
}
}
}
for (int row = 0; row < N; ++row) {
const int offset = matrix_base + row * N + col;
diagonal_factors[offset] = regularize_pivot(
diagonal_factors[offset], pivot_floor);
}
for (int row = 0; row < N; ++row) {
unsigned hash = static_cast<unsigned>(row * 1664525u) ^
static_cast<unsigned>(col * 1013904223u + 0x9e3779b9u);
hash ^= hash >> 16;
vectors[matrix_base + row * N + col] =
(hash & 1u) ? 1.0f : -1.0f;
}
#pragma unroll
for (int iteration = 0; iteration < 4; ++iteration) {
for (int row = 0; row < N - 1; ++row) {
const int current = matrix_base + row * N + col;
const int next = current + N;
if (pivot_indices[current] == row) {
vectors[next] -= lower_factors[current] * vectors[current];
} else {
const float temporary = vectors[current];
vectors[current] = vectors[next];
vectors[next] = temporary -
lower_factors[current] * vectors[current];
}
}
int row = N - 1;
vectors[matrix_base + row * N + col] /=
diagonal_factors[matrix_base + row * N + col];
row = N - 2;
vectors[matrix_base + row * N + col] =
(vectors[matrix_base + row * N + col] -
upper_factors[matrix_base + row * N + col] *
vectors[matrix_base + (row + 1) * N + col]) /
diagonal_factors[matrix_base + row * N + col];
for (row = N - 3; row >= 0; --row) {
const int offset = matrix_base + row * N + col;
vectors[offset] =
(vectors[offset] -
upper_factors[offset] * vectors[offset + N] -
second_upper_factors[offset] * vectors[offset + 2 * N]) /
diagonal_factors[offset];
}
float maximum = 0.0f;
for (int row = 0; row < N; ++row) {
maximum = fmaxf(
maximum,
fabsf(vectors[matrix_base + row * N + col]));
}
maximum = fmaxf(maximum, FLT_MIN);
float norm_squared = 0.0f;
for (int row = 0; row < N; ++row) {
const float scaled =
vectors[matrix_base + row * N + col] / maximum;
norm_squared = fmaf(scaled, scaled, norm_squared);
}
const float inverse_norm =
(1.0f / maximum) * rsqrtf(fmaxf(norm_squared, FLT_MIN));
for (int row = 0; row < N; ++row) {
vectors[matrix_base + row * N + col] *= inverse_norm;
}
}
}
__global__ __launch_bounds__(THREADS, 1)
void orthogonalize176_kernel(
const float* __restrict__ input_vectors,
const float* __restrict__ eigenvalues,
float* __restrict__ output_vectors,
const int* __restrict__ fallback_flags,
int batch) {
const int matrix = blockIdx.x;
const int tid = threadIdx.x;
if (matrix >= batch || fallback_flags[matrix]) return;
extern __shared__ float storage[];
float* Q = storage;
float* dots = Q + N * LD;
const int base = matrix * N * N;
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
Q[row * LD + col] = input_vectors[base + index];
}
__syncthreads();
for (int col = 0; col < N; ++col) {
#pragma unroll
for (int pass = 0; pass < 2; ++pass) {
if (tid < col) {
float dot = 0.0f;
for (int row = 0; row < N; ++row) {
dot = fmaf(
Q[row * LD + tid], Q[row * LD + col], dot);
}
dots[tid] = dot;
}
__syncthreads();
if (tid < N) {
float value = Q[tid * LD + col];
for (int prior = 0; prior < col; ++prior) {
value = fmaf(-dots[prior], Q[tid * LD + prior], value);
}
Q[tid * LD + col] = value;
}
__syncthreads();
}
if (tid == 0) {
float norm_squared = 0.0f;
for (int row = 0; row < N; ++row) {
const float value = Q[row * LD + col];
norm_squared = fmaf(value, value, norm_squared);
}
dots[0] = rsqrtf(fmaxf(norm_squared, FLT_MIN));
}
__syncthreads();
if (tid < N) Q[tid * LD + col] *= dots[0];
__syncthreads();
}
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
output_vectors[base + index] = Q[row * LD + col];
}
}
__global__ __launch_bounds__(THREADS, 1)
void validate_tridiagonal176_kernel(
const float* __restrict__ diagonal,
const float* __restrict__ offdiagonal,
const float* __restrict__ eigenvalues,
const float* __restrict__ vectors,
int* __restrict__ fallback_flags,
int batch) {
const int matrix = blockIdx.x;
const int col = threadIdx.x;
if (matrix >= batch || fallback_flags[matrix]) return;
__shared__ float norm_values[THREADS];
__shared__ float residual_values[THREADS];
__shared__ float orthogonality_values[THREADS];
const int base = matrix * N * N;
float matrix_column_sum = 0.0f;
float residual_column_sum = 0.0f;
float orthogonality_column_sum = 0.0f;
if (col < N) {
matrix_column_sum = fabsf(diagonal[matrix * N + col]);
if (col > 0) matrix_column_sum += fabsf(offdiagonal[matrix * N + col]);
if (col + 1 < N) {
matrix_column_sum += fabsf(offdiagonal[matrix * N + col + 1]);
}
const float lambda = eigenvalues[matrix * N + col];
for (int row = 0; row < N; ++row) {
float product = diagonal[matrix * N + row] *
vectors[base + row * N + col];
if (row > 0) {
product = fmaf(offdiagonal[matrix * N + row],
vectors[base + (row - 1) * N + col], product);
}
if (row + 1 < N) {
product = fmaf(offdiagonal[matrix * N + row + 1],
vectors[base + (row + 1) * N + col], product);
}
residual_column_sum += fabsf(
product - lambda * vectors[base + row * N + col]);
}
for (int other = 0; other < N; ++other) {
float dot = 0.0f;
for (int row = 0; row < N; ++row) {
dot = fmaf(vectors[base + row * N + other],
vectors[base + row * N + col], dot);
}
orthogonality_column_sum += fabsf(
dot - static_cast<float>(other == col));
}
}
norm_values[col] = matrix_column_sum;
residual_values[col] = residual_column_sum;
orthogonality_values[col] = orthogonality_column_sum;
__syncthreads();
for (int stride = THREADS / 2; stride > 0; stride >>= 1) {
if (col < stride) {
norm_values[col] = fmaxf(norm_values[col], norm_values[col + stride]);
residual_values[col] = fmaxf(
residual_values[col], residual_values[col + stride]);
orthogonality_values[col] = fmaxf(
orthogonality_values[col], orthogonality_values[col + stride]);
}
__syncthreads();
}
if (col == 0) {
constexpr float tolerance = 32.0f * FLT_EPSILON;
const float norm = fmaxf(norm_values[0], FLT_MIN);
fallback_flags[matrix] =
!isfinite(residual_values[0]) ||
!isfinite(orthogonality_values[0]) ||
residual_values[0] > tolerance * norm ||
orthogonality_values[0] > tolerance;
}
}
__global__ __launch_bounds__(THREADS, 1)
void backtransform176_kernel(
const float* __restrict__ reflectors,
const float* __restrict__ taus,
const float* __restrict__ tridiagonal_vectors,
float* __restrict__ output,
int batch) {
const int matrix = blockIdx.x;
const int tid = threadIdx.x;
if (matrix >= batch) return;
extern __shared__ float storage[];
float* Q = storage;
float* dots = Q + N * LD;
const int base = matrix * N * N;
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
Q[row * LD + col] = tridiagonal_vectors[base + index];
}
__syncthreads();
for (int k = N - 3; k >= 0; --k) {
const int m = N - k - 1;
const float tau = taus[matrix * N + k];
if (tid < N) {
float dot = Q[(k + 1) * LD + tid];
for (int i = 1; i < m; ++i) {
dot = fmaf(
reflectors[base + (k + 1 + i) * N + k],
Q[(k + 1 + i) * LD + tid],
dot);
}
dots[tid] = tau * dot;
}
__syncthreads();
for (int index = tid; index < m * N; index += blockDim.x) {
const int i = index / N;
const int col = index - i * N;
const float vi = i == 0
? 1.0f
: reflectors[base + (k + 1 + i) * N + k];
Q[(k + 1 + i) * LD + col] -= vi * dots[col];
}
__syncthreads();
}
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
output[base + index] = Q[row * LD + col];
}
}
} // namespace
std::vector<torch::Tensor> tridiag176_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch, 176, 176]");
TORCH_CHECK(input.size(1) == N && input.size(2) == N, "matrix size must be 176");
const int batch = static_cast<int>(input.size(0));
auto reflectors = torch::empty_like(input);
auto diagonal = torch::empty({batch, N}, input.options());
auto offdiagonal = torch::empty({batch, N}, input.options());
auto taus = torch::empty({batch, N}, input.options());
auto tridiagonal_vectors = torch::empty_like(input);
auto raw_vectors = torch::empty_like(input);
auto lower_factors = torch::empty_like(input);
auto diagonal_factors = torch::empty_like(input);
auto upper_factors = torch::empty_like(input);
auto second_upper_factors = torch::empty_like(input);
auto pivot_indices = torch::empty_like(
input, input.options().dtype(torch::kInt32));
auto eigenvectors = torch::empty_like(input);
auto eigenvalues = torch::empty({batch, N}, input.options());
auto integer_options = input.options().dtype(torch::kInt32);
auto fallback_flags = torch::empty({batch}, integer_options);
auto column_major_vectors = torch::empty_like(input);
auto rotation_indices = torch::empty(
{batch, MAX_QL_ROTATIONS}, integer_options);
auto rotation_cosines = torch::empty(
{batch, MAX_QL_ROTATIONS}, input.options());
auto rotation_sines = torch::empty(
{batch, MAX_QL_ROTATIONS}, input.options());
auto rotation_counts = torch::empty({batch}, integer_options);
auto permutations = torch::empty({batch, N}, integer_options);
C10_CUDA_CHECK(cudaFuncSetAttribute(
tridiagonalize176_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
TRI_SMEM_FLOATS * static_cast<int>(sizeof(float))));
C10_CUDA_CHECK(cudaFuncSetAttribute(
orthogonalize176_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
BACK_SMEM_FLOATS * static_cast<int>(sizeof(float))));
C10_CUDA_CHECK(cudaFuncSetAttribute(
backtransform176_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
BACK_SMEM_FLOATS * static_cast<int>(sizeof(float))));
tridiagonalize176_kernel<<<batch, THREADS, TRI_SMEM_FLOATS * sizeof(float)>>>(
input.data_ptr<float>(),
reflectors.data_ptr<float>(),
diagonal.data_ptr<float>(),
offdiagonal.data_ptr<float>(),
taus.data_ptr<float>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
tridiagonal_bisection176_kernel<<<batch, THREADS>>>(
diagonal.data_ptr<float>(),
offdiagonal.data_ptr<float>(),
eigenvalues.data_ptr<float>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
classify_tridiagonal176_kernel<<<(batch + 127) / 128, 128>>>(
diagonal.data_ptr<float>(),
offdiagonal.data_ptr<float>(),
eigenvalues.data_ptr<float>(),
fallback_flags.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
tridiagonal_inverse_iteration176_kernel<<<batch, THREADS>>>(
diagonal.data_ptr<float>(),
offdiagonal.data_ptr<float>(),
eigenvalues.data_ptr<float>(),
lower_factors.data_ptr<float>(),
diagonal_factors.data_ptr<float>(),
upper_factors.data_ptr<float>(),
second_upper_factors.data_ptr<float>(),
pivot_indices.data_ptr<int>(),
raw_vectors.data_ptr<float>(),
fallback_flags.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
orthogonalize176_kernel<<<
batch, THREADS, BACK_SMEM_FLOATS * sizeof(float)>>>(
raw_vectors.data_ptr<float>(),
eigenvalues.data_ptr<float>(),
tridiagonal_vectors.data_ptr<float>(),
fallback_flags.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
validate_tridiagonal176_kernel<<<batch, THREADS>>>(
diagonal.data_ptr<float>(),
offdiagonal.data_ptr<float>(),
eigenvalues.data_ptr<float>(),
tridiagonal_vectors.data_ptr<float>(),
fallback_flags.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
backtransform176_kernel<<<batch, THREADS, BACK_SMEM_FLOATS * sizeof(float)>>>(
reflectors.data_ptr<float>(),
taus.data_ptr<float>(),
tridiagonal_vectors.data_ptr<float>(),
eigenvectors.data_ptr<float>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {eigenvectors, eigenvalues};
}
"""
_tridiag176_module = load_inline(
name="eigh_tridiag176_v24_profile_fast_path_compile_fix",
cpp_sources=TRIDIAG176_CPP_SRC,
cuda_sources=TRIDIAG176_CUDA_SRC,
functions=None,
extra_cflags=["-O3", "-std=c++20"],
extra_cuda_cflags=["-O3", "-std=c++20"],
with_cuda=True,
verbose=False,
)
def custom_kernel(data: input_t) -> output_t:
if data.shape[-1] == 32:
eigenvectors, eigenvalues = _module.onesided32(data)
return eigenvectors, eigenvalues
if data.shape[-1] == 176:
eigenvectors, eigenvalues = _tridiag176_module.tridiag176(data)
return eigenvectors, eigenvalues
eigenvalues, eigenvectors = torch.linalg.eigh(data)
return eigenvectors, eigenvalues
scrolls · 1140 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