Skip to content
KernelIndex
Search⌘K

submission 888730

Rohith-Rongali · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-888730?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.59ms
#189 of 337
2026-07-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:da73db882a536f27017055a6133774c63bb18228ca0f236226bb3d1ec2fe7351
license declaredunknown
license concludedunknown
authorsRohith-Rongali
imported2026-08-26

Techniques

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

fused-epilogue- Fused fp16 / bf16x9 trailing SYRK via a cuBLASLt epilogue op
num-warps = 1_cholesky32_kernel[(data.shape[0],)](data, output, 32 * 32, num_warps=1)
shared-memoryextern __shared__ float smem[];

Kernel source

submission.py387 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
"""
Batched dense Cholesky: A = L @ L.T

Routing (all numbers measured on B200):
  - n in {32, 64}: custom CUDA kernel, one CTA per matrix, matrix resident in
    shared memory, in-place left-looking factorization with a parallelized
    diagonal. cuSOLVER's potrfBatched is far off the memory-bound floor
    (16 MiB in + out ~= 4 us) on these tiny matrices, so a single-launch smem
    kernel wins big: n=32 63 us (vs 113 cuSOLVER, 1.8x), n=64 89 us (1.2x).
    (n=128 also works but loses to cuSOLVER, occupancy-starved -> not routed.)
  - blocked right-looking Cholesky with a TF32 tensor-core trailing update for
    the shapes where it beats cuSOLVER: 512 high-batch, 1024 high-batch,
    2048 x<=2, 4096 x>=2, and every single matrix n>=8192.
    For n>=8192 and n==512 the panel solve is done as a small (block x block)
    triangular inverse + a big TF32 GEMM instead of one big FP32 triangular
    solve, moving the panel onto tensor cores. cuSOLVER runs its trailing SYRK
    in FP32, so the TF32 trailing is the win (measured, per-op via profiler):
    32768x1 79.6->59 ms, 16384x1 19.3->15.8 ms, 640x512 3.80->3.18 ms.
    For batch=1 n>=8192 the trailing Schur update is computed lower-triangle
    only, as a tiled TF32 GEMM over T=8192 tiles that skips the strictly-upper
    tiles (the trailing is symmetric and only its lower triangle is read by
    later steps) -- ~0.56x the FLOPs, a net win over the monolithic baddbmm:
    n=32768 58.4->50.0 ms (-14%), n=16384 -2.6% (harness).
  - otherwise: torch.linalg.cholesky_ex (cuSOLVER), still best for n in
    {128, 256}, low-batch 512/1024, single large-n potrf up to 4096.

The reconstruction gate is scaled_residual <= 20; the TF32 route lands near
0.6 at n=2048 (~33x margin), so the low-precision trailing update is safe.

Dead ends (measured, kept out):
  - CUDA graphs on single-kernel routines (no launch latency to hide; the
    input copy is pure overhead).
  - FP16 trailing update: same 10-bit mantissa as TF32 but needs a separate
    cast + bmm + sub_, adding memory passes over the huge trailing matrix.
  - Fused fp16 / bf16x9 trailing SYRK via a cuBLASLt epilogue op
    (D = beta*C + alpha*(A@B) in one call, out=, from the qr_v2 corpus): builds
    and is numerically fine, but does NOT beat torch's TF32 baddbmm -- geomean
    tf32 1925, fp16 1939, bf16x9 1967 us. The trailing GEMM has arithmetic
    intensity ~256 FLOP/byte (right at the B200 tensor-core knee), so it is
    memory-bound on the fp32 C accumulator: changing the A@B precision doesn't
    move it, and bf16x9's ~9 emulation passes only add cost. (bf16x9 is for
    replacing true FP32; our trailing was already a cheap 1-pass TF32.)
  - blocked route at n=512 and low-batch n=1024: cuSOLVER is competitive there
    and the FP32 diagonal + solve_triangular overhead loses.
"""
from __future__ import annotations

import os

import torch
import triton
import triton.language as tl

from task import input_t, output_t

# cuSOLVER's small/mid batched potrf carries real host-side overhead (workspace
# setup + several internal launches). Capturing the call in a CUDA graph once per
# shape and replaying it erases that overhead: 512x16 757->578, 1024x4 1628->1245,
# 256x64 363->274, 128x256 197->162 us (all measured). potrf's kernel sequence is
# data-independent for a fixed shape, so replaying on new values is exact.
# CHOL_GRAPH overrides the mode: "cusolver"(default), "off", or "blocked" (expt).
_GRAPH_MODE = os.environ.get("CHOL_GRAPH", "cusolver")
_GRAPH_CACHE: dict = {}
_GRAPH_SKIP: set = set()

# Graph-replay the whole blocked route (not just the cuSOLVER fallback) for the
# batched/mid shapes, where many small per-block launches carry host overhead.
_GRAPH_BLOCKED = os.environ.get("CHOL_GRAPH_BLOCKED", "1") != "0"

# Trailing update for batch=1 large n: a lower-triangle-only *tiled* TF32 GEMM
# (skip strictly-upper T x T tiles, ~0.56x the FLOPs) beats the monolithic full
# baddbmm because the trailing is symmetric and only its lower triangle is read
# by later steps. T=8192 tiles are big enough to run near the monolithic gemm's
# efficiency, so the triangle-skip is a net win: n=32768 58.4->50.0 ms (-14%),
# n=16384 -2.6%, n=8192 neutral (its trailing < one tile -> one full gemm).
# (Stock cuBLAS TF32 Ssyrk does NOT help: it dispatches to the same full gemm
# kernel and discards the upper triangle -- see docs/EXPERIMENTS.md.)
# Set CHOL_TRAILING=baddbmm to force the old monolithic path; CHOL_TILE picks T.
_TRAILING = os.environ.get("CHOL_TRAILING", "tiled")
_TILE = int(os.environ.get("CHOL_TILE", "8192"))


def _tiled_triangle_syrk(trailing: torch.Tensor, l_pk: torch.Tensor, tile: int):
    """C[lower] -= L_pk @ L_pk^T via lower-triangle-only T x T tiles (batch=1).

    Skips strictly-upper block-tiles (~half the FLOPs). Diagonal tiles compute
    the full T x T (only the lower part is read by later steps; the upper write
    is harmless). Operates on strided views of `work` (row-stride = full n).
    """
    m = trailing.shape[-1]           # trailing: (1, m, m) strided view of work
    prev = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    for i0 in range(0, m, tile):
        i1 = min(i0 + tile, m)
        ai = l_pk[:, i0:i1, :]
        for j0 in range(0, i1, tile):   # j0 < i1 -> lower incl. diagonal tile
            j1 = min(j0 + tile, m)
            sub = trailing[:, i0:i1, j0:j1]
            torch.baddbmm(
                sub, ai, l_pk[:, j0:j1, :].transpose(-1, -2),
                beta=1.0, alpha=-1.0, out=sub,
            )
    torch.backends.cuda.matmul.allow_tf32 = prev


def _graphed(fn, data: torch.Tensor) -> torch.Tensor:
    """Capture `fn(static_in)` once per shape, then replay on new data.

    Falls back to eager `fn` if capture is unsupported for this shape.
    """
    key = tuple(data.shape)
    if key in _GRAPH_SKIP:
        return fn(data)
    entry = _GRAPH_CACHE.get(key)
    if entry is None:
        try:
            static_in = data.clone()
            for _ in range(3):  # warmup (allocator/cublas plans) before capture
                fn(static_in)
            torch.cuda.synchronize()
            g = torch.cuda.CUDAGraph()
            with torch.cuda.graph(g):
                static_out = fn(static_in)
            entry = (g, static_in, static_out)
            _GRAPH_CACHE[key] = entry
        except Exception:
            _GRAPH_SKIP.add(key)
            return fn(data)
    g, static_in, static_out = entry
    static_in.copy_(data)
    g.replay()
    return static_out.clone()

# ---------------------------------------------------------------------------
# Custom CUDA batched Cholesky for tiny n (one CTA per matrix, smem-resident).
# ---------------------------------------------------------------------------
_CPP_SRC = "void chol_batched(torch::Tensor input, torch::Tensor output);"

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

// In-place left-looking Cholesky, one thread block per matrix.
// The lower triangle is factored in shared memory; each column is read once
// (as the original A) then overwritten with its final L values, and later
// columns only ever read already-finalized columns to the left.
template <int N, int T>
__global__ void chol_kernel(const float* __restrict__ Ain_all,
                            float* __restrict__ Aout_all) {
    extern __shared__ float smem[];
    float* __restrict__ sL = smem;          // N*N factored tile
    float* __restrict__ col = smem + N * N; // scratch for the current column
    const int m = blockIdx.x;
    const int tid = threadIdx.x;
    const float* __restrict__ Ain = Ain_all + (long)m * N * N;
    float* __restrict__ Aout = Aout_all + (long)m * N * N;

    for (int idx = tid; idx < N * N; idx += T) sL[idx] = Ain[idx];
    __syncthreads();

    for (int j = 0; j < N; ++j) {
        // Compute the whole column below+including the diagonal in parallel:
        //   col[i] = A[i][j] - sum_{p<j} L[i][p] * L[j][p]
        // Handling i==j here means the diagonal's dot product is parallelized
        // too (no serial thread-0 critical path).
        for (int i = j + tid; i < N; i += T) {
            float s = sL[i * N + j];
            for (int p = 0; p < j; ++p) {
                s -= sL[i * N + p] * sL[j * N + p];
            }
            col[i] = s;
        }
        __syncthreads();
        const float diag = sqrtf(col[j] > 0.f ? col[j] : 1e-30f);
        const float inv = 1.0f / diag;
        for (int i = j + tid; i < N; i += T) {
            sL[i * N + j] = (i == j) ? diag : col[i] * inv;
        }
        __syncthreads();
    }

    for (int idx = tid; idx < N * N; idx += T) {
        int i = idx / N;
        int j = idx - i * N;
        Aout[idx] = (j <= i) ? sL[idx] : 0.0f;
    }
}

void chol_batched(torch::Tensor input, torch::Tensor output) {
    const int batch = input.size(0);
    const int N = input.size(1);
    const float* in = input.data_ptr<float>();
    float* out = output.data_ptr<float>();
    const size_t smem = (size_t)(N * N + N) * sizeof(float);
    const dim3 grid(batch);

#define LAUNCH(NN, TT)                                                        \
    do {                                                                      \
        cudaFuncSetAttribute(chol_kernel<NN, TT>,                            \
                             cudaFuncAttributeMaxDynamicSharedMemorySize,     \
                             (int)smem);                                      \
        chol_kernel<NN, TT><<<grid, TT, smem>>>(in, out);                    \
    } while (0)

    if (N == 32)       LAUNCH(32, 32);
    else if (N == 64)  LAUNCH(64, 64);
    else if (N == 128) LAUNCH(128, 128);
    else TORCH_CHECK(false, "chol_batched: unsupported N=", N);
#undef LAUNCH

    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "chol_batched launch: ", cudaGetErrorString(err));
}
"""

try:
    from torch.utils.cpp_extension import load_inline

    _CHOL_EXT = load_inline(
        name="chol_ext",
        cpp_sources=[_CPP_SRC],
        cuda_sources=[_CUDA_SRC],
        functions=["chol_batched"],
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        verbose=False,
    )
except Exception:
    _CHOL_EXT = None

# n=128 also builds/passes but loses to cuSOLVER (178 vs 152 us): only 256
# matrices with 64 KB smem/block leaves the GPU occupancy-starved.
_CUDA_NS = (32, 64)


def _cuda_cholesky(data: torch.Tensor) -> torch.Tensor:
    output = torch.empty_like(data)
    _CHOL_EXT.chol_batched(data, output)
    return output


# ---------------------------------------------------------------------------
# Triton fallback for n == 32 (used only if the CUDA extension fails to build).
# ---------------------------------------------------------------------------
@triton.jit
def _cholesky32_kernel(input_ptr, output_ptr, matrix_stride: tl.constexpr):
    matrix = tl.program_id(0)
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]
    offsets = matrix * matrix_stride + rows * 32 + cols
    values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)
    for k in range(32):
        row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
        diagonal = tl.sum(tl.where(col_ids == k, row, 0.0), axis=0)
        diagonal -= tl.sum(tl.where(col_ids < k, row * row, 0.0), axis=0)
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
        column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
        products = tl.where(cols < k, values * row[None, :], 0.0)
        column = (column - tl.sum(products, axis=1)) / diagonal
        values = tl.where((rows == k) & (cols == k), diagonal, values)
        values = tl.where((rows > k) & (cols == k), column[:, None], values)
    tl.store(output_ptr + offsets, values)


def _triton_cholesky32(data: torch.Tensor) -> torch.Tensor:
    output = torch.empty_like(data)
    _cholesky32_kernel[(data.shape[0],)](data, output, 32 * 32, num_warps=1)
    return output


# ---------------------------------------------------------------------------
# Blocked right-looking Cholesky with a TF32 tensor-core trailing update.
# ---------------------------------------------------------------------------
_BLOCK_FOR_N = {
    128: 64,
    256: 128,
    512: 128,
    1024: 512,
    2048: 512,
    4096: 512,
    8192: 1024,
    16384: 1024,
    32768: 1024,
}


def _use_blocked(n: int, batch: int) -> bool:
    if n == 512:
        return batch >= 32  # high-batch 512 (blocked TF32, inverse+GEMM panel)
    if n == 1024:
        return batch >= 32
    if n == 2048:
        return batch <= 2
    if n == 4096:
        return batch >= 2
    if n >= 8192:
        return True
    return False


def _blocked_cholesky(data: torch.Tensor, block: int, use_inv: bool) -> torch.Tensor:
    """Right-looking blocked Cholesky with a TF32 trailing Schur update.

    When `use_inv` (large single matrices), the panel solve
    L_pk = A_pk @ L_kk^{-T} is done as a *small* triangular inverse of the
    diagonal block (once per step, FP32) followed by a *big* TF32 GEMM,
    moving the O(n^2 * block) panel work onto tensor cores. For small batched
    matrices the plain FP32 triangular solve is cheaper (no inverse overhead).
    """
    batch, n, _ = data.shape
    work = data.clone()
    out = torch.zeros_like(data)
    if use_inv:
        eye = torch.eye(block, device=data.device, dtype=data.dtype).expand(
            batch, block, block
        )

    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    for k in range(0, n, block):
        e = min(k + block, n)
        akk = work[:, k:e, k:e]
        lkk = torch.linalg.cholesky_ex(akk, check_errors=False).L
        out[:, k:e, k:e] = lkk

        if e < n:
            a_pk = work[:, e:, k:e]
            if use_inv:
                lkk_inv = torch.linalg.solve_triangular(lkk, eye, upper=False)
                torch.backends.cuda.matmul.allow_tf32 = True
                l_pk = torch.bmm(a_pk, lkk_inv.transpose(-1, -2))
                torch.backends.cuda.matmul.allow_tf32 = False
            else:
                # L_pk @ L_kk^T = A_pk  ->  L_pk = A_pk @ L_kk^{-T}
                l_pk = torch.linalg.solve_triangular(
                    lkk.transpose(-1, -2), a_pk, upper=True, left=False
                )
            out[:, e:, k:e] = l_pk
            trailing = work[:, e:, e:]
            if _TRAILING == "tiled" and batch == 1:
                _tiled_triangle_syrk(trailing, l_pk, _TILE)
            else:
                torch.backends.cuda.matmul.allow_tf32 = True
                torch.baddbmm(
                    trailing, l_pk, l_pk.transpose(-1, -2),
                    beta=1.0, alpha=-1.0, out=trailing,
                )
                torch.backends.cuda.matmul.allow_tf32 = False
    torch.backends.cuda.matmul.allow_tf32 = prev_tf32
    return out


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n in _CUDA_NS and _CHOL_EXT is not None:
        return _cuda_cholesky(data)
    if n == 32:
        return _triton_cholesky32(data)
    if _GRAPH_MODE == "blocked" and n in (128, 256):
        return _graphed(
            lambda d: _blocked_cholesky(d, _BLOCK_FOR_N[n], use_inv=False), data
        )
    if _use_blocked(n, batch):
        blk = _BLOCK_FOR_N[n]
        inv = (n >= 8192 or n == 512)
        # The blocked route is a fixed (data-independent) sequence of many small
        # per-block launches for a given shape; graph-replaying it erases the
        # host-launch overhead. Grader-confirmed wins: 2048x2 2.83->2.70 ms
        # (-4.6%), 4096x2 6.26->6.01 ms (-4%). n>=8192 has few large launches (no
        # host overhead to hide); n=1024x60 improved in the (host-heavy) harness
        # but slightly regressed on the grader; n=512 has a big-batch/big-data
        # footprint whose graph copy costs more than it saves -> all stay eager.
        if _GRAPH_BLOCKED and 2048 <= n < 8192:
            return _graphed(lambda d: _blocked_cholesky(d, blk, use_inv=inv), data)
        return _blocked_cholesky(data, blk, use_inv=inv)
    # cuSOLVER fallback. Graph-replay the small/mid shapes where potrf's host
    # overhead dominates; n>=4096 single potrf has none to hide (graph copy only
    # adds cost), so leave it eager.
    if _GRAPH_MODE == "cusolver" and n <= 2048:
        return _graphed(
            lambda d: torch.linalg.cholesky_ex(d, check_errors=False).L, data
        )
    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 387 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