Skip to content
KernelIndex
Search⌘K

submission 916068

monish devineni · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

cholesky_hybrid_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-916068?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
1.25ms
#144 of 337
2026-07-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ad393fcf8b6ebffc8292368f7376273134ddbe24b898efda275ba4f796e47981
license declaredunknown
license concludedunknown
authorsmonish devineni
imported2026-08-26

Techniques

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

shared-memory__shared__ float smem[MPB][N * S];
vector-width = float4const float4 *rt = reinterpret_cast<const float4 *>(s + t * S);

Kernel source

cholesky_hybrid_submission.py224 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

import sys
import torch
from torch.utils.cpp_extension import load_inline

cuda_src = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>

template <int N, int MPB>
__global__ void chol_warp(const float *__restrict__ A, float *__restrict__ L, int batch) {
    constexpr int S = N + 4;
    __shared__ float smem[MPB][N * S];
    __shared__ float sdiag[MPB][N];

    const int t = threadIdx.x;
    const int w = threadIdx.y;
    const int m = blockIdx.x * MPB + w;
    if (m >= batch) return;

    float *s = smem[w];
    float *sd = sdiag[w];
    const float *Ab = A + (size_t)m * N * N;
    float *Lb = L + (size_t)m * N * N;

    #pragma unroll 8
    for (int r = 0; r < N; ++r) s[r * S + t] = Ab[r * N + t];
    __syncwarp();

    const float4 *rt = reinterpret_cast<const float4 *>(s + t * S);
    for (int j = 0; j < N; ++j) {
        float acc = 0.0f;
        if (t >= j) {
            const float4 *rj = reinterpret_cast<const float4 *>(s + j * S);
            float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
            const int q = j >> 2;
            for (int g = 0; g < q; ++g) {
                float4 u = rt[g], v = rj[g];
                a0 = fmaf(-u.x, v.x, a0);
                a1 = fmaf(-u.y, v.y, a1);
                a2 = fmaf(-u.z, v.z, a2);
                a3 = fmaf(-u.w, v.w, a3);
            }
            acc = s[t * S + j] + ((a0 + a1) + (a2 + a3));
            for (int p = q << 2; p < j; ++p) acc -= s[t * S + p] * s[j * S + p];
        }
        if (t == j) sd[j] = sqrtf(acc);
        __syncwarp();
        if (t > j) s[t * S + j] = acc * (1.0f / sd[j]);
        __syncwarp();
    }

    #pragma unroll 8
    for (int r = 0; r < N; ++r)
        Lb[r * N + t] = (t < r) ? s[r * S + t] : ((t == r) ? sd[r] : 0.0f);
}

template <int N>
__global__ void chol_block(const float *__restrict__ A, float *__restrict__ L) {
    constexpr int S = N + 4;
    extern __shared__ float raw[];
    float *s = raw;
    float *sd = raw + (size_t)N * S;

    const int t = threadIdx.x;
    const float *Ab = A + (size_t)blockIdx.x * N * N;
    float *Lb = L + (size_t)blockIdx.x * N * N;

    #pragma unroll 8
    for (int r = 0; r < N; ++r) s[r * S + t] = Ab[(size_t)r * N + t];
    __syncthreads();

    const float4 *rt = reinterpret_cast<const float4 *>(s + t * S);
    for (int j = 0; j < N; ++j) {
        float acc = 0.0f;
        if (t >= j) {
            const float4 *rj = reinterpret_cast<const float4 *>(s + j * S);
            float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
            const int q = j >> 2;
            for (int g = 0; g < q; ++g) {
                float4 u = rt[g], v = rj[g];
                a0 = fmaf(-u.x, v.x, a0);
                a1 = fmaf(-u.y, v.y, a1);
                a2 = fmaf(-u.z, v.z, a2);
                a3 = fmaf(-u.w, v.w, a3);
            }
            acc = s[t * S + j] + ((a0 + a1) + (a2 + a3));
            for (int p = q << 2; p < j; ++p) acc -= s[t * S + p] * s[j * S + p];
        }
        if (t == j) sd[j] = sqrtf(acc);
        __syncthreads();
        if (t > j) s[t * S + j] = acc * (1.0f / sd[j]);
        __syncthreads();
    }

    #pragma unroll 8
    for (int r = 0; r < N; ++r)
        Lb[(size_t)r * N + t] = (t < r) ? s[r * S + t] : ((t == r) ? sd[r] : 0.0f);
}

template <int N>
static void launch_block(const float *a, float *l, int batch) {
    const size_t bytes = ((size_t)N * (N + 4) + N) * sizeof(float);
    cudaFuncSetAttribute(chol_block<N>, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)bytes);
    chol_block<N><<<batch, N, bytes>>>(a, l);
}

torch::Tensor cholesky_small(torch::Tensor A) {
    TORCH_CHECK(A.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(A.scalar_type() == torch::kFloat32, "input must be float32");
    TORCH_CHECK(A.dim() == 3, "input must be (batch, n, n)");
    A = A.contiguous();

    const int batch = A.size(0);
    const int n = A.size(1);
    auto L = torch::empty_like(A);
    const float *a = A.data_ptr<float>();
    float *l = L.data_ptr<float>();

    if (n == 32) {
        constexpr int MPB = 8;
        chol_warp<32, MPB><<<(batch + MPB - 1) / MPB, dim3(32, MPB)>>>(a, l, batch);
    } else if (n == 64) {
        launch_block<64>(a, l, batch);
    } else if (n == 128) {
        launch_block<128>(a, l, batch);
    } else {
        TORCH_CHECK(false, "cholesky_small: unsupported n=", n);
    }

    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "cholesky kernel failed: ", cudaGetErrorString(err));
    return L;
}
"""

cpp_src = "torch::Tensor cholesky_small(torch::Tensor A);"

_extra_cuda_cflags = ["-Xcompiler", "/Zc:preprocessor"] if sys.platform == "win32" else []

_module = load_inline(
    name="cholesky_hybrid_cuda",
    cpp_sources=cpp_src,
    cuda_sources=cuda_src,
    functions=["cholesky_small"],
    extra_cuda_cflags=_extra_cuda_cflags,
    verbose=False,
)

_SMALL_N = (32, 64, 128)


def _chol_diag(A):
    if A.size(-1) in _SMALL_N:
        return _module.cholesky_small(A)
    return torch.linalg.cholesky_ex(A, check_errors=False).L


def _block_size(n):
    if n <= 256:
        return 128
    if n <= 2048:
        return 256
    if n <= 8192:
        return 512
    return 1024


def _blocked_chol(A, tf32=True):
    n = A.size(-1)
    B = _block_size(n)
    old = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = tf32
    try:
        L = torch.zeros_like(A)
        eye = None
        for k in range(0, n, B):
            kb = min(B, n - k)
            panel = A[:, k:, k:k + kb].clone()
            if k:
                panel.baddbmm_(L[:, k:, :k], L[:, k:k + kb, :k].transpose(-1, -2),
                               beta=1.0, alpha=-1.0)
            Lkk = _chol_diag(panel[:, :kb, :].contiguous())
            L[:, k:k + kb, k:k + kb] = Lkk
            if n - k - kb:
                if n >= 1024:
                    if eye is None or eye.size(-1) != kb:
                        eye = torch.eye(kb, device=A.device, dtype=A.dtype).expand(
                            A.size(0), kb, kb)
                    inv = torch.linalg.solve_triangular(Lkk, eye, upper=False)
                    L[:, k + kb:, k:k + kb] = panel[:, kb:, :] @ inv.transpose(-1, -2)
                else:
                    L[:, k + kb:, k:k + kb] = torch.linalg.solve_triangular(
                        Lkk.transpose(-1, -2), panel[:, kb:, :], upper=True, left=False)
        return L
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old


def custom_kernel(data):
    A = data
    if A.dim() != 3 or not A.is_cuda or A.dtype != torch.float32:
        return torch.linalg.cholesky_ex(A, check_errors=False).L

    b, n = A.size(0), A.size(-1)

    if n in _SMALL_N:
        return _module.cholesky_small(A)

    if n >= 8192 or (n >= 1024 and b >= 16):
        L = _blocked_chol(A)
        if bool((torch.diagonal(L, dim1=-2, dim2=-1) > 0).all()):
            return L

    if b > 1 and n >= 2048 and b <= 4:
        L = torch.empty_like(A)
        for i in range(b):
            L[i] = torch.linalg.cholesky_ex(A[i], check_errors=False).L
        return L

    return torch.linalg.cholesky_ex(A, check_errors=False).L
scrolls · 224 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