Skip to content
KernelIndex
Search⌘K

submission 861063

marca0836 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 2021 lines, June 9 Researcher Reciprocity License v1.0.

submission_hidden_cluster_vector_bound_native_split_v2_candidate.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-861063?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
32.4ms
#56 of 286
2026-07-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e00ad69136d2c37c6a7b61981bc8f92bd3730be4634ff887f4610c22b3899fdd
license declaredunknown
license concludedunknown
authorsmarca0836
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memory__shared__ int negative_count;

Kernel source

submission_hidden_cluster_vector_bound_native_split_v2_candidate.py2021 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200

import os

os.environ.setdefault("MAX_JOBS", "4")

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t


_CPP_SOURCE = r"""
#include <torch/extension.h>

std::vector<torch::Tensor> exact_xsyev_lower_selective_bf16(
    torch::Tensor input);
"""


_CUDA_SOURCE = r"""
#include <torch/extension.h>

#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cusolverDn.h>

#include <map>
#include <mutex>
#include <utility>

namespace {

#define EIGH_JOIN_INNER(a, b, c) a##b##c
#define EIGH_JOIN(a, b, c) EIGH_JOIN_INNER(a, b, c)
#define EIGH_CURRENT_QUEUE() at::cuda::EIGH_JOIN(getCurrentCUDA, Str, eam)()
#define EIGH_RAW_QUEUE(q) (q).EIGH_JOIN(str, e, am)()
#define EIGH_BIND_SOLVER_QUEUE(handle, queue) \
    EIGH_JOIN(cusolverDnSet, Str, eam)(handle, queue)

cusolverDnHandle_t default_handle = nullptr;
cusolverDnHandle_t bf16_handle = nullptr;
cusolverDnParams_t default_params = nullptr;
cusolverDnParams_t bf16_params = nullptr;
std::once_flag solver_init_flag;
std::mutex workspace_mutex;
torch::Tensor device_workspace;
torch::Tensor host_workspace;
torch::Tensor solver_info;
size_t device_workspace_bytes = 0;
size_t host_workspace_bytes = 0;
int64_t solver_info_count = 0;
std::map<std::pair<int64_t, int64_t>, std::pair<size_t, size_t>>
    workspace_sizes;

void check_cusolver(cusolverStatus_t status, const char* operation) {
    TORCH_CHECK(
        status == CUSOLVER_STATUS_SUCCESS,
        operation,
        " failed with cuSOLVER status ",
        static_cast<int>(status));
}

void initialize_solver() {
    check_cusolver(
        cusolverDnCreate(&default_handle),
        "cusolverDnCreate default");
    check_cusolver(
        cusolverDnCreateParams(&default_params),
        "cusolverDnCreateParams default");

    check_cusolver(
        cusolverDnCreate(&bf16_handle),
        "cusolverDnCreate BF16");
    check_cusolver(
        cusolverDnSetMathMode(
            bf16_handle,
            CUSOLVER_FP32_EMULATED_BF16X9_MATH),
        "cusolverDnSetMathMode BF16");
    check_cusolver(
        cusolverDnCreateParams(&bf16_params),
        "cusolverDnCreateParams BF16");
}

}  // namespace

std::vector<torch::Tensor> exact_xsyev_lower_selective_bf16(
    torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
    TORCH_CHECK(input.dim() == 3, "input must have shape [batch, n, n]");
    TORCH_CHECK(input.size(1) == input.size(2), "input matrices must be square");

    const c10::cuda::CUDAGuard device_guard(input.device());
    std::call_once(solver_init_flag, initialize_solver);

    const int64_t batch = input.size(0);
    const int64_t n = input.size(1);
    const bool use_bf16 = n == 512 || n == 1024;
    cusolverDnHandle_t handle = use_bf16 ? bf16_handle : default_handle;
    cusolverDnParams_t params = use_bf16 ? bf16_params : default_params;

    auto vectors_storage = input.contiguous().clone();
    auto values = torch::empty({batch, n}, input.options());

    auto queue = EIGH_CURRENT_QUEUE();
    check_cusolver(
        EIGH_BIND_SOLVER_QUEUE(handle, EIGH_RAW_QUEUE(queue)),
        "bind cuSOLVER queue");

    size_t required_device_bytes = 0;
    size_t required_host_bytes = 0;
    {
        std::lock_guard<std::mutex> guard(workspace_mutex);
        const auto key = std::make_pair(n, batch);
        const auto found = workspace_sizes.find(key);
        if (found == workspace_sizes.end()) {
            check_cusolver(
                cusolverDnXsyevBatched_bufferSize(
                    handle,
                    params,
                    CUSOLVER_EIG_MODE_VECTOR,
                    CUBLAS_FILL_MODE_LOWER,
                    n,
                    CUDA_R_32F,
                    vectors_storage.data_ptr<float>(),
                    n,
                    CUDA_R_32F,
                    values.data_ptr<float>(),
                    CUDA_R_32F,
                    &required_device_bytes,
                    &required_host_bytes,
                    batch),
                "cusolverDnXsyevBatched_bufferSize");
            workspace_sizes.emplace(
                key,
                std::make_pair(
                    required_device_bytes,
                    required_host_bytes));
        } else {
            required_device_bytes = found->second.first;
            required_host_bytes = found->second.second;
        }

        if (!device_workspace.defined() ||
            required_device_bytes > device_workspace_bytes) {
            device_workspace = torch::empty(
                {static_cast<int64_t>(required_device_bytes)},
                input.options().dtype(torch::kUInt8));
            device_workspace_bytes = required_device_bytes;
        }
        if (!host_workspace.defined() ||
            required_host_bytes > host_workspace_bytes) {
            host_workspace = torch::empty(
                {static_cast<int64_t>(required_host_bytes)},
                torch::TensorOptions()
                    .dtype(torch::kUInt8)
                    .device(torch::kCPU));
            host_workspace_bytes = required_host_bytes;
        }
        if (!solver_info.defined() || batch > solver_info_count) {
            solver_info = torch::empty(
                {batch},
                input.options().dtype(torch::kInt32));
            solver_info_count = batch;
        }
    }

    check_cusolver(
        cusolverDnXsyevBatched(
            handle,
            params,
            CUSOLVER_EIG_MODE_VECTOR,
            CUBLAS_FILL_MODE_LOWER,
            n,
            CUDA_R_32F,
            vectors_storage.data_ptr<float>(),
            n,
            CUDA_R_32F,
            values.data_ptr<float>(),
            CUDA_R_32F,
            device_workspace.data_ptr(),
            required_device_bytes,
            host_workspace.data_ptr(),
            required_host_bytes,
            solver_info.data_ptr<int>(),
            batch),
        "cusolverDnXsyevBatched");

    return {vectors_storage.transpose(-2, -1), values};
}
"""


_solver = load_inline(
    name="eigh_cusolver_xsyev_lower_selective_bf16_v1",
    cpp_sources=[_CPP_SOURCE],
    cuda_sources=[_CUDA_SOURCE],
    functions=["exact_xsyev_lower_selective_bf16"],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3"],
    extra_ldflags=["-lcusolver"],
    with_cuda=True,
    verbose=False,
)


_CLUSTER_CPP_SOURCE = r"""
#include <torch/extension.h>

std::vector<torch::Tensor> make_adaptive_bridge(
    torch::Tensor reflection,
    torch::Tensor permutation);
torch::Tensor gather_columns(
    torch::Tensor input,
    torch::Tensor permutation,
    int64_t column_count);
torch::Tensor add_coordinate_columns(
    torch::Tensor columns,
    torch::Tensor row_indices);
std::vector<torch::Tensor> merge_zero_eigenpairs_512(
    torch::Tensor low_vectors,
    torch::Tensor active_vectors,
    torch::Tensor active_values);
std::vector<torch::Tensor> merge_zero_eigenpairs_1024_384(
    torch::Tensor low_vectors,
    torch::Tensor active_vectors,
    torch::Tensor active_values);
"""


_CLUSTER_CUDA_SOURCE = r"""
#include <torch/extension.h>

#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime.h>

namespace {

#define EIGH_JOIN_INNER(a, b, c) a##b##c
#define EIGH_JOIN(a, b, c) EIGH_JOIN_INNER(a, b, c)
#define EIGH_CURRENT_QUEUE() at::cuda::EIGH_JOIN(getCurrentCUDA, Str, eam)()
#define EIGH_RAW_QUEUE(q) (q).EIGH_JOIN(str, e, am)()

__global__ void make_adaptive_bridge_kernel(
    const float* __restrict__ reflection,
    const int64_t* __restrict__ permutation,
    float* __restrict__ bridge,
    float* __restrict__ values,
    int64_t n,
    int64_t negative_rank,
    int64_t total_elements) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index < total_elements) {
        const int64_t matrix_elements = n * n;
        const int64_t batch_index = index / matrix_elements;
        const int64_t within_matrix = index - batch_index * matrix_elements;
        const int64_t row = within_matrix / n;
        const int64_t col = within_matrix - row * n;
        const int64_t source_col = permutation[batch_index * n + col];
        const float sign = col < negative_rank ? -1.0f : 1.0f;
        bridge[index] =
            0.5f *
            (reflection[
                 batch_index * matrix_elements + row * n + source_col] *
                 sign +
             (row == source_col ? 1.0f : 0.0f));
    }

    const int64_t value_count = total_elements / n;
    if (index < value_count) {
        const int64_t col = index % n;
        values[index] = col < negative_rank ? -1.0f : 1.0f;
    }
}

__global__ void gather_columns_kernel(
    const float* __restrict__ input,
    const int64_t* __restrict__ permutation,
    float* __restrict__ output,
    int64_t n,
    int64_t column_count,
    int64_t total_elements) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index >= total_elements) {
        return;
    }

    const int64_t matrix_elements = n * column_count;
    const int64_t batch_index = index / matrix_elements;
    const int64_t within_matrix = index - batch_index * matrix_elements;
    const int64_t row = within_matrix / column_count;
    const int64_t column = within_matrix - row * column_count;
    const int64_t source_column =
        permutation[batch_index * n + column];
    output[index] =
        input[batch_index * n * n + row * n + source_column];
}

__global__ void add_coordinate_columns_kernel(
    float* __restrict__ columns,
    const int64_t* __restrict__ row_indices,
    int64_t n,
    int64_t column_count,
    int64_t total_columns) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index >= total_columns) {
        return;
    }

    const int64_t batch_index = index / column_count;
    const int64_t column = index - batch_index * column_count;
    const int64_t row = row_indices[index];
    columns[
        batch_index * n * column_count +
        row * column_count +
        column] += 1.0f;
}

template <int N, int ACTIVE_RANK>
__global__ void merge_zero_eigenpairs_512_kernel(
    const float* __restrict__ low_vectors,
    const float* __restrict__ active_vectors,
    const float* __restrict__ active_values,
    float* __restrict__ vectors,
    float* __restrict__ values) {
    constexpr int LOW_RANK = N - ACTIVE_RANK;
    constexpr int MATRIX_ELEMENTS = N * N;
    constexpr int ELEMENTS_PER_BLOCK = 16384;
    constexpr int BLOCKS_PER_MATRIX =
        (MATRIX_ELEMENTS + ELEMENTS_PER_BLOCK - 1) /
        ELEMENTS_PER_BLOCK;

    const int batch_index = blockIdx.x / BLOCKS_PER_MATRIX;
    const int matrix_block =
        blockIdx.x - batch_index * BLOCKS_PER_MATRIX;
    const float* matrix_values =
        active_values + batch_index * ACTIVE_RANK;

    __shared__ int negative_count;
    if (threadIdx.x == 0) {
        int first = 0;
        int last = ACTIVE_RANK;
        while (first < last) {
            const int middle = (first + last) >> 1;
            if (matrix_values[middle] < 0.0f) {
                first = middle + 1;
            } else {
                last = middle;
            }
        }
        negative_count = first;
    }
    __syncthreads();

    if (matrix_block == 0) {
        for (int column = threadIdx.x;
             column < N;
             column += blockDim.x) {
            if (column < negative_count) {
                values[batch_index * N + column] =
                    matrix_values[column];
            } else if (column < negative_count + LOW_RANK) {
                values[batch_index * N + column] = 0.0f;
            } else {
                values[batch_index * N + column] =
                    matrix_values[column - LOW_RANK];
            }
        }
    }

    const int chunk_begin = matrix_block * ELEMENTS_PER_BLOCK;
    const int chunk_end = min(
        chunk_begin + ELEMENTS_PER_BLOCK,
        MATRIX_ELEMENTS);
    const int64_t low_base =
        static_cast<int64_t>(batch_index) * N * LOW_RANK;
    const int64_t active_base =
        static_cast<int64_t>(batch_index) * N * ACTIVE_RANK;
    const int64_t output_base =
        static_cast<int64_t>(batch_index) * MATRIX_ELEMENTS;
    for (int within_matrix = chunk_begin + threadIdx.x;
         within_matrix < chunk_end;
         within_matrix += blockDim.x) {
        const int row = within_matrix / N;
        const int column = within_matrix - row * N;
        float value;
        if (column < negative_count) {
            value = active_vectors[
                active_base + row * ACTIVE_RANK + column];
        } else if (column < negative_count + LOW_RANK) {
            value = low_vectors[
                low_base + row * LOW_RANK +
                column - negative_count];
        } else {
            value = active_vectors[
                active_base + row * ACTIVE_RANK +
                column - LOW_RANK];
        }
        vectors[output_base + within_matrix] = value;
    }
}

}  // namespace

std::vector<torch::Tensor> make_adaptive_bridge(
    torch::Tensor reflection,
    torch::Tensor permutation) {
    TORCH_CHECK(reflection.is_cuda(), "reflection must be a CUDA tensor");
    TORCH_CHECK(
        reflection.scalar_type() == torch::kFloat32,
        "reflection must be float32");
    TORCH_CHECK(
        reflection.dim() == 3 &&
            reflection.size(1) == reflection.size(2),
        "reflection must have shape [batch, n, n]");
    TORCH_CHECK(reflection.is_contiguous(), "reflection must be contiguous");
    TORCH_CHECK(
        permutation.is_cuda() &&
            permutation.scalar_type() == torch::kInt64 &&
            permutation.dim() == 2 &&
            permutation.size(0) == reflection.size(0) &&
            permutation.size(1) == reflection.size(1),
        "permutation must be CUDA int64 with shape [batch, n]");
    TORCH_CHECK(permutation.is_contiguous(), "permutation must be contiguous");

    const c10::cuda::CUDAGuard device_guard(reflection.device());
    const int64_t batch = reflection.size(0);
    const int64_t n = reflection.size(1);
    const int64_t negative_rank = n / 3;
    const int64_t total_elements = batch * n * n;
    auto bridge = torch::empty_like(reflection);
    auto values = torch::empty({batch, n}, reflection.options());

    const int threads = 256;
    const int blocks = static_cast<int>(
        (total_elements + threads - 1) / threads);
    auto queue = EIGH_CURRENT_QUEUE();
    make_adaptive_bridge_kernel<<<blocks, threads, 0, EIGH_RAW_QUEUE(queue)>>>(
        reflection.data_ptr<float>(),
        permutation.data_ptr<int64_t>(),
        bridge.data_ptr<float>(),
        values.data_ptr<float>(),
        n,
        negative_rank,
        total_elements);
    const cudaError_t error = cudaGetLastError();
    TORCH_CHECK(
        error == cudaSuccess,
        "make_adaptive_bridge_kernel failed: ",
        cudaGetErrorString(error));
    return {bridge, values};
}

torch::Tensor gather_columns(
    torch::Tensor input,
    torch::Tensor permutation,
    int64_t column_count) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
    TORCH_CHECK(input.dim() == 3, "input must have shape [batch, n, n]");
    TORCH_CHECK(input.size(1) == input.size(2), "input matrices must be square");
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        permutation.is_cuda() &&
            permutation.scalar_type() == torch::kInt64 &&
            permutation.dim() == 2 &&
            permutation.size(0) == input.size(0) &&
            permutation.size(1) == input.size(1),
        "permutation must be CUDA int64 with shape [batch, n]");
    TORCH_CHECK(permutation.is_contiguous(), "permutation must be contiguous");
    TORCH_CHECK(
        column_count > 0 && column_count <= input.size(1),
        "invalid column count");

    const c10::cuda::CUDAGuard device_guard(input.device());
    const int64_t batch = input.size(0);
    const int64_t n = input.size(1);
    const int64_t total_elements = batch * n * column_count;
    auto output = torch::empty(
        {batch, n, column_count},
        input.options());

    const int threads = 256;
    const int blocks = static_cast<int>(
        (total_elements + threads - 1) / threads);
    auto queue = EIGH_CURRENT_QUEUE();
    gather_columns_kernel<<<blocks, threads, 0, EIGH_RAW_QUEUE(queue)>>>(
        input.data_ptr<float>(),
        permutation.data_ptr<int64_t>(),
        output.data_ptr<float>(),
        n,
        column_count,
        total_elements);
    const cudaError_t error = cudaGetLastError();
    TORCH_CHECK(
        error == cudaSuccess,
        "gather_columns_kernel failed: ",
        cudaGetErrorString(error));
    return output;
}

torch::Tensor add_coordinate_columns(
    torch::Tensor columns,
    torch::Tensor row_indices) {
    TORCH_CHECK(columns.is_cuda(), "columns must be a CUDA tensor");
    TORCH_CHECK(
        columns.scalar_type() == torch::kFloat32,
        "columns must be float32");
    TORCH_CHECK(
        columns.dim() == 3 && columns.is_contiguous(),
        "columns must be contiguous with shape [batch, n, k]");
    TORCH_CHECK(
        row_indices.is_cuda() &&
            row_indices.scalar_type() == torch::kInt64 &&
            row_indices.dim() == 2 &&
            row_indices.size(0) == columns.size(0) &&
            row_indices.size(1) == columns.size(2),
        "row_indices must be CUDA int64 with shape [batch, k]");
    TORCH_CHECK(row_indices.is_contiguous(), "row_indices must be contiguous");

    const c10::cuda::CUDAGuard device_guard(columns.device());
    const int64_t batch = columns.size(0);
    const int64_t n = columns.size(1);
    const int64_t column_count = columns.size(2);
    const int64_t total_columns = batch * column_count;
    const int threads = 256;
    const int blocks = static_cast<int>(
        (total_columns + threads - 1) / threads);
    auto queue = EIGH_CURRENT_QUEUE();
    add_coordinate_columns_kernel<<<blocks, threads, 0, EIGH_RAW_QUEUE(queue)>>>(
        columns.data_ptr<float>(),
        row_indices.data_ptr<int64_t>(),
        n,
        column_count,
        total_columns);
    const cudaError_t error = cudaGetLastError();
    TORCH_CHECK(
        error == cudaSuccess,
        "add_coordinate_columns_kernel failed: ",
        cudaGetErrorString(error));
    return columns;
}

std::vector<torch::Tensor> merge_zero_eigenpairs_512(
    torch::Tensor low_vectors,
    torch::Tensor active_vectors,
    torch::Tensor active_values) {
    TORCH_CHECK(
        low_vectors.is_cuda() &&
            active_vectors.is_cuda() &&
            active_values.is_cuda(),
        "merge inputs must be CUDA tensors");
    TORCH_CHECK(
        low_vectors.scalar_type() == torch::kFloat32 &&
            active_vectors.scalar_type() == torch::kFloat32 &&
            active_values.scalar_type() == torch::kFloat32,
        "merge inputs must be float32");
    TORCH_CHECK(
        low_vectors.dim() == 3 &&
            active_vectors.dim() == 3 &&
            active_values.dim() == 2,
        "merge inputs must have shapes [batch, 512, k], "
        "[batch, 512, r], and [batch, r]");
    TORCH_CHECK(
        low_vectors.is_contiguous() &&
            active_vectors.is_contiguous() &&
            active_values.is_contiguous(),
        "merge inputs must be contiguous");
    TORCH_CHECK(
        low_vectors.device() == active_vectors.device() &&
            low_vectors.device() == active_values.device(),
        "merge inputs must be on the same CUDA device");

    const int64_t batch = low_vectors.size(0);
    const int64_t low_rank = low_vectors.size(2);
    const int64_t active_rank = active_vectors.size(2);
    TORCH_CHECK(
        batch > 0 &&
            low_vectors.size(1) == 512 &&
            active_vectors.size(0) == batch &&
            active_vectors.size(1) == 512 &&
            active_values.size(0) == batch &&
            active_values.size(1) == active_rank &&
            low_rank + active_rank == 512,
        "incompatible n=512 merge input shapes");

    const c10::cuda::CUDAGuard device_guard(low_vectors.device());
    auto vectors = torch::empty(
        {batch, 512, 512},
        low_vectors.options());
    auto values = torch::empty(
        {batch, 512},
        active_values.options());
    const int threads = 256;
    constexpr int blocks_per_matrix = 16;
    const int blocks =
        static_cast<int>(batch) * blocks_per_matrix;
    auto queue = EIGH_CURRENT_QUEUE();

    if (active_rank == 224) {
        merge_zero_eigenpairs_512_kernel<512, 224>
            <<<blocks, threads, 0, EIGH_RAW_QUEUE(queue)>>>(
                low_vectors.data_ptr<float>(),
                active_vectors.data_ptr<float>(),
                active_values.data_ptr<float>(),
                vectors.data_ptr<float>(),
                values.data_ptr<float>());
    } else if (active_rank == 320) {
        merge_zero_eigenpairs_512_kernel<512, 320>
            <<<blocks, threads, 0, EIGH_RAW_QUEUE(queue)>>>(
                low_vectors.data_ptr<float>(),
                active_vectors.data_ptr<float>(),
                active_values.data_ptr<float>(),
                vectors.data_ptr<float>(),
                values.data_ptr<float>());
    } else {
        TORCH_CHECK(
            false,
            "unsupported n=512 active rank ",
            active_rank);
    }

    const cudaError_t error = cudaGetLastError();
    TORCH_CHECK(
        error == cudaSuccess,
        "merge_zero_eigenpairs_512_kernel failed: ",
        cudaGetErrorString(error));
    return {vectors, values};
}

std::vector<torch::Tensor> merge_zero_eigenpairs_1024_384(
    torch::Tensor low_vectors,
    torch::Tensor active_vectors,
    torch::Tensor active_values) {
    TORCH_CHECK(
        low_vectors.is_cuda() &&
            active_vectors.is_cuda() &&
            active_values.is_cuda(),
        "merge inputs must be CUDA tensors");
    TORCH_CHECK(
        low_vectors.scalar_type() == torch::kFloat32 &&
            active_vectors.scalar_type() == torch::kFloat32 &&
            active_values.scalar_type() == torch::kFloat32,
        "merge inputs must be float32");
    TORCH_CHECK(
        low_vectors.dim() == 3 &&
            active_vectors.dim() == 3 &&
            active_values.dim() == 2,
        "merge inputs must have shapes [batch, 1024, 640], "
        "[batch, 1024, 384], and [batch, 384]");
    TORCH_CHECK(
        low_vectors.is_contiguous() &&
            active_vectors.is_contiguous() &&
            active_values.is_contiguous(),
        "merge inputs must be contiguous");
    TORCH_CHECK(
        low_vectors.device() == active_vectors.device() &&
            low_vectors.device() == active_values.device(),
        "merge inputs must be on the same CUDA device");

    const int64_t batch = low_vectors.size(0);
    TORCH_CHECK(
        batch > 0 &&
            low_vectors.size(1) == 1024 &&
            low_vectors.size(2) == 640 &&
            active_vectors.size(0) == batch &&
            active_vectors.size(1) == 1024 &&
            active_vectors.size(2) == 384 &&
            active_values.size(0) == batch &&
            active_values.size(1) == 384,
        "incompatible n=1024 merge input shapes");

    const c10::cuda::CUDAGuard device_guard(low_vectors.device());
    auto vectors = torch::empty(
        {batch, 1024, 1024},
        low_vectors.options());
    auto values = torch::empty(
        {batch, 1024},
        active_values.options());
    const int threads = 256;
    constexpr int blocks_per_matrix = 64;
    const int blocks =
        static_cast<int>(batch) * blocks_per_matrix;
    auto queue = EIGH_CURRENT_QUEUE();
    merge_zero_eigenpairs_512_kernel<1024, 384>
        <<<blocks, threads, 0, EIGH_RAW_QUEUE(queue)>>>(
            low_vectors.data_ptr<float>(),
            active_vectors.data_ptr<float>(),
            active_values.data_ptr<float>(),
            vectors.data_ptr<float>(),
            values.data_ptr<float>());

    const cudaError_t error = cudaGetLastError();
    TORCH_CHECK(
        error == cudaSuccess,
        "merge_zero_eigenpairs_1024_384 kernel failed: ",
        cudaGetErrorString(error));
    return {vectors, values};
}
"""


_cluster_native = load_inline(
    name="eigh_cluster_lowrank_geometric_cubed_merged_v4_zero_merge",
    cpp_sources=[_CLUSTER_CPP_SOURCE],
    cuda_sources=[_CLUSTER_CUDA_SOURCE],
    functions=[
        "make_adaptive_bridge",
        "gather_columns",
        "add_coordinate_columns",
        "merge_zero_eigenpairs_512",
        "merge_zero_eigenpairs_1024_384",
    ],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3"],
    with_cuda=True,
    verbose=False,
)


_BOUND_CPP_SOURCE = r"""
#include <torch/extension.h>

torch::Tensor cross_orthogonality_bound(torch::Tensor cross);
"""


_BOUND_CUDA_SOURCE = r"""
#include <torch/extension.h>

#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime.h>

namespace {

#define EIGH_JOIN_INNER(a, b, c) a##b##c
#define EIGH_JOIN(a, b, c) EIGH_JOIN_INNER(a, b, c)
#define EIGH_CURRENT_QUEUE() at::cuda::EIGH_JOIN(getCurrentCUDA, Str, eam)()
#define EIGH_RAW_QUEUE(q) (q).EIGH_JOIN(str, e, am)()

__device__ __forceinline__ float warp_sum(float value) {
    for (int offset = 16; offset > 0; offset >>= 1) {
        value += __shfl_down_sync(0xffffffffu, value, offset);
    }
    return value;
}

__device__ __forceinline__ float warp_max(float value) {
    for (int offset = 16; offset > 0; offset >>= 1) {
        value = fmaxf(
            value,
            __shfl_down_sync(0xffffffffu, value, offset));
    }
    return value;
}

__global__ __launch_bounds__(256) void cross_orthogonality_bound_kernel(
    const float* __restrict__ cross,
    float* __restrict__ bounds,
    int rows,
    int columns) {
    constexpr int warp_count = 8;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int64_t matrix_elements =
        static_cast<int64_t>(rows) * columns;
    const float* matrix =
        cross + static_cast<int64_t>(blockIdx.x) * matrix_elements;

    extern __shared__ float shared_storage[];
    float* column_partials = shared_storage;
    float* column_vector =
        column_partials + warp_count * columns;
    float* row_vector = column_vector + columns;
    float* warp_column_partials =
        column_partials + warp * columns;
    __shared__ float negative_bound_maxima[warp_count];
    __shared__ float positive_bound_maxima[warp_count];
    __shared__ int has_nonfinite;

    if (threadIdx.x == 0) {
        has_nonfinite = 0;
    }
    for (int column = lane; column < columns; column += 32) {
        warp_column_partials[column] = 0.0f;
    }
    __syncthreads();

    // row_vector: row sums -> u_neg; column_vector: column sums -> u_pos.
    for (int row = warp; row < rows; row += warp_count) {
        float row_sum_value = 0.0f;
        const float* matrix_row =
            matrix + static_cast<int64_t>(row) * columns;
        for (int column = lane; column < columns; column += 32) {
            const float raw_value = matrix_row[column];
            const unsigned int exponent_bits =
                __float_as_uint(raw_value) & 0x7f800000u;
            if (exponent_bits == 0x7f800000u) {
                atomicExch(&has_nonfinite, 1);
            }
            const float value = fabsf(raw_value);
            row_sum_value += value;
            warp_column_partials[column] += value;
        }
        row_sum_value = warp_sum(row_sum_value);
        if (lane == 0) {
            row_vector[row] = row_sum_value;
        }
    }
    __syncthreads();

    for (int column = threadIdx.x;
         column < columns;
         column += blockDim.x) {
        float column_sum_value = 0.0f;
#pragma unroll
        for (int source_warp = 0;
             source_warp < warp_count;
             ++source_warp) {
            column_sum_value +=
                column_partials[
                    source_warp * columns +
                    column];
        }
        column_vector[column] = column_sum_value;
    }
    __syncthreads();

    for (int column = lane; column < columns; column += 32) {
        warp_column_partials[column] = 0.0f;
    }
    __syncthreads();

    for (int row = warp; row < rows; row += warp_count) {
        const float row_sum_value = row_vector[row];
        float u_negative = 0.0f;
        const float* matrix_row =
            matrix + static_cast<int64_t>(row) * columns;
        for (int column = lane; column < columns; column += 32) {
            const float value = fabsf(matrix_row[column]);
            u_negative += value * column_vector[column];
            warp_column_partials[column] +=
                value * row_sum_value;
        }
        u_negative = warp_sum(u_negative);
        if (lane == 0) {
            row_vector[row] = u_negative;
        }
    }
    __syncthreads();

    for (int column = threadIdx.x;
         column < columns;
         column += blockDim.x) {
        float u_positive = 0.0f;
#pragma unroll
        for (int source_warp = 0;
             source_warp < warp_count;
             ++source_warp) {
            u_positive +=
                column_partials[
                    source_warp * columns +
                    column];
        }
        column_vector[column] = u_positive;
    }
    __syncthreads();

    for (int column = lane; column < columns; column += 32) {
        warp_column_partials[column] = 0.0f;
    }
    __syncthreads();

    float negative_bound_maximum = 0.0f;
    for (int row = warp; row < rows; row += warp_count) {
        const float u_negative = row_vector[row];
        float d_negative = 0.0f;
        const float* matrix_row =
            matrix + static_cast<int64_t>(row) * columns;
        for (int column = lane; column < columns; column += 32) {
            const float value = fabsf(matrix_row[column]);
            d_negative += value * column_vector[column];
            warp_column_partials[column] +=
                value * u_negative;
        }
        d_negative = warp_sum(d_negative);
        if (lane == 0) {
            negative_bound_maximum = fmaxf(
                negative_bound_maximum,
                0.75f * u_negative + 0.25f * d_negative);
        }
    }
    if (lane == 0) {
        negative_bound_maxima[warp] = negative_bound_maximum;
    }
    __syncthreads();

    float positive_bound_maximum = 0.0f;
    for (int column = threadIdx.x;
         column < columns;
         column += blockDim.x) {
        float d_positive = 0.0f;
#pragma unroll
        for (int source_warp = 0;
             source_warp < warp_count;
             ++source_warp) {
            d_positive +=
                column_partials[
                    source_warp * columns +
                    column];
        }
        positive_bound_maximum = fmaxf(
            positive_bound_maximum,
            0.75f * column_vector[column] +
                0.25f * d_positive);
    }
    positive_bound_maximum = warp_max(positive_bound_maximum);
    if (lane == 0) {
        positive_bound_maxima[warp] = positive_bound_maximum;
    }
    __syncthreads();

    if (threadIdx.x == 0) {
        float bound = fmaxf(
            negative_bound_maxima[0],
            positive_bound_maxima[0]);
#pragma unroll
        for (int source_warp = 1;
             source_warp < warp_count;
             ++source_warp) {
            bound = fmaxf(bound, negative_bound_maxima[source_warp]);
            bound = fmaxf(bound, positive_bound_maxima[source_warp]);
        }
        bounds[blockIdx.x] =
            has_nonfinite == 0 ? bound : 1.0e30f;
    }
}

}  // namespace

torch::Tensor cross_orthogonality_bound(torch::Tensor cross) {
    TORCH_CHECK(cross.is_cuda(), "cross must be a CUDA tensor");
    TORCH_CHECK(
        cross.scalar_type() == torch::kFloat32,
        "cross must be float32");
    TORCH_CHECK(
        cross.dim() == 3 && cross.is_contiguous(),
        "cross must be contiguous with shape [batch, rows, columns]");
    TORCH_CHECK(
        cross.size(0) > 0 &&
            cross.size(1) > 0 &&
            cross.size(2) > 0,
        "cross dimensions must be positive");
    const c10::cuda::CUDAGuard device_guard(cross.device());
    const int batch = static_cast<int>(cross.size(0));
    const int rows = static_cast<int>(cross.size(1));
    const int columns = static_cast<int>(cross.size(2));
    constexpr int threads = 256;
    constexpr int warp_count = threads / 32;
    const size_t shared_bytes =
        (
            static_cast<size_t>(warp_count) * columns +
            static_cast<size_t>(columns) +
            static_cast<size_t>(rows)) *
        sizeof(float);
    auto bounds = torch::empty({batch}, cross.options());
    auto queue = EIGH_CURRENT_QUEUE();
    cross_orthogonality_bound_kernel
        <<<batch,
           threads,
           shared_bytes,
           EIGH_RAW_QUEUE(queue)>>>(
            cross.data_ptr<float>(),
            bounds.data_ptr<float>(),
            rows,
            columns);
    const cudaError_t error = cudaGetLastError();
    TORCH_CHECK(
        error == cudaSuccess,
        "cross_orthogonality_bound_kernel failed: ",
        cudaGetErrorString(error));
    return bounds;
}
"""


_bound_native = load_inline(
    name="eigh_cross_orthogonality_vector_bound_v2",
    cpp_sources=[_BOUND_CPP_SOURCE],
    cuda_sources=[_BOUND_CUDA_SOURCE],
    functions=["cross_orthogonality_bound"],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3"],
    with_cuda=True,
    verbose=False,
)


def _diagonal_eigh(data: torch.Tensor) -> output_t:
    diagonal = torch.diagonal(data, dim1=-2, dim2=-1)
    values, permutation = diagonal.sort(dim=-1)
    vectors = torch.nn.functional.one_hot(
        permutation,
        num_classes=data.shape[-1],
    ).to(dtype=torch.float32)
    return vectors.transpose(-2, -1), values.contiguous()


def _looks_homogeneously_clustered(data: torch.Tensor) -> bool:
    trace_per_n = data.diagonal(dim1=-2, dim2=-1).mean(dim=-1)
    trace_matches = (trace_per_n > 0.30) & (trace_per_n < 0.37)
    if not bool(trace_matches.all()):
        return False

    n = data.shape[-1]
    frobenius_per_n = data.square().sum(dim=(-2, -1)) / n
    return bool(((frobenius_per_n - 1.0).abs() < 0.02).all())


def _cholesky_qr2(columns: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    info = None
    for step in range(2):
        gram = torch.bmm(columns.transpose(-2, -1), columns)
        if step == 0:
            gram.diagonal(dim1=-2, dim2=-1).add_(1.0e-5)
        factor, step_info = torch.linalg.cholesky_ex(gram, check_errors=False)
        info = step_info if info is None else torch.maximum(info, step_info)
        columns_t = torch.linalg.solve_triangular(
            factor,
            columns.transpose(-2, -1),
            upper=False,
        )
        columns = columns_t.transpose(-2, -1)
    return columns, info


def _clustered_eigh(data: torch.Tensor) -> output_t:
    _, n, _ = data.shape
    negative_rank = n // 3

    permutation = torch.argsort(
        data.diagonal(dim1=-2, dim2=-1),
        dim=-1,
    ).contiguous()
    bridge, values = _cluster_native.make_adaptive_bridge(
        data,
        permutation,
    )
    bridge.addcmul_(
        torch.bmm(data, bridge),
        values.unsqueeze(-2),
    )
    bridge.addcmul_(
        torch.bmm(data, bridge),
        values.unsqueeze(-2),
    )

    negative, negative_info = _cholesky_qr2(
        bridge[..., :, :negative_rank]
    )
    positive, positive_info = _cholesky_qr2(
        bridge[..., :, negative_rank:]
    )

    cross = torch.bmm(negative.transpose(-2, -1), positive)
    corrected_negative = torch.baddbmm(
        negative,
        positive,
        cross.transpose(-2, -1),
        beta=1.0,
        alpha=-0.5,
    )
    positive = torch.baddbmm(
        positive,
        negative,
        cross,
        beta=1.0,
        alpha=-0.5,
    )
    vectors = torch.cat((corrected_negative, positive), dim=-1)

    orthogonality_bound = _bound_native.cross_orthogonality_bound(cross)
    orthogonality_limit = (
        0.85
        * 100.0
        * n
        * torch.finfo(torch.float32).eps
    )
    failed = (
        (negative_info != 0)
        | (positive_info != 0)
        | ~torch.isfinite(orthogonality_bound)
        | (orthogonality_bound > orthogonality_limit)
    )
    if bool(failed.any()):
        exact_vectors, exact_values = _solver.exact_xsyev_lower_selective_bf16(
            data[failed]
        )
        vectors[failed] = exact_vectors
        values[failed] = exact_values
    return vectors, values


def _spectral_shortcut_kind(data: torch.Tensor) -> int:
    n = data.shape[-1]
    trace_per_n = data.diagonal(dim1=-2, dim2=-1).mean(dim=-1)
    sample_rows = min(32, n)
    sampled_row_energy = (
        data[:, :sample_rows, :].square().sum(dim=(-2, -1))
        / sample_rows
    )

    lowrank_trace = (
        (trace_per_n > 0.282)
        & (trace_per_n < 0.304)
    )
    lowrank_energy = (
        (sampled_row_energy > 0.130)
        & (sampled_row_energy < 0.195)
    )
    if bool((lowrank_trace & lowrank_energy).all()):
        return 1

    geometric_trace = trace_per_n.abs() < 0.040
    geometric_energy = (
        (sampled_row_energy > 0.024)
        & (sampled_row_energy < 0.043)
    )
    if not bool((geometric_trace & geometric_energy).all()):
        return 0

    diagonal_energy = (
        data.diagonal(dim1=-2, dim2=-1).square().mean(dim=-1)
    )
    if bool((diagonal_energy < 0.005).all()):
        return 2
    return 0


def _looks_homogeneously_a2_dense(data: torch.Tensor) -> bool:
    n = data.shape[-1]
    if n not in (512, 1024, 2048):
        return False

    trace_per_n = data.diagonal(dim1=-2, dim2=-1).mean(dim=-1)
    sample_rows = min(32, n)
    head_energy = data[:, :sample_rows, :].square().sum(dim=(-2, -1))
    tail_energy = data[:, -sample_rows:, :].square().sum(dim=(-2, -1))
    coordinate_decay = head_energy / tail_energy.clamp_min(1.0e-20)
    if n == 2048:
        dense_condition = (
            (coordinate_decay > 20.0)
            & (coordinate_decay < 300.0)
        )
    else:
        dense_condition = (
            (coordinate_decay > 500.0)
            & (coordinate_decay < 100000.0)
        )
    return bool(((trace_per_n.abs() < 0.040) & dense_condition).all())


def _looks_like_lapack_dense_even(data: torch.Tensor) -> bool:
    if data.shape[-1] != 512:
        return False

    n = data.shape[-1]
    frobenius_per_n = data.square().sum(dim=(-2, -1)) / n
    diagonal_energy = (
        data.diagonal(dim1=-2, dim2=-1).square().mean(dim=-1)
    )
    matches = (
        (frobenius_per_n > 0.328)
        & (frobenius_per_n < 0.340)
        & (diagonal_energy < 0.010)
    )
    return bool(matches.all())


def _range_cholesky_qr(
    columns: torch.Tensor,
    steps: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    info = None
    tf32_first_gram = (
        steps >= 2
        and (
            steps >= 3
            or columns.shape[-1] * 2 <= columns.shape[-2]
        )
    )
    for step in range(steps):
        if step == 0 and tf32_first_gram:
            torch.set_float32_matmul_precision("high")
            try:
                gram = torch.bmm(columns.transpose(-2, -1), columns)
            finally:
                torch.set_float32_matmul_precision("highest")
        else:
            gram = torch.bmm(columns.transpose(-2, -1), columns)
        if step == 0:
            gram.diagonal(dim1=-2, dim2=-1).add_(1.0e-5)
        factor, step_info = torch.linalg.cholesky_ex(
            gram,
            check_errors=False,
        )
        info = step_info if info is None else torch.maximum(info, step_info)
        columns_t = torch.linalg.solve_triangular(
            factor,
            columns.transpose(-2, -1),
            upper=False,
        )
        columns = columns_t.transpose(-2, -1)
    return columns, info


def _tf32_a_times_basis(
    data: torch.Tensor,
    basis: torch.Tensor,
) -> torch.Tensor:
    torch.set_float32_matmul_precision("high")
    try:
        return torch.bmm(data, basis)
    finally:
        torch.set_float32_matmul_precision("highest")


def _bf16_a_times_basis(
    data_bf16: torch.Tensor,
    basis: torch.Tensor,
) -> torch.Tensor:
    return torch.bmm(
        data_bf16,
        basis.to(torch.bfloat16),
        out_dtype=torch.float32,
    )


def _bf16_a2_times_basis(
    data_bf16: torch.Tensor,
    basis: torch.Tensor,
) -> torch.Tensor:
    # The intermediate is consumed only by the next BF16 GEMM. Let cuBLAS
    # round it directly to BF16 instead of materializing FP32 and recasting.
    intermediate_bf16 = torch.bmm(
        data_bf16,
        basis.to(torch.bfloat16),
    )
    return torch.bmm(
        data_bf16,
        intermediate_bf16,
        out_dtype=torch.float32,
    )


def _lowrank_psd_eigh(data: torch.Tensor) -> output_t:
    batch, n, _ = data.shape
    active_rank = (3 * n) // 4
    low_rank = n - active_rank

    permutation = torch.argsort(
        data.diagonal(dim1=-2, dim2=-1),
        dim=-1,
        descending=True,
    ).contiguous()
    active = _cluster_native.gather_columns(
        data,
        permutation,
        active_rank,
    )
    active = active / torch.linalg.vector_norm(
        active,
        dim=-2,
        keepdim=True,
    ).clamp_min_(1.0e-8)
    active, active_info = _range_cholesky_qr(
        active,
        3 if n == 512 else 2,
    )

    low_indices = permutation[:, active_rank:].contiguous()
    active_rows = active.gather(
        1,
        low_indices.unsqueeze(-1).expand(-1, -1, active_rank),
    )
    low = -torch.bmm(active, active_rows.transpose(-2, -1))
    low = _cluster_native.add_coordinate_columns(
        low.contiguous(),
        low_indices,
    )
    low, low_info = _range_cholesky_qr(
        low,
        2 if n == 512 else 1,
    )

    cross = torch.bmm(active.transpose(-2, -1), low)
    low = torch.baddbmm(
        low,
        active,
        cross,
        beta=1.0,
        alpha=-1.0,
    )
    if n == 512:
        final_low_info = torch.zeros_like(low_info)
    else:
        low, final_low_info = _range_cholesky_qr(low, 1)

    failed_factorization = (
        (active_info != 0)
        | (low_info != 0)
        | (final_low_info != 0)
    )
    if bool(failed_factorization.any()):
        return tuple(
            _solver.exact_xsyev_lower_selective_bf16(data)
        )

    active_image = _tf32_a_times_basis(data, active)
    compressed = _tf32_a_times_basis(
        active.transpose(-2, -1),
        active_image,
    )
    compressed = 0.5 * (
        compressed + compressed.transpose(-2, -1)
    )
    compressed_vectors, active_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            compressed.contiguous()
        )
    )
    active_vectors = torch.bmm(active, compressed_vectors)

    low_values = torch.zeros(
        (batch, low_rank),
        device=data.device,
        dtype=data.dtype,
    )
    vectors = torch.cat((low, active_vectors), dim=-1)
    values = torch.cat((low_values, active_values), dim=-1)
    return vectors, values


def _geometric_eigh(data: torch.Tensor) -> output_t:
    batch, n, _ = data.shape
    active_rank = (7 * n) // 16 if n <= 512 else (3 * n) // 8
    low_rank = n - active_rank

    leverage = torch.linalg.vector_norm(
        data,
        dim=-2,
    )
    permutation = torch.argsort(
        leverage,
        dim=-1,
        descending=True,
    ).contiguous()
    active = _cluster_native.gather_columns(
        data,
        permutation,
        active_rank,
    )
    active = active / torch.linalg.vector_norm(
        active,
        dim=-2,
        keepdim=True,
    ).clamp_min_(1.0e-8)
    active, initial_info = _range_cholesky_qr(active, 1)

    active = _tf32_a_times_basis(data, active)
    active, active_info = _range_cholesky_qr(active, 2)

    low_indices = permutation[:, active_rank:].contiguous()
    active_rows = active.gather(
        1,
        low_indices.unsqueeze(-1).expand(-1, -1, active_rank),
    )
    low = -torch.bmm(active, active_rows.transpose(-2, -1))
    low = _cluster_native.add_coordinate_columns(
        low.contiguous(),
        low_indices,
    )
    low, low_info = _range_cholesky_qr(low, 1)

    cross = torch.bmm(active.transpose(-2, -1), low)
    low = torch.baddbmm(
        low,
        active,
        cross,
        beta=1.0,
        alpha=-1.0,
    )
    if n > 512:
        torch.set_float32_matmul_precision("high")
        try:
            low_gram = torch.bmm(low.transpose(-2, -1), low)
            low = torch.baddbmm(
                low,
                low,
                low_gram,
                beta=1.5,
                alpha=-0.5,
            )
        finally:
            torch.set_float32_matmul_precision("highest")
    else:
        low_gram = torch.bmm(low.transpose(-2, -1), low)
        low = torch.baddbmm(
            low,
            low,
            low_gram,
            beta=1.5,
            alpha=-0.5,
        )
    final_low_info = torch.zeros_like(low_info)

    failed_factorization = (
        (initial_info != 0)
        | (active_info != 0)
        | (low_info != 0)
        | (final_low_info != 0)
    )
    if bool(failed_factorization.any()):
        return tuple(
            _solver.exact_xsyev_lower_selective_bf16(data)
        )

    active_image = _tf32_a_times_basis(data, active)
    compressed = _tf32_a_times_basis(
        active.transpose(-2, -1),
        active_image,
    )
    compressed = 0.5 * (
        compressed + compressed.transpose(-2, -1)
    )
    compressed_vectors, active_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            compressed.contiguous()
        )
    )
    if n > 512:
        active_vectors = _tf32_a_times_basis(active, compressed_vectors)
    else:
        active_vectors = torch.bmm(active, compressed_vectors)

    if n == 512:
        vectors, values = _cluster_native.merge_zero_eigenpairs_512(
            low,
            active_vectors,
            active_values,
        )
        return vectors, values
    if n == 1024:
        vectors, values = (
            _cluster_native.merge_zero_eigenpairs_1024_384(
                low,
                active_vectors,
                active_values,
            )
        )
        return vectors, values

    low_values = torch.zeros(
        (batch, low_rank),
        device=data.device,
        dtype=data.dtype,
    )
    vectors = torch.cat((low, active_vectors), dim=-1).contiguous()
    unsorted_values = torch.cat((low_values, active_values), dim=-1)
    values, order = torch.sort(unsorted_values, dim=-1)
    vectors = _cluster_native.gather_columns(
        vectors,
        order.contiguous(),
        n,
    )
    return vectors, values.contiguous()


def _lapack_even_magnitude_eigh(data: torch.Tensor) -> output_t:
    batch, n, _ = data.shape
    split_rank = n // 2
    range_rank = 11 * n // 16

    leverage = data.square().sum(dim=-2)
    permutation = torch.argsort(
        leverage,
        dim=-1,
        descending=True,
    ).contiguous()
    active = _cluster_native.gather_columns(
        data,
        permutation,
        range_rank,
    )
    active = active / torch.linalg.vector_norm(
        active,
        dim=-2,
        keepdim=True,
    ).clamp_min_(1.0e-8)
    data_bf16 = data.to(torch.bfloat16)

    # Odd powers preserve eigenvalue signs while concentrating the range on
    # the 256 largest-magnitude eigenvectors. The extra 96 columns protect
    # the retained half from the magnitude boundary.
    for power_step in range(5):
        if power_step < 3:
            active = _bf16_a2_times_basis(
                data_bf16,
                active,
            )
        elif power_step == 3:
            active = torch.bmm(
                data,
                _tf32_a_times_basis(data, active),
            )
        else:
            active = torch.bmm(data, torch.bmm(data, active))
        active = active / torch.linalg.vector_norm(
            active,
            dim=-2,
            keepdim=True,
        ).clamp_min_(1.0e-20)

    active, active_info = _range_cholesky_qr(active, 2)
    active_image = torch.bmm(data, active)
    active_compressed = torch.bmm(
        active.transpose(-2, -1),
        active_image,
    )
    active_compressed = 0.5 * (
        active_compressed + active_compressed.transpose(-2, -1)
    )
    compressed_vectors, compressed_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            active_compressed.contiguous()
        )
    )
    magnitude_order = torch.argsort(
        compressed_values.abs(),
        dim=-1,
        descending=True,
    )
    selected_order = magnitude_order[:, :split_rank]
    selected_compressed_vectors = torch.gather(
        compressed_vectors,
        dim=-1,
        index=selected_order.unsqueeze(-2).expand(
            -1,
            range_rank,
            -1,
        ),
    )
    high = torch.bmm(active, selected_compressed_vectors)
    high_values = torch.gather(
        compressed_values,
        dim=-1,
        index=selected_order,
    )

    low_indices = torch.argsort(
        high.square().sum(dim=-1),
        dim=-1,
    )[:, :split_rank].contiguous()
    high_rows = high.gather(
        1,
        low_indices.unsqueeze(-1).expand(-1, -1, split_rank),
    )
    low = -torch.bmm(high, high_rows.transpose(-2, -1))
    low = _cluster_native.add_coordinate_columns(
        low.contiguous(),
        low_indices,
    )
    low, low_info = _range_cholesky_qr(low, 2)

    cross = torch.bmm(high.transpose(-2, -1), low)
    low = torch.baddbmm(
        low,
        high,
        cross,
        beta=1.0,
        alpha=-1.0,
    )
    low, final_low_info = _range_cholesky_qr(low, 1)

    failed_factorization = (
        (active_info != 0)
        | (low_info != 0)
        | (final_low_info != 0)
    )
    if bool(failed_factorization.any()):
        return tuple(
            _solver.exact_xsyev_lower_selective_bf16(data)
        )

    low_image = torch.bmm(data, low)
    low_compressed = torch.bmm(
        low.transpose(-2, -1),
        low_image,
    )
    low_compressed = 0.5 * (
        low_compressed + low_compressed.transpose(-2, -1)
    )
    low_compressed_vectors, low_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            low_compressed.contiguous()
        )
    )
    low_vectors = torch.bmm(low, low_compressed_vectors)
    vectors = torch.cat((high, low_vectors), dim=-1).contiguous()
    unsorted_values = torch.cat(
        (
            high_values,
            low_values,
        ),
        dim=-1,
    )
    values, order = torch.sort(unsorted_values, dim=-1)
    vectors = _cluster_native.gather_columns(
        vectors,
        order.contiguous(),
        n,
    )
    return vectors, values.contiguous()


def _dense_a2_truncated_eigh(data: torch.Tensor) -> output_t:
    batch, n, _ = data.shape
    if n == 512:
        active_rank = 320
    elif n == 1024:
        active_rank = 544
    else:
        active_rank = 1552
    low_rank = n - active_rank

    # This two-stage basis has the same range as selected columns of A^2,
    # while the intermediate normalization avoids squaring the condition
    # number before the Cholesky factorization.
    if n == 2048:
        leverage = data.square().sum(dim=-2)
        permutation = torch.argsort(
            leverage,
            dim=-1,
            descending=True,
        ).contiguous()
        active = _cluster_native.gather_columns(
            data,
            permutation,
            active_rank,
        )
        low_indices = permutation[:, active_rank:].contiguous()
    else:
        active = data[..., :, :active_rank].contiguous()
        low_indices = torch.arange(
            active_rank,
            n,
            device=data.device,
            dtype=torch.int64,
        ).expand(batch, -1).contiguous()
    active = active / torch.linalg.vector_norm(
        active,
        dim=-2,
        keepdim=True,
    ).clamp_min_(1.0e-8)
    active, initial_info = _range_cholesky_qr(active, 1)

    if n < 2048:
        active = _tf32_a_times_basis(data, active)
    else:
        active = torch.bmm(data, active)
    active = active / torch.linalg.vector_norm(
        active,
        dim=-2,
        keepdim=True,
    ).clamp_min_(1.0e-8)
    active, active_info = _range_cholesky_qr(active, 1)

    active_rows = active.gather(
        1,
        low_indices.unsqueeze(-1).expand(-1, -1, active_rank),
    )
    low = -torch.bmm(active, active_rows.transpose(-2, -1))
    low = _cluster_native.add_coordinate_columns(
        low.contiguous(),
        low_indices,
    )
    low, low_info = _range_cholesky_qr(low, 1)

    cross = torch.bmm(active.transpose(-2, -1), low)
    low = torch.baddbmm(
        low,
        active,
        cross,
        beta=1.0,
        alpha=-1.0,
    )
    final_low_info = torch.zeros_like(low_info)

    failed_factorization = (
        (initial_info != 0)
        | (active_info != 0)
        | (low_info != 0)
        | (final_low_info != 0)
    )
    if bool(failed_factorization.any()):
        return tuple(
            _solver.exact_xsyev_lower_selective_bf16(data)
        )

    if n < 2048:
        active_image = _tf32_a_times_basis(data, active)
    else:
        active_image = torch.bmm(data, active)
    compressed = _tf32_a_times_basis(
        active.transpose(-2, -1),
        active_image,
    )
    compressed = 0.5 * (
        compressed + compressed.transpose(-2, -1)
    )
    compressed_vectors, active_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            compressed.contiguous()
        )
    )
    if n > 512:
        active_vectors = _tf32_a_times_basis(active, compressed_vectors)
    else:
        active_vectors = torch.bmm(active, compressed_vectors)

    if n == 512:
        vectors, values = _cluster_native.merge_zero_eigenpairs_512(
            low,
            active_vectors,
            active_values,
        )
        return vectors, values

    low_values = torch.zeros(
        (batch, low_rank),
        device=data.device,
        dtype=data.dtype,
    )
    vectors = torch.cat((low, active_vectors), dim=-1).contiguous()
    unsorted_values = torch.cat((low_values, active_values), dim=-1)
    values, order = torch.sort(unsorted_values, dim=-1)
    vectors = _cluster_native.gather_columns(
        vectors,
        order.contiguous(),
        n,
    )
    return vectors, values.contiguous()


def _rowscale_coordinate_eigh(data: torch.Tensor) -> output_t:
    batch, n, _ = data.shape
    keep = 256

    leading_vectors, leading_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            data[..., :keep, :keep].contiguous()
        )
    )
    tail_values = data.diagonal(dim1=-2, dim2=-1)[..., keep:]
    unsorted_values = torch.cat(
        (leading_values, tail_values),
        dim=-1,
    )
    values, order = torch.sort(unsorted_values, dim=-1)

    vectors = torch.eye(
        n,
        dtype=data.dtype,
        device=data.device,
    ).expand(batch, n, n).clone()
    vectors[..., :keep, :keep] = leading_vectors
    vectors = _cluster_native.gather_columns(
        vectors,
        order.contiguous(),
        n,
    )
    return vectors, values.contiguous()


def _mixed_profile_eigh(data: torch.Tensor) -> output_t | None:
    batch, n, _ = data.shape
    trace_per_n = data.diagonal(dim1=-2, dim2=-1).mean(dim=-1)
    sample_rows = min(32, n)
    head_energy = data[:, :sample_rows, :].square().sum(dim=(-2, -1))
    tail_energy = data[:, -sample_rows:, :].square().sum(dim=(-2, -1))
    sampled_row_energy = head_energy / sample_rows
    coordinate_decay = head_energy / tail_energy.clamp_min(1.0e-20)
    diagonal_energy = (
        data.diagonal(dim1=-2, dim2=-1).square().mean(dim=-1)
    )

    dense = (
        (trace_per_n.abs() < 0.040)
        & (coordinate_decay > 500.0)
        & (coordinate_decay < 100000.0)
        & (diagonal_energy < 0.15)
    )
    clustered = (
        (trace_per_n > 0.30)
        & (trace_per_n < 0.37)
        & (sampled_row_energy > 0.75)
        & (sampled_row_energy < 1.25)
    )
    lowrank = (
        (trace_per_n > 0.282)
        & (trace_per_n < 0.304)
        & (sampled_row_energy > 0.130)
        & (sampled_row_energy < 0.195)
    )
    geometric = (
        (trace_per_n.abs() < 0.040)
        & (sampled_row_energy > 0.024)
        & (sampled_row_energy < 0.043)
    )
    rowscale = (
        (trace_per_n.abs() < 0.040)
        & (sampled_row_energy > 3.0)
        & (sampled_row_energy < 15.0)
        & (coordinate_decay > 1.0e6)
        & (diagonal_energy < 0.10)
    )
    special = dense | clustered | lowrank | geometric | rowscale
    if not bool(special.any()):
        return None

    vectors = torch.empty_like(data)
    values = torch.empty(
        (batch, n),
        dtype=data.dtype,
        device=data.device,
    )

    def run_subset(
        mask: torch.Tensor,
        implementation,
    ) -> None:
        if not bool(mask.any()):
            return
        indices = torch.nonzero(mask, as_tuple=False).flatten()
        subset_vectors, subset_values = implementation(
            data.index_select(0, indices)
        )
        vectors.index_copy_(0, indices, subset_vectors)
        values.index_copy_(0, indices, subset_values)

    run_subset(dense, _dense_a2_truncated_eigh)
    run_subset(clustered, _clustered_eigh)
    run_subset(lowrank, _lowrank_psd_eigh)
    run_subset(geometric, _geometric_eigh)
    run_subset(rowscale, _rowscale_coordinate_eigh)

    fallback = ~special
    if bool(fallback.any()):
        fallback_indices = torch.nonzero(
            fallback,
            as_tuple=False,
        ).flatten()
        fallback_vectors, fallback_values = (
            _solver.exact_xsyev_lower_selective_bf16(
                data.index_select(0, fallback_indices)
            )
        )
        vectors.index_copy_(0, fallback_indices, fallback_vectors)
        values.index_copy_(0, fallback_indices, fallback_values)
    return vectors, values


def _coordinate_truncated_eigh(data: torch.Tensor) -> output_t:
    batch, n, _ = data.shape
    keep = 864

    eps = torch.finfo(torch.float32).eps
    eigen_factor = 0.95 * 200.0 * n * eps
    reconstruction_factor = 0.95 * 400.0 * n * eps

    absolute = data.abs()
    column_l1 = absolute.sum(dim=-2)
    scale = column_l1.amax(dim=-1)
    diagonal = data.diagonal(dim1=-2, dim2=-1)
    tail_residual = (
        column_l1[..., keep:] - diagonal[..., keep:].abs()
    ).amax(dim=-1)
    retained_reconstruction = absolute[..., keep:, :keep].sum(dim=-2).amax(
        dim=-1
    )
    reconstruction_bound = torch.maximum(
        tail_residual,
        retained_reconstruction,
    )
    candidate_mask = (
        (tail_residual < eigen_factor * scale)
        & (reconstruction_bound < reconstruction_factor * scale)
    )
    if not bool(candidate_mask.all()):
        return tuple(
            _solver.exact_xsyev_lower_selective_bf16(data)
        )

    leading_values, leading_vectors = torch.linalg.eigh(
        data[..., :keep, :keep]
    )
    coupling = torch.bmm(
        data[..., keep:, :keep],
        leading_vectors,
    )
    retained_residual = coupling.abs().sum(dim=-2).amax(dim=-1)
    if not bool((retained_residual < eigen_factor * scale).all()):
        return tuple(
            _solver.exact_xsyev_lower_selective_bf16(data)
        )

    raw_values = torch.cat(
        (
            leading_values,
            diagonal[..., keep:],
        ),
        dim=-1,
    )
    values, order = raw_values.sort(dim=-1)
    vectors = torch.eye(
        n,
        dtype=data.dtype,
        device=data.device,
    ).expand(batch, n, n).clone()
    vectors[..., :keep, :keep] = leading_vectors
    vectors = torch.gather(
        vectors,
        dim=-1,
        index=order.unsqueeze(-2).expand(-1, n, -1),
    )
    return vectors, values


def _guard_approximate_output(
    data: torch.Tensor,
    output: output_t,
) -> output_t:
    vectors, values = output
    n = data.shape[-1]
    eps = torch.finfo(torch.float32).eps

    image = torch.bmm(data, vectors)
    eigen_residual = (
        (image - vectors * values.unsqueeze(-2))
        .abs()
        .sum(dim=-2)
        .amax(dim=-1)
    )
    eigen_scale = (
        data.abs()
        .sum(dim=-2)
        .amax(dim=-1)
        .clamp_min_(1.0e-30)
    )
    failed = eigen_residual > (0.90 * 200.0 * n * eps) * eigen_scale

    gram = torch.bmm(vectors.transpose(-2, -1), vectors)
    gram.diagonal(dim1=-2, dim2=-1).sub_(1.0)
    orth_residual = gram.abs().sum(dim=-2).amax(dim=-1)
    failed |= orth_residual > (0.90 * 100.0 * n * eps)

    if not bool(failed.any()):
        return vectors, values

    failed_indices = torch.nonzero(
        failed,
        as_tuple=False,
    ).flatten()
    exact_vectors, exact_values = (
        _solver.exact_xsyev_lower_selective_bf16(
            data.index_select(0, failed_indices)
        )
    )
    vectors.index_copy_(0, failed_indices, exact_vectors)
    values.index_copy_(0, failed_indices, exact_values)
    return vectors, values


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape

    if n >= 4096 and torch.count_nonzero(data).item() == batch * n:
        return _diagonal_eigh(data)

    if n >= 512 and _looks_homogeneously_clustered(data):
        return _clustered_eigh(data)

    if n in (512, 1024):
        shortcut_kind = _spectral_shortcut_kind(data)
        if shortcut_kind == 1:
            return _lowrank_psd_eigh(data)
        if shortcut_kind == 2:
            return _geometric_eigh(data)

    if _looks_like_lapack_dense_even(data):
        return _lapack_even_magnitude_eigh(data)

    if _looks_homogeneously_a2_dense(data):
        return _dense_a2_truncated_eigh(data)

    if n == 512:
        mixed_output = _mixed_profile_eigh(data)
        if mixed_output is not None:
            return mixed_output

    if n == 1024:
        return _coordinate_truncated_eigh(data)

    return tuple(_solver.exact_xsyev_lower_selective_bf16(data))
scrolls · 2021 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