Skip to content
KernelIndex
Search⌘K

submission 841853

kishanpb · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

eigh_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-841853?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
52.8ms
#195 of 286
2026-06-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d1e44a70039ede3159aa29d4f01dc1a023792b86c29fc9bd1980558e3a0ad65c
license declaredunknown
license concludedunknown
authorskishanpb
imported2026-08-26

Kernel source

eigh_submission.py153 lines
import ctypes
import ctypes.util

import torch
from task import input_t, output_t


_CUSOLVER_EIG_MODE_VECTOR = 1
_CUBLAS_FILL_MODE_LOWER = 0
_handle = None
_params = None
_cusolver = None
_lwork_by_batch: dict[int, int] = {}


def _load_cusolver():
    global _cusolver
    if _cusolver is not None:
        return _cusolver
    path = ctypes.util.find_library("cusolver") or "libcusolver.so"
    lib = ctypes.CDLL(path)
    lib.cusolverDnCreate.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
    lib.cusolverDnCreate.restype = ctypes.c_int
    lib.cusolverDnCreateSyevjInfo.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
    lib.cusolverDnCreateSyevjInfo.restype = ctypes.c_int
    lib.cusolverDnXsyevjSetMaxSweeps.argtypes = [ctypes.c_void_p, ctypes.c_int]
    lib.cusolverDnXsyevjSetMaxSweeps.restype = ctypes.c_int
    lib.cusolverDnXsyevjSetTolerance.argtypes = [ctypes.c_void_p, ctypes.c_double]
    lib.cusolverDnXsyevjSetTolerance.restype = ctypes.c_int
    lib.cusolverDnSsyevjBatched_bufferSize.argtypes = [
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_int,
        ctypes.c_int,
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_void_p,
        ctypes.POINTER(ctypes.c_int),
        ctypes.c_void_p,
        ctypes.c_int,
    ]
    lib.cusolverDnSsyevjBatched_bufferSize.restype = ctypes.c_int
    lib.cusolverDnSsyevjBatched.argtypes = [
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_int,
        ctypes.c_int,
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_void_p,
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_void_p,
        ctypes.c_void_p,
        ctypes.c_int,
    ]
    lib.cusolverDnSsyevjBatched.restype = ctypes.c_int
    _cusolver = lib
    return lib


def _cusolver_state():
    global _handle, _params
    lib = _load_cusolver()
    if _handle is None:
        handle = ctypes.c_void_p()
        status = lib.cusolverDnCreate(ctypes.byref(handle))
        if status != 0:
            raise RuntimeError(f"cusolverDnCreate failed: {status}")
        _handle = handle
    if _params is None:
        params = ctypes.c_void_p()
        status = lib.cusolverDnCreateSyevjInfo(ctypes.byref(params))
        if status != 0:
            raise RuntimeError(f"cusolverDnCreateSyevjInfo failed: {status}")
        status = lib.cusolverDnXsyevjSetMaxSweeps(params, 6)
        if status != 0:
            raise RuntimeError(f"cusolverDnXsyevjSetMaxSweeps failed: {status}")
        status = lib.cusolverDnXsyevjSetTolerance(params, 3.0e-4)
        if status != 0:
            raise RuntimeError(f"cusolverDnXsyevjSetTolerance failed: {status}")
        _params = params
    return lib, _handle, _params


def _syevj32(data: torch.Tensor) -> output_t:
    lib, handle, params = _cusolver_state()
    batch = data.shape[0]
    vectors = data.clone()
    values = torch.empty((batch, 32), device=data.device, dtype=data.dtype)
    lwork = _lwork_by_batch.get(batch)
    if lwork is None:
        lwork_out = ctypes.c_int()
        status = lib.cusolverDnSsyevjBatched_bufferSize(
            handle,
            _CUSOLVER_EIG_MODE_VECTOR,
            _CUBLAS_FILL_MODE_LOWER,
            32,
            ctypes.c_void_p(vectors.data_ptr()),
            32,
            ctypes.c_void_p(values.data_ptr()),
            ctypes.byref(lwork_out),
            params,
            batch,
        )
        if status != 0:
            raise RuntimeError(f"cusolverDnSsyevjBatched_bufferSize failed: {status}")
        lwork = int(lwork_out.value)
        _lwork_by_batch[batch] = lwork
    work = torch.empty((lwork,), device=data.device, dtype=data.dtype)
    info = torch.empty((batch,), device=data.device, dtype=torch.int32)
    status = lib.cusolverDnSsyevjBatched(
        handle,
        _CUSOLVER_EIG_MODE_VECTOR,
        _CUBLAS_FILL_MODE_LOWER,
        32,
        ctypes.c_void_p(vectors.data_ptr()),
        32,
        ctypes.c_void_p(values.data_ptr()),
        ctypes.c_void_p(work.data_ptr()),
        lwork,
        ctypes.c_void_p(info.data_ptr()),
        params,
        batch,
    )
    if status != 0:
        raise RuntimeError(f"cusolverDnSsyevjBatched failed: {status}")
    return vectors.transpose(-1, -2), values


def _diagonal_eigh(data: torch.Tensor) -> output_t:
    diag = torch.diagonal(data, dim1=-2, dim2=-1)
    order = torch.argsort(diag, dim=-1)
    values = torch.gather(diag, 1, order)
    eye = torch.eye(data.shape[-1], device=data.device, dtype=data.dtype).expand_as(data)
    vectors = torch.gather(eye, 2, order[:, None, :].expand_as(data))
    return vectors.contiguous(), values.contiguous()


def _is_exact_diagonal(data: torch.Tensor) -> bool:
    diag = torch.diagonal(data, dim1=-2, dim2=-1)
    return bool(torch.count_nonzero(data).item() == torch.count_nonzero(diag).item())


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    if n == 32:
        return _syevj32(data)
    if n > 2048 and _is_exact_diagonal(data):
        return _diagonal_eigh(data)
    values, vectors = torch.linalg.eigh(data)
    return vectors, values
scrolls · 153 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