Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
45.7ms
#100 of 286
2026-07-12

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