Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
1.41ms
#167 of 337
2026-07-17

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.

mmapanel = tl.dot(panel, tl.trans(inverse), input_precision=IP)
num-warps = 8num_warps=8,
stages = 3for reduction_start in tl.range(0, NB, BK, num_stages=3):

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