submission 887406
ikudrautsau · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 242 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-887406?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:46af2fa80f1a8e7a7fbc147931b7be665a5c17958cdc85a59318c7fca2b29996
license declaredunknown
license concludedunknown
authorsikudrautsau
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
update = tl.dot(left, tl.trans(right), input_precision="tf32x3")num-warps = 1
num_warps=1,stages = 2
num_stages=2,Kernel source
submission.py242 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
BLOCK = 32
PANEL_GROUP = 4
@triton.jit
def _cholesky32_kernel(input_ptr, output_ptr, matrix_stride: tl.constexpr):
matrix = tl.program_id(0)
ids = tl.arange(0, 32)
rows = ids[:, None]
cols = ids[None, :]
offsets = matrix * matrix_stride + rows * 32 + cols
values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)
for k in tl.static_range(0, 32):
row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
diagonal -= tl.sum(tl.where(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)
@triton.jit
def _copy_lower_kernel(input_ptr, output_ptr, n: tl.constexpr, total_elements):
offsets = tl.program_id(0) * 512 + tl.arange(0, 512)
mask = offsets < total_elements
matrix_offsets = offsets % (n * n)
rows = matrix_offsets // n
cols = matrix_offsets % n
values = tl.load(input_ptr + offsets, mask=mask, other=0.0)
tl.store(output_ptr + offsets, tl.where(rows >= cols, values, 0.0), mask=mask)
@triton.jit
def _factor_diagonal_kernel(
output_ptr,
n: tl.constexpr,
start,
block_size,
BLOCK_SIZE: tl.constexpr,
):
batch = tl.program_id(0)
ids = tl.arange(0, BLOCK_SIZE)
rows = ids[:, None]
cols = ids[None, :]
matrix = output_ptr + batch * n * n
offsets = (start + rows) * n + start + cols
valid = (rows < block_size) & (cols < block_size)
values = tl.load(matrix + offsets, mask=valid & (rows >= cols), other=0.0)
for k in tl.static_range(0, BLOCK_SIZE):
row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
diagonal -= tl.sum(tl.where(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(matrix + offsets, values, mask=valid & (rows >= cols))
@triton.jit
def _solve_panel_kernel(
output_ptr,
n: tl.constexpr,
start,
block_size,
panel_start,
panel_rows,
BLOCK_SIZE: tl.constexpr,
PANEL_ROWS: tl.constexpr,
):
program = tl.program_id(0)
row_programs = tl.cdiv(panel_rows, PANEL_ROWS)
batch = program // row_programs
row_group = program % row_programs
row_ids = tl.arange(0, PANEL_ROWS)
col_ids = tl.arange(0, BLOCK_SIZE)
panel_row = panel_start + row_group * PANEL_ROWS + row_ids
matrix = output_ptr + batch * n * n
values = tl.load(
matrix + panel_row[:, None] * n + start + col_ids[None, :],
mask=(panel_row[:, None] < panel_start + panel_rows)
& (col_ids[None, :] < block_size),
other=0.0,
)
for j in tl.static_range(0, BLOCK_SIZE):
diagonal_row = tl.load(
matrix + (start + j) * n + start + col_ids,
mask=(j < block_size) & (col_ids <= j),
other=0.0,
)
value = tl.sum(
tl.where(col_ids[None, :] == j, values, 0.0), axis=1
)
value -= tl.sum(
tl.where(
col_ids[None, :] < j,
values * diagonal_row[None, :],
0.0,
),
axis=1,
)
diagonal = tl.sum(tl.where(col_ids == j, diagonal_row, 0.0), axis=0)
values = tl.where(col_ids[None, :] == j, value[:, None] / diagonal, values)
tl.store(
matrix + panel_row[:, None] * n + start + col_ids[None, :],
values,
mask=(panel_row[:, None] < panel_start + panel_rows)
& (col_ids[None, :] < block_size),
)
@triton.jit
def _update_trailing_kernel(
output_ptr,
n: tl.constexpr,
panel_start,
panel_width,
triangular_tiles,
BLOCK_SIZE: tl.constexpr,
):
program = tl.program_id(0)
batch = program // triangular_tiles
triangular_id = program % triangular_tiles
row_tile = tl.floor(
(tl.sqrt(8.0 * triangular_id.to(tl.float32) + 1.0) - 1.0) * 0.5
).to(tl.int32)
base = row_tile * (row_tile + 1) // 2
row_tile = tl.where(base > triangular_id, row_tile - 1, row_tile)
next_base = (row_tile + 1) * (row_tile + 2) // 2
row_tile = tl.where(next_base <= triangular_id, row_tile + 1, row_tile)
column_tile = triangular_id - row_tile * (row_tile + 1) // 2
rows = panel_start + row_tile * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
cols = panel_start + column_tile * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
reduction = tl.arange(0, BLOCK_SIZE)
matrix = output_ptr + batch * n * n
left = tl.load(
matrix + rows[:, None] * n + panel_start - panel_width + reduction[None, :],
mask=(rows[:, None] < n) & (reduction[None, :] < panel_width),
other=0.0,
)
right = tl.load(
matrix + cols[:, None] * n + panel_start - panel_width + reduction[None, :],
mask=(cols[:, None] < n) & (reduction[None, :] < panel_width),
other=0.0,
)
update = tl.dot(left, tl.trans(right), input_precision="tf32x3")
offsets = rows[:, None] * n + cols[None, :]
mask = (rows[:, None] < n) & (cols[None, :] < n)
mask = mask & (rows[:, None] >= cols[None, :])
values = tl.load(matrix + offsets, mask=mask, other=0.0)
tl.store(matrix + offsets, values - update, mask=mask)
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
block = 64 if n >= 16384 else BLOCK
factor_warps = 4 if n >= 16384 else 2
if n <= 512:
panel_group = 8
elif 4096 <= n <= 8192:
panel_group = 2
else:
panel_group = PANEL_GROUP
update_warps = 2 if n == 512 and batch >= 640 else 4
output = torch.empty_like(data)
if n == 32:
_cholesky32_kernel[(batch,)](
data,
output,
32 * 32,
num_warps=1,
)
return output
total_elements = data.numel()
_copy_lower_kernel[(triton.cdiv(total_elements, 512),)](
data,
output,
n,
total_elements,
num_warps=4,
)
for start in range(0, n, block):
block_size = min(block, n - start)
_factor_diagonal_kernel[(batch,)](
output,
n,
start,
block_size,
BLOCK_SIZE=block,
num_warps=factor_warps,
)
panel_start = start + block_size
remaining = n - panel_start
if remaining == 0:
continue
panel_programs = triton.cdiv(remaining, panel_group)
_solve_panel_kernel[(batch * panel_programs,)](
output,
n,
start,
block_size,
panel_start,
remaining,
BLOCK_SIZE=block,
PANEL_ROWS=panel_group,
num_warps=1,
)
trailing_tiles = triton.cdiv(remaining, block)
triangular_tiles = trailing_tiles * (trailing_tiles + 1) // 2
_update_trailing_kernel[(batch * triangular_tiles,)](
output,
n,
panel_start,
block_size,
triangular_tiles,
BLOCK_SIZE=block,
num_warps=update_warps,
num_stages=2,
)
return output
scrolls · 242 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