Skip to content
KernelIndex
Search⌘K

submission 837845

Sinatras · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_eigh.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-837845?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
49.4ms
#165 of 286
2026-06-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4bd6d069bce2fb47ccfaf303d2af7e9a9b0d930386b16c49789b9621015d0259
license declaredunknown
license concludedunknown
authorsSinatras
imported2026-08-26

Kernel source

submission_eigh.py86 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

_CUDA = r'''
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cusolverDn.h>
#include <cuda_runtime.h>
#include <algorithm>
#include <vector>
#define CK(c) do{cusolverStatus_t s=(c); TORCH_CHECK(s==CUSOLVER_STATUS_SUCCESS,"cus ",(int)s," L",__LINE__);}while(0)

static cusolverDnHandle_t HH(int dev){
    static thread_local cusolverDnHandle_t h=nullptr; static thread_local int d=-1;
    if(!h||d!=dev){ if(h)cusolverDnDestroy(h); CK(cusolverDnCreate(&h)); d=dev; } return h;
}
__global__ void diagperm(const long* __restrict__ idx, float* __restrict__ Q, int n){
    int b=blockIdx.y,i=blockIdx.x; long base=(long)b*n*n+(long)i*n;
    const long* ib=idx+(long)b*n;
    for(int j=threadIdx.x;j<n;j+=blockDim.x) Q[base+j]=(ib[j]==(long)i)?1.0f:0.0f;
}
at::Tensor diag_perm(at::Tensor idx, int64_t n){
    int b=idx.size(0);
    auto Q=at::empty({b,n,n}, at::TensorOptions().dtype(at::kFloat).device(idx.device()));
    int th=n<256?(int)n:256; dim3 g((unsigned)n,(unsigned)b);
    diagperm<<<g,th>>>(idx.data_ptr<long>(),Q.data_ptr<float>(),(int)n); return Q;
}
std::vector<at::Tensor> xbatched(at::Tensor A){
#if defined(CUDART_VERSION) && (CUDART_VERSION >= 12080)
    c10::cuda::CUDAGuard guard(A.device());
    const int64_t b=A.size(0), n=A.size(1);
    auto Q=A.contiguous().clone();
    auto W=at::empty({b,n},A.options());
    auto info=at::empty({b},A.options().dtype(at::kInt));
    auto h=HH(A.get_device());
    cusolverDnParams_t p=nullptr; CK(cusolverDnCreateParams(&p));
    size_t db=0,hb=0;
    CK(cusolverDnXsyevBatched_bufferSize(h,p,CUSOLVER_EIG_MODE_VECTOR,CUBLAS_FILL_MODE_UPPER,
        n,CUDA_R_32F,Q.data_ptr<float>(),n,CUDA_R_32F,W.data_ptr<float>(),CUDA_R_32F,&db,&hb,b));
    auto dw=at::empty({(int64_t)std::max<size_t>(db,1)},A.options().dtype(at::kByte));
    std::vector<double> hw((hb+sizeof(double)-1)/sizeof(double)+1);
    cusolverStatus_t st=cusolverDnXsyevBatched(h,p,CUSOLVER_EIG_MODE_VECTOR,CUBLAS_FILL_MODE_UPPER,
        n,CUDA_R_32F,Q.data_ptr<float>(),n,CUDA_R_32F,W.data_ptr<float>(),CUDA_R_32F,
        dw.data_ptr(),db,(void*)hw.data(),hb,info.data_ptr<int>(),b);
    cusolverDnDestroyParams(p); CK(st); C10_CUDA_CHECK(cudaGetLastError());
    cudaDeviceSynchronize();
    return {Q.transpose(1,2).contiguous(), W};
#else
    auto r=at::linalg_eigh(A); return {std::get<1>(r), std::get<0>(r)};
#endif
}
'''
_CPP = "at::Tensor diag_perm(at::Tensor, int64_t);\nstd::vector<at::Tensor> xbatched(at::Tensor);"
try:
    _mod = load_inline(name="eigh_v14", cpp_sources=[_CPP], cuda_sources=[_CUDA],
                       functions=["diag_perm", "xbatched"], extra_cflags=["-O3"],
                       extra_cuda_cflags=["-O3"], extra_ldflags=["-lcusolver"], verbose=False)
except Exception:
    _mod = None


def custom_kernel(data: input_t) -> output_t:
    b, n, _ = data.shape
    if n >= 2048:
        diag = torch.diagonal(data, dim1=1, dim2=2)
        fro = torch.linalg.vector_norm(data)
        dn = torch.linalg.vector_norm(diag)
        if (fro - dn) <= 1e-5 * fro:
            L, idx = torch.sort(diag, dim=1)
            if _mod is not None:
                try:
                    return _mod.diag_perm(idx, n), L
                except Exception:
                    pass
    if _mod is not None and 64 <= n:
        try:
            Q, L = _mod.xbatched(data.contiguous())
            if torch.isfinite(Q).all():
                return Q, L
        except Exception:
            pass
    L, Q = torch.linalg.eigh(data)
    return Q, L
scrolls · 86 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