submission 879555
amandeepsp · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 758 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-879555?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:a68dc0dc01e6a1d7e61bec0f1e2759c0c0e605193616787ae3671d2f78389d6f
license declaredunknown
license concludedunknown
authorsamandeepsp
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
accumulator += tl.dot(left, right_t, input_precision=dot_precision)num-warps = 1
num_warps = 1 if n <= 32 else 2persistent-kernel
def _persistent_cholesky_kernel(Kernel source
submission.py758 lines
#!POPCORN leaderboard cholesky
from __future__ import annotations
import torch
import triton
import triton.language as tl
input_t = torch.Tensor
output_t = torch.Tensor
@triton.jit
def _register_cholesky_kernel(
a,
out,
n: tl.constexpr,
block_n: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, block_n)
cols = tl.arange(0, block_n)
base = batch * n * n
valid = (rows[:, None] < n) & (cols[None, :] < n)
tile = tl.load(
a + base + rows[:, None] * n + cols[None, :],
mask=valid,
other=0.0,
)
# Same register-tile pattern that worked well for the QR panel: extract one
# column, normalize it, then apply its rank-1 Schur update without touching
# global memory until the complete factor is ready.
for j in range(n):
column = tl.sum(tl.where(cols[None, :] == j, tile, 0.0), axis=1)
diagonal_value = tl.sum(tl.where(rows == j, column, 0.0))
diagonal = tl.sqrt(tl.maximum(diagonal_value, 1.0e-6))
factor_column = tl.where(rows == j, diagonal, column / diagonal)
active_column = tl.where(rows >= j, factor_column, 0.0)
tile = tl.where(cols[None, :] == j, active_column[:, None], tile)
tile = tl.where(
(rows[:, None] > j) & (cols[None, :] > j),
tile - active_column[:, None] * active_column[None, :],
tile,
)
result = tl.where(rows[:, None] >= cols[None, :], tile, 0.0)
tl.store(out + base + rows[:, None] * n + cols[None, :], result, mask=valid)
def _register_cholesky(data: torch.Tensor) -> torch.Tensor:
batch, n, _ = data.shape
out = torch.empty_like(data)
block_n = triton.next_power_of_2(n)
num_warps = 1 if n <= 32 else 2
_register_cholesky_kernel[(batch,)](
data,
out,
n=n,
block_n=block_n,
num_warps=num_warps,
)
return out
@triton.jit
def _initialize_lower_kernel(a, out, total: tl.constexpr, n: tl.constexpr, block: tl.constexpr):
offsets = tl.program_id(0) * block + tl.arange(0, block)
valid = offsets < total
within_matrix = offsets % (n * n)
row = within_matrix // n
col = within_matrix - row * n
value = tl.load(a + offsets, mask=valid & (col <= row), other=0.0)
tl.store(out + offsets, value, mask=valid)
@triton.jit
def _copy_first_panel_kernel(
a,
out,
matrix_stride,
n: tl.constexpr,
block_size: tl.constexpr,
block_m: tl.constexpr,
):
batch = tl.program_id(1)
rows = tl.program_id(0) * block_m + tl.arange(0, block_m)
cols = tl.arange(0, block_size)
pointers = batch * matrix_stride + rows[:, None] * n + cols[None, :]
mask = (rows[:, None] < n) & (rows[:, None] >= cols[None, :])
value = tl.load(a + pointers, mask=mask, other=0.0)
tl.store(out + pointers, value, mask=mask)
@triton.jit
def _zero_upper_kernel(
out,
matrix_stride,
n: tl.constexpr,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
batch = tl.program_id(2)
rows = tl.program_id(0) * block_m + tl.arange(0, block_m)
cols = tl.program_id(1) * block_n + tl.arange(0, block_n)
mask = (rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > rows[:, None])
tl.store(
out + batch * matrix_stride + rows[:, None] * n + cols[None, :],
0.0,
mask=mask,
)
@triton.jit
def _factor_diagonal_block_kernel(
out,
matrix_stride,
panel_start,
n: tl.constexpr,
block_size: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, block_size)
base = batch * matrix_stride + panel_start * n + panel_start
j = 0
while j < block_size:
diagonal_sum = 0.0
k = 0
while k < j:
value = tl.load(out + base + j * n + k)
diagonal_sum += value * value
k += 1
ajj = tl.load(out + base + j * n + j)
diagonal = tl.sqrt(tl.maximum(ajj - diagonal_sum, 1.0e-6))
row_sum = tl.zeros((block_size,), dtype=tl.float32)
k = 0
while k < j:
lik = tl.load(out + base + rows * n + k)
ljk = tl.load(out + base + j * n + k)
row_sum += lik * ljk
k += 1
aij = tl.load(out + base + rows * n + j)
value = (aij - row_sum) / diagonal
value = tl.where(rows == j, diagonal, value)
tl.store(out + base + rows * n + j, value, mask=rows >= j)
tl.debug_barrier()
j += 1
@triton.jit
def _factor_diagonal_register_kernel(
out,
matrix_stride,
panel_start,
n: tl.constexpr,
block_size: tl.constexpr,
pivot_floor: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, block_size)
cols = tl.arange(0, block_size)
base = batch * matrix_stride + panel_start * n + panel_start
lower_row = tl.maximum(rows[:, None], cols[None, :])
lower_col = tl.minimum(rows[:, None], cols[None, :])
tile = tl.load(out + base + lower_row * n + lower_col)
for j in range(block_size):
column = tl.sum(tl.where(cols[None, :] == j, tile, 0.0), axis=1)
diagonal_value = tl.sum(tl.where(rows == j, column, 0.0))
diagonal = tl.sqrt(tl.maximum(diagonal_value, pivot_floor))
factor_column = tl.where(rows == j, diagonal, column / diagonal)
active_column = tl.where(rows >= j, factor_column, 0.0)
tile = tl.where(cols[None, :] == j, active_column[:, None], tile)
tile = tl.where(
(rows[:, None] > j) & (cols[None, :] > j),
tile - active_column[:, None] * active_column[None, :],
tile,
)
tl.store(
out + base + rows[:, None] * n + cols[None, :],
tile,
mask=rows[:, None] >= cols[None, :],
)
@triton.jit
def _left_looking_panel_update_kernel(
out,
matrix_stride,
panel_start,
n: tl.constexpr,
block_size: tl.constexpr,
block_m: tl.constexpr,
reduction_block: tl.constexpr,
dot_precision: tl.constexpr,
):
batch = tl.program_id(1)
rows = panel_start + tl.program_id(0) * block_m + tl.arange(0, block_m)
cols = panel_start + tl.arange(0, block_size)
reduction = tl.arange(0, reduction_block)
matrix = batch * matrix_stride
accumulator = tl.zeros((block_m, block_size), dtype=tl.float32)
history = 0
while history < panel_start:
left = tl.load(
out + matrix + rows[:, None] * n + history + reduction[None, :],
mask=(rows[:, None] < n) & (history + reduction[None, :] < panel_start),
other=0.0,
)
right_t = tl.load(
out + matrix + cols[None, :] * n + history + reduction[:, None],
mask=(history + reduction[:, None] < panel_start),
other=0.0,
)
accumulator += tl.dot(left, right_t, input_precision=dot_precision)
history += reduction_block
pointers = out + matrix + rows[:, None] * n + cols[None, :]
mask = (rows[:, None] < n) & (rows[:, None] >= cols[None, :])
current = tl.load(pointers, mask=mask, other=0.0)
tl.store(pointers, current - accumulator, mask=mask)
@triton.jit
def _panel_trsm_kernel(
out,
matrix_stride,
panel_start,
trailing_start,
n: tl.constexpr,
block_size: tl.constexpr,
block_m: tl.constexpr,
):
batch = tl.program_id(1)
rows = trailing_start + tl.program_id(0) * block_m + tl.arange(0, block_m)
row_mask = rows < n
matrix = batch * matrix_stride
diagonal_base = matrix + panel_start * n + panel_start
j = 0
while j < block_size:
value = tl.load(
out + matrix + rows * n + panel_start + j,
mask=row_mask,
other=0.0,
)
k = 0
while k < j:
previous = tl.load(
out + matrix + rows * n + panel_start + k,
mask=row_mask,
other=0.0,
)
factor = tl.load(out + diagonal_base + j * n + k)
value -= previous * factor
k += 1
diagonal = tl.load(out + diagonal_base + j * n + j)
tl.store(
out + matrix + rows * n + panel_start + j,
value / diagonal,
mask=row_mask,
)
j += 1
@triton.jit
def _panel_trsm_register_kernel(
out,
matrix_stride,
panel_start,
trailing_start,
n: tl.constexpr,
block_size: tl.constexpr,
block_m: tl.constexpr,
):
batch = tl.program_id(1)
rows = trailing_start + tl.program_id(0) * block_m + tl.arange(0, block_m)
cols = tl.arange(0, block_size)
row_mask = rows < n
matrix = batch * matrix_stride
tile = tl.load(
out + matrix + rows[:, None] * n + panel_start + cols[None, :],
mask=row_mask[:, None],
other=0.0,
)
diagonal = tl.load(
out
+ matrix
+ (panel_start + cols[:, None]) * n
+ panel_start
+ cols[None, :]
)
# Solve X L^T = A one column at a time, but keep the entire row tile in
# registers. Each solved column immediately updates all remaining RHS
# columns, eliminating the repeated global loads in the scalar version.
for j in range(block_size):
current = tl.sum(tl.where(cols[None, :] == j, tile, 0.0), axis=1)
diagonal_j = tl.sum(
tl.where((cols[:, None] == j) & (cols[None, :] == j), diagonal, 0.0)
)
solved = current / diagonal_j
diagonal_column = tl.sum(
tl.where(cols[None, :] == j, diagonal, 0.0), axis=1
)
tile = tl.where(
cols[None, :] == j,
solved[:, None],
tl.where(
cols[None, :] > j,
tile - solved[:, None] * diagonal_column[None, :],
tile,
),
)
tl.store(
out + matrix + rows[:, None] * n + panel_start + cols[None, :],
tile,
mask=row_mask[:, None],
)
@triton.jit
def _factor_trsm_fused_kernel(
out,
matrix_stride,
panel_start,
trailing_start,
n: tl.constexpr,
block_size: tl.constexpr,
block_m: tl.constexpr,
pivot_floor: tl.constexpr,
):
row_tile = tl.program_id(0)
batch = tl.program_id(1)
indices = tl.arange(0, block_size)
matrix = batch * matrix_stride
diagonal_base = matrix + panel_start * n + panel_start
lower_row = tl.maximum(indices[:, None], indices[None, :])
lower_col = tl.minimum(indices[:, None], indices[None, :])
diagonal = tl.load(out + diagonal_base + lower_row * n + lower_col)
for j in range(block_size):
column = tl.sum(
tl.where(indices[None, :] == j, diagonal, 0.0), axis=1
)
diagonal_value = tl.sum(tl.where(indices == j, column, 0.0))
pivot = tl.sqrt(tl.maximum(diagonal_value, pivot_floor))
factor_column = tl.where(indices == j, pivot, column / pivot)
active_column = tl.where(indices >= j, factor_column, 0.0)
diagonal = tl.where(
indices[None, :] == j, active_column[:, None], diagonal
)
diagonal = tl.where(
(indices[:, None] > j) & (indices[None, :] > j),
diagonal - active_column[:, None] * active_column[None, :],
diagonal,
)
rows = trailing_start + row_tile * block_m + tl.arange(0, block_m)
row_mask = rows < n
tile = tl.load(
out + matrix + rows[:, None] * n + panel_start + indices[None, :],
mask=row_mask[:, None],
other=0.0,
)
for j in range(block_size):
current = tl.sum(
tl.where(indices[None, :] == j, tile, 0.0), axis=1
)
pivot = tl.sum(
tl.where(
(indices[:, None] == j) & (indices[None, :] == j),
diagonal,
0.0,
)
)
solved = current / pivot
factor_column = tl.sum(
tl.where(indices[None, :] == j, diagonal, 0.0), axis=1
)
tile = tl.where(
indices[None, :] == j,
solved[:, None],
tl.where(
indices[None, :] > j,
tile - solved[:, None] * factor_column[None, :],
tile,
),
)
tl.store(
out + matrix + rows[:, None] * n + panel_start + indices[None, :],
tile,
mask=row_mask[:, None],
)
if row_tile == 0:
tl.store(
out
+ diagonal_base
+ indices[:, None] * n
+ indices[None, :],
diagonal,
mask=indices[:, None] >= indices[None, :],
)
@triton.jit
def _trailing_update_kernel(
out,
source,
matrix_stride,
panel_start,
trailing_start,
n: tl.constexpr,
block_size: tl.constexpr,
block_m: tl.constexpr,
block_n: tl.constexpr,
dot_precision: tl.constexpr,
load_from_source: tl.constexpr,
):
tile_m = tl.program_id(0)
tile_n = tl.program_id(1)
if tile_m * block_m + block_m <= tile_n * block_n:
return
batch = tl.program_id(2)
rows = trailing_start + tile_m * block_m + tl.arange(0, block_m)
cols = trailing_start + tile_n * block_n + tl.arange(0, block_n)
reduction = tl.arange(0, block_size)
matrix = batch * matrix_stride
left = tl.load(
out + matrix + rows[:, None] * n + panel_start + reduction[None, :],
mask=(rows[:, None] < n),
other=0.0,
)
right_t = tl.load(
out + matrix + cols[None, :] * n + panel_start + reduction[:, None],
mask=(cols[None, :] < n),
other=0.0,
)
if dot_precision == "f16x2" or dot_precision == "f16x3":
left_hi = left.to(tl.float16)
right_hi = right_t.to(tl.float16)
left_lo = (left - left_hi.to(tl.float32)).to(tl.float16)
right_lo = (right_t - right_hi.to(tl.float32)).to(tl.float16)
update = tl.dot(left_hi, right_hi, out_dtype=tl.float32)
update = tl.dot(left_hi, right_lo, acc=update, out_dtype=tl.float32)
if dot_precision == "f16x3":
update = tl.dot(left_lo, right_hi, acc=update, out_dtype=tl.float32)
else:
update = tl.dot(left, right_t, input_precision=dot_precision)
pointers = out + matrix + rows[:, None] * n + cols[None, :]
mask = (rows[:, None] < n) & (cols[None, :] < n) & (rows[:, None] >= cols[None, :])
if load_from_source:
current = tl.load(
source + matrix + rows[:, None] * n + cols[None, :],
mask=mask,
other=0.0,
)
else:
current = tl.load(pointers, mask=mask, other=0.0)
tl.store(pointers, current - update, mask=mask)
@triton.jit
def _persistent_cholesky_kernel(
a,
out,
n: tl.constexpr,
block_n: tl.constexpr,
):
"""One persistent Triton program per matrix.
Lanes own output rows. The outer loop advances one Cholesky column at a
time; a CTA barrier makes the newly written column visible before it is
consumed by the following iteration. Keeping the matrix batch dimension
in the launch grid exposes abundant parallelism for the small benchmark
shapes without paying one kernel launch per column.
"""
batch = tl.program_id(0)
rows = tl.arange(0, block_n)
valid_row = rows < n
matrix_offset = batch * n * n
# Zero the strict upper triangle while the factorization writes the lower
# triangle. Doing this here avoids a separate torch.tril launch.
cols = tl.arange(0, block_n)
upper_ptrs = out + matrix_offset + rows[:, None] * n + cols[None, :]
upper_mask = valid_row[:, None] & (cols[None, :] < n) & (cols[None, :] > rows[:, None])
tl.store(upper_ptrs, 0.0, mask=upper_mask)
j = 0
while j < n:
# Every lane evaluates the diagonal recurrence. This deliberately
# duplicates only O(n^2) scalar work per CTA and avoids a broadcast or
# scratch allocation; the row update below is the O(n^3) term.
diagonal_sum = 0.0
k = 0
while k < j:
value = tl.load(out + matrix_offset + j * n + k)
diagonal_sum += value * value
k += 1
ajj = tl.load(a + matrix_offset + j * n + j)
diagonal = tl.sqrt(tl.maximum(ajj - diagonal_sum, 1.0e-6))
row_sum = tl.zeros((block_n,), dtype=tl.float32)
k = 0
while k < j:
lik = tl.load(
out + matrix_offset + rows * n + k,
mask=valid_row,
other=0.0,
)
ljk = tl.load(out + matrix_offset + j * n + k)
row_sum += lik * ljk
k += 1
aij = tl.load(
a + matrix_offset + rows * n + j,
mask=valid_row,
other=0.0,
)
value = (aij - row_sum) / diagonal
value = tl.where(rows == j, diagonal, value)
tl.store(
out + matrix_offset + rows * n + j,
value,
mask=valid_row & (rows >= j),
)
tl.debug_barrier()
j += 1
def _persistent_cholesky(data: torch.Tensor) -> torch.Tensor:
batch, n, _ = data.shape
out = torch.empty_like(data)
block_n = triton.next_power_of_2(n)
if block_n <= 64:
num_warps = 2
elif block_n <= 128:
num_warps = 4
else:
num_warps = 8
_persistent_cholesky_kernel[(batch,)](
data,
out,
n=n,
block_n=block_n,
num_warps=num_warps,
)
return out
def _blocked_cholesky(
data: torch.Tensor,
dot_precision: str = "tf32x3",
update_shape: tuple[int, int] | None = None,
update_warps: int = 4,
panel_m: int = 32,
panel_warps: int = 4,
fuse_factor_trsm: bool = False,
) -> torch.Tensor:
batch, n, _ = data.shape
out = torch.empty_like(data)
matrix_stride = n * n
_copy_first_panel_kernel[(triton.cdiv(n, 128), batch)](
data,
out,
matrix_stride,
n=n,
block_size=32,
block_m=128,
num_warps=4,
)
block_size = 32
if update_shape is None:
update_m = 128 if n >= 1024 else 32
update_n = 16 if n >= 1024 else 32
else:
update_m, update_n = update_shape
pivot_floor = 1.0e-4 if n >= 512 else 1.0e-6
for panel_start in range(0, n, block_size):
trailing_start = panel_start + block_size
if trailing_start == n:
_factor_diagonal_register_kernel[(batch,)](
out,
matrix_stride,
panel_start,
n=n,
block_size=block_size,
pivot_floor=pivot_floor,
num_warps=1,
)
continue
trailing = n - trailing_start
if fuse_factor_trsm:
_factor_trsm_fused_kernel[(triton.cdiv(trailing, panel_m), batch)](
out,
matrix_stride,
panel_start,
trailing_start,
n=n,
block_size=block_size,
block_m=panel_m,
pivot_floor=pivot_floor,
num_warps=panel_warps,
)
else:
_factor_diagonal_register_kernel[(batch,)](
out,
matrix_stride,
panel_start,
n=n,
block_size=block_size,
pivot_floor=pivot_floor,
num_warps=1,
)
_panel_trsm_register_kernel[(triton.cdiv(trailing, panel_m), batch)](
out,
matrix_stride,
panel_start,
trailing_start,
n=n,
block_size=block_size,
block_m=panel_m,
num_warps=panel_warps,
)
tiles_m = triton.cdiv(trailing, update_m)
tiles_n = triton.cdiv(trailing, update_n)
_trailing_update_kernel[(tiles_m, tiles_n, batch)](
out,
data,
matrix_stride,
panel_start,
trailing_start,
n=n,
block_size=block_size,
block_m=update_m,
block_n=update_n,
dot_precision=dot_precision,
load_from_source=panel_start == 0,
num_warps=update_warps,
)
zero_m = 32
zero_n = 128
_zero_upper_kernel[
(triton.cdiv(n, zero_m), triton.cdiv(n, zero_n), batch)
](
out,
matrix_stride,
n=n,
block_m=zero_m,
block_n=zero_n,
num_warps=4,
)
return out
def _left_looking_cholesky(data: torch.Tensor) -> torch.Tensor:
batch, n, _ = data.shape
out = torch.empty_like(data)
total = data.numel()
init_block = 256
_initialize_lower_kernel[(triton.cdiv(total, init_block),)](
data,
out,
total=total,
n=n,
block=init_block,
num_warps=8,
)
block_size = 32
block_m = 128
matrix_stride = n * n
pivot_floor = 1.0e-4 if n >= 1024 else 1.0e-6
for panel_start in range(0, n, block_size):
if panel_start:
_left_looking_panel_update_kernel[
(triton.cdiv(n - panel_start, block_m), batch)
](
out,
matrix_stride,
panel_start,
n=n,
block_size=block_size,
block_m=block_m,
reduction_block=32,
dot_precision="tf32x3",
num_warps=4,
)
_factor_diagonal_register_kernel[(batch,)](
out,
matrix_stride,
panel_start,
n=n,
block_size=block_size,
pivot_floor=pivot_floor,
num_warps=1,
)
trailing_start = panel_start + block_size
if trailing_start == n:
continue
_panel_trsm_register_kernel[
(triton.cdiv(n - trailing_start, 32), batch)
](
out,
matrix_stride,
panel_start,
trailing_start,
n=n,
block_size=block_size,
block_m=32,
num_warps=4,
)
return out
def custom_kernel(data: input_t) -> output_t:
batch = data.shape[0]
n = data.shape[-1]
if n <= 32:
return _register_cholesky(data)
if n == 64:
return _blocked_cholesky(data, "f16x3", (32, 32), 4)
if n == 512:
if batch >= 128:
return _blocked_cholesky(
data, "f16x3", (32, 128), 4, 128, 4, True
)
return _blocked_cholesky(data, "f16x3", (64, 64), 4, 32, 4, True)
if n == 1024:
panel_m = 128 if batch >= 32 else 32
return _blocked_cholesky(
data, "f16x3", (32, 128), 4, panel_m, 4, True
)
if n == 2048:
shape = (64, 64) if batch >= 8 else (32, 128)
return _blocked_cholesky(data, "f16x3", shape, 4, 32, 4, True)
if n == 4096 and batch >= 2:
return _blocked_cholesky(data, "f16x3", (64, 64), 4, 16, 2)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 758 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