Skip to content
KernelIndex
Search⌘K

submission 909141

swb0387 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 244 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-909141?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.82ms
#232 of 337
2026-07-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3b2047f99a772f8ece173e4969db8eea6b7b7a7f873ff8cefe40a5438f97fb40
license declaredunknown
license concludedunknown
authorsswb0387
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 8num_warps=8 if tile_size == 128 else 4,
stages = 3num_stages=3,

Kernel source

submission.py244 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl

from task import input_t, output_t


@triton.jit
def _cholesky_small_kernel(
    data,
    output,
    n: tl.constexpr,
    block_n: tl.constexpr,
    block_rows: tl.constexpr,
):
    """One Triton program factors one matrix."""
    batch = tl.program_id(0).to(tl.int64)
    base = batch * n * n
    k = tl.arange(0, block_n)

    for column in tl.range(0, n):
        pivot_row = tl.load(
            output + base + column * n + k,
            mask=k < column,
            other=0.0,
        ).to(tl.float32)
        diagonal_square = (
            tl.load(data + base + column * n + column).to(tl.float32)
            - tl.sum(pivot_row * pivot_row)
        )
        diagonal = tl.sqrt(tl.maximum(diagonal_square, 1.0e-20))
        tl.store(output + base + column * n + column, diagonal)

        for row_start in tl.range(column + 1, n, block_rows):
            rows = row_start + tl.arange(0, block_rows)
            row_mask = rows < n
            row_values = tl.load(
                output + base + rows[:, None] * n + k[None, :],
                mask=row_mask[:, None] & (k[None, :] < column),
                other=0.0,
            ).to(tl.float32)
            dot = tl.sum(row_values * pivot_row[None, :], axis=1)
            source = tl.load(
                data + base + rows * n + column,
                mask=row_mask,
                other=0.0,
            ).to(tl.float32)
            tl.store(
                output + base + rows * n + column,
                (source - dot) / diagonal,
                mask=row_mask,
            )

        upper_rows = tl.arange(0, block_n)
        tl.store(
            output + base + upper_rows * n + column,
            0.0,
            mask=upper_rows < column,
        )


@triton.jit
def _syrk_update_kernel(
    matrix,
    n,
    panel_start,
    panel_end,
    block_m: tl.constexpr,
    block_n: tl.constexpr,
    block_k: tl.constexpr,
):
    """Lower-triangular A22 -= L21 @ L21.T using Tensor Cores."""
    batch = tl.program_id(0).to(tl.int64)
    tile_row = tl.program_id(1)
    tile_column = tl.program_id(2)

    if tile_column <= tile_row:
        rows = panel_end + tile_row * block_m + tl.arange(0, block_m)
        columns = (
            panel_end + tile_column * block_n + tl.arange(0, block_n)
        )
        accumulator = tl.zeros((block_m, block_n), dtype=tl.float32)

        for k_start in tl.range(
            panel_start,
            panel_end,
            block_k,
            num_stages=3,
        ):
            k = k_start + tl.arange(0, block_k)
            left = tl.load(
                matrix
                + batch * n * n
                + rows[:, None] * n
                + k[None, :],
                mask=(rows[:, None] < n) & (k[None, :] < panel_end),
                other=0.0,
            )
            right = tl.load(
                matrix
                + batch * n * n
                + k[:, None]
                + columns[None, :] * n,
                mask=(k[:, None] < panel_end) & (columns[None, :] < n),
                other=0.0,
            )
            accumulator += tl.dot(
                left,
                right,
                input_precision="tf32",
            )

        addresses = (
            matrix
            + batch * n * n
            + rows[:, None] * n
            + columns[None, :]
        )
        mask = (
            (rows[:, None] < n)
            & (columns[None, :] < n)
            & (rows[:, None] >= columns[None, :])
        )
        current = tl.load(addresses, mask=mask, other=0.0)
        tl.store(addresses, current - accumulator, mask=mask)


@triton.jit
def _zero_upper_kernel(
    matrix,
    elements,
    n: tl.constexpr,
    block: tl.constexpr,
):
    offsets = tl.program_id(0).to(tl.int64) * block + tl.arange(0, block)
    matrix_offset = offsets % (n * n)
    row = matrix_offset // n
    column = matrix_offset - row * n
    tl.store(
        matrix + offsets,
        0.0,
        mask=(offsets < elements) & (column > row),
    )


def _blocked_cholesky(
    data: torch.Tensor,
    block_size: int = 512,
    tile_size: int = 64,
) -> torch.Tensor:
    """Right-looking Cholesky with a custom Tensor Core trailing update."""
    batch, n, _ = data.shape
    output = data.clone()

    for panel_start in range(0, n, block_size):
        panel_end = min(panel_start + block_size, n)
        diagonal_view = output[
            :, panel_start:panel_end, panel_start:panel_end
        ]
        diagonal = torch.linalg.cholesky_ex(
            diagonal_view,
            check_errors=False,
        ).L
        diagonal_view.copy_(diagonal)

        if panel_end == n:
            break

        right_hand_side = output[
            :, panel_end:, panel_start:panel_end
        ].transpose(1, 2)
        panel = torch.linalg.solve_triangular(
            diagonal,
            right_hand_side,
            upper=False,
            left=True,
        ).transpose(1, 2)
        output[:, panel_end:, panel_start:panel_end].copy_(panel)

        active = n - panel_end
        grid = (
            batch,
            triton.cdiv(active, tile_size),
            triton.cdiv(active, tile_size),
        )
        _syrk_update_kernel[grid](
            output,
            n,
            panel_start,
            panel_end,
            block_m=tile_size,
            block_n=tile_size,
            block_k=32,
            num_warps=8 if tile_size == 128 else 4,
            num_stages=3,
        )

    elements = output.numel()
    _zero_upper_kernel[(triton.cdiv(elements, 1024),)](
        output,
        elements,
        n=n,
        block=1024,
        num_warps=8,
    )
    return output


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n == 32:
        output = torch.empty_like(data)
        _cholesky_small_kernel[(batch,)](
            data,
            output,
            n=n,
            block_n=32,
            block_rows=32,
            num_warps=1,
            num_stages=1,
        )
        return output
    if n == 1024 and batch == 60:
        return _blocked_cholesky(
            data,
            block_size=256,
            tile_size=64,
        )
    if (
        (n == 2048 and batch == 2)
        or (n == 4096 and batch == 2)
    ):
        return _blocked_cholesky(data)
    if n >= 16384:
        return _blocked_cholesky(
            data,
            block_size=1024,
            tile_size=128,
        )
    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 244 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