Skip to content
KernelIndex
Search⌘K

submission 849740

kevinniechen_12917 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-849740?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.4ms
#142 of 286
2026-07-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cbfdef6ee8303abd35b0149e927c392b34b18586aa952c47cd97270f057f5487
license declaredunknown
license concludedunknown
authorskevinniechen_12917
imported2026-08-26

Kernel source

submission.py120 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> xsyev_batched(torch::Tensor a);
"""

CUDA_SRC = r"""
#include <torch/extension.h>
#include <cusolverDn.h>
#include <stdexcept>
#include <vector>

#define CK(expr)                                                          \
  do {                                                                    \
    cusolverStatus_t st_ = (expr);                                        \
    if (st_ != CUSOLVER_STATUS_SUCCESS) {                                 \
      throw std::runtime_error(std::string("cusolver err ") +            \
                               std::to_string((int)st_));                 \
    }                                                                     \
  } while (0)

namespace {
cusolverDnHandle_t sol_handle() {
  static cusolverDnHandle_t h = nullptr;
  if (!h) CK(cusolverDnCreate(&h));
  return h;
}
cusolverDnParams_t sol_params() {
  static cusolverDnParams_t p = nullptr;
  if (!p) CK(cusolverDnCreateParams(&p));
  return p;
}
}  // namespace

std::vector<torch::Tensor> xsyev_batched(torch::Tensor a) {
  int64_t batch = a.size(0), n = a.size(1);
  torch::Tensor w = torch::empty({batch, n}, a.options());
  torch::Tensor info = torch::empty({batch}, a.options().dtype(torch::kInt32));
  size_t dev_b = 0, host_b = 0;
  CK(cusolverDnXsyevBatched_bufferSize(sol_handle(), sol_params(),
      CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_LOWER, n, CUDA_R_32F,
      a.data_ptr(), n, CUDA_R_32F, w.data_ptr(), CUDA_R_32F, &dev_b, &host_b,
      batch));
  torch::Tensor dws = torch::empty({(int64_t)std::max<size_t>(dev_b, 16)},
                                   a.options().dtype(torch::kUInt8));
  static std::vector<uint8_t> hws;
  if (hws.size() < host_b) hws.resize(host_b);
  CK(cusolverDnXsyevBatched(sol_handle(), sol_params(),
      CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_LOWER, n, CUDA_R_32F,
      a.data_ptr(), n, CUDA_R_32F, w.data_ptr(), CUDA_R_32F, dws.data_ptr(),
      dev_b, hws.empty() ? nullptr : (void*)hws.data(), host_b,
      info.data_ptr<int>(), batch));
  return {a, w};
}
"""

module = load_inline(
    name="eigh_v8",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["xsyev_batched"],
    with_cuda=True,
    extra_ldflags=["-lcusolver"],
    verbose=False,
)


def _xsyev(a):
    buf, w = module.xsyev_batched(a.contiguous())
    return buf.transpose(-1, -2), w


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    dev = data.device
    if n <= 32:
        w, q = torch.linalg.eigh(data)
        return q, w

    if n < 512:
        buf, w = _xsyev(data.clone())
        return buf, w

    # diagonal fast path: sorted diagonal is the exact answer; shrinks the
    # cusolver batch for mixed inputs. Detect on unit-scaled data to avoid
    # fp32 overflow for high-magnitude inputs.
    amax = data.abs().amax(dim=(-2, -1), keepdim=True).clamp_min(1e-30)
    As = data / amax
    d2 = As * As
    total = d2.sum(dim=(-2, -1))
    diag2 = torch.diagonal(d2, dim1=-2, dim2=-1).sum(-1)
    is_diag = (total - diag2) <= 1e-12 * total

    if not bool(is_diag.any()):
        buf, w = _xsyev(data.clone())
        return buf, w

    Q = torch.empty(batch, n, n, device=dev)
    L = torch.empty(batch, n, device=dev)
    i = torch.nonzero(is_diag).flatten()
    dvals = torch.diagonal(data[i], dim1=-2, dim2=-1)
    ds, si = torch.sort(dvals, dim=-1)
    L[i] = ds
    eye = torch.eye(n, device=dev)
    Q[i] = torch.gather(eye.expand(len(i), n, n), -1,
                        si.unsqueeze(-2).expand(len(i), n, n))
    j = torch.nonzero(~is_diag).flatten()
    if len(j) > 0:
        qb, wb = _xsyev(data[j].clone())
        Q[j] = qb.contiguous()
        L[j] = wb
    return Q, L
scrolls · 120 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