Skip to content
KernelIndex
Search⌘K

submission 923062

Shaman Shetty · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-923062?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
3.78ms
#330 of 337
2026-07-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:52057f386ef74dfb040a0b303df07449051356b8e8c8223ee7e711a6c7e91505
license declaredunknown
license concludedunknown
authorsShaman Shetty
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memoryextern __shared__ float smem[];

Kernel source

submission2.py301 lines
import torch

try:
    from task import input_t, output_t
except Exception:
    input_t = torch.Tensor
    output_t = torch.Tensor

try:
    torch.backends.cuda.matmul.allow_tf32 = False
    torch.set_float32_matmul_precision("highest")
except Exception:
    pass

_EXT = None
_BASE = 128
_TRSM_OUT = None
_PLAN = {}

_CPP_DECL = "void bchol(torch::Tensor dst, torch::Tensor src, bool zero_upper);"

_CUDA_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>


__global__ void bchol_kernel(
    float* __restrict__ dst, const float* __restrict__ src,
    const long batch, const int n,
    const long d_srow, const long d_sbatch,
    const long s_srow, const long s_sbatch,
    const int ld, const int zero_upper)
{
    extern __shared__ float smem[];
    float* a   = smem;
    float* dia = smem + (size_t)n * ld;

    const long b = blockIdx.x;
    if (b >= batch) return;
    const float* S = src + b * s_sbatch;
    float*       D = dst + b * d_sbatch;
    const int t = threadIdx.x;
    const int T = blockDim.x;

    for (int idx = t; idx < n * n; idx += T) {
        const int r = idx / n;
        const int c = idx - r * n;
        a[r * ld + c] = S[(long)r * s_srow + c];
    }
    __syncthreads();

    const int i = t;  
    for (int j = 0; j < n; ++j) {
        const float ajj = a[j * ld + j];
        const float dj  = sqrtf(ajj);
        const float inv = 1.0f / dj;
        if (i == j) dia[j] = dj;
        if (i > j && i < n) a[i * ld + j] *= inv;
        __syncthreads();
        if (i > j && i < n) {
            const float lij = a[i * ld + j];
            float* ai = a + i * ld;
            #pragma unroll 4
            for (int k = j + 1; k <= i; ++k)
                ai[k] -= lij * a[k * ld + j];
        }
        __syncthreads();
    }

    for (int idx = t; idx < n * n; idx += T) {
        const int r = idx / n;
        const int c = idx - r * n;
        if (c < r)            D[(long)r * d_srow + c] = a[r * ld + c];
        else if (c == r)      D[(long)r * d_srow + c] = dia[r];
        else if (zero_upper)  D[(long)r * d_srow + c] = 0.0f;
    }
}

void bchol(torch::Tensor dst, torch::Tensor src, bool zero_upper)
{
    TORCH_CHECK(dst.is_cuda() && src.is_cuda(), "bchol: CUDA tensors required");
    TORCH_CHECK(dst.scalar_type() == at::kFloat && src.scalar_type() == at::kFloat,
                "bchol: float32 required");
    TORCH_CHECK(dst.dim() == 3 && src.dim() == 3, "bchol: (B,n,n) required");
    const long batch = dst.size(0);
    const long n_l   = dst.size(1);
    TORCH_CHECK(src.size(0) == batch && src.size(1) == n_l &&
                src.size(2) == n_l && dst.size(2) == n_l, "bchol: shape mismatch");
    TORCH_CHECK(n_l >= 1 && n_l <= 128, "bchol: 1 <= n <= 128 required");
    TORCH_CHECK(dst.stride(2) == 1 && src.stride(2) == 1,
                "bchol: innermost stride must be 1");
    if (batch == 0) return;

    const int n = (int)n_l;
    int ld = n + 1;
    if ((ld & 1) == 0) ld += 1;              
    const int T = ((n + 31) / 32) * 32;
    const size_t smem = ((size_t)n * ld + n) * sizeof(float);
    if (smem > 48 * 1024) {
        static size_t smem_set = 0;
        if (smem > smem_set) {
            cudaError_t err = cudaFuncSetAttribute(
                bchol_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
            TORCH_CHECK(err == cudaSuccess, "bchol: shared memory opt-in failed");
            smem_set = smem;
        }
    }
    const at::cuda::OptionalCUDAGuard guard(at::device_of(dst));
    bchol_kernel<<<dim3((unsigned int)batch), dim3(T), smem>>>(
        dst.data_ptr<float>(), src.data_ptr<float>(),
        batch, n,
        dst.stride(1), dst.stride(0),
        src.stride(1), src.stride(0),
        ld, zero_upper ? 1 : 0);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}
"""


def _base_chol(V, zero_upper):
    if _EXT is not None:
        _EXT.bchol(V, V, zero_upper)
    else:
        V.copy_(torch.linalg.cholesky_ex(V, check_errors=False).L)


def _panel_nb(n):
    if n <= 2048:
        return _BASE
    if n <= 4096:
        return 256
    return 512


def _probe_trsm_out(device):
    try:
        g = torch.Generator(device="cpu").manual_seed(0)
        L = torch.randn(2, 5, 5, generator=g).tril() + 5.0 * torch.eye(5)
        Bf = torch.randn(2, 7, 9, generator=g)
        L, Bf = L.to(device), Bf.to(device)
        B = Bf[:, :, 2:7]
        U = L.transpose(-1, -2)
        ref = torch.linalg.solve_triangular(U, B, upper=True, left=False)
        torch.linalg.solve_triangular(U, B, upper=True, left=False, out=B)
        return bool(torch.allclose(B, ref, rtol=1e-4, atol=1e-5))
    except Exception:
        return False


def _trsm_right(L11, B):
    global _TRSM_OUT
    U = L11.transpose(-1, -2)
    if _TRSM_OUT is None:
        _TRSM_OUT = _probe_trsm_out(B.device)
    if _TRSM_OUT:
        torch.linalg.solve_triangular(U, B, upper=True, left=False, out=B)
    else:
        B.copy_(torch.linalg.solve_triangular(U, B, upper=True, left=False))


def _update_trailing(A22, P):
    m = A22.size(-1)
    Pt = P.transpose(-1, -2)
    if m <= 256:
        A22.baddbmm_(P, Pt, alpha=-1.0)
        return
    if m <= 1024:
        cb = 256
    elif m <= 4096:
        cb = 512
    elif m <= 16384:
        cb = 1024
    else:
        cb = 2048
    for c0 in range(0, m, cb):
        c1 = min(c0 + cb, m)
        A22[:, c0:, c0:c1].baddbmm_(P[:, c0:, :], Pt[:, :, c0:c1], alpha=-1.0)


def _chol_ip(A):
    n = A.size(-1)
    if n <= _BASE:
        _base_chol(A, False)
        return
    nb0 = _panel_nb(n)
    j = 0
    while j < n:
        nb = min(nb0, n - j)
        k1 = j + nb
        D = A[:, j:k1, j:k1]
        if nb <= _BASE:
            _base_chol(D, False)
        else:
            _chol_ip(D)
        if k1 < n:
            P = A[:, k1:, j:k1]
            _trsm_right(D, P)
            _update_trailing(A[:, k1:, k1:], P)
        j = k1


def _run_custom(data):
    n = data.size(-1)
    if n <= _BASE and _EXT is not None and data.stride(-1) == 1:
        L = torch.empty(data.shape, dtype=data.dtype, device=data.device)
        _EXT.bchol(L, data, True)
        return L
    A = data.contiguous()
    L = A.clone() if A is data else A
    _chol_ip(L)
    return L.tril_()


def _run_torch(data):
    return torch.linalg.cholesky_ex(data, check_errors=False).L


def _quick_ok(A, L):
    try:
        if not bool(torch.isfinite(L).all()):
            return False
        n = A.size(-1)
        s = min(n, 128)
        idx = torch.randperm(n, device=A.device)[:s]
        rec = torch.matmul(L[:, idx, :], L.transpose(-1, -2))
        num = (rec - A[:, idx, :]).norm()
        den = A[:, idx, :].norm().clamp_min(1e-30)
        return bool((num / den) < 2e-3)
    except Exception:
        return False


def _select(A):
    try:
        Lc = _run_custom(A)
        ok = _quick_ok(A, Lc)
        del Lc
        if ok:
            return _run_custom
    except Exception:
        pass
    return _run_torch


def custom_kernel(data: input_t) -> output_t:
    A = data
    if not torch.is_tensor(A):
        return _run_torch(A)
    if A.dim() == 2:
        return custom_kernel(A.unsqueeze(0)).squeeze(0)
    if A.dim() != 3 or (not A.is_cuda) or A.dtype != torch.float32:
        return _run_torch(A)
    key = (A.size(0), A.size(1), A.device.index)
    fn = _PLAN.get(key)
    if fn is None:
        fn = _select(A)
        _PLAN[key] = fn
    return fn(A)


def _init():
    global _EXT, _BASE
    if not torch.cuda.is_available():
        return
    try:
        from torch.utils.cpp_extension import load_inline
        ext = load_inline(
            name="bchol_ext_v1",
            cpp_sources=_CPP_DECL,
            cuda_sources=_CUDA_SRC,
            functions=["bchol"],
            extra_cuda_cflags=["-O3"],
            verbose=False,
        )
    except Exception:
        return
    dev = torch.device("cuda", torch.cuda.current_device())
    for base in (128, 96):
        try:
            g = torch.Generator(device=dev).manual_seed(0)
            R = torch.randn(3, base, base, device=dev, generator=g)
            A = R @ R.transpose(-1, -2) + base * torch.eye(base, device=dev)
            L = torch.empty_like(A)
            ext.bchol(L, A, True)
            Ld = L.double()
            resid = (Ld @ Ld.transpose(-1, -2) - A.double()).norm() / A.double().norm()
            ok = (bool(torch.isfinite(L).all())
                  and float(resid) < 1e-4
                  and float(L.triu(1).abs().max()) == 0.0
                  and bool((L.diagonal(dim1=-2, dim2=-1) > 0).all()))
            if ok:
                _EXT, _BASE = ext, base
                return
        except Exception:
            continue


_init()
scrolls · 301 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