Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
1.84ms
#238 of 337
2026-07-19

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.

mmaupdate = tl.dot(left, tl.trans(right), input_precision="tf32x3")
num-warps = 1num_warps=1,
stages = 2num_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