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