Skip to content
KernelIndex
Search⌘K

submission 877050

rd9000 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_final.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-877050?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
48.2ms
#139 of 286
2026-07-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f7a2f567f6eea48d4e7e72d69531ef1f398de9baef8272d8d4f487eb3b718210
license declaredunknown
license concludedunknown
authorsrd9000
imported2026-08-26

Kernel source

submission_final.py106 lines
"""Batched symmetric eigh.

torch.linalg.eigh on CUDA already dispatches to batched cuSOLVER (Xsyev) for
every n here, which is the measured performance frontier for n >= 176 among
cuSOLVER-based approaches. For n <= 32 the batched Jacobi solver
(syevjBatched) with a task-appropriate tolerance is ~1.5x faster than the
default path (Jacobi rotations keep Q orthogonal by construction; 16 sweeps
at tol 1e-6 sit ~100x inside the n-scaled task gates).
"""

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 <cusolverDn.h>

#define CUSOLVER_CHECK(expr)                                              \
  do {                                                                    \
    cusolverStatus_t st_ = (expr);                                        \
    TORCH_CHECK(st_ == CUSOLVER_STATUS_SUCCESS, "cusolver error ",        \
                static_cast<int>(st_), " at ", __FILE__, ":", __LINE__);  \
  } while (0)

static cusolverDnHandle_t get_handle() {
  static cusolverDnHandle_t handle = [] {
    cusolverDnHandle_t h;
    CUSOLVER_CHECK(cusolverDnCreate(&h));
    return h;
  }();
  return handle;
}

// A: (batch, n, n) fp32 contiguous, symmetric. Overwritten with eigenvectors
// (column-major -> the row-major view is Q^T). Returns W (batch, n) ascending.
torch::Tensor syevj_batched(torch::Tensor A, double tol, int64_t max_sweeps) {
  TORCH_CHECK(A.is_cuda() && A.dtype() == torch::kFloat32 && A.is_contiguous());
  int64_t batch = A.size(0), n = A.size(1);
  auto handle = get_handle();

  static syevjInfo_t jparams = [&] {
    syevjInfo_t p;
    CUSOLVER_CHECK(cusolverDnCreateSyevjInfo(&p));
    CUSOLVER_CHECK(cusolverDnXsyevjSetSortEig(p, 1));
    return p;
  }();
  CUSOLVER_CHECK(cusolverDnXsyevjSetTolerance(jparams, tol));
  CUSOLVER_CHECK(cusolverDnXsyevjSetMaxSweeps(jparams, (int)max_sweeps));

  auto W = torch::empty({batch, n}, A.options());
  int lwork = 0;
  CUSOLVER_CHECK(cusolverDnSsyevjBatched_bufferSize(
      handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_LOWER, (int)n,
      A.data_ptr<float>(), (int)n, W.data_ptr<float>(), &lwork, jparams,
      (int)batch));
  auto work = torch::empty({(int64_t)lwork}, A.options());
  auto info = torch::empty({batch}, A.options().dtype(torch::kInt32));

  CUSOLVER_CHECK(cusolverDnSsyevjBatched(
      handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_LOWER, (int)n,
      A.data_ptr<float>(), (int)n, W.data_ptr<float>(), work.data_ptr<float>(),
      lwork, info.data_ptr<int>(), jparams, (int)batch));
  return W;
}
"""

_ext = None


def _get_ext():
    global _ext
    if _ext is None:
        _ext = load_inline(
            name="eigh_syevj_ext",
            cpp_sources=[_CPP_SRC],
            functions=["syevj_batched"],
            with_cuda=True,
            extra_ldflags=["-lcusolver"],
            verbose=False,
        )
    return _ext


def _l1(x: torch.Tensor) -> torch.Tensor:
    return x.abs().sum(dim=-2).amax(dim=-1)


def custom_kernel(data: input_t) -> output_t:
    b, n, _ = data.shape
    if n > 32:
        values, vectors = torch.linalg.eigh(data)
        return vectors, values

    try:
        ext = _get_ext()
        a = data.clone()  # solver overwrites its input with eigenvectors
        w = ext.syevj_batched(a, 1e-6, 16)
        q = a.transpose(-1, -2)

        return q, w
    except Exception:  # blanket safety net
        values, vectors = torch.linalg.eigh(data)
        return vectors, values
scrolls · 106 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