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
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