Skip to content
KernelIndex
Search⌘K

submission 930424

gpuseed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-930424?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.06ms
#123 of 337
2026-07-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:daac65de2f36db1d5584f3a55d0982d7c4366c53e48ae2bf330c50140e985c0f
license declaredunknown
license concludedunknown
authorsgpuseed
imported2026-08-26

Techniques

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

num-warps = 4num_warps=4,
shared-memory__shared__ float panels[1][2][32][32];

Kernel source

submission_1.py2629 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = True
torch.set_float32_matmul_precision("high")


_cuda_n32_source = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>

__global__ void cholesky32_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch
) {
    constexpr unsigned mask = 0xffffffffu;
    const int lane = threadIdx.x & 31;
    const int matrix = blockIdx.x * 4 + (threadIdx.x >> 5);
    if (matrix >= batch) {
        return;
    }

    const float* matrix_input = input + matrix * 32 * 32;
    float* matrix_output = output + matrix * 32 * 32;
    float row[32];

    #pragma unroll
    for (int col = 0; col < 32; ++col) {
        row[col] = matrix_input[lane * 32 + col];
    }

    #pragma unroll
    for (int pivot = 0; pivot < 32; ++pivot) {
        float value = row[pivot];
        #pragma unroll
        for (int col = 0; col < pivot; ++col) {
            value -= row[col] * __shfl_sync(mask, row[col], pivot);
        }

        if (lane == pivot) {
            row[pivot] = sqrtf(value);
        }
        __syncwarp(mask);

        const float diagonal = __shfl_sync(mask, row[pivot], pivot);
        if (lane > pivot) {
            row[pivot] = value / diagonal;
        }
        __syncwarp(mask);
    }

    #pragma unroll
    for (int col = 0; col < 32; ++col) {
        matrix_output[lane * 32 + col] = lane >= col ? row[col] : 0.0f;
    }
}

torch::Tensor cholesky32_cuda(torch::Tensor input, torch::Tensor output) {
    const int batch = static_cast<int>(input.size(0));
    const int blocks = (batch + 3) / 4;
    cholesky32_kernel<<<blocks, 128>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        batch
    );
    return output;
}

__global__ void cholesky64_recursive_kernel(
    const float* __restrict__ input,
    float* __restrict__ output
) {
    constexpr unsigned mask = 0xffffffffu;
    const int lane = threadIdx.x & 31;
    const int warp_in_matrix = (threadIdx.x >> 5) & 1;
    const int local_matrix = threadIdx.x >> 6;
    const int matrix = blockIdx.x + local_matrix;
    const float* matrix_input = input + matrix * 64 * 64;
    float* matrix_output = output + matrix * 64 * 64;

    // [column][row] keeps simultaneous row accesses bank-conflict free.
    __shared__ float panels[1][2][32][32];

    if (warp_in_matrix == 0) {
        float row[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            row[col] = matrix_input[lane * 64 + col];
        }

        #pragma unroll
        for (int pivot = 0; pivot < 32; ++pivot) {
            float value = row[pivot];
            #pragma unroll
            for (int col = 0; col < pivot; ++col) {
                value -= row[col] * __shfl_sync(mask, row[col], pivot);
            }
            if (lane == pivot) {
                row[pivot] = sqrtf(value);
            }
            __syncwarp(mask);
            const float diagonal = __shfl_sync(mask, row[pivot], pivot);
            if (lane > pivot) {
                row[pivot] = value / diagonal;
            }
            __syncwarp(mask);
        }

        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            const float value = lane >= col ? row[col] : 0.0f;
            panels[local_matrix][0][col][lane] = value;
            matrix_output[lane * 64 + col] = value;
            matrix_output[lane * 64 + col + 32] = 0.0f;
        }
    }
    __syncthreads();

    if (warp_in_matrix == 1) {
        float panel[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            float value = matrix_input[(lane + 32) * 64 + col];
            #pragma unroll
            for (int inner = 0; inner < col; ++inner) {
                value -= (
                    panel[inner]
                    * panels[local_matrix][0][inner][col]
                );
            }
            panel[col] = (
                value / panels[local_matrix][0][col][col]
            );
            panels[local_matrix][1][col][lane] = panel[col];
            matrix_output[(lane + 32) * 64 + col] = panel[col];
        }
    }
    __syncthreads();

    if (warp_in_matrix == 1) {
        float row[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            float value = matrix_input[(lane + 32) * 64 + col + 32];
            #pragma unroll
            for (int inner = 0; inner < 32; ++inner) {
                value -= (
                    panels[local_matrix][1][inner][lane]
                    * panels[local_matrix][1][inner][col]
                );
            }
            row[col] = lane >= col ? value : 0.0f;
        }

        #pragma unroll
        for (int pivot = 0; pivot < 32; ++pivot) {
            float value = row[pivot];
            #pragma unroll
            for (int col = 0; col < pivot; ++col) {
                value -= row[col] * __shfl_sync(mask, row[col], pivot);
            }
            if (lane == pivot) {
                row[pivot] = sqrtf(value);
            }
            __syncwarp(mask);
            const float diagonal = __shfl_sync(mask, row[pivot], pivot);
            if (lane > pivot) {
                row[pivot] = value / diagonal;
            }
            __syncwarp(mask);
        }

        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            matrix_output[(lane + 32) * 64 + col + 32] = (
                lane >= col ? row[col] : 0.0f
            );
        }
    }
}

torch::Tensor cholesky64_recursive_cuda(
    torch::Tensor input,
    torch::Tensor output
) {
    const int batch = static_cast<int>(input.size(0));
    cholesky64_recursive_kernel<<<batch, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>()
    );
    return output;
}

template <int stride>
__global__ void factor64_strided_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int tile_start
) {
    constexpr unsigned mask = 0xffffffffu;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int matrix = blockIdx.x;
    const float* matrix_input = input + matrix * stride * stride;
    float* matrix_output = output + matrix * stride * stride;
    __shared__ float panels[2][32][32];

    if (warp == 0) {
        // Lane is the column during global transfer, so each row is one
        // coalesced transaction. Shared storage is [column][row].
        #pragma unroll
        for (int row_idx = 0; row_idx < 32; ++row_idx) {
            panels[0][lane][row_idx] = matrix_input[
                (tile_start + row_idx) * stride + tile_start + lane
            ];
        }
    } else {
        #pragma unroll
        for (int row_idx = 0; row_idx < 32; ++row_idx) {
            panels[1][lane][row_idx] = matrix_input[
                (tile_start + 32 + row_idx) * stride
                + tile_start + lane
            ];
        }
    }
    __syncthreads();

    if (warp == 0) {
        float factor_row[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            factor_row[col] = panels[0][col][lane];
        }
        #pragma unroll
        for (int pivot = 0; pivot < 32; ++pivot) {
            float value = factor_row[pivot];
            #pragma unroll
            for (int col = 0; col < pivot; ++col) {
                value -= (
                    factor_row[col]
                    * __shfl_sync(mask, factor_row[col], pivot)
                );
            }
            if (lane == pivot) {
                factor_row[pivot] = sqrtf(value);
            }
            __syncwarp(mask);
            const float diagonal = __shfl_sync(
                mask, factor_row[pivot], pivot
            );
            if (lane > pivot) {
                factor_row[pivot] = value / diagonal;
            }
            __syncwarp(mask);
        }
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            panels[0][col][lane] = (
                lane >= col ? factor_row[col] : 0.0f
            );
        }
    }
    __syncthreads();

    if (warp == 1) {
        float panel[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            float value = panels[1][col][lane];
            #pragma unroll
            for (int inner = 0; inner < col; ++inner) {
                value -= panel[inner] * panels[0][inner][col];
            }
            panel[col] = value / panels[0][col][col];
            panels[1][col][lane] = panel[col];
        }
    } else {
        // Store L00 and its upper-right zero quadrant coalescently.
        #pragma unroll
        for (int row_idx = 0; row_idx < 32; ++row_idx) {
            matrix_output[
                (tile_start + row_idx) * stride + tile_start + lane
            ] = panels[0][lane][row_idx];
            matrix_output[
                (tile_start + row_idx) * stride
                + tile_start + 32 + lane
            ] = 0.0f;
            if (tile_start == 0) {
                matrix_output[row_idx * stride + 64 + lane] = 0.0f;
                matrix_output[row_idx * stride + 96 + lane] = 0.0f;
            }
        }
    }
    __syncthreads();

    if (warp == 1) {
        // Store L10 and reuse the first tile buffer for A11.
        #pragma unroll
        for (int row_idx = 0; row_idx < 32; ++row_idx) {
            matrix_output[
                (tile_start + 32 + row_idx) * stride
                + tile_start + lane
            ] = panels[1][lane][row_idx];
            panels[0][lane][row_idx] = matrix_input[
                (tile_start + 32 + row_idx) * stride
                + tile_start + 32 + lane
            ];
        }
    }
    __syncthreads();

    if (warp == 1) {
        float factor_row[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            float value = panels[0][col][lane];
            #pragma unroll
            for (int inner = 0; inner < 32; ++inner) {
                value -= panels[1][inner][lane] * panels[1][inner][col];
            }
            factor_row[col] = lane >= col ? value : 0.0f;
        }
        #pragma unroll
        for (int pivot = 0; pivot < 32; ++pivot) {
            float value = factor_row[pivot];
            #pragma unroll
            for (int col = 0; col < pivot; ++col) {
                value -= (
                    factor_row[col]
                    * __shfl_sync(mask, factor_row[col], pivot)
                );
            }
            if (lane == pivot) {
                factor_row[pivot] = sqrtf(value);
            }
            __syncwarp(mask);
            const float diagonal = __shfl_sync(
                mask, factor_row[pivot], pivot
            );
            if (lane > pivot) {
                factor_row[pivot] = value / diagonal;
            }
            __syncwarp(mask);
        }
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            panels[0][col][lane] = (
                lane >= col ? factor_row[col] : 0.0f
            );
        }
        __syncwarp(mask);
        #pragma unroll
        for (int row_idx = 0; row_idx < 32; ++row_idx) {
            matrix_output[
                (tile_start + 32 + row_idx) * stride
                + tile_start + 32 + lane
            ] = panels[0][lane][row_idx];
            if (tile_start == 0) {
                matrix_output[
                    (row_idx + 32) * stride + 64 + lane
                ] = 0.0f;
                matrix_output[
                    (row_idx + 32) * stride + 96 + lane
                ] = 0.0f;
            }
        }
    }
}

torch::Tensor factor64_strided128_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start
) {
    const int batch = static_cast<int>(input.size(0));
    const int stride = static_cast<int>(input.size(1));
    if (stride == 128) {
        factor64_strided_kernel<128><<<batch, 64>>>(
            input.data_ptr<float>(),
            output.data_ptr<float>(),
            static_cast<int>(tile_start)
        );
    } else if (stride == 256) {
        factor64_strided_kernel<256><<<batch, 64>>>(
            input.data_ptr<float>(),
            output.data_ptr<float>(),
            static_cast<int>(tile_start)
        );
    } else if (stride == 512) {
        factor64_strided_kernel<512><<<batch, 64>>>(
            input.data_ptr<float>(),
            output.data_ptr<float>(),
            static_cast<int>(tile_start)
        );
    } else if (stride == 1024) {
        factor64_strided_kernel<1024><<<batch, 64>>>(
            input.data_ptr<float>(),
            output.data_ptr<float>(),
            static_cast<int>(tile_start)
        );
    } else {
        factor64_strided_kernel<2048><<<batch, 64>>>(
            input.data_ptr<float>(),
            output.data_ptr<float>(),
            static_cast<int>(tile_start)
        );
    }
    return output;
}

template <int stride>
__global__ void trsm_panel64_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int tile_start,
    int panel_count,
    int panel_base
) {
    const int thread = threadIdx.x;
    const int matrix = blockIdx.x / panel_count;
    const int panel_id = blockIdx.x - matrix * panel_count;
    const int panel_start = panel_base + panel_id * 64;
    const float* matrix_input = input + matrix * stride * stride;
    float* matrix_output = output + matrix * stride * stride;
    __shared__ float lower[64][64];
    __shared__ float panel[64][64];

    // Thread is the column during transfer, giving coalesced rows.
    #pragma unroll
    for (int row = 0; row < 64; ++row) {
        lower[thread][row] = matrix_output[
            (tile_start + row) * stride + tile_start + thread
        ];
        panel[thread][row] = matrix_input[
            (panel_start + row) * stride + tile_start + thread
        ];
    }
    __syncthreads();

    // Thread becomes the panel row. [column][row] keeps row-parallel
    // accesses conflict free throughout the forward substitution.
    #pragma unroll
    for (int col = 0; col < 64; ++col) {
        float value = panel[col][thread];
        #pragma unroll
        for (int inner = 0; inner < col; ++inner) {
            value -= panel[inner][thread] * lower[inner][col];
        }
        panel[col][thread] = value / lower[col][col];
    }
    __syncthreads();

    // Reinterpret thread as the column for a coalesced global store.
    #pragma unroll
    for (int row = 0; row < 64; ++row) {
        matrix_output[
            (panel_start + row) * stride + tile_start + thread
        ] = panel[thread][row];
        matrix_output[
            (tile_start + row) * stride + panel_start + thread
        ] = 0.0f;
    }
}

torch::Tensor trsm128_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<128><<<batch, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        0,
        1,
        64
    );
    return output;
}

torch::Tensor trsm256_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<256><<<batch * panel_count, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        static_cast<int>(tile_start),
        static_cast<int>(panel_count),
        static_cast<int>(tile_start + 64)
    );
    return output;
}

torch::Tensor trsm512_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<512><<<batch * panel_count, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        static_cast<int>(tile_start),
        static_cast<int>(panel_count),
        static_cast<int>(tile_start + 64)
    );
    return output;
}

torch::Tensor trsm1024_external_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<1024><<<batch * 8, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        static_cast<int>(tile_start),
        8,
        512
    );
    return output;
}

torch::Tensor trsm1024_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<1024><<<batch * panel_count, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        static_cast<int>(tile_start),
        static_cast<int>(panel_count),
        static_cast<int>(tile_start + 64)
    );
    return output;
}

torch::Tensor trsm2048_external_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<2048><<<batch * 16, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        static_cast<int>(tile_start),
        16,
        1024
    );
    return output;
}

torch::Tensor trsm2048_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
) {
    const int batch = static_cast<int>(input.size(0));
    trsm_panel64_kernel<2048><<<batch * panel_count, 64>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        static_cast<int>(tile_start),
        static_cast<int>(panel_count),
        static_cast<int>(tile_start + 64)
    );
    return output;
}
"""

_cuda_n32_cpp = r"""
#include <torch/extension.h>
torch::Tensor cholesky32_cuda(torch::Tensor input, torch::Tensor output);
torch::Tensor cholesky64_recursive_cuda(
    torch::Tensor input,
    torch::Tensor output
);
torch::Tensor factor64_strided128_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start
);
torch::Tensor trsm128_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output
);
torch::Tensor trsm256_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
);
torch::Tensor trsm512_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
);
torch::Tensor trsm1024_external_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start
);
torch::Tensor trsm1024_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
);
torch::Tensor trsm2048_external_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start
);
torch::Tensor trsm2048_panel64_cuda(
    torch::Tensor input,
    torch::Tensor output,
    int64_t tile_start,
    int64_t panel_count
);
"""

_cuda_n32_module = load_inline(
    name="cholesky_cuda_direct2048_v5",
    cpp_sources=_cuda_n32_cpp,
    cuda_sources=_cuda_n32_source,
    functions=[
        "cholesky32_cuda",
        "cholesky64_recursive_cuda",
        "factor64_strided128_cuda",
        "trsm128_panel64_cuda",
        "trsm256_panel64_cuda",
        "trsm512_panel64_cuda",
        "trsm1024_external_panel64_cuda",
        "trsm1024_panel64_cuda",
        "trsm2048_external_panel64_cuda",
        "trsm2048_panel64_cuda",
    ],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3"],
    verbose=False,
)


@triton.jit
def _cholesky_banachiewicz_32(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    """Factor one 32x32 SPD matrix per Triton program."""
    matrix_id = tl.program_id(0)
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]
    offsets = matrix_id * matrix_stride + rows * 32 + cols

    # Keep the lower triangle in registers. At step k, previously computed
    # columns contain L and the remaining columns still contain A.
    values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)

    for k in range(32):
        pivot_row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)

        diagonal = tl.sum(
            tl.where(col_ids == k, pivot_row, 0.0),
            axis=0,
        )
        diagonal -= tl.sum(
            tl.where(col_ids < k, pivot_row * pivot_row, 0.0),
            axis=0,
        )
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))

        current_column = tl.sum(
            tl.where(cols == k, values, 0.0),
            axis=1,
        )
        dot_products = tl.sum(
            tl.where(cols < k, values * pivot_row[None, :], 0.0),
            axis=1,
        )
        current_column = (current_column - dot_products) / diagonal

        values = tl.where(
            (rows == k) & (cols == k),
            diagonal,
            values,
        )
        values = tl.where(
            (rows > k) & (cols == k),
            current_column[:, None],
            values,
        )

    tl.store(output_ptr + offsets, values)


@triton.jit
def _factor_lower_tile32(values):
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]

    for k in range(32):
        pivot_row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)

        diagonal = tl.sum(
            tl.where(col_ids == k, pivot_row, 0.0),
            axis=0,
        )
        diagonal -= tl.sum(
            tl.where(col_ids < k, pivot_row * pivot_row, 0.0),
            axis=0,
        )
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))

        current_column = tl.sum(
            tl.where(cols == k, values, 0.0),
            axis=1,
        )
        dot_products = tl.sum(
            tl.where(cols < k, values * pivot_row[None, :], 0.0),
            axis=1,
        )
        current_column = (current_column - dot_products) / diagonal

        values = tl.where(
            (rows == k) & (cols == k),
            diagonal,
            values,
        )
        values = tl.where(
            (rows > k) & (cols == k),
            current_column[:, None],
            values,
        )

    return values


@triton.jit
def _potrf64_diagonal_tile(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]
    base = matrix_id * matrix_stride
    offsets = (
        base
        + (rows + tile_start) * 64
        + (cols + tile_start)
    )

    values = tl.where(
        rows >= cols,
        tl.load(input_ptr + offsets),
        0.0,
    )
    values = _factor_lower_tile32(values)
    tl.store(output_ptr + offsets, values)

    if tile_start == 0:
        upper_offsets = base + rows * 64 + (cols + 32)
        tl.store(output_ptr + upper_offsets, 0.0)


@triton.jit
def _trsm64_panel(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]
    base = matrix_id * matrix_stride
    offsets_00 = base + rows * 64 + cols
    offsets_10 = base + (rows + 32) * 64 + cols

    l00 = tl.load(output_ptr + offsets_00)
    l10 = tl.load(input_ptr + offsets_10)

    for k in range(32):
        pivot_row = tl.sum(tl.where(rows == k, l00, 0.0), axis=0)
        diagonal = tl.sum(
            tl.where(col_ids == k, pivot_row, 0.0),
            axis=0,
        )
        current_column = tl.sum(
            tl.where(cols == k, l10, 0.0),
            axis=1,
        )
        dot_products = tl.sum(
            tl.where(cols < k, l10 * pivot_row[None, :], 0.0),
            axis=1,
        )
        solved_column = (current_column - dot_products) / diagonal
        l10 = tl.where(cols == k, solved_column[:, None], l10)

    tl.store(output_ptr + offsets_10, l10)


@triton.jit
def _syrk64_trailing_update(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]
    base = matrix_id * matrix_stride
    offsets_10 = base + (rows + 32) * 64 + cols
    offsets_11 = base + (rows + 32) * 64 + (cols + 32)

    l10 = tl.load(output_ptr + offsets_10)
    a11 = tl.load(input_ptr + offsets_11)
    update = tl.dot(
        l10,
        tl.trans(l10),
        input_precision="ieee",
        out_dtype=tl.float32,
    )
    schur = tl.where(rows >= cols, a11 - update, 0.0)
    tl.store(output_ptr + offsets_11, schur)


@triton.jit
def _right_trsm32(panel, lower):
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]

    for k in range(32):
        pivot_row = tl.sum(tl.where(rows == k, lower, 0.0), axis=0)
        diagonal = tl.sum(
            tl.where(col_ids == k, pivot_row, 0.0),
            axis=0,
        )
        current_column = tl.sum(
            tl.where(cols == k, panel, 0.0),
            axis=1,
        )
        dot_products = tl.sum(
            tl.where(cols < k, panel * pivot_row[None, :], 0.0),
            axis=1,
        )
        solved_column = (current_column - dot_products) / diagonal
        panel = tl.where(cols == k, solved_column[:, None], panel)
    return panel


@triton.jit
def _cholesky64_fused(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    lower_mask = rows >= cols
    base = matrix_id * matrix_stride
    offsets_00 = base + rows * 64 + cols
    offsets_01 = offsets_00 + 32
    offsets_10 = offsets_00 + 32 * 64
    offsets_11 = offsets_10 + 32

    l00 = _factor_lower_tile32(
        tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
    )
    l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
    a11 = tl.where(
        lower_mask,
        tl.load(input_ptr + offsets_11),
        0.0,
    )
    update = tl.dot(
        l10,
        tl.trans(l10),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    l11 = _factor_lower_tile32(
        tl.where(lower_mask, a11 - update, 0.0)
    )

    tl.store(output_ptr + offsets_00, l00)
    tl.store(output_ptr + offsets_01, 0.0)
    tl.store(output_ptr + offsets_10, l10)
    tl.store(output_ptr + offsets_11, l11)


@triton.jit
def _potrf128_tile64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    lower_mask = rows >= cols
    base = matrix_id * matrix_stride

    offsets_00 = base + (rows + tile_start) * 128 + cols + tile_start
    offsets_01 = offsets_00 + 32
    offsets_10 = offsets_00 + 32 * 128
    offsets_11 = offsets_10 + 32

    l00 = tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
    l00 = _factor_lower_tile32(l00)

    l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
    a11 = tl.where(lower_mask, tl.load(input_ptr + offsets_11), 0.0)
    update = tl.dot(
        l10,
        tl.trans(l10),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    l11 = _factor_lower_tile32(
        tl.where(lower_mask, a11 - update, 0.0)
    )

    tl.store(output_ptr + offsets_00, l00)
    tl.store(output_ptr + offsets_01, 0.0)
    tl.store(output_ptr + offsets_10, l10)
    tl.store(output_ptr + offsets_11, l11)

    if tile_start == 0:
        top_right = base + rows * 128 + cols + 64
        tl.store(output_ptr + top_right, 0.0)
        tl.store(output_ptr + top_right + 32, 0.0)
        tl.store(output_ptr + top_right + 32 * 128, 0.0)
        tl.store(output_ptr + top_right + 32 * 128 + 32, 0.0)


@triton.jit
def _trsm128_panel64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // 2
    row_tile = program_id % 2
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    row_start = 64 + row_tile * 32

    diagonal_00 = base + rows * 128 + cols
    diagonal_10 = diagonal_00 + 32 * 128
    diagonal_11 = diagonal_10 + 32
    panel_0 = base + (rows + row_start) * 128 + cols
    panel_1 = panel_0 + 32

    l00 = tl.load(output_ptr + diagonal_00)
    l10 = tl.load(output_ptr + diagonal_10)
    l11 = tl.load(output_ptr + diagonal_11)
    solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
    rhs_1 = tl.load(input_ptr + panel_1)
    rhs_1 -= tl.dot(
        solved_0,
        tl.trans(l10),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    solved_1 = _right_trsm32(rhs_1, l11)

    tl.store(output_ptr + panel_0, solved_0)
    tl.store(output_ptr + panel_1, solved_1)


@triton.jit
def _syrk128_trailing64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // 3
    tile_id = program_id % 3
    row_tile = tl.where(tile_id == 0, 0, 1)
    col_tile = tl.where(tile_id == 2, 1, 0)

    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    row_start = 64 + row_tile * 32
    col_start = 64 + col_tile * 32
    left_0 = base + (rows + row_start) * 128 + cols
    left_1 = left_0 + 32
    right_0 = base + (rows + col_start) * 128 + cols
    right_1 = right_0 + 32
    trailing = base + (rows + row_start) * 128 + cols + col_start

    update = tl.dot(
        tl.load(output_ptr + left_0),
        tl.trans(tl.load(output_ptr + right_0)),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    update += tl.dot(
        tl.load(output_ptr + left_1),
        tl.trans(tl.load(output_ptr + right_1)),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    values = tl.load(input_ptr + trailing) - update
    values = tl.where(
        (row_tile > col_tile) | (rows >= cols),
        values,
        0.0,
    )
    tl.store(output_ptr + trailing, values)


@triton.jit
def _potrf256_tile64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    lower_mask = rows >= cols
    base = matrix_id * matrix_stride
    offsets_00 = base + (rows + tile_start) * 256 + cols + tile_start
    offsets_01 = offsets_00 + 32
    offsets_10 = offsets_00 + 32 * 256
    offsets_11 = offsets_10 + 32

    l00 = _factor_lower_tile32(
        tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
    )
    l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
    a11 = tl.where(lower_mask, tl.load(input_ptr + offsets_11), 0.0)
    update = tl.dot(
        l10,
        tl.trans(l10),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    l11 = _factor_lower_tile32(
        tl.where(lower_mask, a11 - update, 0.0)
    )
    tl.store(output_ptr + offsets_00, l00)
    tl.store(output_ptr + offsets_01, 0.0)
    tl.store(output_ptr + offsets_10, l10)
    tl.store(output_ptr + offsets_11, l11)


@triton.jit
def _trsm256_panel64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    panel_count: tl.constexpr,
):
    rows_per_matrix: tl.constexpr = panel_count * 2
    program_id = tl.program_id(0)
    matrix_id = program_id // rows_per_matrix
    local_id = program_id % rows_per_matrix
    panel_id = local_id // 2
    row_tile = local_id % 2
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    row_start = tile_start + 64 + panel_id * 64 + row_tile * 32
    diagonal_00 = (
        base + (rows + tile_start) * 256 + cols + tile_start
    )
    diagonal_10 = diagonal_00 + 32 * 256
    diagonal_11 = diagonal_10 + 32
    panel_0 = base + (rows + row_start) * 256 + cols + tile_start
    panel_1 = panel_0 + 32

    l00 = tl.load(output_ptr + diagonal_00)
    l10 = tl.load(output_ptr + diagonal_10)
    l11 = tl.load(output_ptr + diagonal_11)
    solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
    rhs_1 = tl.load(input_ptr + panel_1)
    rhs_1 -= tl.dot(
        solved_0,
        tl.trans(l10),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    solved_1 = _right_trsm32(rhs_1, l11)
    tl.store(output_ptr + panel_0, solved_0)
    tl.store(output_ptr + panel_1, solved_1)

    upper_0 = base + (rows + tile_start) * 256 + cols + row_start
    upper_1 = upper_0 + 32 * 256
    tl.store(output_ptr + upper_0, 0.0)
    tl.store(output_ptr + upper_1, 0.0)

@triton.jit
def _update256_tiles64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    remaining_tiles: tl.constexpr,
):
    pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
    programs_per_matrix: tl.constexpr = pair_count
    program_id = tl.program_id(0)
    matrix_id = program_id // programs_per_matrix
    local_id = program_id % programs_per_matrix
    pair_id = local_id

    if remaining_tiles == 3:
        row_tile = tl.where(pair_id == 0, 0, tl.where(pair_id < 3, 1, 2))
        col_tile = tl.where(
            pair_id == 0,
            0,
            tl.where(
                pair_id == 1,
                0,
                tl.where(pair_id == 2, 1, pair_id - 3),
            ),
        )
    elif remaining_tiles == 2:
        row_tile = tl.where(pair_id == 0, 0, 1)
        col_tile = tl.where(pair_id == 2, 1, 0)
    else:
        row_tile = 0
        col_tile = 0

    row_start = tile_start + 64 + row_tile * 64
    col_start = tile_start + 64 + col_tile * 64
    rows = tl.arange(0, 64)[:, None]
    cols = tl.arange(0, 64)[None, :]
    base = matrix_id * matrix_stride
    left_0 = base + (rows + row_start) * 256 + cols + tile_start
    right_0 = base + (rows + col_start) * 256 + cols + tile_start
    trailing = base + (rows + row_start) * 256 + cols + col_start

    update = tl.dot(
        tl.load(output_ptr + left_0),
        tl.trans(tl.load(output_ptr + right_0)),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    values = tl.load(input_ptr + trailing) - update
    lower_tile = row_tile > col_tile
    tl.store(
        output_ptr + trailing,
        tl.where(lower_tile | (rows >= cols), values, 0.0),
    )


@triton.jit
def _zero256_upper(output_ptr, matrix_stride: tl.constexpr):
    program_id = tl.program_id(0)
    matrix_id = program_id // 24
    local_id = program_id % 24
    pair_id = local_id // 4
    block_id = local_id % 4
    row_tile = tl.where(pair_id < 3, 0, tl.where(pair_id < 5, 1, 2))
    col_tile = tl.where(
        pair_id == 0,
        1,
        tl.where(
            pair_id == 1,
            2,
            tl.where(pair_id == 2, 3, tl.where(pair_id == 3, 2, 3)),
        ),
    )
    row_subtile = block_id // 2
    col_subtile = block_id % 2
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    offsets = (
        base
        + (rows + row_tile * 64 + row_subtile * 32) * 256
        + cols
        + col_tile * 64
        + col_subtile * 32
    )
    tl.store(output_ptr + offsets, 0.0)


@triton.jit
def _potrf512_tile64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    dot_precision: tl.constexpr,
):
    matrix_id = tl.program_id(0)
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    lower_mask = rows >= cols
    base = matrix_id * matrix_stride
    offsets_00 = base + (rows + tile_start) * 512 + cols + tile_start
    offsets_01 = offsets_00 + 32
    offsets_10 = offsets_00 + 32 * 512
    offsets_11 = offsets_10 + 32

    l00 = _factor_lower_tile32(
        tl.where(lower_mask, tl.load(input_ptr + offsets_00), 0.0)
    )
    l10 = _right_trsm32(tl.load(input_ptr + offsets_10), l00)
    a11 = tl.where(lower_mask, tl.load(input_ptr + offsets_11), 0.0)
    update = tl.dot(
        l10,
        tl.trans(l10),
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    l11 = _factor_lower_tile32(
        tl.where(lower_mask, a11 - update, 0.0)
    )
    tl.store(output_ptr + offsets_00, l00)
    tl.store(output_ptr + offsets_01, 0.0)
    tl.store(output_ptr + offsets_10, l10)
    tl.store(output_ptr + offsets_11, l11)


@triton.jit
def _trsm512_panel64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    panel_count: tl.constexpr,
    dot_precision: tl.constexpr,
    use_2d_grid: tl.constexpr,
):
    rows_per_matrix: tl.constexpr = panel_count * 2
    if use_2d_grid:
        matrix_id = tl.program_id(1)
        local_id = tl.program_id(0)
    else:
        program_id = tl.program_id(0)
        matrix_id = program_id // rows_per_matrix
        local_id = program_id % rows_per_matrix
    panel_id = local_id // 2
    row_tile = local_id % 2
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    row_start = tile_start + 64 + panel_id * 64 + row_tile * 32
    diagonal_00 = (
        base + (rows + tile_start) * 512 + cols + tile_start
    )
    diagonal_10 = diagonal_00 + 32 * 512
    diagonal_11 = diagonal_10 + 32
    panel_0 = base + (rows + row_start) * 512 + cols + tile_start
    panel_1 = panel_0 + 32

    l00 = tl.load(output_ptr + diagonal_00)
    l10 = tl.load(output_ptr + diagonal_10)
    l11 = tl.load(output_ptr + diagonal_11)
    solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
    rhs_1 = tl.load(input_ptr + panel_1)
    rhs_1 -= tl.dot(
        solved_0,
        tl.trans(l10),
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    solved_1 = _right_trsm32(rhs_1, l11)
    tl.store(output_ptr + panel_0, solved_0)
    tl.store(output_ptr + panel_1, solved_1)

    upper_0 = base + (rows + tile_start) * 512 + cols + row_start
    upper_1 = upper_0 + 32 * 512
    tl.store(output_ptr + upper_0, 0.0)
    tl.store(output_ptr + upper_1, 0.0)


@triton.jit
def _update512_tiles32(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    remaining_tiles: tl.constexpr,
    dot_precision: tl.constexpr,
    use_2d_grid: tl.constexpr,
):
    pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
    programs_per_matrix: tl.constexpr = pair_count * 4
    if use_2d_grid:
        matrix_id = tl.program_id(1)
        local_id = tl.program_id(0)
    else:
        program_id = tl.program_id(0)
        matrix_id = program_id // programs_per_matrix
        local_id = program_id % programs_per_matrix
    pair_id = local_id // 4
    block_id = local_id % 4

    row_tile = tl.cast(
        (tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = pair_id - row_tile * (row_tile + 1) // 2
    row_subtile = block_id // 2
    col_subtile = block_id % 2
    row_start = tile_start + 64 + row_tile * 64 + row_subtile * 32
    col_start = tile_start + 64 + col_tile * 64 + col_subtile * 32
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    left_0 = base + (rows + row_start) * 512 + cols + tile_start
    left_1 = left_0 + 32
    right_0 = base + (rows + col_start) * 512 + cols + tile_start
    right_1 = right_0 + 32
    trailing = base + (rows + row_start) * 512 + cols + col_start

    update = tl.dot(
        tl.load(output_ptr + left_0),
        tl.trans(tl.load(output_ptr + right_0)),
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    update += tl.dot(
        tl.load(output_ptr + left_1),
        tl.trans(tl.load(output_ptr + right_1)),
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    values = tl.load(input_ptr + trailing) - update
    lower_tile = row_tile > col_tile
    lower_block = (row_subtile > col_subtile) | (
        (row_subtile == col_subtile) & (rows >= cols)
    )
    tl.store(
        output_ptr + trailing,
        tl.where(lower_tile | lower_block, values, 0.0),
    )


@triton.jit
def _update512_tiles64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    remaining_tiles: tl.constexpr,
    dot_precision: tl.constexpr,
    use_2d_grid: tl.constexpr,
):
    pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
    programs_per_matrix: tl.constexpr = pair_count
    if use_2d_grid:
        matrix_id = tl.program_id(1)
        local_id = tl.program_id(0)
    else:
        program_id = tl.program_id(0)
        matrix_id = program_id // programs_per_matrix
        local_id = program_id % programs_per_matrix
    pair_id = local_id

    row_tile = tl.cast(
        (tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = pair_id - row_tile * (row_tile + 1) // 2
    row_start = tile_start + 64 + row_tile * 64
    col_start = tile_start + 64 + col_tile * 64
    rows = tl.arange(0, 64)[:, None]
    cols = tl.arange(0, 64)[None, :]
    base = matrix_id * matrix_stride
    left_0 = base + (rows + row_start) * 512 + cols + tile_start
    right_0 = base + (rows + col_start) * 512 + cols + tile_start
    trailing = base + (rows + row_start) * 512 + cols + col_start

    update = tl.dot(
        tl.load(output_ptr + left_0),
        tl.trans(tl.load(output_ptr + right_0)),
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    values = tl.load(input_ptr + trailing) - update
    lower_tile = row_tile > col_tile
    tl.store(
        output_ptr + trailing,
        tl.where(lower_tile | (rows >= cols), values, 0.0),
    )


@triton.jit
def _trsm512_second64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    panel_count: tl.constexpr,
):
    programs_per_matrix: tl.constexpr = panel_count * 4
    program_id = tl.program_id(0)
    matrix_id = program_id // programs_per_matrix
    row_block = program_id % programs_per_matrix
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    row_start = tile_start + 128 + row_block * 32

    diagonal_00 = (
        base + (rows + tile_start) * 512 + cols + tile_start
    )
    diagonal_20 = diagonal_00 + 64 * 512
    diagonal_21 = diagonal_20 + 32
    diagonal_22 = diagonal_20 + 64
    diagonal_30 = diagonal_20 + 32 * 512
    diagonal_31 = diagonal_30 + 32
    diagonal_32 = diagonal_30 + 64
    diagonal_33 = diagonal_30 + 96

    panel_0 = base + (rows + row_start) * 512 + cols + tile_start
    panel_1 = panel_0 + 32
    panel_2 = panel_0 + 64
    panel_3 = panel_0 + 96

    solved_0 = tl.load(output_ptr + panel_0)
    solved_1 = tl.load(output_ptr + panel_1)
    l20 = tl.load(output_ptr + diagonal_20)
    l21 = tl.load(output_ptr + diagonal_21)
    l22 = tl.load(output_ptr + diagonal_22)
    l30 = tl.load(output_ptr + diagonal_30)
    l31 = tl.load(output_ptr + diagonal_31)
    l32 = tl.load(output_ptr + diagonal_32)
    l33 = tl.load(output_ptr + diagonal_33)

    rhs_2 = tl.load(input_ptr + panel_2)
    rhs_2 -= tl.dot(
        solved_0,
        tl.trans(l20),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    rhs_2 -= tl.dot(
        solved_1,
        tl.trans(l21),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    solved_2 = _right_trsm32(rhs_2, l22)

    rhs_3 = tl.load(input_ptr + panel_3)
    rhs_3 -= tl.dot(
        solved_0,
        tl.trans(l30),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    rhs_3 -= tl.dot(
        solved_1,
        tl.trans(l31),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    rhs_3 -= tl.dot(
        solved_2,
        tl.trans(l32),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    solved_3 = _right_trsm32(rhs_3, l33)

    tl.store(output_ptr + panel_2, solved_2)
    tl.store(output_ptr + panel_3, solved_3)

    upper_2 = (
        base + (rows + tile_start + 64) * 512 + cols + row_start
    )
    tl.store(output_ptr + upper_2, 0.0)
    tl.store(output_ptr + upper_2 + 32 * 512, 0.0)


@triton.jit
def _trsm512_panel128(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    panel_count: tl.constexpr,
):
    programs_per_matrix: tl.constexpr = panel_count * 4
    program_id = tl.program_id(0)
    matrix_id = program_id // programs_per_matrix
    local_id = program_id % programs_per_matrix
    panel_id = local_id // 4
    row_block = local_id % 4
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    row_start = tile_start + 128 + panel_id * 128 + row_block * 32

    diagonal_00 = (
        base + (rows + tile_start) * 512 + cols + tile_start
    )
    diagonal_10 = diagonal_00 + 32 * 512
    diagonal_11 = diagonal_10 + 32
    diagonal_20 = diagonal_00 + 64 * 512
    diagonal_21 = diagonal_20 + 32
    diagonal_22 = diagonal_20 + 64
    diagonal_30 = diagonal_20 + 32 * 512
    diagonal_31 = diagonal_30 + 32
    diagonal_32 = diagonal_30 + 64
    diagonal_33 = diagonal_30 + 96

    panel_0 = base + (rows + row_start) * 512 + cols + tile_start
    panel_1 = panel_0 + 32
    panel_2 = panel_0 + 64
    panel_3 = panel_0 + 96

    l00 = tl.load(output_ptr + diagonal_00)
    l10 = tl.load(output_ptr + diagonal_10)
    l11 = tl.load(output_ptr + diagonal_11)
    l20 = tl.load(output_ptr + diagonal_20)
    l21 = tl.load(output_ptr + diagonal_21)
    l22 = tl.load(output_ptr + diagonal_22)
    l30 = tl.load(output_ptr + diagonal_30)
    l31 = tl.load(output_ptr + diagonal_31)
    l32 = tl.load(output_ptr + diagonal_32)
    l33 = tl.load(output_ptr + diagonal_33)

    solved_0 = _right_trsm32(tl.load(input_ptr + panel_0), l00)
    rhs_1 = tl.load(input_ptr + panel_1)
    rhs_1 -= tl.dot(
        solved_0,
        tl.trans(l10),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    solved_1 = _right_trsm32(rhs_1, l11)

    rhs_2 = tl.load(input_ptr + panel_2)
    rhs_2 -= tl.dot(
        solved_0,
        tl.trans(l20),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    rhs_2 -= tl.dot(
        solved_1,
        tl.trans(l21),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    solved_2 = _right_trsm32(rhs_2, l22)

    rhs_3 = tl.load(input_ptr + panel_3)
    rhs_3 -= tl.dot(
        solved_0,
        tl.trans(l30),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    rhs_3 -= tl.dot(
        solved_1,
        tl.trans(l31),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    rhs_3 -= tl.dot(
        solved_2,
        tl.trans(l32),
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    solved_3 = _right_trsm32(rhs_3, l33)

    tl.store(output_ptr + panel_0, solved_0)
    tl.store(output_ptr + panel_1, solved_1)
    tl.store(output_ptr + panel_2, solved_2)
    tl.store(output_ptr + panel_3, solved_3)

    upper_0 = base + (rows + tile_start) * 512 + cols + row_start
    tl.store(output_ptr + upper_0, 0.0)
    tl.store(output_ptr + upper_0 + 32 * 512, 0.0)
    tl.store(output_ptr + upper_0 + 64 * 512, 0.0)
    tl.store(output_ptr + upper_0 + 96 * 512, 0.0)


@triton.jit
def _update512_second_panel64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    panel_count: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // panel_count
    row_tile = program_id % panel_count
    rows = tl.arange(0, 64)
    cols = tl.arange(0, 64)
    k_ids = tl.arange(0, 64)
    row_start = tile_start + 64 + row_tile * 64
    col_start = tile_start + 64
    base = matrix_id * matrix_stride

    left = tl.load(
        output_ptr
        + base
        + (rows[:, None] + row_start) * 512
        + k_ids[None, :]
        + tile_start
    )
    right = tl.load(
        output_ptr
        + base
        + (cols[None, :] + col_start) * 512
        + k_ids[:, None]
        + tile_start
    )
    update = tl.dot(
        left,
        right,
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    panel = (
        base
        + (rows[:, None] + row_start) * 512
        + cols[None, :]
        + col_start
    )
    values = tl.load(input_ptr + panel) - update
    lower_mask = (row_tile > 0) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(
        output_ptr + panel,
        tl.where(lower_mask, values, 0.0),
    )


@triton.jit
def _update512_tiles128(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    remaining_tiles: tl.constexpr,
):
    pair_count: tl.constexpr = remaining_tiles * (remaining_tiles + 1) // 2
    program_id = tl.program_id(0)
    matrix_id = program_id // pair_count
    pair_id = program_id % pair_count
    row_tile = tl.cast(
        (tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = pair_id - row_tile * (row_tile + 1) // 2
    rows = tl.arange(0, 64)
    cols = tl.arange(0, 64)
    k_ids = tl.arange(0, 128)
    row_start = tile_start + 128 + row_tile * 64
    col_start = tile_start + 128 + col_tile * 64
    base = matrix_id * matrix_stride

    left = tl.load(
        output_ptr
        + base
        + (rows[:, None] + row_start) * 512
        + k_ids[None, :]
        + tile_start
    )
    right = tl.load(
        output_ptr
        + base
        + (cols[None, :] + col_start) * 512
        + k_ids[:, None]
        + tile_start
    )
    update = tl.dot(
        left,
        right,
        input_precision="tf32",
        out_dtype=tl.float32,
    )
    trailing = (
        base
        + (rows[:, None] + row_start) * 512
        + cols[None, :]
        + col_start
    )
    values = tl.load(input_ptr + trailing) - update
    lower_mask = (row_tile > col_tile) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(
        output_ptr + trailing,
        tl.where(lower_mask, values, 0.0),
    )


@triton.jit
def _update1024_second_panel64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start,
    panel_count,
    dot_precision: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // panel_count
    row_tile = program_id % panel_count
    rows = tl.arange(0, 64)
    cols = tl.arange(0, 64)
    k_ids = tl.arange(0, 64)
    row_start = tile_start + 64 + row_tile * 64
    col_start = tile_start + 64
    base = matrix_id * matrix_stride
    left = tl.load(
        output_ptr
        + base
        + (rows[:, None] + row_start) * 1024
        + k_ids[None, :]
        + tile_start
    )
    right = tl.load(
        output_ptr
        + base
        + (cols[None, :] + col_start) * 1024
        + k_ids[:, None]
        + tile_start
    )
    update = tl.dot(
        left,
        right,
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    panel = (
        base
        + (rows[:, None] + row_start) * 1024
        + cols[None, :]
        + col_start
    )
    values = tl.load(input_ptr + panel) - update
    lower_mask = (row_tile > 0) | (rows[:, None] >= cols[None, :])
    tl.store(
        output_ptr + panel,
        tl.where(lower_mask, values, 0.0),
    )


@triton.jit
def _update1024_tiles128(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start,
    remaining_tiles,
    dot_precision: tl.constexpr,
):
    pair_count = remaining_tiles * (remaining_tiles + 1) // 2
    program_id = tl.program_id(0)
    matrix_id = program_id // pair_count
    pair_id = program_id % pair_count
    row_tile = tl.cast(
        (tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = pair_id - row_tile * (row_tile + 1) // 2
    rows = tl.arange(0, 64)
    cols = tl.arange(0, 64)
    k_ids = tl.arange(0, 128)
    row_start = tile_start + 128 + row_tile * 64
    col_start = tile_start + 128 + col_tile * 64
    base = matrix_id * matrix_stride
    left = tl.load(
        output_ptr
        + base
        + (rows[:, None] + row_start) * 1024
        + k_ids[None, :]
        + tile_start
    )
    right = tl.load(
        output_ptr
        + base
        + (cols[None, :] + col_start) * 1024
        + k_ids[:, None]
        + tile_start
    )
    update = tl.dot(
        left,
        right,
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    trailing = (
        base
        + (rows[:, None] + row_start) * 1024
        + cols[None, :]
        + col_start
    )
    values = tl.load(input_ptr + trailing) - update
    lower_mask = (row_tile > col_tile) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(
        output_ptr + trailing,
        tl.where(lower_mask, values, 0.0),
    )


@triton.jit
def _update2048_second_panel64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start,
    panel_count,
    dot_precision: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // panel_count
    row_tile = program_id % panel_count
    rows = tl.arange(0, 64)
    cols = tl.arange(0, 64)
    k_ids = tl.arange(0, 64)
    row_start = tile_start + 64 + row_tile * 64
    col_start = tile_start + 64
    base = matrix_id * matrix_stride
    left = tl.load(
        output_ptr
        + base
        + (rows[:, None] + row_start) * 2048
        + k_ids[None, :]
        + tile_start
    )
    right = tl.load(
        output_ptr
        + base
        + (cols[None, :] + col_start) * 2048
        + k_ids[:, None]
        + tile_start
    )
    update = tl.dot(
        left,
        right,
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    panel = (
        base
        + (rows[:, None] + row_start) * 2048
        + cols[None, :]
        + col_start
    )
    values = tl.load(input_ptr + panel) - update
    lower_mask = (row_tile > 0) | (rows[:, None] >= cols[None, :])
    tl.store(
        output_ptr + panel,
        tl.where(lower_mask, values, 0.0),
    )


@triton.jit
def _update2048_tiles128(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start,
    remaining_tiles,
    dot_precision: tl.constexpr,
):
    pair_count = remaining_tiles * (remaining_tiles + 1) // 2
    program_id = tl.program_id(0)
    matrix_id = program_id // pair_count
    pair_id = program_id % pair_count
    row_tile = tl.cast(
        (tl.sqrt(8.0 * pair_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = pair_id - row_tile * (row_tile + 1) // 2
    rows = tl.arange(0, 64)
    cols = tl.arange(0, 64)
    k_ids = tl.arange(0, 128)
    row_start = tile_start + 128 + row_tile * 64
    col_start = tile_start + 128 + col_tile * 64
    base = matrix_id * matrix_stride
    left = tl.load(
        output_ptr
        + base
        + (rows[:, None] + row_start) * 2048
        + k_ids[None, :]
        + tile_start
    )
    right = tl.load(
        output_ptr
        + base
        + (cols[None, :] + col_start) * 2048
        + k_ids[:, None]
        + tile_start
    )
    update = tl.dot(
        left,
        right,
        input_precision=dot_precision,
        out_dtype=tl.float32,
    )
    trailing = (
        base
        + (rows[:, None] + row_start) * 2048
        + cols[None, :]
        + col_start
    )
    values = tl.load(input_ptr + trailing) - update
    lower_mask = (row_tile > col_tile) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(
        output_ptr + trailing,
        tl.where(lower_mask, values, 0.0),
    )


@triton.jit
def _zero512_upper(output_ptr, matrix_stride: tl.constexpr):
    program_id = tl.program_id(0)
    matrix_id = program_id // 120
    pair_id = program_id % 120
    col_block = tl.cast(
        (1.0 + tl.sqrt(1.0 + 8.0 * pair_id)) * 0.5,
        tl.int32,
    )
    row_block = pair_id - col_block * (col_block - 1) // 2
    rows = tl.arange(0, 32)[:, None]
    cols = tl.arange(0, 32)[None, :]
    base = matrix_id * matrix_stride
    offsets = (
        base + (rows + row_block * 32) * 512 + cols + col_block * 32
    )
    tl.store(output_ptr + offsets, 0.0)


@triton.jit
def _gather_lower_tiles(
    input_ptr,
    output_ptr,
    input_batch_stride: tl.constexpr,
    input_ld: tl.constexpr,
    output_n: tl.constexpr,
    tile_count: tl.constexpr,
):
    pair_count: tl.constexpr = tile_count * (tile_count + 1) // 2
    program_id = tl.program_id(0)
    matrix_id = program_id // pair_count
    tile_id = program_id % pair_count
    row_tile = tl.cast(
        (tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = tile_id - row_tile * (row_tile + 1) // 2
    rows = row_tile * 64 + tl.arange(0, 64)[:, None]
    cols = col_tile * 64 + tl.arange(0, 64)[None, :]
    input_offsets = (
        matrix_id * input_batch_stride
        + rows * input_ld
        + cols
    )
    output_offsets = (
        matrix_id * output_n * output_n
        + rows * output_n
        + cols
    )
    tl.store(output_ptr + output_offsets, tl.load(input_ptr + input_offsets))


@triton.jit
def _update1024_trsm_tiles64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    remaining_tiles: tl.constexpr,
):
    programs_per_matrix: tl.constexpr = 8 * remaining_tiles
    program_id = tl.program_id(0)
    matrix_id = program_id // programs_per_matrix
    local_id = program_id % programs_per_matrix
    row_tile = local_id // remaining_tiles
    col_tile = local_id % remaining_tiles
    row_start = 512 + row_tile * 64
    col_start = tile_start + 64 + col_tile * 64
    rows = tl.arange(0, 64)[:, None]
    cols = tl.arange(0, 64)[None, :]
    base = matrix_id * matrix_stride
    solved = base + (rows + row_start) * 1024 + cols + tile_start
    factor = base + (rows + col_start) * 1024 + cols + tile_start
    trailing = base + (rows + row_start) * 1024 + cols + col_start
    update = tl.dot(
        tl.load(output_ptr + solved),
        tl.trans(tl.load(output_ptr + factor)),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    values = tl.load(input_ptr + trailing) - update
    tl.store(output_ptr + trailing, values)


@triton.jit
def _update2048_trsm_tiles64(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
    tile_start: tl.constexpr,
    remaining_tiles: tl.constexpr,
):
    programs_per_matrix: tl.constexpr = 16 * remaining_tiles
    program_id = tl.program_id(0)
    matrix_id = program_id // programs_per_matrix
    local_id = program_id % programs_per_matrix
    row_tile = local_id // remaining_tiles
    col_tile = local_id % remaining_tiles
    row_start = 1024 + row_tile * 64
    col_start = tile_start + 64 + col_tile * 64
    rows = tl.arange(0, 64)[:, None]
    cols = tl.arange(0, 64)[None, :]
    base = matrix_id * matrix_stride
    solved = base + (rows + row_start) * 2048 + cols + tile_start
    factor = base + (rows + col_start) * 2048 + cols + tile_start
    trailing = base + (rows + row_start) * 2048 + cols + col_start
    update = tl.dot(
        tl.load(output_ptr + solved),
        tl.trans(tl.load(output_ptr + factor)),
        input_precision="tf32x3",
        out_dtype=tl.float32,
    )
    values = tl.load(input_ptr + trailing) - update
    tl.store(output_ptr + trailing, values)


@triton.jit
def _schur1024_split512(
    input_ptr,
    l10_ptr,
    schur_ptr,
    matrix_stride: tl.constexpr,
    l10_batch_stride: tl.constexpr,
    l10_row_stride: tl.constexpr,
    l10_col_stride: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // 36
    tile_id = program_id % 36
    row_tile = tl.cast(
        (tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = tile_id - row_tile * (row_tile + 1) // 2
    rows = row_tile * 64 + tl.arange(0, 64)
    cols = col_tile * 64 + tl.arange(0, 64)
    accumulator = tl.zeros((64, 64), tl.float32)
    l10_base = matrix_id * l10_batch_stride

    for k_start in range(0, 512, 64):
        k_ids = k_start + tl.arange(0, 64)
        left = tl.load(
            l10_ptr
            + l10_base
            + rows[:, None] * l10_row_stride
            + k_ids[None, :] * l10_col_stride
        )
        right = tl.load(
            l10_ptr
            + l10_base
            + k_ids[:, None] * l10_col_stride
            + cols[None, :] * l10_row_stride
        )
        accumulator += tl.dot(
            left,
            right,
            input_precision="tf32x3",
            out_dtype=tl.float32,
        )

    input_offsets = (
        matrix_id * matrix_stride
        + (rows[:, None] + 512) * 1024
        + cols[None, :]
        + 512
    )
    schur_offsets = (
        matrix_id * 512 * 512 + rows[:, None] * 512 + cols[None, :]
    )
    values = tl.load(input_ptr + input_offsets) - accumulator
    lower_mask = (row_tile > col_tile) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(schur_ptr + schur_offsets, tl.where(lower_mask, values, 0.0))


@triton.jit
def _schur2048_split1024(
    input_ptr,
    l10_ptr,
    schur_ptr,
    matrix_stride: tl.constexpr,
    l10_batch_stride: tl.constexpr,
    l10_row_stride: tl.constexpr,
    l10_col_stride: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // 136
    tile_id = program_id % 136
    row_tile = tl.cast(
        (tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = tile_id - row_tile * (row_tile + 1) // 2
    rows = row_tile * 64 + tl.arange(0, 64)
    cols = col_tile * 64 + tl.arange(0, 64)
    accumulator = tl.zeros((64, 64), tl.float32)
    l10_base = matrix_id * l10_batch_stride

    for k_start in range(0, 1024, 64):
        k_ids = k_start + tl.arange(0, 64)
        left = tl.load(
            l10_ptr
            + l10_base
            + rows[:, None] * l10_row_stride
            + k_ids[None, :] * l10_col_stride
        )
        right = tl.load(
            l10_ptr
            + l10_base
            + k_ids[:, None] * l10_col_stride
            + cols[None, :] * l10_row_stride
        )
        accumulator += tl.dot(
            left,
            right,
            input_precision="tf32x3",
            out_dtype=tl.float32,
        )

    input_offsets = (
        matrix_id * matrix_stride
        + (rows[:, None] + 1024) * 2048
        + cols[None, :]
        + 1024
    )
    schur_offsets = (
        matrix_id * 1024 * 1024 + rows[:, None] * 1024 + cols[None, :]
    )
    values = tl.load(input_ptr + input_offsets) - accumulator
    lower_mask = (row_tile > col_tile) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(schur_ptr + schur_offsets, tl.where(lower_mask, values, 0.0))


@triton.jit
def _schur4096_split2048(
    input_ptr,
    l10_ptr,
    schur_ptr,
    matrix_stride: tl.constexpr,
    l10_batch_stride: tl.constexpr,
    l10_row_stride: tl.constexpr,
    l10_col_stride: tl.constexpr,
):
    program_id = tl.program_id(0)
    matrix_id = program_id // 528
    tile_id = program_id % 528
    row_tile = tl.cast(
        (tl.sqrt(8.0 * tile_id + 1.0) - 1.0) * 0.5,
        tl.int32,
    )
    col_tile = tile_id - row_tile * (row_tile + 1) // 2
    rows = row_tile * 64 + tl.arange(0, 64)
    cols = col_tile * 64 + tl.arange(0, 64)
    accumulator = tl.zeros((64, 64), tl.float32)
    l10_base = matrix_id * l10_batch_stride

    for k_start in range(0, 2048, 64):
        k_ids = k_start + tl.arange(0, 64)
        left = tl.load(
            l10_ptr
            + l10_base
            + rows[:, None] * l10_row_stride
            + k_ids[None, :] * l10_col_stride
        )
        right = tl.load(
            l10_ptr
            + l10_base
            + k_ids[:, None] * l10_col_stride
            + cols[None, :] * l10_row_stride
        )
        accumulator += tl.dot(
            left,
            right,
            input_precision="tf32x3",
            out_dtype=tl.float32,
        )

    input_offsets = (
        matrix_id * matrix_stride
        + (rows[:, None] + 2048) * 4096
        + cols[None, :]
        + 2048
    )
    schur_offsets = (
        matrix_id * 2048 * 2048 + rows[:, None] * 2048 + cols[None, :]
    )
    values = tl.load(input_ptr + input_offsets) - accumulator
    lower_mask = (row_tile > col_tile) | (
        rows[:, None] >= cols[None, :]
    )
    tl.store(schur_ptr + schur_offsets, tl.where(lower_mask, values, 0.0))


def _serial_cholesky(data):
    output = torch.empty_like(data)
    info = torch.empty(
        (data.shape[0],),
        dtype=torch.int32,
        device=data.device,
    )
    for index in range(data.shape[0]):
        torch.linalg.cholesky_ex(
            data[index],
            check_errors=False,
            out=(output[index], info[index]),
        )
    return output


def _pure_deferred_large_blocks(data: torch.Tensor, block: int) -> torch.Tensor:
    _, n, _ = data.shape
    output = torch.zeros_like(data)

    for tile_start in range(0, n, block):
        tile_end = tile_start + block
        if tile_start == 0:
            source = data
        else:
            active_history = output[:, tile_start:, :tile_start]
            pivot_history = output[:, tile_start:tile_end, :tile_start]
            torch.baddbmm(
                data[:, tile_start:, tile_start:tile_end],
                active_history,
                pivot_history.transpose(1, 2),
                beta=1.0,
                alpha=-1.0,
                out=output[:, tile_start:, tile_start:tile_end],
            )
            source = output

        diagonal = source[:, tile_start:tile_end, tile_start:tile_end]
        diagonal_factor = torch.linalg.cholesky_ex(
            diagonal,
            check_errors=False,
        ).L
        output[:, tile_start:tile_end, tile_start:tile_end].copy_(
            diagonal_factor
        )

        if tile_end < n:
            panel = source[:, tile_end:, tile_start:tile_end]
            solved = torch.linalg.solve_triangular(
                diagonal_factor,
                panel.transpose(-1, -2),
                upper=False,
            )
            output[:, tile_end:, tile_start:tile_end].copy_(
                solved.transpose(-1, -2)
            )

    diagonal = torch.diagonal(output, dim1=-2, dim2=-1)
    valid = torch.isfinite(diagonal).all() & (diagonal > 0.0).all()
    if not bool(valid.item()):
        return torch.linalg.cholesky_ex(data, check_errors=False).L
    return output


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if (n == 2048 and batch == 2) or (n == 4096 and batch == 2):
        return _serial_cholesky(data)

    if n == 2048 and batch >= 8:
        output = torch.empty_like(data)
        for tile_start in range(0, 2048, 128):
            remaining_blocks = (1920 - tile_start) // 128
            if tile_start == 0:
                source = data
            else:
                active_history = output[:, tile_start:, :tile_start]
                pivot_history = output[
                    :,
                    tile_start:tile_start + 128,
                    :tile_start,
                ]
                torch.baddbmm(
                    data[:, tile_start:, tile_start:tile_start + 128],
                    active_history,
                    pivot_history.transpose(1, 2),
                    beta=1.0,
                    alpha=-1.0,
                    out=output[
                        :,
                        tile_start:,
                        tile_start:tile_start + 128,
                    ],
                )
                source = output
            _cuda_n32_module.factor64_strided128_cuda(
                source,
                output,
                tile_start,
            )
            first_half_panels = 1 + remaining_blocks * 2
            _cuda_n32_module.trsm2048_panel64_cuda(
                source,
                output,
                tile_start,
                first_half_panels,
            )
            _update2048_second_panel64[
                (batch * first_half_panels,)
            ](
                source,
                output,
                n * n,
                tile_start=tile_start,
                panel_count=first_half_panels,
                dot_precision="tf32x3",
                num_warps=4,
            )
            _cuda_n32_module.factor64_strided128_cuda(
                output,
                output,
                tile_start + 64,
            )
            remaining_tiles = remaining_blocks * 2
            if remaining_tiles:
                _cuda_n32_module.trsm2048_panel64_cuda(
                    output,
                    output,
                    tile_start + 64,
                    remaining_tiles,
                )
        diagonal = torch.diagonal(output, dim1=-2, dim2=-1)
        valid = torch.isfinite(diagonal).all() & (diagonal > 0.0).all()
        if not bool(valid.item()):
            return torch.linalg.cholesky_ex(data, check_errors=False).L
        return output

    if n == 1024 and batch >= 4:
        output = torch.empty_like(data)
        for tile_start in range(0, 1024, 128):
            source = data if tile_start == 0 else output
            remaining_blocks = (896 - tile_start) // 128
            _cuda_n32_module.factor64_strided128_cuda(
                source,
                output,
                tile_start,
            )
            first_half_panels = 1 + remaining_blocks * 2
            _cuda_n32_module.trsm1024_panel64_cuda(
                source,
                output,
                tile_start,
                first_half_panels,
            )
            _update1024_second_panel64[
                (batch * first_half_panels,)
            ](
                source,
                output,
                n * n,
                tile_start=tile_start,
                panel_count=first_half_panels,
                dot_precision="tf32x3",
                num_warps=4,
            )
            _cuda_n32_module.factor64_strided128_cuda(
                output,
                output,
                tile_start + 64,
            )
            remaining_tiles = remaining_blocks * 2
            if remaining_tiles:
                _cuda_n32_module.trsm1024_panel64_cuda(
                    output,
                    output,
                    tile_start + 64,
                    remaining_tiles,
                )
                pair_count = remaining_tiles * (remaining_tiles + 1) // 2
                _update1024_tiles128[
                    (batch * pair_count,)
                ](
                    source,
                    output,
                    n * n,
                    tile_start=tile_start,
                    remaining_tiles=remaining_tiles,
                    dot_precision="tf32x3",
                    num_warps=4,
                )
        return output

    if batch == 1 and n in (8192, 16384, 32768):
        return _pure_deferred_large_blocks(data, 2048)

    if batch == 1 and n == 4096:
        return torch.linalg.cholesky_ex(data, check_errors=False).L

    if n not in (32, 64, 128, 256, 512):
        return torch.linalg.cholesky_ex(data, check_errors=False).L

    output = torch.empty_like(data)
    if n == 32:
        _cuda_n32_module.cholesky32_cuda(data, output)
    elif n == 64:
        if batch % 4 == 0:
            _cuda_n32_module.cholesky64_recursive_cuda(data, output)
        else:
            _cholesky64_fused[(batch,)](
                data,
                output,
                n * n,
                num_warps=1,
            )
    elif n == 128:
        _cuda_n32_module.factor64_strided128_cuda(data, output, 0)
        _cuda_n32_module.trsm128_panel64_cuda(data, output)
        _syrk128_trailing64[(batch * 3,)](
            data,
            output,
            n * n,
            num_warps=2,
        )
        _cuda_n32_module.factor64_strided128_cuda(output, output, 64)
    elif n == 256:
        for tile_start, remaining_tiles in ((0, 3), (64, 2), (128, 1)):
            source = data if tile_start == 0 else output
            _cuda_n32_module.factor64_strided128_cuda(
                source,
                output,
                tile_start,
            )
            _cuda_n32_module.trsm256_panel64_cuda(
                source,
                output,
                tile_start,
                remaining_tiles,
            )
            pair_count = remaining_tiles * (remaining_tiles + 1) // 2
            _update256_tiles64[(batch * pair_count,)](
                source,
                output,
                n * n,
                tile_start=tile_start,
                remaining_tiles=remaining_tiles,
                num_warps=4,
            )
        _cuda_n32_module.factor64_strided128_cuda(output, output, 192)
    elif batch == 60 or batch >= 128:
        for tile_start, remaining_blocks in ((0, 3), (128, 2), (256, 1)):
            source = data if tile_start == 0 else output
            _cuda_n32_module.factor64_strided128_cuda(
                source,
                output,
                tile_start,
            )
            first_half_panels = 1 + remaining_blocks * 2
            _cuda_n32_module.trsm512_panel64_cuda(
                source,
                output,
                tile_start,
                first_half_panels,
            )
            _update512_second_panel64[(batch * first_half_panels,)](
                source,
                output,
                n * n,
                tile_start=tile_start,
                panel_count=first_half_panels,
                num_warps=4,
            )
            _cuda_n32_module.factor64_strided128_cuda(
                output,
                output,
                tile_start + 64,
            )
            remaining_tiles = remaining_blocks * 2
            _cuda_n32_module.trsm512_panel64_cuda(
                output,
                output,
                tile_start + 64,
                remaining_tiles,
            )
            pair_count = remaining_tiles * (remaining_tiles + 1) // 2
            _update512_tiles128[(batch * pair_count,)](
                source,
                output,
                n * n,
                tile_start=tile_start,
                remaining_tiles=remaining_tiles,
                num_warps=4 if batch == 60 else 8,
            )

        _cuda_n32_module.factor64_strided128_cuda(
            output,
            output,
            384,
        )
        _cuda_n32_module.trsm512_panel64_cuda(
            output,
            output,
            384,
            1,
        )
        _update512_second_panel64[(batch,)](
            output,
            output,
            n * n,
            tile_start=384,
            panel_count=1,
            num_warps=4,
        )
        _cuda_n32_module.factor64_strided128_cuda(
            output,
            output,
            448,
        )
    else:
        trsm_warps = 1 if batch >= 128 else 2
        dot_precision = "tf32" if batch >= 128 else "tf32x3"
        use_2d_grid = batch < 128
        panel_steps = (
            (0, 7),
            (64, 6),
            (128, 5),
            (192, 4),
            (256, 3),
            (320, 2),
            (384, 1),
        )
        for tile_start, remaining_tiles in panel_steps:
            source = data if tile_start == 0 else output
            _cuda_n32_module.factor64_strided128_cuda(
                source,
                output,
                tile_start,
            )
            _cuda_n32_module.trsm512_panel64_cuda(
                source,
                output,
                tile_start,
                remaining_tiles,
            )
            pair_count = remaining_tiles * (remaining_tiles + 1) // 2
            if batch >= 128:
                _update512_tiles64[(batch * pair_count,)](
                    source,
                    output,
                    n * n,
                    tile_start=tile_start,
                    remaining_tiles=remaining_tiles,
                    dot_precision=dot_precision,
                    use_2d_grid=False,
                    num_warps=4,
                )
            else:
                _update512_tiles32[(pair_count * 4, batch)](
                    source,
                    output,
                    n * n,
                    tile_start=tile_start,
                    remaining_tiles=remaining_tiles,
                    dot_precision=dot_precision,
                    use_2d_grid=True,
                    num_warps=2,
                )
        _cuda_n32_module.factor64_strided128_cuda(
            output,
            output,
            448,
        )
    return output
scrolls · 2629 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