Skip to content
KernelIndex
Search⌘K

submission 905488

oogway8030 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_final.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-905488?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.95ms
#257 of 337
2026-07-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f8462352420cb5670907e1719215cdca2ba444d728eee9361d6605633ab4e371
license declaredunknown
license concludedunknown
authorsoogway8030
imported2026-08-26

Techniques

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

num-warps = 1_chol32_kernel[(batch,)](data, out, 32 * 32, 32, 0, num_warps=1)

Kernel source

submission_final.py133 lines
#!POPCORN leaderboard cholesky

import torch
import triton
import triton.language as tl
from task import input_t, output_t


@triton.jit
def _chol32_kernel(in_ptr, out_ptr, batch_stride, row_stride, base):
    b = tl.program_id(0)
    rn = tl.arange(0, 32)
    rows = rn[:, None]
    cols = rn[None, :]
    offs = base + b * batch_stride + rows * row_stride + cols
    values = tl.where(rows >= cols, tl.load(in_ptr + offs), 0.0)
    for k in range(32):
        colk = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
        dk = tl.sum(tl.where(rn == k, colk, 0.0), axis=0)
        s = tl.sqrt(tl.maximum(dk, 1e-30))
        colk = tl.where(rn >= k, colk / s, 0.0)
        values = tl.where(cols == k, colk[:, None], values)
        values = tl.where(
            (cols > k) & (rows >= cols),
            values - colk[:, None] * colk[None, :],
            values,
        )
    tl.store(out_ptr + offs, values)


_NB = 32


def _small32(data, batch):
    out = torch.empty_like(data)
    _chol32_kernel[(batch,)](data, out, 32 * 32, 32, 0, num_warps=1)
    return out


def _blocked(data, batch, n):
    W = data.clone()
    nn = n * n
    for k0 in range(0, n, _NB):
        k1 = k0 + _NB
        _chol32_kernel[(batch,)](W, W, nn, n, k0 * n + k0, num_warps=1)
        if k1 < n:
            L11 = W[:, k0:k1, k0:k1]
            A21 = W[:, k1:, k0:k1]
            X = torch.linalg.solve_triangular(
                L11.transpose(-1, -2), A21, upper=True, left=False
            )
            A21.copy_(X)
            W[:, k1:, k1:].baddbmm_(X, X.transpose(-1, -2), beta=1, alpha=-1)
    return torch.tril(W)


def _use_blocked(batch, n):
    return n in (2048, 4096) and batch >= 2


def _compute(data, batch, n):
    if n == 32:
        return _small32(data, batch)
    if _use_blocked(batch, n):
        return _blocked(data, batch, n)
    return torch.linalg.cholesky_ex(data, check_errors=False).L


def _count(batch, n):
    return max(1, min(50, (256 << 20) // (batch * n * n * 4)))


_slots = {}  # (ptr, batch, n) -> list[(graph, out)] (small) or (graph, out) (blocked)
_ring_idx = {}
_warm_shapes = set()
_graphs_ok = [True]


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    small = n == 32
    eligible = (
        _graphs_ok[0]
        and data.is_contiguous()
        and (small or _use_blocked(batch, n))
    )
    shape = (batch, n)
    if not eligible or shape not in _warm_shapes:
        _warm_shapes.add(shape)
        return _compute(data, batch, n)

    key = (data.data_ptr(), batch, n)
    if small:
        # ring of graphs: cheap captures (single kernel), zero-copy returns
        entry = _slots.get(key)
        if entry is None:
            entry = []
            _slots[key] = entry
            _ring_idx[key] = 0
        if len(entry) < _count(batch, n) + 1:
            try:
                g = torch.cuda.CUDAGraph()
                with torch.cuda.graph(g):
                    out = _compute(data, batch, n)
                g.replay()
                entry.append((g, out))
                return out
            except Exception:
                _graphs_ok[0] = False
                return _compute(data, batch, n)
        g, out = entry[_ring_idx[key] % len(entry)]
        _ring_idx[key] += 1
        g.replay()
        return out

    # blocked shapes: one graph per key, return a clone (graphs are expensive
    # to capture, clones are cheap relative to the factorization itself)
    entry = _slots.get(key)
    if entry is None:
        try:
            g = torch.cuda.CUDAGraph()
            with torch.cuda.graph(g):
                out = _compute(data, batch, n)
            g.replay()
            _slots[key] = (g, out)
            return out.clone()
        except Exception:
            _graphs_ok[0] = False
            return _compute(data, batch, n)
    g, out = entry
    g.replay()
    return out.clone()
scrolls · 133 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