Skip to content
KernelIndex
Search⌘K

submission 894455

aresmaniii · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f008034c85ee15770ac42baa148ffdeb7796d9868a4ac69cb6753e629ec80a61
license declaredunknown
license concludedunknown
authorsaresmaniii
imported2026-08-26

Techniques

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

split-kdef _split_kernel(

Kernel source

submission.py385 lines
"""Batched FP32 Cholesky for the GPU MODE `cholesky` leaderboard (B200).

Ranking is the geometric mean over 15 benchmark entries, which weights the
tiny ones exactly as heavily as the 32768 x 32768 one. Those two ends want
opposite things, so the file has two engines:

  * Small n is bound by launch latency, not arithmetic. Every low-batch entry
    from n=32 to n=1024 carries exactly 2^22 FP32 elements, so the whole
    factorisation is ~4M*n/3 flops -- microseconds of B200 compute. n <= 32
    therefore runs as a single fused Triton kernel, one program per matrix,
    fully unrolled so every pivot index is a compile-time constant. That turns
    the column extraction from a masked full-tile reduction into a broadcast
    and is worth 2.2x; the kernel then runs at ~86% of peak bandwidth.

  * Large n is bound by GEMM throughput, and FP32 GEMM leaves most of a
    Blackwell on the floor. The blocked driver can run its panel update as
    three BF16 GEMMs on a hi/lo split of each operand, which costs 3x the
    tensor-core flops but is ~5x cheaper per flop -- and lands within 10x of
    FP32 accuracy, against a checker that allows ~10000x. Plain TF32 is 750x
    off FP32 and diverges outright on near-singular inputs; the split is the
    thing that makes low precision usable here.

Between them sits a left-looking blocked driver whose panel solve is a GEMM
rather than a TRSM, because the tile kernel hands back each diagonal block's
inverse along with its factor.

Thresholds are measured crossovers, but they were measured on an RTX 5070
(sm_120); see README.md for how to re-sweep them on B200.
"""

import os

import torch
import triton
import triton.language as tl

from task import input_t, output_t


# --------------------------------------------------------------------------
# Tuning knobs. All overridable from the environment; see tools/sweep.py.
# --------------------------------------------------------------------------

def _knob(name, default):
    return int(os.environ.get(name, default))


# The tile kernel below is fully unrolled, so its IR is O(N^3) -- N steps over
# an N x N tile -- and `ptxas` memory use follows. Measured, sm_120:
#
#     N=32   0.8GB /  3s      N=128   1.4GB / 23s
#     N=64   0.9GB /  6s      N=256   >8GB, does not finish
#
# 256 is off a cliff, not on a curve: it took a 29GB machine down three times
# before this clamp existed. Nothing here may instantiate the kernel past 128,
# whatever the environment says. Raising this needs `tl.range` in place of
# `tl.static_range` above the clamp, which costs the 2.2x the unroll buys.
_UNROLL_MAX = 128

# Largest n handled entirely inside one Triton program.
TILE_MAX = min(_knob("CHOL_TILE_MAX", 32), _UNROLL_MAX)

# Panel width of the blocked driver. 0 = use the table below. The tile kernel
# costs ~O(width^3) per matrix, so wide panels only pay off once the GEMMs
# they feed are big enough to dominate their launch count.
BASE = min(_knob("CHOL_BASE", 0), _UNROLL_MAX)
# Measured, sm_120 (geomean over entries 4-9,11): width 64 beats 32 by 1.13x
# and beats 128 by 2.0x. 128 loses because the tile kernel is O(width^3) of
# strictly serial work per matrix, so it outgrows the GEMMs it saves.
_BASE_BY_N = ((512, 32), (1 << 30, 64))

# Run the panel update as three BF16 GEMMs on a hi/lo split once the update is
# at least this many flops. Below it the extra launches cost more than the
# tensor cores save.
#
# Off by default: measured on sm_120 it loses on every entry it fires on
# (640x512: 9.0ms -> 14.4ms; 60x1024: 4.27 -> 4.95; 2x4096: 10.2 -> 11.7).
# Three BF16 GEMMs plus the hi/lo conversion traffic cost more than this card's
# BF16 tensor path saves over its FP32 path. The premise -- ~5x cheaper per
# flop -- is a B200 number, so the path is kept and re-enabled with
# CHOL_SPLIT_FLOPS=2e10 for re-measurement there.
SPLIT_MIN_FLOPS = float(os.environ.get("CHOL_SPLIT_FLOPS", "inf"))

# A lone matrix this large goes to cuSOLVER. It is well tuned for the
# single-matrix case and degrades sharply the moment there is more than one
# (15ms for batch=2 at n=4096 vs 3.7ms for batch=1), so the cutoff is on both.
CUSOLVER_MIN_N = _knob("CHOL_CUSOLVER_N", 4096)
CUSOLVER_MAX_BATCH = _knob("CHOL_CUSOLVER_B", 1)

_WARPS = {8: 1, 16: 1, 32: 1, 64: 4, 128: 8}


# --------------------------------------------------------------------------
# In-CTA batched Cholesky.
# --------------------------------------------------------------------------

@triton.jit
def _tile_chol_kernel(
    a_ptr, l_ptr, v_ptr,
    a_sb, a_sr, a_sc,
    l_sb, l_sr, l_sc,
    v_sb, v_sr, v_sc,
    n,
    N: tl.constexpr, EXACT: tl.constexpr, WANT_INV: tl.constexpr,
):
    pid = tl.program_id(0)
    r = tl.arange(0, N)
    c = tl.arange(0, N)
    rr = r[:, None]
    cc = c[None, :]
    lower = rr >= cc

    a_off = pid * a_sb + rr * a_sr + cc * a_sc
    if EXACT:
        a = tl.where(lower, tl.load(a_ptr + a_off), 0.0)
    else:
        inside = (rr < n) & (cc < n)
        a = tl.load(a_ptr + a_off, mask=inside, other=0.0)
        a = tl.where(lower & inside, a, 0.0)
        # Identity padding keeps the padded pivots well defined, so the loop
        # can always run the full N steps and stay statically unrollable.
        a = tl.where((rr == cc) & (rr >= n), 1.0, a)

    # Relative pivot floor: never triggers on an honestly SPD input, but keeps
    # the output finite and the diagonal strictly positive if one is only
    # semidefinite to FP32 (the `lowrank` and `diagonal` cases get close).
    floor = tl.max(tl.where(rr == cc, a, 0.0)) * 1e-14 + 1e-38

    # static_range, not range: it makes k a compile-time constant, so `cc == k`
    # is a known mask and the column extraction folds into a broadcast rather
    # than a masked reduction over the whole tile. Worth 2.2x at N=32.
    for k in tl.static_range(N):
        col = tl.sum(tl.where(cc == k, a, 0.0), axis=1)
        akk = tl.sum(tl.where(r == k, col, 0.0), axis=0)
        ok = akk > floor
        d = tl.sqrt(tl.where(ok, akk, floor))
        rd = tl.where(ok, 1.0 / d, 0.0)
        col = tl.where(r == k, d, tl.where(r > k, col * rd, 0.0))
        a = tl.where(cc == k, col[:, None], a)
        a -= tl.where(lower & (rr > k) & (cc > k), col[:, None] * col[None, :], 0.0)

    if WANT_INV:
        # Forward substitution on the identity, all right-hand sides at once,
        # so the driver's panel solve becomes a GEMM: cuBLAS batched TRSM
        # measures ~3x slower than the equivalent GEMM at these shapes.
        # Interleaving this into the sweep above was tried and is slower --
        # it lengthens the dependent chain the factorisation is bound on.
        y = tl.where(rr == cc, 1.0, 0.0)
        for k in tl.static_range(N):
            lcol = tl.sum(tl.where(cc == k, a, 0.0), axis=1)
            lkk = tl.sum(tl.where(r == k, lcol, 0.0), axis=0)
            yrow = tl.sum(tl.where(rr == k, y, 0.0), axis=0) / lkk
            y = tl.where(rr == k, yrow[None, :], y)
            y -= tl.where(rr > k, lcol[:, None] * yrow[None, :], 0.0)

    l_off = pid * l_sb + rr * l_sr + cc * l_sc
    if EXACT:
        tl.store(l_ptr + l_off, a)
    else:
        tl.store(l_ptr + l_off, a, mask=(rr < n) & (cc < n))

    if WANT_INV:
        v_off = pid * v_sb + rr * v_sr + cc * v_sc
        if EXACT:
            tl.store(v_ptr + v_off, y)
        else:
            tl.store(v_ptr + v_off, y, mask=(rr < n) & (cc < n))


def _tile_chol(src: torch.Tensor, dst: torch.Tensor,
               inv: torch.Tensor | None = None) -> None:
    """Factor every matrix in `src` (b, n, n) into `dst`; `src`/`dst` may alias.

    When `inv` is given it receives the inverse of each factor.
    """
    b, n, _ = src.shape
    tile = max(8, triton.next_power_of_2(n))
    ref = inv if inv is not None else dst
    _tile_chol_kernel[(b,)](
        src, dst, ref,
        src.stride(0), src.stride(1), src.stride(2),
        dst.stride(0), dst.stride(1), dst.stride(2),
        ref.stride(0), ref.stride(1), ref.stride(2),
        n,
        N=tile, EXACT=(tile == n), WANT_INV=inv is not None,
        num_warps=_WARPS.get(tile, 8),
    )


# --------------------------------------------------------------------------
# BF16 hi/lo split, one launch per panel.
# --------------------------------------------------------------------------

@triton.jit
def _split_kernel(
    src_ptr, hi_ptr, lo_ptr,
    s_b, s_r, s_c, h_b, h_r, h_c,
    rows, cols,
    BR: tl.constexpr, BC: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_r = tl.program_id(1)
    rr = (pid_r * BR + tl.arange(0, BR))[:, None]
    cc = tl.arange(0, BC)[None, :]
    mask = (rr < rows) & (cc < cols)

    x = tl.load(src_ptr + pid_b * s_b + rr * s_r + cc * s_c, mask=mask, other=0.0)
    hi = x.to(tl.bfloat16)
    lo = (x - hi.to(tl.float32)).to(tl.bfloat16)
    off = pid_b * h_b + rr * h_r + cc * h_c
    tl.store(hi_ptr + off, hi, mask=mask)
    tl.store(lo_ptr + off, lo, mask=mask)


def _write_split(src: torch.Tensor, hi: torch.Tensor, lo: torch.Tensor) -> None:
    b, rows, cols = src.shape
    bc = max(16, triton.next_power_of_2(cols))
    br = max(1, min(64, 4096 // bc))
    _split_kernel[(b, triton.cdiv(rows, br))](
        src, hi, lo,
        src.stride(0), src.stride(1), src.stride(2),
        hi.stride(0), hi.stride(1), hi.stride(2),
        rows, cols, BR=br, BC=bc,
    )


# --------------------------------------------------------------------------
# Blocked driver.
# --------------------------------------------------------------------------

def _panel_width(n: int) -> int:
    if BASE:
        return BASE
    for limit, width in _BASE_BY_N:
        if n <= limit:
            return width
    return _BASE_BY_N[-1][1]


def _chol_blocked(a: torch.Tensor, split: bool) -> torch.Tensor:
    """Left-looking blocked Cholesky.

    Left-looking rather than right-looking because it hits the n^3/3 flop count
    exactly: a right-looking trailing update either does 2x the work or needs a
    launch per block column to exploit symmetry. Each panel is brought up to
    date by one GEMM against everything to its left, so the launch count is
    O(n/width) and independent of batch size, and the output never needs a
    final tril -- L is seeded with zeros and only its lower blocks are written.

    With `split`, that GEMM becomes three BF16 GEMMs over a hi/lo decomposition
    of L. The decomposition is maintained incrementally, one panel at a time,
    so it costs O(n^2) of conversion rather than O(n^3/width).
    """
    b, n, _ = a.shape
    width = _panel_width(n)
    out = torch.zeros_like(a)
    if split:
        hi = torch.zeros((b, n, n), dtype=torch.bfloat16, device=a.device)
        lo = torch.zeros((b, n, n), dtype=torch.bfloat16, device=a.device)

    for k in range(0, n, width):
        kb = min(width, n - k)
        rows = n - k
        panel = torch.empty((b, rows, kb), dtype=a.dtype, device=a.device)
        src = a[:, k:, k:k + kb]

        if k == 0:
            panel.copy_(src)
        elif not split:
            torch.baddbmm(
                src, out[:, k:, :k], out[:, k:k + kb, :k].transpose(-1, -2),
                beta=1.0, alpha=-1.0, out=panel,
            )
        else:
            xh, xl = hi[:, k:, :k], lo[:, k:, :k]
            yh = hi[:, k:k + kb, :k].transpose(-1, -2)
            yl = lo[:, k:k + kb, :k].transpose(-1, -2)
            # (xh + xl) @ (yh + yl), dropping xl@yl -- it is below the FP32
            # rounding of the result anyway.
            torch.baddbmm(src, xh, yh, beta=1.0, alpha=-1.0,
                          out_dtype=torch.float32, out=panel)
            torch.baddbmm(panel, xh, yl, beta=1.0, alpha=-1.0,
                          out_dtype=torch.float32, out=panel)
            torch.baddbmm(panel, xl, yh, beta=1.0, alpha=-1.0,
                          out_dtype=torch.float32, out=panel)

        diag = out[:, k:k + kb, k:k + kb]
        if rows > kb:
            inv = torch.empty((b, kb, kb), dtype=a.dtype, device=a.device)
            _tile_chol(panel[:, :kb, :], diag, inv=inv)
            # L21 = A21 @ inv(L11).T
            torch.bmm(panel[:, kb:, :], inv.transpose(-1, -2),
                      out=out[:, k + kb:, k:k + kb])
        else:
            _tile_chol(panel[:, :kb, :], diag)

        if split and k + kb < n:
            _write_split(out[:, k:, k:k + kb], hi[:, k:, k:k + kb],
                         lo[:, k:, k:k + kb])

    return out


# --------------------------------------------------------------------------

def _reference(data: torch.Tensor) -> torch.Tensor:
    return torch.linalg.cholesky_ex(data, check_errors=False).L


def _use_split(b: int, n: int, width: int) -> bool:
    # Flops in the left-looking updates: sum over panels of 2*(n-k)*width*k.
    flops = 2.0 * b * sum((n - k) * min(width, n - k) * k
                          for k in range(0, n, width))
    return flops >= SPLIT_MIN_FLOPS


def _dispatch(data: torch.Tensor) -> torch.Tensor:
    b, n, _ = data.shape

    if n <= TILE_MAX:
        out = torch.empty_like(data)
        _tile_chol(data, out)
        return out

    if n >= CUSOLVER_MIN_N and b <= CUSOLVER_MAX_BATCH:
        return _reference(data)

    return _chol_blocked(data, _use_split(b, n, _panel_width(n)))


_fast_path: bool | None = None


def _fast_path_works() -> bool:
    """Check the fast path against cuSOLVER once, on the first call.

    The leaderboard runs on hardware and library versions that cannot be
    reproduced here, and a ranked run that fails validation scores nothing at
    all. This exercises the tile kernel, the blocked driver including its
    strided `out=` GEMMs, and the BF16 split path, on inputs small enough to
    be free, and permanently falls back to cuSOLVER if anything will not
    compile, an ATen signature differs, or the answer is wrong. Worst case the
    submission scores like the baseline instead of not at all.
    """
    global _fast_path
    if _fast_path is not None:
        return _fast_path

    _fast_path = False
    try:
        gen = torch.Generator(device="cuda")
        gen.manual_seed(0)
        widths = sorted({_panel_width(96), _panel_width(1 << 20)})
        sizes = [(3, 16), (2, TILE_MAX + 1)] + [(2, 2 * w + 5) for w in widths]
        for b, n in sizes:
            m = torch.randn((b, n, n), device="cuda", dtype=torch.float32,
                            generator=gen)
            a = m @ m.transpose(-1, -2) / n
            a.diagonal(dim1=-2, dim2=-1).add_(1.0)
            a = (0.5 * (a + a.transpose(-1, -2))).contiguous()
            want = _reference(a)

            candidates = [_dispatch(a)]
            if n > TILE_MAX:
                candidates.append(_chol_blocked(a, True))
            for got in candidates:
                if got.shape != a.shape or got.dtype != a.dtype:
                    return False
                if not torch.isfinite(got).all().item():
                    return False
                if torch.triu(got, 1).abs().max().item() != 0.0:
                    return False
                if (got - want).abs().max().item() > 1e-3 * want.abs().max().item():
                    return False
        _fast_path = True
    except Exception:
        _fast_path = False
    return _fast_path


def custom_kernel(data: input_t) -> output_t:
    if not _fast_path_works():
        return _reference(data)
    return _dispatch(data)
scrolls · 385 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