submission 882512
Harshwardhan · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 849 lines, June 9 Researcher Reciprocity License v1.0.
submission_32_n256_parallel_updates.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-882512?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:09a9a06f5a76deb0222a1e6643ece6bebe72d3829155188ab344a1dd4d15bb65
license declaredunknown
license concludedunknown
authorsHarshwardhan
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
panel = tl.dot(panel, tl.trans(inverse), input_precision=IP)num-warps = 8
num_warps=8,stages = 3
for reduction_start in tl.range(0, NB, BK, num_stages=3):tile-k = 64
BK=64,tile-m = 128
BM=128,tile-n = 128
BN=128,Kernel source
submission_32_n256_parallel_updates.py849 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _chol_fused(A, L, n, BLOCK: tl.constexpr):
batch_index = tl.program_id(0).to(tl.int64)
offsets = tl.arange(0, BLOCK)
rows = offsets[:, None]
cols = offsets[None, :]
base = batch_index * n * n
values = tl.load(A + base + rows * n + cols)
for k in range(n):
column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
pivot = tl.sum(tl.where(offsets == k, column, 0.0), axis=0)
column = tl.where(offsets >= k, column / tl.sqrt(pivot), 0.0)
values = tl.where(cols == k, column[:, None], values)
values = tl.where(
cols > k,
values - column[:, None] * column[None, :],
values,
)
values = tl.where(rows >= cols, values, 0.0)
tl.store(L + base + rows * n + cols, values)
@triton.jit
def _chol_small_blocked(
A,
L,
n,
BS: tl.constexpr,
NSTEPS: tl.constexpr,
IP: tl.constexpr,
):
batch_index = tl.program_id(0).to(tl.int64)
offsets = tl.arange(0, BS)
local_rows = offsets[:, None]
local_cols = offsets[None, :]
base = batch_index * n * n
for tile_row in tl.static_range(0, NSTEPS):
matrix_rows = tile_row * BS + local_rows
for tile_col in tl.static_range(0, NSTEPS):
matrix_cols = tile_col * BS + local_cols
values = tl.load(A + base + matrix_rows * n + matrix_cols)
values = tl.where(matrix_rows >= matrix_cols, values, 0.0)
tl.store(L + base + matrix_rows * n + matrix_cols, values)
tl.debug_barrier()
for step in tl.static_range(0, NSTEPS):
block_start = step * BS
diagonal_pointer = (
L + base + (block_start + local_rows) * n + block_start + local_cols
)
diagonal = tl.load(diagonal_pointer)
for k in tl.static_range(0, BS):
column = tl.sum(
tl.where(local_cols == k, diagonal, 0.0),
axis=1,
)
pivot = tl.sum(
tl.where(offsets == k, column, 0.0),
axis=0,
)
column = tl.where(
offsets >= k,
column / tl.sqrt(pivot),
0.0,
)
diagonal = tl.where(local_cols == k, column[:, None], diagonal)
diagonal = tl.where(
local_cols > k,
diagonal - column[:, None] * column[None, :],
diagonal,
)
diagonal = tl.where(local_rows >= local_cols, diagonal, 0.0)
tl.store(diagonal_pointer, diagonal)
if step < NSTEPS - 1:
inverse = tl.zeros((BS, BS), dtype=tl.float32)
for k in tl.static_range(0, BS):
diagonal_row = tl.sum(
tl.where(local_rows == k, diagonal, 0.0),
axis=0,
)
pivot = tl.sum(
tl.where(offsets == k, diagonal_row, 0.0),
axis=0,
)
accumulated = tl.sum(
tl.where(offsets < k, diagonal_row, 0.0)[:, None]
* inverse,
axis=0,
)
identity_row = tl.where(offsets == k, 1.0, 0.0)
solved_row = (identity_row - accumulated) / pivot
inverse = tl.where(
local_rows == k,
solved_row[None, :],
inverse,
)
for panel_tile in tl.static_range(step + 1, NSTEPS):
panel_start = panel_tile * BS
panel_pointer = (
L
+ base
+ (panel_start + local_rows) * n
+ block_start
+ local_cols
)
panel = tl.load(panel_pointer)
panel = tl.dot(panel, tl.trans(inverse), input_precision=IP)
tl.store(panel_pointer, panel)
tl.debug_barrier()
for col_tile in tl.static_range(step + 1, NSTEPS):
col_start = col_tile * BS
right_panel = tl.load(
L
+ base
+ (col_start + local_rows) * n
+ block_start
+ local_cols
)
for row_tile in tl.static_range(col_tile, NSTEPS):
row_start = row_tile * BS
left_panel = tl.load(
L
+ base
+ (row_start + local_rows) * n
+ block_start
+ local_cols
)
trailing_pointer = (
L
+ base
+ (row_start + local_rows) * n
+ col_start
+ local_cols
)
trailing = tl.load(trailing_pointer)
trailing -= tl.dot(
left_panel,
tl.trans(right_panel),
input_precision=IP,
)
tl.store(trailing_pointer, trailing)
tl.debug_barrier()
@triton.jit
def _chol_medium_blocked(
A,
L,
n,
BS: tl.constexpr,
NSTEPS: tl.constexpr,
IP: tl.constexpr,
):
"""Compact one-CTA tiled Cholesky for the n=256, batch=64 branch.
Unlike _chol_small_blocked, tile loops stay as compiler-generated loops.
This avoids the explosive code size of the old statically unrolled n=256
probe while preserving the same proven block Cholesky arithmetic.
"""
batch_index = tl.program_id(0).to(tl.int64)
offsets = tl.arange(0, BS)
local_rows = offsets[:, None]
local_cols = offsets[None, :]
base = batch_index * n * n
# Copy and triangularize in the same launch. Flattening the tile traversal
# keeps the control-flow graph small.
for tile_id in tl.range(
0,
NSTEPS * NSTEPS,
loop_unroll_factor=1,
):
tile_row = tile_id // NSTEPS
tile_col = tile_id - tile_row * NSTEPS
matrix_rows = tile_row * BS + local_rows
matrix_cols = tile_col * BS + local_cols
values = tl.load(A + base + matrix_rows * n + matrix_cols)
values = tl.where(matrix_rows >= matrix_cols, values, 0.0)
tl.store(L + base + matrix_rows * n + matrix_cols, values)
tl.debug_barrier()
for step in tl.range(
0,
NSTEPS,
loop_unroll_factor=1,
disable_licm=True,
):
block_start = step * BS
diagonal_pointer = (
L + base + (block_start + local_rows) * n + block_start + local_cols
)
diagonal = tl.load(diagonal_pointer)
# The fixed 32x32 POTRF is intentionally unrolled; only the growing
# tile traversal remains dynamic.
for k in tl.static_range(0, BS):
column = tl.sum(
tl.where(local_cols == k, diagonal, 0.0),
axis=1,
)
pivot = tl.sum(
tl.where(offsets == k, column, 0.0),
axis=0,
)
column = tl.where(
offsets >= k,
column / tl.sqrt(pivot),
0.0,
)
diagonal = tl.where(local_cols == k, column[:, None], diagonal)
diagonal = tl.where(
local_cols > k,
diagonal - column[:, None] * column[None, :],
diagonal,
)
diagonal = tl.where(local_rows >= local_cols, diagonal, 0.0)
tl.store(diagonal_pointer, diagonal)
if step < NSTEPS - 1:
# Form inv(L_kk) once, then use tensor-core GEMMs for the panel
# solve and trailing lower-triangle update.
inverse = tl.zeros((BS, BS), dtype=tl.float32)
for k in tl.static_range(0, BS):
diagonal_row = tl.sum(
tl.where(local_rows == k, diagonal, 0.0),
axis=0,
)
pivot = tl.sum(
tl.where(offsets == k, diagonal_row, 0.0),
axis=0,
)
accumulated = tl.sum(
tl.where(offsets < k, diagonal_row, 0.0)[:, None]
* inverse,
axis=0,
)
identity_row = tl.where(offsets == k, 1.0, 0.0)
solved_row = (identity_row - accumulated) / pivot
inverse = tl.where(
local_rows == k,
solved_row[None, :],
inverse,
)
for panel_tile in tl.range(
step + 1,
NSTEPS,
loop_unroll_factor=1,
):
panel_start = panel_tile * BS
panel_pointer = (
L
+ base
+ (panel_start + local_rows) * n
+ block_start
+ local_cols
)
panel = tl.load(panel_pointer)
panel = tl.dot(
panel,
tl.trans(inverse),
input_precision=IP,
)
tl.store(panel_pointer, panel)
tl.debug_barrier()
for col_tile in tl.range(
step + 1,
NSTEPS,
loop_unroll_factor=1,
):
col_start = col_tile * BS
right_panel = tl.load(
L
+ base
+ (col_start + local_rows) * n
+ block_start
+ local_cols
)
for row_tile in tl.range(
col_tile,
NSTEPS,
loop_unroll_factor=1,
):
row_start = row_tile * BS
left_panel = tl.load(
L
+ base
+ (row_start + local_rows) * n
+ block_start
+ local_cols
)
trailing_pointer = (
L
+ base
+ (row_start + local_rows) * n
+ col_start
+ local_cols
)
trailing = tl.load(trailing_pointer)
trailing -= tl.dot(
left_panel,
tl.trans(right_panel),
input_precision=IP,
)
tl.store(trailing_pointer, trailing)
tl.debug_barrier()
@triton.jit
def _chol_medium_panel_step(
L,
n,
block_start,
BS: tl.constexpr,
IP: tl.constexpr,
):
"""Factor one diagonal tile and solve its complete panel."""
batch_index = tl.program_id(0).to(tl.int64)
offsets = tl.arange(0, BS)
local_rows = offsets[:, None]
local_cols = offsets[None, :]
base = batch_index * n * n
diagonal_pointer = (
L + base + (block_start + local_rows) * n + block_start + local_cols
)
diagonal = tl.load(diagonal_pointer)
for k in tl.static_range(0, BS):
column = tl.sum(
tl.where(local_cols == k, diagonal, 0.0),
axis=1,
)
pivot = tl.sum(
tl.where(offsets == k, column, 0.0),
axis=0,
)
column = tl.where(
offsets >= k,
column / tl.sqrt(pivot),
0.0,
)
diagonal = tl.where(local_cols == k, column[:, None], diagonal)
diagonal = tl.where(
local_cols > k,
diagonal - column[:, None] * column[None, :],
diagonal,
)
diagonal = tl.where(local_rows >= local_cols, diagonal, 0.0)
tl.store(diagonal_pointer, diagonal)
if block_start + BS < n:
inverse = tl.zeros((BS, BS), dtype=tl.float32)
for k in tl.static_range(0, BS):
diagonal_row = tl.sum(
tl.where(local_rows == k, diagonal, 0.0),
axis=0,
)
pivot = tl.sum(
tl.where(offsets == k, diagonal_row, 0.0),
axis=0,
)
accumulated = tl.sum(
tl.where(offsets < k, diagonal_row, 0.0)[:, None]
* inverse,
axis=0,
)
identity_row = tl.where(offsets == k, 1.0, 0.0)
solved_row = (identity_row - accumulated) / pivot
inverse = tl.where(
local_rows == k,
solved_row[None, :],
inverse,
)
for panel_start in tl.range(
block_start + BS,
n,
BS,
loop_unroll_factor=1,
):
panel_pointer = (
L
+ base
+ (panel_start + local_rows) * n
+ block_start
+ local_cols
)
panel = tl.load(panel_pointer)
panel = tl.dot(
panel,
tl.trans(inverse),
input_precision=IP,
)
tl.store(panel_pointer, panel)
@triton.jit
def _chol_medium_update_step(
L,
n,
block_start,
BS: tl.constexpr,
IP: tl.constexpr,
):
"""Update one independent lower-triangular trailing tile."""
row_tile = tl.program_id(0)
col_tile = tl.program_id(1)
batch_index = tl.program_id(2).to(tl.int64)
if row_tile < col_tile:
return
offsets = tl.arange(0, BS)
local_rows = offsets[:, None]
local_cols = offsets[None, :]
base = batch_index * n * n
trailing_start = block_start + BS
matrix_rows = trailing_start + row_tile * BS + local_rows
matrix_cols = trailing_start + col_tile * BS + local_cols
left_panel = tl.load(
L + base + matrix_rows * n + block_start + local_cols
)
right_panel = tl.load(
L
+ base
+ (trailing_start + col_tile * BS + local_rows) * n
+ block_start
+ local_cols
)
trailing_pointer = L + base + matrix_rows * n + matrix_cols
trailing = tl.load(trailing_pointer)
trailing -= tl.dot(
left_panel,
tl.trans(right_panel),
input_precision=IP,
)
tl.store(trailing_pointer, trailing)
@triton.jit
def _syrk_wide(
L,
n,
block_start,
NB: tl.constexpr,
BM: tl.constexpr,
BN: tl.constexpr,
BK: tl.constexpr,
IP: tl.constexpr,
):
row_tile = tl.program_id(0)
col_tile = tl.program_id(1)
if row_tile < col_tile:
return
trailing_start = block_start + NB
matrix_rows = trailing_start + row_tile * BM + tl.arange(0, BM)
matrix_cols = trailing_start + col_tile * BN + tl.arange(0, BN)
reduction_offsets = tl.arange(0, BK)
row_pointers = matrix_rows[:, None].to(tl.int64) * n
row_mask = matrix_rows[:, None] < n
col_mask = matrix_cols[:, None] < n
accumulator = tl.zeros((BM, BN), dtype=tl.float32)
for reduction_start in tl.range(0, NB, BK, num_stages=3):
left = tl.load(
L
+ row_pointers
+ block_start
+ reduction_start
+ reduction_offsets[None, :],
mask=row_mask,
other=0.0,
)
right = tl.load(
L
+ matrix_cols[:, None].to(tl.int64) * n
+ block_start
+ reduction_start
+ reduction_offsets[None, :],
mask=col_mask,
other=0.0,
)
accumulator = tl.dot(
left,
tl.trans(right),
accumulator,
input_precision=IP,
)
trailing_pointer = L + row_pointers + matrix_cols[None, :]
store_mask = (
(matrix_rows[:, None] >= matrix_cols[None, :])
& row_mask
& (matrix_cols[None, :] < n)
)
trailing = tl.load(trailing_pointer, mask=store_mask, other=0.0)
tl.store(
trailing_pointer,
trailing - accumulator,
mask=store_mask,
)
@triton.jit
def _syrk_left_diagonal(
L,
n,
block_start,
reduction_size,
NB: tl.constexpr,
BM: tl.constexpr,
BN: tl.constexpr,
BK: tl.constexpr,
IP: tl.constexpr,
):
row_tile = tl.program_id(0)
col_tile = tl.program_id(1)
if row_tile < col_tile:
return
matrix_rows = block_start + row_tile * BM + tl.arange(0, BM)
matrix_cols = block_start + col_tile * BN + tl.arange(0, BN)
reduction_offsets = tl.arange(0, BK)
row_pointers = matrix_rows[:, None].to(tl.int64) * n
row_mask = matrix_rows[:, None] < block_start + NB
col_mask = matrix_cols[:, None] < block_start + NB
accumulator = tl.zeros((BM, BN), dtype=tl.float32)
for reduction_start in tl.range(
0,
reduction_size,
BK,
num_stages=3,
):
left = tl.load(
L
+ row_pointers
+ reduction_start
+ reduction_offsets[None, :],
mask=row_mask,
other=0.0,
)
right = tl.load(
L
+ matrix_cols[:, None].to(tl.int64) * n
+ reduction_start
+ reduction_offsets[None, :],
mask=col_mask,
other=0.0,
)
accumulator = tl.dot(
left,
tl.trans(right),
accumulator,
input_precision=IP,
)
target_pointer = L + row_pointers + matrix_cols[None, :]
store_mask = (
(matrix_rows[:, None] >= matrix_cols[None, :])
& row_mask
& (matrix_cols[None, :] < block_start + NB)
)
target = tl.load(target_pointer, mask=store_mask, other=0.0)
tl.store(
target_pointer,
target - accumulator,
mask=store_mask,
)
def _tf32_addmm_in_place(target, left, right):
target.addmm_(left, right, beta=1.0, alpha=-1.0)
def _run_with_tf32(factor, data, block_size):
previous_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
return factor(data, block_size)
finally:
torch.backends.cuda.matmul.allow_tf32 = previous_tf32
def _right_solve_blocked(panel, diagonal_factor, solve_block):
upper = diagonal_factor.mT
width = panel.shape[1]
for solve_start in range(0, width, solve_block):
solve_end = solve_start + solve_block
current = panel[:, solve_start:solve_end]
diagonal = upper[
solve_start:solve_end,
solve_start:solve_end,
]
torch.linalg.solve_triangular(
diagonal,
current,
upper=True,
left=False,
out=current,
)
if solve_end < width:
remaining = panel[:, solve_end:]
cross = upper[solve_start:solve_end, solve_end:]
_tf32_addmm_in_place(remaining, current, cross)
def _factor_wide_right(data, block_size):
batch, n, _ = data.shape
output = torch.tril(data)
for batch_index in range(batch):
matrix = output[batch_index]
for block_start in range(0, n, block_size):
block_end = block_start + block_size
diagonal_view = matrix[
block_start:block_end,
block_start:block_end,
]
diagonal_factor = torch.linalg.cholesky_ex(
diagonal_view,
check_errors=False,
).L
diagonal_view.copy_(diagonal_factor)
if block_end < n:
panel = matrix[block_end:, block_start:block_end]
_right_solve_blocked(
panel,
diagonal_factor,
512,
)
trailing_size = n - block_end
tiles = triton.cdiv(trailing_size, 128)
_syrk_wide[(tiles, tiles)](
matrix,
n,
block_start,
NB=block_size,
BM=128,
BN=128,
BK=64,
IP="tf32",
num_warps=8,
)
return output
def _factor_wide_left(data, block_size):
batch, n, _ = data.shape
output = torch.tril(data)
for batch_index in range(batch):
matrix = output[batch_index]
for block_start in range(0, n, block_size):
block_end = block_start + block_size
if block_start > 0:
diagonal_tiles = triton.cdiv(block_size, 128)
_syrk_left_diagonal[(diagonal_tiles, diagonal_tiles)](
matrix,
n,
block_start,
block_start,
NB=block_size,
BM=128,
BN=128,
BK=64,
IP="tf32",
num_warps=8,
)
if block_end < n:
panel_view = matrix[
block_end:,
block_start:block_end,
]
left_history = matrix[block_end:, :block_start]
right_history = matrix[
block_start:block_end,
:block_start,
].mT
_tf32_addmm_in_place(
panel_view,
left_history,
right_history,
)
diagonal_view = matrix[
block_start:block_end,
block_start:block_end,
]
diagonal_factor = torch.linalg.cholesky_ex(
diagonal_view,
check_errors=False,
).L
diagonal_view.copy_(diagonal_factor)
if block_end < n:
panel = matrix[block_end:, block_start:block_end]
_right_solve_blocked(
panel,
diagonal_factor,
512,
)
return output
def _factor_medium_parallel(data, block_size):
"""Blocked Cholesky with B200-wide parallel trailing updates."""
batch, n, _ = data.shape
output = torch.tril(data)
for block_start in range(0, n, block_size):
_chol_medium_panel_step[(batch,)](
output,
n,
block_start,
BS=block_size,
IP="tf32x3",
num_warps=4,
)
trailing_size = n - block_start - block_size
if trailing_size > 0:
tiles = triton.cdiv(trailing_size, block_size)
_chol_medium_update_step[(tiles, tiles, batch)](
output,
n,
block_start,
BS=block_size,
IP="tf32x3",
num_warps=4,
)
return output
def _factor_separately(data):
batch = data.shape[0]
output = torch.empty_like(data)
info = torch.empty((batch,), device=data.device, dtype=torch.int32)
for batch_index in range(batch):
torch.linalg.cholesky_ex(
data[batch_index],
check_errors=False,
out=(output[batch_index], info[batch_index]),
)
return output
def custom_kernel(data):
batch, n, _ = data.shape
if batch == 1 and n == 8192:
return _run_with_tf32(_factor_wide_right, data, 4096)
if batch == 1 and n == 16384:
return _run_with_tf32(_factor_wide_left, data, 4096)
if batch == 1 and n == 32768:
return _run_with_tf32(_factor_wide_left, data, 2048)
if n == 32:
output = torch.empty_like(data)
_chol_fused[(batch,)](
data,
output,
n,
BLOCK=32,
num_warps=2,
)
return output
if n == 64:
output = torch.empty_like(data)
_chol_small_blocked[(batch,)](
data,
output,
n,
BS=32,
NSTEPS=2,
IP="tf32x3",
num_warps=4,
)
return output
if n == 128:
output = torch.empty_like(data)
_chol_small_blocked[(batch,)](
data,
output,
n,
BS=32,
NSTEPS=4,
IP="tf32x3",
num_warps=4,
)
return output
if n == 256:
return _factor_medium_parallel(data, 32)
if n == 2048 and batch <= 2:
return _factor_separately(data)
if n == 4096 and batch == 2:
return _factor_separately(data)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 849 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