Skip to content
KernelIndex
Search⌘K

submission 893427

d_lolo_ · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-893427?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.81ms
#231 of 337
2026-07-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b2f6a35fdf6f46b42bff3b883a4630875bf509d38c64964b5c0c3792b057a9e7
license declaredunknown
license concludedunknown
authorsd_lolo_
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float s[]; // N*(N+1)/2 floats

Kernel source

submission.py356 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

"""Batched dense Cholesky (A = L L^T, fp32 SPD -> lower-triangular L).

Per-shape routing to the fastest CORRECT kernel, with a cuSOLVER
(torch.linalg.cholesky_ex) fallback that can never regress or fail:

  n in {32, 64}                      -> fused warp-per-matrix register-resident
                                        CUDA kernel (fp32). One CUDA warp
                                        factorizes one matrix, kills cuSOLVER's
                                        per-matrix launch overhead.
  n == 128                           -> fused one-CTA-per-matrix blocked Cholesky
                                        with a PACKED lower triangle in shared
                                        memory (fp32).
  n in {1024, 2048} AND batch >= 48  -> blocked right-looking Cholesky with a
                                        TF32 tensor-core trailing GEMM (the
                                        O(n^3/3) Schur update is the bulk), fp32
                                        potf2/trsm. Guarded: falls back to
                                        cuSOLVER if TF32 pushes any diagonal
                                        block indefinite (near-singular inputs).
  everything else                    -> cuSOLVER baseline.

Self-contained: CUDA sources are inlined and built lazily via cpp_extension
(so importing this file is cheap and a build failure for one kernel does not
break the others). Every specialized path is wrapped in try/except that falls
back to the cuSOLVER baseline.
"""

import os

# Only needed on dev boxes where CUDA headers/ptxas aren't on the default path.
# No-op on a properly configured runner (these paths won't exist there).
_dev_ptxas = "/usr/local/cuda-13.0/bin/ptxas"
_dev_inc = "/usr/local/cuda-13.0/targets/sbsa-linux/include"
_dev_home = "/usr/local/cuda-13.0"
if os.path.exists(_dev_ptxas):
    os.environ.setdefault("TRITON_PTXAS_PATH", _dev_ptxas)
if os.path.isdir(_dev_home):
    os.environ.setdefault("CUDA_HOME", _dev_home)
if os.path.isdir(_dev_inc) and _dev_inc not in os.environ.get("CPATH", ""):
    os.environ["CPATH"] = _dev_inc + ":" + os.environ.get("CPATH", "")

import torch

from task import input_t, output_t


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


# --------------------------------------------------------------------------- #
# Extension 1: warp-per-matrix register-resident Cholesky for n in {32, 64}   #
# --------------------------------------------------------------------------- #
_WARP_CUDA = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>

// Warp-per-matrix, register-resident right-looking Cholesky.
// N = matrix dim, CPL = columns per lane = N/32.
template<int N, int CPL>
__global__ void chol_warp(const float* __restrict__ A,
                          float* __restrict__ Lout, int batch) {
    const int warp = (blockIdx.x * blockDim.x + threadIdx.x) >> 5;
    const int lane = threadIdx.x & 31;
    if (warp >= batch) return;
    const float* Ab = A + (size_t)warp * N * N;
    float* Lb = Lout + (size_t)warp * N * N;

    float col[CPL][N];
    #pragma unroll
    for (int c = 0; c < CPL; c++) {
        const int j = lane + c * 32;
        #pragma unroll
        for (int i = 0; i < N; i++) col[c][i] = Ab[i * N + j];
    }

    #pragma unroll
    for (int k = 0; k < N; k++) {
        const int kc = k >> 5;      // column-group holding pivot col k
        const int kl = k & 31;      // lane holding it
        const float akk = __shfl_sync(0xffffffffu, col[kc][k], kl);
        const float d = sqrtf(fmaxf(akk, 1e-30f));
        const float inv = 1.0f / d;
        float Lself[CPL];
        #pragma unroll
        for (int c = 0; c < CPL; c++) Lself[c] = 0.0f;
        #pragma unroll
        for (int i = 0; i < N; i++) {
            const float raw = __shfl_sync(0xffffffffu, col[kc][i], kl);
            const float Lki = (i == k) ? d : (i > k ? raw * inv : 0.0f);
            #pragma unroll
            for (int c = 0; c < CPL; c++) {
                const int j = lane + c * 32;
                if (i == j) Lself[c] = Lki;
                if (j > k && i >= j) col[c][i] -= Lki * Lself[c];
                if (j == k) col[c][i] = Lki;     // finalize pivot column
            }
        }
    }

    #pragma unroll
    for (int c = 0; c < CPL; c++) {
        const int j = lane + c * 32;
        #pragma unroll
        for (int i = 0; i < N; i++) Lb[i * N + j] = (i >= j) ? col[c][i] : 0.0f;
    }
}

torch::Tensor chol_warp_launch(torch::Tensor A, int64_t warps_per_block) {
    const int batch = A.size(0);
    const int n = A.size(1);
    auto L = torch::empty_like(A);
    const int threads = (int)warps_per_block * 32;
    const int blocks = (batch + (int)warps_per_block - 1) / (int)warps_per_block;
    const float* a = A.data_ptr<float>();
    float* l = L.data_ptr<float>();
    if (n == 32)
        chol_warp<32, 1><<<blocks, threads>>>(a, l, batch);
    else if (n == 64)
        chol_warp<64, 2><<<blocks, threads>>>(a, l, batch);
    else
        TORCH_CHECK(false, "chol_warp supports n in {32,64}");
    return L;
}
'''
_WARP_CPP = "torch::Tensor chol_warp_launch(torch::Tensor A, int64_t warps_per_block);"
_WPB = {32: 1, 64: 4}   # warps per block, tuned per size


# --------------------------------------------------------------------------- #
# Extension 2: blocked packed-shared-memory Cholesky for n == 128             #
# --------------------------------------------------------------------------- #
_SMEM_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <math.h>

// BLOCKED + PACKED lower triangle (small shared -> high occupancy).
// Packed index p(i,j) = i*(i+1)/2 + j (valid for i>=j). n=128 packed = 33 KB.
// One CTA per matrix; right-looking with potf2 (warp 0) + trsm + syrk.
#define PIDX(i, j) (((i) * ((i) + 1)) >> 1) + (j)
template<int N, int NB>
__global__ void chol_blk_packed_kernel(const float* __restrict__ A,
                                       float* __restrict__ L, int batch) {
    const int m = blockIdx.x;
    if (m >= batch) return;
    extern __shared__ float s[];              // N*(N+1)/2 floats
    const float* Am = A + (long long)m * N * N;
    float* Lm = L + (long long)m * N * N;
    const int tid = threadIdx.x;
    const int nt = blockDim.x;
    const int warp = tid >> 5;
    const int lane = tid & 31;
    const int P = N * (N + 1) / 2;

    for (int p = tid; p < P; p += nt) {
        int i = (int)((sqrtf(8.0f * p + 1.0f) - 1.0f) * 0.5f);
        while ((long long)(i + 1) * (i + 2) / 2 <= p) ++i;
        while ((long long)i * (i + 1) / 2 > p) --i;
        int j = p - ((i * (i + 1)) >> 1);
        s[p] = Am[(long long)i * N + j];
    }
    __syncthreads();

    for (int pc = 0; pc < N; pc += NB) {
        if (warp == 0) {
            for (int c = 0; c < NB; ++c) {
                int dcc = PIDX(pc + c, pc + c);
                float d = sqrtf(s[dcc]);
                if (lane == 0) s[dcc] = d;
                for (int r = lane; r < NB; r += 32)
                    if (r > c) s[PIDX(pc + r, pc + c)] /= d;
                __syncwarp();
                for (int r = lane; r < NB; r += 32) {
                    if (r > c) {
                        float lic = s[PIDX(pc + r, pc + c)];
                        for (int j = c + 1; j <= r; ++j)
                            s[PIDX(pc + r, pc + j)] -= lic * s[PIDX(pc + j, pc + c)];
                    }
                }
                __syncwarp();
            }
        }
        __syncthreads();

        int r0 = pc + NB;
        int rows = N - r0;
        if (rows > 0) {
            for (int i = r0 + tid; i < N; i += nt) {
                int rowbase = (i * (i + 1)) >> 1;
                for (int c = 0; c < NB; ++c) {
                    float acc = s[rowbase + pc + c];
                    for (int cc = 0; cc < c; ++cc)
                        acc -= s[rowbase + pc + cc] * s[PIDX(pc + c, pc + cc)];
                    s[rowbase + pc + c] = acc / s[PIDX(pc + c, pc + c)];
                }
            }
            __syncthreads();

            int tot = rows * rows;
            for (int idx = tid; idx < tot; idx += nt) {
                int di = idx / rows;
                int dj = idx - di * rows;
                if (di >= dj) {
                    int i = r0 + di, j = r0 + dj;
                    int ibase = (i * (i + 1)) >> 1;
                    int jbase = (j * (j + 1)) >> 1;
                    float acc = 0.0f;
                    #pragma unroll
                    for (int c = 0; c < NB; ++c)
                        acc += s[ibase + pc + c] * s[jbase + pc + c];
                    s[ibase + j] -= acc;
                }
            }
            __syncthreads();
        }
    }
    for (int idx = tid; idx < N * N; idx += nt) {
        int i = idx / N, j = idx - i * N;
        Lm[idx] = (i >= j) ? s[((i * (i + 1)) >> 1) + j] : 0.0f;
    }
}

void launch_blkpacked128(torch::Tensor A, torch::Tensor L) {
    int batch = A.size(0);
    const int N = 128;
    size_t shmem = (size_t)N * (N + 1) / 2 * sizeof(float);   // 33 KB -> 3 blk/SM
    cudaFuncSetAttribute(chol_blk_packed_kernel<128, 8>,
        cudaFuncAttributeMaxDynamicSharedMemorySize, shmem);
    chol_blk_packed_kernel<128, 8><<<batch, 256, shmem>>>(
        A.data_ptr<float>(), L.data_ptr<float>(), batch);
}
"""
_SMEM_CPP = "void launch_blkpacked128(torch::Tensor A, torch::Tensor L);"


# --------------------------------------------------------------------------- #
# Lazy extension builders (cached in module globals)                          #
# --------------------------------------------------------------------------- #
_EXT_WARP = None       # None = not tried; False = build failed; else module
_EXT_SMEM = None


def _get_warp():
    global _EXT_WARP
    if _EXT_WARP is None:
        try:
            from torch.utils.cpp_extension import load_inline
            _EXT_WARP = load_inline(
                name="sub_warpshfl",
                cpp_sources=_WARP_CPP,
                cuda_sources=_WARP_CUDA,
                functions=["chol_warp_launch"],
                extra_cuda_cflags=["-O3", "--use_fast_math"],
                verbose=False,
            )
        except Exception:
            _EXT_WARP = False
    return _EXT_WARP


def _get_smem():
    global _EXT_SMEM
    if _EXT_SMEM is None:
        try:
            from torch.utils.cpp_extension import load_inline
            _EXT_SMEM = load_inline(
                name="sub_smemblk",
                cpp_sources=_SMEM_CPP,
                cuda_sources=_SMEM_CUDA,
                functions=["launch_blkpacked128"],
                extra_cuda_cflags=["-O3"],
                verbose=False,
            )
        except Exception:
            _EXT_SMEM = False
    return _EXT_SMEM


# --------------------------------------------------------------------------- #
# tf32-blocked path for high-batch n in {1024, 2048}                          #
# --------------------------------------------------------------------------- #
def _blocked_tf32(A, bs):
    """Right-looking blocked Cholesky, TF32 tensor-core trailing GEMM, fp32
    potf2/trsm. Returns (L, bad) where bad > 0 iff a diagonal block lost
    positive-definiteness (the TF32 breakdown signal on near-singular inputs)."""
    B, n, _ = A.shape
    A = A.clone()
    L = torch.zeros_like(A)
    bad = None
    for j in range(0, n, bs):
        je = min(j + bs, n)
        res = torch.linalg.cholesky_ex(A[:, j:je, j:je], check_errors=False)
        Ljj = res.L
        info = res.info
        bad = info.amax() if bad is None else torch.maximum(bad, info.amax())
        L[:, j:je, j:je] = Ljj
        if je < n:
            A21 = A[:, je:, j:je]
            L21 = torch.linalg.solve_triangular(
                Ljj.transpose(-1, -2), A21, upper=True, left=False
            )
            L[:, je:, j:je] = L21
            A[:, je:, je:] -= L21 @ L21.transpose(-1, -2)
    return L, bad


def _tf32_path(data, n):
    bs = 256 if n == 1024 else 512
    old = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        L, bad = _blocked_tf32(data, bs=bs)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old
    # Robustness guard: bad>0 iff TF32 pushed a diagonal block indefinite.
    if bad is not None and int(bad.item()) == 0:
        return L
    return None  # signal caller to fall back


# --------------------------------------------------------------------------- #
# Router                                                                       #
# --------------------------------------------------------------------------- #
def custom_kernel(data: input_t) -> output_t:
    # Only handle the exact contract; otherwise cuSOLVER baseline.
    if not (isinstance(data, torch.Tensor) and data.is_cuda
            and data.dtype == torch.float32 and data.dim() == 3
            and data.shape[-1] == data.shape[-2]):
        return _baseline(data)

    B, n, _ = data.shape

    try:
        if n in (32, 64):
            ext = _get_warp()
            if ext:
                return ext.chol_warp_launch(data.contiguous(), _WPB[n])
        elif n == 128:
            ext = _get_smem()
            if ext:
                A = data.contiguous()
                L = torch.empty_like(A)
                ext.launch_blkpacked128(A, L)
                return L
        elif n in (1024, 2048) and B >= 48:
            L = _tf32_path(data, n)
            if L is not None:
                return L
    except Exception:
        pass

    return _baseline(data)
scrolls · 356 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