Skip to content
KernelIndex
Search⌘K

submission 882186

techhelp · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-882186?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.66ms
#206 of 337
2026-07-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2e9315560e1913bb68db1c8ae56e413c4634e40a026d028857cf11adb72de611
license declaredunknown
license concludedunknown
authorstechhelp
imported2026-08-26

Techniques

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

fp4_register_prec("fp4", 8192,
num-warps = 1num_warps=1,

Kernel source

submission.py210 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl

from task import input_t, output_t

NB_SMALL = 1024
NB_LARGE = 2048
NB_LARGE_THRESH = 8192
_FP8 = torch.float8_e4m3fn
_FP8_MAX = 448.0
try:
    _FP4 = torch.float4_e2m1fn
    _FP4_MAX = 6.0
except AttributeError:
    _FP4 = None
    _FP4_MAX = 0.0


@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 _cholesky32(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n != 32:
        return torch.linalg.cholesky_ex(data, check_errors=False).L
    output = torch.empty_like(data)
    _cholesky32_kernel[(batch,)](
        data,
        output,
        n * n,
        num_warps=1,
    )
    return output


def _quant_bmm(
    Y: torch.Tensor, Z: torch.Tensor, dtype: torch.dtype, max_val: float
) -> torch.Tensor:
    sy = max_val / Y.abs().amax(dim=(1, 2), keepdim=True).clamp_min(1e-20)
    sz = max_val / Z.abs().amax(dim=(1, 2), keepdim=True).clamp_min(1e-20)
    Yq = (Y * sy).to(dtype)
    Zq = (Z * sz).to(dtype)
    C = torch.matmul(Yq, Zq.mT.contiguous())
    return C / (sy * sz)


# Precision tiers: (mm_key, min_n) — ordered best→worst, first passing quality check wins
# mm_key — key into _MM dict below
# min_n — only use for n >= this threshold
_MM = {}
_TRAIL = {}
_PREC_TIERS = []


def _register_prec(key: str, min_n: int, mm_fn, trail_fn=None):
    _MM[key] = mm_fn
    _TRAIL[key] = trail_fn or (lambda dst, Y, Z: dst.sub_(mm_fn(Y, Z)))
    _PREC_TIERS.append((key, min_n))


_register_prec("fp8", 4096,
    lambda Y, Z: _quant_bmm(Y, Z, _FP8, _FP8_MAX))
_register_prec("fp32", 0,
    lambda Y, Z: Y @ Z.mT,
    lambda dst, Y, Z: torch.baddbmm(dst, Y, Z.mT, beta=1, alpha=-1, out=dst))
if _FP4 is not None:
    _register_prec("fp4", 8192,
        lambda Y, Z: _quant_bmm(Y, Z, _FP4, _FP4_MAX))


def _cholesky_lower(A: torch.Tensor, prec: str, nb: int) -> None:
    trail = _TRAIL[prec]
    b, n = A.shape[0], A.shape[-1]
    info = torch.empty(b, dtype=torch.int32, device='cuda')
    for j in range(0, n, nb):
        jb = min(nb, n - j)

        if j > 0:
            Lc = A[:, j : j + jb, :j]
            trail(A[:, j : j + jb, j : j + jb], Lc, Lc)

        torch.linalg.cholesky_ex(
            A[:, j : j + jb, j : j + jb],
            upper=False, check_errors=False,
            out=(A[:, j : j + jb, j : j + jb], info),
        )

        if j + jb < n:
            Lp = A[:, j + jb : n, :j]
            Lc = A[:, j : j + jb, :j]
            trail(A[:, j + jb : n, j : j + jb], Lp, Lc)
            A[:, j + jb : n, j : j + jb] = torch.linalg.solve_triangular(
                A[:, j : j + jb, j : j + jb].mT,
                A[:, j + jb : n, j : j + jb],
                upper=True,
                left=False,
            )


def _quality_ok(out: torch.Tensor, data: torch.Tensor) -> bool:
    if not torch.isfinite(out).all():
        return False
    od = (out * out).sum(dim=-1)
    dd = data.diagonal(dim1=-2, dim2=-1)
    return bool(
        (
            (od - dd).abs()
            < 10.0
            * data.shape[-1]
            * torch.finfo(torch.float32).eps
            * dd.abs().clamp_min(1e-20).max()
        ).all()
    )


def _run_blocked(A: torch.Tensor, prec: str, nb: int) -> torch.Tensor:
    _cholesky_lower(A, prec, nb)
    return torch.tril(A)


def _pick_prec(n: int, data: torch.Tensor, nb: int) -> str:
    for key, min_n in _PREC_TIERS:
        if n >= min_n:
            if key == "fp32":
                return key
            try:
                A = data.clone()
                _run_blocked(A, key, nb)
                if _quality_ok(torch.tril(A), data):
                    return key
            except Exception:
                continue
    return "fp32"


_GRAPH_CACHE = {}


def _use_blocked(n: int, b: int) -> bool:
    if n >= 16384:
        return True
    if n == 8192 and b <= 1:
        return True
    if n == 4096 and b == 2:
        return True
    if n == 2048:
        return b <= 2
    return False


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    b = data.shape[0]
    if n == 32:
        return _cholesky32(data)
    if not _use_blocked(n, b):
        return torch.linalg.cholesky_ex(data, upper=False, check_errors=False)[0]

    nb = NB_LARGE if n >= NB_LARGE_THRESH else NB_SMALL

    key = (b, n)
    entry = _GRAPH_CACHE.get(key)
    if entry is None:
        prev = torch.backends.cuda.matmul.allow_tf32
        torch.backends.cuda.matmul.allow_tf32 = True
        try:
            prec = _pick_prec(n, data, nb)
            A = data.clone()
            for _ in range(2):
                _run_blocked(A.clone(), prec, nb)
            torch.cuda.synchronize()
            static_in = data.clone()
            g = torch.cuda.CUDAGraph()
            with torch.cuda.graph(g):
                A_lower = _run_blocked(static_in, prec, nb)
            entry = (g, static_in, A_lower, prec)
        finally:
            torch.backends.cuda.matmul.allow_tf32 = prev
        _GRAPH_CACHE[key] = entry
    g, static_in, A_lower, prec = entry
    static_in.copy_(data)
    g.replay()
    out = A_lower.clone()
    return out
scrolls · 210 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