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
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