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
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 = 1
num_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