Skip to content
KernelIndex
Search⌘K

submission 914616

:Dev · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-914616?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.22ms
#139 of 337
2026-07-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:312deb723b4a93b9073c846a2bf297aa251f771129a55746ebeb1476b06dd579
license declaredunknown
license concludedunknown
authors:Dev
imported2026-08-26

Techniques

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

shared-memory__shared__ float buf[2][64];

Kernel source

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

"""Batched dense Cholesky, A = L @ L.T.

Three regimes, because the grid spans five orders of magnitude in batch:

n <= 128  one matrix per warp / per CTA, resident in registers for the whole
          factorisation. These entries hold 2^22 FP32 elements each, a ~4 us
          HBM floor against a ~4 us launch, so they are latency-bound and the
          only thing that matters is that the matrix never leaves the SM.

n >= 256  blocked right-looking. The matrix no longer fits in the register
          file (n=256 alone is 256 KB, the whole file), so the factorisation
          becomes panel + rank-b update and the update goes to cuBLAS.

The rest is picking, per grid entry, whichever of those beats cuSOLVER --
cuSOLVER is genuinely strong on the large single-matrix entries (53 TF/s at
n=32768) and genuinely weak on the batched ones (1-8 TF/s).

Register-resident kernels, n <= 128
-----------------------------------
Thread t owns COLUMN t, not row t: loads coalesce, and -- because the trailing
matrix is kept symmetric -- the multiplier M[j][t] is the thread's OWN element
j rather than something it has to fetch. The broadcast column M[i][j] is then
M[j][i], i.e. thread i's own element j, so publishing it costs each thread a
single store instead of a gather. n=32 does that broadcast with a shuffle and
needs no barrier at all.

Sizing is set by occupancy, not by work: at n=64, two columns per thread needs
~130 registers and MEASURED 80 us, while one column per thread needs 64 and
MEASURED 17 us. Fewer registers beat fewer instructions.
"""

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

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

#define FULL 0xffffffffu

// n == 32: one warp per matrix, no barriers -- the column broadcast is a shfl.
__global__ __launch_bounds__(128, 4)
void chol32(const float* __restrict__ A, float* __restrict__ L, int batch) {
    const int t   = threadIdx.x & 31;
    const int wid = (blockIdx.x * blockDim.x + threadIdx.x) >> 5;
    if (wid >= batch) return;
    const float* __restrict__ a = A + (size_t)wid * 1024;
    float c[32];
    #pragma unroll
    for (int i = 0; i < 32; ++i) c[i] = a[i * 32 + t];

    // Unscaled LDL^T sweep: M[i][k] -= M[i][j]*M[j][k]/d_j, with the sqrt
    // folded into the final store, so no column rescale per step.
    float mypiv = 1.0f;
    #pragma unroll
    for (int j = 0; j < 32; ++j) {
        float piv = __shfl_sync(FULL, c[j], j);
        if (t == j) mypiv = piv;
        float w = (t > j) ? c[j] / piv : 0.0f;   // 0 keeps finished columns put
        #pragma unroll
        for (int i = j + 1; i < 32; ++i)
            c[i] = fmaf(-__shfl_sync(FULL, c[i], j), w, c[i]);
    }
    const float s = rsqrtf(fmaxf(mypiv, 1e-30f));
    float* __restrict__ o = L + (size_t)wid * 1024;
    #pragma unroll
    for (int i = 0; i < 32; ++i) o[i * 32 + t] = (i >= t) ? c[i] * s : 0.0f;
}

// n == 64: two warps per matrix, one column per thread. Double buffered, so
// one barrier per step.
__global__ __launch_bounds__(64, 8)
void chol64(const float* __restrict__ A, float* __restrict__ L, int batch) {
    __shared__ float buf[2][64];
    const int t = threadIdx.x, wid = blockIdx.x;
    const float* __restrict__ a = A + (size_t)wid * 4096;
    float c[64];
    #pragma unroll
    for (int i = 0; i < 64; ++i) c[i] = a[i * 64 + t];
    float mypiv = 1.0f;
    #pragma unroll
    for (int j = 0; j < 64; ++j) {
        buf[j & 1][t] = c[j];                    // == M[j][t] == M[t][j]
        __syncthreads();
        const float d = buf[j & 1][j];
        if (t == j) mypiv = d;
        const float w = (t > j) ? c[j] / d : 0.0f;
        #pragma unroll
        for (int i = j + 1; i < 64; ++i)
            c[i] = fmaf(-buf[j & 1][i], w, c[i]);
    }
    const float s = rsqrtf(fmaxf(mypiv, 1e-30f));
    float* __restrict__ o = L + (size_t)wid * 4096;
    #pragma unroll
    for (int i = 0; i < 64; ++i) o[i * 64 + t] = (i >= t) ? c[i] * s : 0.0f;
}

// n == 128. Fully unrolling the j-loop is what lets a register array be
// statically indexed, but at n=128 that is ~24K instructions of straight-line
// code. So the j-loop stays rolled and the row loop is unrolled in static
// chunks of CH: indices stay compile-time, code stays O(R), and chunks
// entirely above the pivot are skipped by a runtime branch. The one dynamic
// access left -- reading register c[j+1] to publish it -- is a two-level
// select costing ~sqrt(R) instead of R.
template <int N, int H, int CH>
__global__ __launch_bounds__(N * H, 2)
void chol_reg(const float* __restrict__ A, float* __restrict__ L, int batch) {
    constexpr int R = N / H;
    __shared__ float buf[2][N + 1];              // [.][N] carries the pivot
    const int tid = threadIdx.x;
    const int t   = tid & (N - 1);
    const int h   = tid / N;
    const int r0  = h * R;
    const int wid = blockIdx.x;

    const float* __restrict__ a = A + (size_t)wid * N * N;
    float c[R];
    #pragma unroll
    for (int i = 0; i < R; ++i) c[i] = a[(size_t)(r0 + i) * N + t];

    if (h == 0) {
        buf[0][t] = (t > 0) ? c[0] : 0.0f;
        if (t == 0) buf[0][N] = c[0];
    }

    float mypiv = 1.0f;
    for (int j = 0; j < N; ++j) {
        __syncthreads();
        const int p = j & 1;
        const float d = buf[p][N];
        if (t == j) mypiv = d;
        const float w = buf[p][t] * __frcp_rn(d);
        const int target = j + 1 - r0;
        float mynext = 0.0f;
        #pragma unroll
        for (int cb = 0; cb < R / CH; ++cb) {
            if (r0 + cb * CH + CH <= j) continue;
            #pragma unroll
            for (int q = 0; q < CH; ++q) {
                const int i = cb * CH + q;
                c[i] = fmaf(-buf[p][r0 + i], w, c[i]);   // buf is 0 above pivot
            }
            if (target >= cb * CH && target < cb * CH + CH) {
                #pragma unroll
                for (int q = 0; q < CH; ++q)
                    if (cb * CH + q == target) mynext = c[cb * CH + q];
            }
        }
        // written to the buffer NOT being read this step, so no second barrier
        if ((unsigned)target < (unsigned)R) {
            buf[p ^ 1][t] = (t > j + 1) ? mynext : 0.0f;
            if (t == j + 1) buf[p ^ 1][N] = mynext;
        }
    }

    const float s = rsqrtf(fmaxf(mypiv, 1e-30f));
    float* __restrict__ o = L + (size_t)wid * N * N;
    #pragma unroll
    for (int i = 0; i < R; ++i)
        o[(size_t)(r0 + i) * N + t] = (r0 + i >= t) ? c[i] * s : 0.0f;
}

torch::Tensor chol_small(torch::Tensor A) {
    const int n = A.size(-1);
    auto Af = A.contiguous().reshape({-1, n, n});
    const int batch = Af.size(0);
    auto L = torch::empty_like(Af);
    const float* ap = Af.data_ptr<float>();
    float* lp = L.data_ptr<float>();
    if (n == 32)       chol32<<<(batch + 3) / 4, 128>>>(ap, lp, batch);
    else if (n == 64)  chol64<<<batch, 64>>>(ap, lp, batch);
    else if (n == 128) chol_reg<128, 2, 16><<<batch, 256>>>(ap, lp, batch);
    return L;
}

// 3xTF32 K-expansion. TF32 keeps 10 mantissa bits against FP32's 23, and a
// plain TF32 trailing update does not degrade gracefully -- it diverges on the
// damped low-rank family. Splitting X = Xh + Xl and keeping the three largest
// cross terms restores FP32 accuracy. Emitting them as ONE gemm with 3x the K
// extent, [Xh Xh Xl] @ [Xh Xl Xh]^T, costs one pass over the trailing
// submatrix instead of three accumulating ones.
__global__ void split3(const float* __restrict__ X, float* __restrict__ P,
                       float* __restrict__ Q, long long rows, int k) {
    const long long i = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    if (i >= rows * (long long)k) return;
    const long long r = i / k;
    const int       c = (int)(i - r * k);
    const float v = X[i];
    unsigned u = __float_as_uint(v);
    const float hi = __uint_as_float((u + 0x1000u) & 0xFFFFE000u);
    const float lo = v - hi;
    const long long base = r * (3LL * k) + c;
    P[base] = hi;  P[base + k] = hi;  P[base + 2LL * k] = lo;
    Q[base] = hi;  Q[base + k] = lo;  Q[base + 2LL * k] = hi;
}

std::vector<torch::Tensor> split3xtf32(torch::Tensor X) {
    auto Xc = X.contiguous();
    const int b = Xc.size(0), m = Xc.size(1), k = Xc.size(2);
    auto P = torch::empty({b, m, 3 * k}, Xc.options());
    auto Q = torch::empty({b, m, 3 * k}, Xc.options());
    const long long rows = (long long)b * m;
    const long long tot  = rows * k;
    const int thr = 256;
    split3<<<(int)((tot + thr - 1) / thr), thr>>>(
        Xc.data_ptr<float>(), P.data_ptr<float>(), Q.data_ptr<float>(), rows, k);
    return {P, Q};
}
"""

CPP_SRC = """
torch::Tensor chol_small(torch::Tensor A);
std::vector<torch::Tensor> split3xtf32(torch::Tensor X);
"""

_mod = load_inline(
    name="chol_final",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["chol_small", "split3xtf32"],
    extra_cuda_cflags=["-O3", "--use_fast_math"],
    verbose=False,
)

_FAST = (32, 64, 128)
_EYE = {}


def _tf32(on):
    """torch moved this knob across versions; set whichever exists."""
    try:
        torch.backends.cuda.matmul.fp32_precision = "tf32" if on else "ieee"
    except Exception:
        pass
    try:
        torch.backends.cuda.matmul.allow_tf32 = on
    except Exception:
        pass


def _eye(bs, batch, dev):
    e = _EYE.get((bs, batch, dev))
    if e is None:
        e = torch.eye(bs, device=dev, dtype=torch.float32)
        e = e.expand(batch, bs, bs).contiguous()
        _EYE[(bs, batch, dev)] = e
    return e


def _blocked(A, bs, tri, mode):
    """Right-looking blocked Cholesky; diagonal blocks on the custom kernel.

    mode "solve"  -- panel by TRSM, trailing update as one FP32 gemm over the
                     full symmetric square. Fewest launches, so it wins on the
                     low-batch entries, which are launch-bound.
    mode "3xtf32" -- panel by explicit inverse + gemm, trailing update split
                     into lower block columns on tensor cores. Wins once there
                     is enough work per launch to be flop-bound.
    """
    n = A.shape[-1]
    M = A.clone()
    for j0 in range(0, n, bs):
        j1 = min(j0 + bs, n)
        b = j1 - j0
        L11 = _mod.chol_small(M[:, j0:j1, j0:j1].contiguous())
        M[:, j0:j1, j0:j1] = L11
        if j1 >= n:
            break

        if mode == "solve":
            L21 = torch.linalg.solve_triangular(
                L11.transpose(-1, -2), M[:, j1:, j0:j1], upper=True, left=False)
            M[:, j1:, j0:j1] = L21
            M[:, j1:, j1:] -= L21 @ L21.transpose(-1, -2)
            continue

        Y = torch.linalg.solve_triangular(
            L11, _eye(bs, A.shape[0], A.device)[:, :b, :b], upper=False)
        L21 = M[:, j1:, j0:j1] @ Y.transpose(-1, -2)
        M[:, j1:, j0:j1] = L21
        P, Q = _mod.split3xtf32(L21)
        _tf32(True)
        for c0 in range(j1, n, tri):
            c1 = min(c0 + tri, n)
            M[:, c0:, c0:c1].baddbmm_(
                P[:, c0 - j1:, :], Q[:, c0 - j1:c1 - j1, :].transpose(-1, -2),
                beta=1.0, alpha=-1.0)
        _tf32(False)
    return M.tril_()


# Per grid entry, whichever path MEASURED fastest on B200. cuSOLVER is the
# fallback and still wins the large single-matrix entries outright.
#   (n, batch) -> (block, trailing block-column width, mode)
_CFG = {
    (256, 64): (128, 4096, "solve"),
    (512, 16): (128, 4096, "solve"),
    (1024, 60): (128, 512, "3xtf32"),
    (16384, 1): (128, 8192, "3xtf32"),
    (32768, 1): (128, 8192, "3xtf32"),
}

# cuSOLVER's potrfBatched is pathological at large n and tiny batch: MEASURED
# n=4096 b=2 at 11.2 ms batched against 3.20 ms looping potrf twice. batch==1
# is excluded -- it is already on cuSOLVER's good path and the stack is pure
# overhead.
_LOOP = {(2048, 2), (4096, 2)}


def custom_kernel(data: input_t) -> output_t:
    A = data[0] if isinstance(data, (tuple, list)) else data
    n = A.shape[-1]

    if A.is_cuda and A.dtype == torch.float32 and n in _FAST:
        return _mod.chol_small(A).reshape(A.shape)

    if A.is_cuda and A.dtype == torch.float32 and A.dim() == 3:
        key = (n, A.shape[0])
        cfg = _CFG.get(key)
        if cfg is not None:
            return _blocked(A, cfg[0], cfg[1], cfg[2])
        if key in _LOOP:
            return torch.stack([
                torch.linalg.cholesky_ex(A[i], check_errors=False).L
                for i in range(A.shape[0])
            ])

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