Skip to content
KernelIndex
Search⌘K

submission 914518

Deepon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-914518?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.11ms
#128 of 337
2026-07-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b36a00debfc5b74c14cdfdaa7d5e1f0a69544b16360a05b197be746f1505f5c1
license declaredunknown
license concludedunknown
authorsDeepon
imported2026-08-26

Techniques

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

mmanamespace wmma = nvcuda::wmma;
num-warps = 8num_warps=8,
persistent-kernel_PERSISTENT_WMMA_SOURCE = r"""
shared-memory__shared__ float tiles[kWarpsPerBlock][32 * 33];
stages = 4num_stages=4,
tile-k = 32BLOCK_K=32,

Kernel source

submission.py2756 lines
import torch
import triton
import triton.language as tl
from functools import lru_cache
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


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

namespace {

constexpr int kWarpsPerBlock = 4;

__global__ void cholesky32_warp_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch
) {
    __shared__ float tiles[kWarpsPerBlock][32 * 33];

    const int warp_in_block = threadIdx.x >> 5;
    const int lane = threadIdx.x & 31;
    const int matrix = blockIdx.x * kWarpsPerBlock + warp_in_block;
    if (matrix >= batch) {
        return;
    }

    float* tile = tiles[warp_in_block];
    const float* matrix_input = input + static_cast<long long>(matrix) * 1024;
    float* matrix_output = output + static_cast<long long>(matrix) * 1024;

    #pragma unroll
    for (int linear = lane; linear < 1024; linear += 32) {
        const int row = linear >> 5;
        const int col = linear & 31;
        tile[row * 33 + col] =
            row >= col ? matrix_input[linear] : 0.0f;
    }
    __syncwarp();

    float row_values[32];
    #pragma unroll
    for (int col = 0; col < 32; ++col) {
        row_values[col] = tile[lane * 33 + col];
    }

    constexpr unsigned kMask = 0xffffffffu;
    #pragma unroll
    for (int pivot = 0; pivot < 32; ++pivot) {
        float dot = 0.0f;
        #pragma unroll
        for (int col = 0; col < pivot; ++col) {
            const float pivot_value =
                __shfl_sync(kMask, row_values[col], pivot);
            dot = fmaf(row_values[col], pivot_value, dot);
        }

        float diagonal = 0.0f;
        if (lane == pivot) {
            diagonal = sqrtf(fmaxf(row_values[pivot] - dot, 0.0f));
        }
        diagonal = __shfl_sync(kMask, diagonal, pivot);

        if (lane == pivot) {
            row_values[pivot] = diagonal;
        } else if (lane > pivot) {
            row_values[pivot] =
                (row_values[pivot] - dot) / diagonal;
        }
    }

    #pragma unroll
    for (int col = 0; col < 32; ++col) {
        tile[lane * 33 + col] = row_values[col];
    }
    __syncwarp();

    #pragma unroll
    for (int linear = lane; linear < 1024; linear += 32) {
        const int row = linear >> 5;
        const int col = linear & 31;
        matrix_output[linear] = tile[row * 33 + col];
    }
}

}  // namespace

torch::Tensor cholesky32_cuda(torch::Tensor input) {
    auto output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int blocks = (batch + kWarpsPerBlock - 1) / kWarpsPerBlock;
    cholesky32_warp_kernel<<<
        blocks,
        kWarpsPerBlock * 32
    >>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        batch
    );
    return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def("cholesky32", &cholesky32_cuda);
}
"""


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

namespace {

constexpr int kWarpsPerBlock = 2;
constexpr int kN = 64;

template <bool kMakeInverse>
__global__ void cholesky64_register_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    float* __restrict__ inverse,
    int batch
) {
    __shared__ float tiles[kWarpsPerBlock][kN * (kN + 1)];

    const int warp_in_block = threadIdx.x >> 5;
    const int lane = threadIdx.x & 31;
    const int matrix = blockIdx.x * kWarpsPerBlock + warp_in_block;
    if (matrix >= batch) {
        return;
    }

    const float* matrix_input =
        input + static_cast<long long>(matrix) * kN * kN;
    float* matrix_output =
        output + static_cast<long long>(matrix) * kN * kN;
    float* matrix_inverse = kMakeInverse
        ? inverse + static_cast<long long>(matrix) * kN * kN
        : nullptr;
    float* tile = tiles[warp_in_block];
    const int row0 = lane;
    const int row1 = lane + 32;

    #pragma unroll 4
    for (int linear = lane; linear < kN * kN; linear += 32) {
        const int row = linear >> 6;
        const int col = linear & 63;
        tile[row * (kN + 1) + col] =
            row >= col ? matrix_input[linear] : 0.0f;
    }
    __syncwarp();

    float values0[kN];
    float values1[kN];
    #pragma unroll
    for (int col = 0; col < kN; ++col) {
        values0[col] = tile[row0 * (kN + 1) + col];
        values1[col] = tile[row1 * (kN + 1) + col];
    }

    constexpr unsigned kMask = 0xffffffffu;
    #pragma unroll
    for (int pivot = 0; pivot < kN; ++pivot) {
        float dot0 = 0.0f;
        float dot1 = 0.0f;
        #pragma unroll
        for (int col = 0; col < pivot; ++col) {
            const float source =
                pivot < 32 ? values0[col] : values1[col];
            const float pivot_value =
                __shfl_sync(kMask, source, pivot & 31);
            dot0 = fmaf(values0[col], pivot_value, dot0);
            dot1 = fmaf(values1[col], pivot_value, dot1);
        }

        float diagonal = 0.0f;
        if (row0 == pivot) {
            diagonal = sqrtf(
                fmaxf(values0[pivot] - dot0, 0.0f)
            );
        } else if (row1 == pivot) {
            diagonal = sqrtf(
                fmaxf(values1[pivot] - dot1, 0.0f)
            );
        }
        diagonal = __shfl_sync(kMask, diagonal, pivot & 31);

        if (row0 == pivot) {
            values0[pivot] = diagonal;
        } else if (row0 > pivot) {
            values0[pivot] = (values0[pivot] - dot0) / diagonal;
        }
        if (row1 == pivot) {
            values1[pivot] = diagonal;
        } else if (row1 > pivot) {
            values1[pivot] = (values1[pivot] - dot1) / diagonal;
        }
    }

    #pragma unroll
    for (int col = 0; col < kN; ++col) {
        tile[row0 * (kN + 1) + col] = values0[col];
        tile[row1 * (kN + 1) + col] = values1[col];
    }
    __syncwarp();

    #pragma unroll 4
    for (int linear = lane; linear < kN * kN; linear += 32) {
        const int row = linear >> 6;
        const int col = linear & 63;
        matrix_output[linear] = tile[row * (kN + 1) + col];
    }

    if constexpr (kMakeInverse) {
        #pragma unroll
        for (int which = 0; which < 2; ++which) {
            const int row = which == 0 ? row0 : row1;
            const float inverse_diagonal =
                1.0f / tile[row * (kN + 1) + row];
            #pragma unroll 1
            for (int col = row - 1; col >= 0; --col) {
                float value = 0.0f;
                #pragma unroll 4
                for (int inner = col + 1; inner <= row; ++inner) {
                    const float inverse_item =
                        inner == row
                            ? inverse_diagonal
                            : tile[inner * (kN + 1) + row];
                    value = fmaf(
                        inverse_item,
                        tile[inner * (kN + 1) + col],
                        value
                    );
                }
                tile[col * (kN + 1) + row] =
                    -value / tile[col * (kN + 1) + col];
            }
        }
        __syncwarp();

        #pragma unroll 4
        for (int linear = lane; linear < kN * kN; linear += 32) {
            const int row = linear >> 6;
            const int col = linear & 63;
            matrix_inverse[linear] =
                col < row
                    ? tile[col * (kN + 1) + row]
                    : (
                        col == row
                            ? 1.0f / tile[row * (kN + 1) + row]
                            : 0.0f
                    );
        }
    }
}

}  // namespace

torch::Tensor cholesky64_cuda(torch::Tensor input) {
    auto output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int blocks = (batch + kWarpsPerBlock - 1) / kWarpsPerBlock;
    cholesky64_register_kernel<false><<<
        blocks,
        kWarpsPerBlock * 32
    >>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        nullptr,
        batch
    );
    return output;
}

std::vector<torch::Tensor> cholesky64_inverse_cuda(torch::Tensor input) {
    auto output = torch::empty_like(input);
    auto inverse = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int blocks = (batch + kWarpsPerBlock - 1) / kWarpsPerBlock;
    cholesky64_register_kernel<true><<<
        blocks,
        kWarpsPerBlock * 32
    >>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        inverse.data_ptr<float>(),
        batch
    );
    return {output, inverse};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def("cholesky64", &cholesky64_cuda);
    module.def("cholesky64_inverse", &cholesky64_inverse_cuda);
}
"""


_CUDA128_SOURCE = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <c10/cuda/CUDAException.h>

namespace {

constexpr int kN = 128;
constexpr int kPitch = kN + 1;
constexpr int kThreads = 256;
constexpr int kSharedBytes = kN * kPitch * sizeof(float);

__global__ void cholesky128_panel_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch
) {
    constexpr int kPanel = 32;
    constexpr int kMatrixElements = kN * kN;
    extern __shared__ float tile[];

    const int matrix = static_cast<int>(blockIdx.x);
    if (matrix >= batch) {
        return;
    }

    const int thread = static_cast<int>(threadIdx.x);
    const int lane = thread & 31;
    const int warp = thread >> 5;
    const long long matrix_offset =
        static_cast<long long>(matrix) * kMatrixElements;
    const float* matrix_input = input + matrix_offset;
    float* matrix_output = output + matrix_offset;

    for (
        int element = thread;
        element < kMatrixElements;
        element += kThreads
    ) {
        const int row = element / kN;
        const int col = element - row * kN;
        tile[row * kPitch + col] =
            row >= col ? matrix_input[element] : 0.0f;
    }
    __syncthreads();

    #pragma unroll 1
    for (int start = 0; start < kN; start += kPanel) {
        if (warp == 0) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                if (lane == local_col) {
                    const int diagonal = start + local_col;
                    float value = tile[diagonal * kPitch + diagonal];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        const float item = tile[
                            diagonal * kPitch + start + previous
                        ];
                        value = fmaf(-item, item, value);
                    }
                    tile[diagonal * kPitch + diagonal] =
                        sqrtf(fmaxf(value, 0.0f));
                }
                __syncwarp();

                if (lane > local_col) {
                    const int row = start + lane;
                    const int col = start + local_col;
                    float value = tile[row * kPitch + col];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        value = fmaf(
                            -tile[row * kPitch + start + previous],
                            tile[col * kPitch + start + previous],
                            value
                        );
                    }
                    tile[row * kPitch + col] =
                        value / tile[col * kPitch + col];
                }
                __syncwarp();
            }
        }
        __syncthreads();

        const int panel_end = start + kPanel;
        for (
            int row = panel_end + thread;
            row < kN;
            row += kThreads
        ) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                const int col = start + local_col;
                float value = tile[row * kPitch + col];
                #pragma unroll
                for (
                    int previous = 0;
                    previous < local_col;
                    ++previous
                ) {
                    value = fmaf(
                        -tile[row * kPitch + start + previous],
                        tile[col * kPitch + start + previous],
                        value
                    );
                }
                tile[row * kPitch + col] =
                    value / tile[col * kPitch + col];
            }
        }
        __syncthreads();

        const int trailing = kN - panel_end;
        const int trailing_elements = trailing * trailing;
        for (
            int local = thread;
            local < trailing_elements;
            local += kThreads
        ) {
            const int local_row = local / trailing;
            const int local_col = local - local_row * trailing;
            if (local_row >= local_col) {
                const int row = panel_end + local_row;
                const int col = panel_end + local_col;
                float value = tile[row * kPitch + col];
                #pragma unroll
                for (int previous = 0; previous < kPanel; ++previous) {
                    value = fmaf(
                        -tile[row * kPitch + start + previous],
                        tile[col * kPitch + start + previous],
                        value
                    );
                }
                tile[row * kPitch + col] = value;
            }
        }
        __syncthreads();
    }

    for (
        int element = thread;
        element < kMatrixElements;
        element += kThreads
    ) {
        const int row = element / kN;
        const int col = element - row * kN;
        matrix_output[element] =
            row >= col ? tile[row * kPitch + col] : 0.0f;
    }
}

bool configure_cholesky128() {
    const cudaError_t shared_result = cudaFuncSetAttribute(
        cholesky128_panel_kernel,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        kSharedBytes
    );
    const cudaError_t carveout_result = cudaFuncSetAttribute(
        cholesky128_panel_kernel,
        cudaFuncAttributePreferredSharedMemoryCarveout,
        100
    );
    return shared_result == cudaSuccess && carveout_result == cudaSuccess;
}

}  // namespace

torch::Tensor cholesky128_cuda(torch::Tensor input) {
    static const bool configured = configure_cholesky128();
    TORCH_CHECK(configured, "failed to configure shared memory");

    auto output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    cholesky128_panel_kernel<<<
        batch,
        kThreads,
        kSharedBytes
    >>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        batch
    );
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def("cholesky128", &cholesky128_cuda);
}
"""


_CUDA128_INVERSE_SOURCE = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <c10/cuda/CUDAException.h>
#include <vector>

namespace {

constexpr int kN = 128;
constexpr int kPitch = kN + 1;
constexpr int kPanel = 32;
constexpr int kThreads = 256;
constexpr int kMatrixElements = kN * kN;
constexpr int kSharedBytes = kN * kPitch * sizeof(float);

__global__ void cholesky128_inverse_kernel(
    const float* __restrict__ input,
    float* __restrict__ factor_output,
    float* __restrict__ inverse_output,
    int batch
) {
    extern __shared__ float tile[];

    const int matrix = static_cast<int>(blockIdx.x);
    if (matrix >= batch) {
        return;
    }

    const int thread = static_cast<int>(threadIdx.x);
    const int lane = thread & 31;
    const int warp = thread >> 5;
    const long long matrix_offset =
        static_cast<long long>(matrix) * kMatrixElements;
    const float* matrix_input = input + matrix_offset;
    float* matrix_factor = factor_output + matrix_offset;
    float* matrix_inverse = inverse_output + matrix_offset;

    for (
        int element = thread;
        element < kMatrixElements;
        element += kThreads
    ) {
        const int row = element / kN;
        const int col = element - row * kN;
        tile[row * kPitch + col] =
            row >= col ? matrix_input[element] : 0.0f;
    }
    __syncthreads();

    #pragma unroll 1
    for (int start = 0; start < kN; start += kPanel) {
        if (warp == 0) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                if (lane == local_col) {
                    const int diagonal = start + local_col;
                    float value =
                        tile[diagonal * kPitch + diagonal];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        const float item = tile[
                            diagonal * kPitch + start + previous
                        ];
                        value = fmaf(-item, item, value);
                    }
                    tile[diagonal * kPitch + diagonal] =
                        sqrtf(fmaxf(value, 0.0f));
                }
                __syncwarp();

                if (lane > local_col) {
                    const int row = start + lane;
                    const int col = start + local_col;
                    float value = tile[row * kPitch + col];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        value = fmaf(
                            -tile[row * kPitch + start + previous],
                            tile[col * kPitch + start + previous],
                            value
                        );
                    }
                    tile[row * kPitch + col] =
                        value / tile[col * kPitch + col];
                }
                __syncwarp();
            }
        }
        __syncthreads();

        const int panel_end = start + kPanel;
        for (
            int row = panel_end + thread;
            row < kN;
            row += kThreads
        ) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                const int col = start + local_col;
                float value = tile[row * kPitch + col];
                #pragma unroll
                for (
                    int previous = 0;
                    previous < local_col;
                    ++previous
                ) {
                    value = fmaf(
                        -tile[row * kPitch + start + previous],
                        tile[col * kPitch + start + previous],
                        value
                    );
                }
                tile[row * kPitch + col] =
                    value / tile[col * kPitch + col];
            }
        }
        __syncthreads();

        const int trailing = kN - panel_end;
        const int trailing_elements = trailing * trailing;
        for (
            int local = thread;
            local < trailing_elements;
            local += kThreads
        ) {
            const int local_row = local / trailing;
            const int local_col = local - local_row * trailing;
            if (local_row >= local_col) {
                const int row = panel_end + local_row;
                const int col = panel_end + local_col;
                float value = tile[row * kPitch + col];
                #pragma unroll
                for (int previous = 0; previous < kPanel; ++previous) {
                    value = fmaf(
                        -tile[row * kPitch + start + previous],
                        tile[col * kPitch + start + previous],
                        value
                    );
                }
                tile[row * kPitch + col] = value;
            }
        }
        __syncthreads();
    }

    if (thread < kN) {
        const int row = thread;
        const float inverse_diagonal =
            1.0f / tile[row * kPitch + row];
        #pragma unroll 1
        for (int col = row - 1; col >= 0; --col) {
            float value = 0.0f;
            #pragma unroll 4
            for (int inner = col + 1; inner <= row; ++inner) {
                const float inverse_item =
                    inner == row
                        ? inverse_diagonal
                        : tile[inner * kPitch + row];
                value = fmaf(
                    inverse_item,
                    tile[inner * kPitch + col],
                    value
                );
            }
            tile[col * kPitch + row] =
                -value / tile[col * kPitch + col];
        }
    }
    __syncthreads();

    for (
        int element = thread;
        element < kMatrixElements;
        element += kThreads
    ) {
        const int row = element / kN;
        const int col = element - row * kN;
        matrix_factor[element] =
            row >= col ? tile[row * kPitch + col] : 0.0f;
        matrix_inverse[element] =
            col < row
                ? tile[col * kPitch + row]
                : (
                    col == row
                        ? 1.0f / tile[row * kPitch + row]
                        : 0.0f
                );
    }
}

bool configure_cholesky128_inverse() {
    const cudaError_t shared_result = cudaFuncSetAttribute(
        cholesky128_inverse_kernel,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        kSharedBytes
    );
    const cudaError_t carveout_result = cudaFuncSetAttribute(
        cholesky128_inverse_kernel,
        cudaFuncAttributePreferredSharedMemoryCarveout,
        100
    );
    return shared_result == cudaSuccess && carveout_result == cudaSuccess;
}

}  // namespace

std::vector<torch::Tensor> cholesky128_inverse_cuda(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be FP32");
    TORCH_CHECK(input.dim() == 3, "input rank must be three");
    TORCH_CHECK(
        input.size(1) == kN && input.size(2) == kN,
        "matrix dimensions must be 128x128"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");

    static const bool configured = configure_cholesky128_inverse();
    TORCH_CHECK(configured, "failed to configure shared memory");

    auto factor = torch::empty_like(input);
    auto inverse = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    cholesky128_inverse_kernel<<<
        batch,
        kThreads,
        kSharedBytes
    >>>(
        input.data_ptr<float>(),
        factor.data_ptr<float>(),
        inverse.data_ptr<float>(),
        batch
    );
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return {factor, inverse};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def("cholesky128_inverse", &cholesky128_inverse_cuda);
}
"""


_CUDA256_PACKED_SOURCE = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <c10/cuda/CUDAException.h>

namespace {

constexpr int kN = 256;
constexpr int kPanel = 32;
constexpr int kThreads = 256;
constexpr int kMatrixElements = kN * kN;
constexpr int kPackedElements = kN * (kN + 1) / 2;
constexpr int kSharedBytes = kPackedElements * sizeof(float);

__device__ __forceinline__ int packed_offset(int row, int col) {
    return row * (row + 1) / 2 + col;
}

__global__ void cholesky256_packed_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch
) {
    extern __shared__ float packed[];

    const int matrix = static_cast<int>(blockIdx.x);
    if (matrix >= batch) {
        return;
    }

    const int thread = static_cast<int>(threadIdx.x);
    const int lane = thread & 31;
    const int warp = thread >> 5;
    const long long matrix_offset =
        static_cast<long long>(matrix) * kMatrixElements;
    const float* matrix_input = input + matrix_offset;
    float* matrix_output = output + matrix_offset;

    for (
        int element = thread;
        element < kMatrixElements;
        element += kThreads
    ) {
        const int row = element >> 8;
        const int col = element & (kN - 1);
        if (row >= col) {
            packed[packed_offset(row, col)] = matrix_input[element];
        }
    }
    __syncthreads();

    #pragma unroll 1
    for (int start = 0; start < kN; start += kPanel) {
        if (warp == 0) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                if (lane == local_col) {
                    const int diagonal = start + local_col;
                    float value =
                        packed[packed_offset(diagonal, diagonal)];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        const float item = packed[
                            packed_offset(
                                diagonal,
                                start + previous
                            )
                        ];
                        value = fmaf(-item, item, value);
                    }
                    packed[packed_offset(diagonal, diagonal)] =
                        sqrtf(fmaxf(value, 0.0f));
                }
                __syncwarp();

                if (lane > local_col) {
                    const int row = start + lane;
                    const int col = start + local_col;
                    float value = packed[packed_offset(row, col)];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        value = fmaf(
                            -packed[
                                packed_offset(
                                    row,
                                    start + previous
                                )
                            ],
                            packed[
                                packed_offset(
                                    col,
                                    start + previous
                                )
                            ],
                            value
                        );
                    }
                    packed[packed_offset(row, col)] =
                        value / packed[packed_offset(col, col)];
                }
                __syncwarp();
            }
        }
        __syncthreads();

        const int panel_end = start + kPanel;
        for (
            int row = panel_end + thread;
            row < kN;
            row += kThreads
        ) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                const int col = start + local_col;
                float value = packed[packed_offset(row, col)];
                #pragma unroll
                for (
                    int previous = 0;
                    previous < local_col;
                    ++previous
                ) {
                    value = fmaf(
                        -packed[
                            packed_offset(
                                row,
                                start + previous
                            )
                        ],
                        packed[
                            packed_offset(
                                col,
                                start + previous
                            )
                        ],
                        value
                    );
                }
                packed[packed_offset(row, col)] =
                    value / packed[packed_offset(col, col)];
            }
        }
        __syncthreads();

        const int trailing = kN - panel_end;
        const int trailing_elements = trailing * trailing;
        for (
            int local = thread;
            local < trailing_elements;
            local += kThreads
        ) {
            const int local_row = local / trailing;
            const int local_col = local - local_row * trailing;
            if (local_row >= local_col) {
                const int row = panel_end + local_row;
                const int col = panel_end + local_col;
                float value = packed[packed_offset(row, col)];
                #pragma unroll
                for (int previous = 0; previous < kPanel; ++previous) {
                    value = fmaf(
                        -packed[
                            packed_offset(
                                row,
                                start + previous
                            )
                        ],
                        packed[
                            packed_offset(
                                col,
                                start + previous
                            )
                        ],
                        value
                    );
                }
                packed[packed_offset(row, col)] = value;
            }
        }
        __syncthreads();
    }

    for (
        int element = thread;
        element < kMatrixElements;
        element += kThreads
    ) {
        const int row = element >> 8;
        const int col = element & (kN - 1);
        matrix_output[element] =
            row >= col ? packed[packed_offset(row, col)] : 0.0f;
    }
}

bool configure_cholesky256_packed() {
    const cudaError_t shared_result = cudaFuncSetAttribute(
        cholesky256_packed_kernel,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        kSharedBytes
    );
    const cudaError_t carveout_result = cudaFuncSetAttribute(
        cholesky256_packed_kernel,
        cudaFuncAttributePreferredSharedMemoryCarveout,
        100
    );
    return shared_result == cudaSuccess && carveout_result == cudaSuccess;
}

}  // namespace

torch::Tensor cholesky256_packed_cuda(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be FP32");
    TORCH_CHECK(input.dim() == 3, "input rank must be three");
    TORCH_CHECK(
        input.size(1) == kN && input.size(2) == kN,
        "matrix dimensions must be 256x256"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");

    static const bool configured = configure_cholesky256_packed();
    TORCH_CHECK(configured, "failed to configure shared memory");

    auto output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    cholesky256_packed_kernel<<<
        batch,
        kThreads,
        kSharedBytes
    >>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        batch
    );
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def("cholesky256_packed", &cholesky256_packed_cuda);
}
"""


_PERSISTENT_WMMA_SOURCE = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <mma.h>
#include <c10/cuda/CUDAException.h>

namespace {

namespace wmma = nvcuda::wmma;

constexpr int kPanel = 32;
constexpr int kThreads = 256;
constexpr int kWarps = kThreads / 32;

template <int N>
__global__ void persistent_wmma_cholesky_kernel(
    const float* __restrict__ input,
    half* __restrict__ history,
    float* __restrict__ output,
    int batch
) {
    extern __shared__ float panel[];

    const int matrix = static_cast<int>(blockIdx.x);
    if (matrix >= batch) {
        return;
    }

    const int thread = static_cast<int>(threadIdx.x);
    const int lane = thread & 31;
    const int warp = thread >> 5;
    const long long matrix_offset =
        static_cast<long long>(matrix) * N * N;
    const float* matrix_input = input + matrix_offset;
    half* matrix_history = history + matrix_offset;
    float* matrix_output = output + matrix_offset;

    #pragma unroll 1
    for (int start = 0; start < N; start += kPanel) {
        const int row_tiles = (N - start) / 16;
        const int tile_jobs = row_tiles * 2;

        for (int job = warp; job < tile_jobs; job += kWarps) {
            const int row_tile = job >> 1;
            const int col_tile = job & 1;
            const int row_relative = row_tile * 16;
            const int row_global = start + row_relative;
            const int col_relative = col_tile * 16;

            wmma::fragment<
                wmma::accumulator,
                16,
                16,
                16,
                float
            > accumulator;
            wmma::load_matrix_sync(
                accumulator,
                matrix_input
                    + static_cast<long long>(row_global) * N
                    + start
                    + col_relative,
                N,
                wmma::mem_row_major
            );
            #pragma unroll
            for (
                int element = 0;
                element < accumulator.num_elements;
                ++element
            ) {
                accumulator.x[element] = -accumulator.x[element];
            }

            for (int inner = 0; inner < start; inner += 16) {
                wmma::fragment<
                    wmma::matrix_a,
                    16,
                    16,
                    16,
                    half,
                    wmma::row_major
                > left_fragment;
                wmma::fragment<
                    wmma::matrix_b,
                    16,
                    16,
                    16,
                    half,
                    wmma::col_major
                > right_fragment;
                wmma::load_matrix_sync(
                    left_fragment,
                    matrix_history
                        + static_cast<long long>(row_global) * N
                        + inner,
                    N
                );
                wmma::load_matrix_sync(
                    right_fragment,
                    matrix_history
                        + static_cast<long long>(
                            start + col_relative
                        ) * N
                        + inner,
                    N
                );
                wmma::mma_sync(
                    accumulator,
                    left_fragment,
                    right_fragment,
                    accumulator
                );
            }

            #pragma unroll
            for (
                int element = 0;
                element < accumulator.num_elements;
                ++element
            ) {
                accumulator.x[element] = -accumulator.x[element];
            }
            wmma::store_matrix_sync(
                panel + row_relative * kPanel + col_relative,
                accumulator,
                kPanel,
                wmma::mem_row_major
            );
        }
        __syncthreads();

        if (warp == 0) {
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                if (lane == local_col) {
                    float value =
                        panel[local_col * kPanel + local_col];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        const float item =
                            panel[local_col * kPanel + previous];
                        value = fmaf(-item, item, value);
                    }
                    const float diagonal = sqrtf(fmaxf(value, 0.0f));
                    const half quantized = __float2half_rn(diagonal);
                    panel[local_col * kPanel + local_col] =
                        __half2float(quantized);
                    matrix_history[
                        static_cast<long long>(start + local_col) * N
                        + start
                        + local_col
                    ] = quantized;
                }
                __syncwarp();

                if (lane > local_col) {
                    float value = panel[lane * kPanel + local_col];
                    #pragma unroll
                    for (
                        int previous = 0;
                        previous < local_col;
                        ++previous
                    ) {
                        value = fmaf(
                            -panel[lane * kPanel + previous],
                            panel[local_col * kPanel + previous],
                            value
                        );
                    }
                    value /= panel[
                        local_col * kPanel + local_col
                    ];
                    const half quantized = __float2half_rn(value);
                    panel[lane * kPanel + local_col] =
                        __half2float(quantized);
                    matrix_history[
                        static_cast<long long>(start + lane) * N
                        + start
                        + local_col
                    ] = quantized;
                }
                __syncwarp();
            }

            for (int local_col = 0; local_col <= lane; ++local_col) {
                matrix_output[
                    static_cast<long long>(start + lane) * N
                    + start
                    + local_col
                ] = panel[lane * kPanel + local_col];
            }
        }
        __syncthreads();

        for (
            int row = start + kPanel + thread;
            row < N;
            row += kThreads
        ) {
            const int row_relative = row - start;
            float* row_values = panel + row_relative * kPanel;
            #pragma unroll
            for (int local_col = 0; local_col < kPanel; ++local_col) {
                float value = row_values[local_col];
                #pragma unroll
                for (
                    int previous = 0;
                    previous < local_col;
                    ++previous
                ) {
                    value = fmaf(
                        -row_values[previous],
                        panel[local_col * kPanel + previous],
                        value
                    );
                }
                value /= panel[local_col * kPanel + local_col];
                const half quantized = __float2half_rn(value);
                const float quantized_float = __half2float(quantized);
                row_values[local_col] = quantized_float;
                matrix_history[
                    static_cast<long long>(row) * N
                    + start
                    + local_col
                ] = quantized;
                matrix_output[
                    static_cast<long long>(row) * N
                    + start
                    + local_col
                ] = quantized_float;
            }
        }
        __syncthreads();
    }

    for (
        int element = thread;
        element < N * N;
        element += kThreads
    ) {
        const int row = element / N;
        const int col = element - row * N;
        if (col > row) {
            matrix_output[element] = 0.0f;
        }
    }
}

template <int N>
bool configure_persistent_kernel() {
    constexpr int kSharedBytes = N * kPanel * sizeof(float);
    const cudaError_t shared_result = cudaFuncSetAttribute(
        persistent_wmma_cholesky_kernel<N>,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        kSharedBytes
    );
    const cudaError_t carveout_result = cudaFuncSetAttribute(
        persistent_wmma_cholesky_kernel<N>,
        cudaFuncAttributePreferredSharedMemoryCarveout,
        100
    );
    return shared_result == cudaSuccess && carveout_result == cudaSuccess;
}

}  // namespace

torch::Tensor persistent_wmma_cholesky_cuda(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be FP32");
    TORCH_CHECK(input.dim() == 3, "input rank must be three");
    TORCH_CHECK(input.size(1) == input.size(2), "input must be square");
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");

    const int batch = static_cast<int>(input.size(0));
    const int n = static_cast<int>(input.size(1));
    auto output = torch::empty_like(input);
    auto history = torch::empty(
        input.sizes(),
        input.options().dtype(at::kHalf)
    );

    if (n == 512) {
        static const bool configured = configure_persistent_kernel<512>();
        TORCH_CHECK(configured, "failed to configure n=512 kernel");
        constexpr int shared_bytes = 512 * kPanel * sizeof(float);
        persistent_wmma_cholesky_kernel<512><<<
            batch,
            kThreads,
            shared_bytes
        >>>(
            input.data_ptr<float>(),
            reinterpret_cast<half*>(history.data_ptr<at::Half>()),
            output.data_ptr<float>(),
            batch
        );
    } else if (n == 1024) {
        static const bool configured = configure_persistent_kernel<1024>();
        TORCH_CHECK(configured, "failed to configure n=1024 kernel");
        constexpr int shared_bytes = 1024 * kPanel * sizeof(float);
        persistent_wmma_cholesky_kernel<1024><<<
            batch,
            kThreads,
            shared_bytes
        >>>(
            input.data_ptr<float>(),
            reinterpret_cast<half*>(history.data_ptr<at::Half>()),
            output.data_ptr<float>(),
            batch
        );
    } else {
        TORCH_CHECK(false, "matrix dimension must be 512 or 1024");
    }

    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def(
        "persistent_wmma_cholesky",
        &persistent_wmma_cholesky_cuda
    );
}
"""


_FAST_UPDATE_SOURCE = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContextLight.h>
#include <cublas_v2.h>

torch::Tensor bf16_update_cuda(
    torch::Tensor left,
    torch::Tensor right,
    torch::Tensor output
) {
    TORCH_CHECK(left.is_cuda(), "left must be CUDA");
    TORCH_CHECK(right.is_cuda(), "right must be CUDA");
    TORCH_CHECK(output.is_cuda(), "output must be CUDA");
    TORCH_CHECK(left.scalar_type() == at::kFloat, "left must be FP32");
    TORCH_CHECK(right.scalar_type() == at::kFloat, "right must be FP32");
    TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be FP32");
    TORCH_CHECK(left.dim() == right.dim(), "rank mismatch");
    TORCH_CHECK(left.dim() == output.dim(), "rank mismatch");
    TORCH_CHECK(left.dim() == 2 || left.dim() == 3, "rank must be 2 or 3");
    TORCH_CHECK(left.stride(-1) == 1, "left inner stride");
    TORCH_CHECK(right.stride(-1) == 1, "right inner stride");
    TORCH_CHECK(output.stride(-1) == 1, "output inner stride");

    const int m = static_cast<int>(left.size(-2));
    const int n = static_cast<int>(right.size(-2));
    const int k = static_cast<int>(left.size(-1));
    TORCH_CHECK(right.size(-1) == k, "contracting dimension mismatch");
    TORCH_CHECK(output.size(-2) == m, "output row mismatch");
    TORCH_CHECK(output.size(-1) == n, "output column mismatch");

    const int batches = left.dim() == 3
        ? static_cast<int>(left.size(0))
        : 1;
    TORCH_CHECK(
        right.dim() == 2 || right.size(0) == batches,
        "right batch mismatch"
    );
    TORCH_CHECK(
        output.dim() == 2 || output.size(0) == batches,
        "output batch mismatch"
    );

    const long long left_batch = left.dim() == 3 ? left.stride(0) : 0;
    const long long right_batch = right.dim() == 3 ? right.stride(0) : 0;
    const long long output_batch = output.dim() == 3 ? output.stride(0) : 0;
    const int left_leading = static_cast<int>(left.stride(-2));
    const int right_leading = static_cast<int>(right.stride(-2));
    const int output_leading = static_cast<int>(output.stride(-2));
    const float alpha = -1.0f;
    const float beta = 1.0f;

    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    const cublasStatus_t status = cublasGemmStridedBatchedEx(
        handle,
        CUBLAS_OP_T,
        CUBLAS_OP_N,
        n,
        m,
        k,
        &alpha,
        right.data_ptr<float>(),
        CUDA_R_32F,
        right_leading,
        right_batch,
        left.data_ptr<float>(),
        CUDA_R_32F,
        left_leading,
        left_batch,
        &beta,
        output.data_ptr<float>(),
        CUDA_R_32F,
        output_leading,
        output_batch,
        batches,
        CUBLAS_COMPUTE_32F_FAST_16F,
        CUBLAS_GEMM_DEFAULT_TENSOR_OP
    );
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "cuBLAS update failed");
    return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def("bf16_update", &bf16_update_cuda);
}
"""


_LOWER_CUBLAS_SOURCE = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContextLight.h>
#include <cublas_v2.h>
#include <algorithm>
#include <cstdint>
#include <limits>

namespace {

void check_status(cublasStatus_t status, const char* operation) {
    TORCH_CHECK(
        status == CUBLAS_STATUS_SUCCESS,
        operation,
        " failed with cuBLAS status ",
        static_cast<int>(status)
    );
}

int checked_int(int64_t value, const char* name) {
    TORCH_CHECK(
        value >= 0
            && value <= static_cast<int64_t>(
                std::numeric_limits<int>::max()
            ),
        name,
        " is outside the cuBLAS integer range"
    );
    return static_cast<int>(value);
}

void validate_inputs(
    const torch::Tensor& c,
    const torch::Tensor& a
) {
    TORCH_CHECK(c.is_cuda(), "c must be a CUDA tensor");
    TORCH_CHECK(a.is_cuda(), "a must be a CUDA tensor");
    TORCH_CHECK(
        c.scalar_type() == at::kFloat,
        "c must have dtype torch.float32"
    );
    TORCH_CHECK(
        a.scalar_type() == at::kFloat,
        "a must have dtype torch.float32"
    );
    TORCH_CHECK(c.dim() == 2, "c must be two-dimensional");
    TORCH_CHECK(a.dim() == 2, "a must be two-dimensional");
    TORCH_CHECK(
        c.device() == a.device(),
        "c and a must be on the same device"
    );
    TORCH_CHECK(
        c.size(0) == c.size(1),
        "c must be square"
    );
    TORCH_CHECK(
        c.size(0) == a.size(0),
        "c and a must have the same row count"
    );
    TORCH_CHECK(
        c.stride(1) == 1,
        "c must have unit column stride"
    );
    TORCH_CHECK(
        a.stride(1) == 1,
        "a must have unit column stride"
    );
    TORCH_CHECK(
        c.stride(0) >= c.size(1),
        "c rows must not overlap"
    );
    TORCH_CHECK(
        a.stride(0) >= a.size(1),
        "a rows must not overlap"
    );
}

class HandleModeGuard {
public:
    explicit HandleModeGuard(cublasHandle_t handle)
        : handle_(handle) {
        check_status(
            cublasGetMathMode(handle_, &previous_),
            "cublasGetMathMode"
        );
        check_status(
            cublasSetMathMode(
                handle_,
                CUBLAS_TF32_TENSOR_OP_MATH
            ),
            "cublasSetMathMode"
        );
        check_status(
            cublasGetPointerMode(
                handle_,
                &previous_pointer_
            ),
            "cublasGetPointerMode"
        );
        check_status(
            cublasSetPointerMode(
                handle_,
                CUBLAS_POINTER_MODE_HOST
            ),
            "cublasSetPointerMode"
        );
    }

    HandleModeGuard(const HandleModeGuard&) = delete;
    HandleModeGuard& operator=(const HandleModeGuard&) = delete;

    ~HandleModeGuard() {
        cublasSetPointerMode(handle_, previous_pointer_);
        cublasSetMathMode(handle_, previous_);
    }

private:
    cublasHandle_t handle_;
    cublasMath_t previous_;
    cublasPointerMode_t previous_pointer_;
};

}  // namespace

torch::Tensor syrk_lower_in_place(
    torch::Tensor c,
    torch::Tensor a
) {
    validate_inputs(c, a);

    const int size = checked_int(c.size(0), "size");
    const int rank = checked_int(a.size(1), "rank");
    const int lda = checked_int(a.stride(0), "a row stride");
    const int ldc = checked_int(c.stride(0), "c row stride");
    if (size == 0 || rank == 0) {
        return c;
    }

    cublasHandle_t handle =
        at::cuda::getCurrentCUDABlasHandle();
    HandleModeGuard mode_guard(handle);
    const float alpha = -1.0f;
    const float beta = 1.0f;

    check_status(
        cublasSsyrk(
            handle,
            CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T,
            size,
            rank,
            &alpha,
            a.data_ptr<float>(),
            lda,
            &beta,
            c.data_ptr<float>(),
            ldc
        ),
        "cublasSsyrk"
    );
    return c;
}

torch::Tensor gemm_lower_blocks_in_place(
    torch::Tensor c,
    torch::Tensor a,
    int64_t row_block
) {
    validate_inputs(c, a);
    TORCH_CHECK(row_block > 0, "row_block must be positive");

    const int size = checked_int(c.size(0), "size");
    const int rank = checked_int(a.size(1), "rank");
    const int lda = checked_int(a.stride(0), "a row stride");
    const int ldc = checked_int(c.stride(0), "c row stride");
    const int block = checked_int(row_block, "row_block");
    if (size == 0 || rank == 0) {
        return c;
    }

    cublasHandle_t handle =
        at::cuda::getCurrentCUDABlasHandle();
    HandleModeGuard mode_guard(handle);
    const float alpha = -1.0f;
    const float beta = 1.0f;

    for (int row_begin = 0; row_begin < size; row_begin += block) {
        const int rows =
            std::min(block, size - row_begin);
        const int columns = row_begin + rows;
        const float* left =
            a.data_ptr<float>()
            + static_cast<int64_t>(row_begin) * lda;
        float* destination =
            c.data_ptr<float>()
            + static_cast<int64_t>(row_begin) * ldc;

        check_status(
            cublasGemmEx(
                handle,
                CUBLAS_OP_T,
                CUBLAS_OP_N,
                columns,
                rows,
                rank,
                &alpha,
                a.data_ptr<float>(),
                CUDA_R_32F,
                lda,
                left,
                CUDA_R_32F,
                lda,
                &beta,
                destination,
                CUDA_R_32F,
                ldc,
                CUBLAS_COMPUTE_32F_FAST_TF32,
                CUBLAS_GEMM_DEFAULT
            ),
            "cublasGemmEx"
        );
    }
    return c;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
    module.def(
        "syrk_lower_in_place",
        &syrk_lower_in_place
    );
    module.def(
        "gemm_lower_blocks_in_place",
        &gemm_lower_blocks_in_place
    );
}
"""


@lru_cache(maxsize=1)
def _cuda32_extension():
    return load_inline(
        name="b200_cholesky32_padded_v3",
        cpp_sources="",
        cuda_sources=_CUDA32_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _cuda64_extension():
    return load_inline(
        name="b200_cholesky64_inverse_v3",
        cpp_sources="",
        cuda_sources=_CUDA64_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _cuda128_extension():
    return load_inline(
        name="b200_cholesky128_panel32_v2",
        cpp_sources="",
        cuda_sources=_CUDA128_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _cuda128_inverse_extension():
    return load_inline(
        name="b200_cholesky128_inverse_panel32_v1",
        cpp_sources="",
        cuda_sources=_CUDA128_INVERSE_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _cuda256_packed_extension():
    return load_inline(
        name="b200_cholesky256_packed_panel32_v1",
        cpp_sources="",
        cuda_sources=_CUDA256_PACKED_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _persistent_wmma_extension():
    return load_inline(
        name="b200_persistent_wmma_cholesky_v1",
        cpp_sources="",
        cuda_sources=_PERSISTENT_WMMA_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _fast_update_extension():
    return load_inline(
        name="b200_bf16_update_v1",
        cpp_sources="",
        cuda_sources=_FAST_UPDATE_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3"],
        extra_ldflags=["-lcublas"],
        with_cuda=True,
        verbose=False,
    )


@lru_cache(maxsize=1)
def _lower_cublas_extension():
    return load_inline(
        name="b200_cublas_lower_v1",
        cpp_sources="",
        cuda_sources=_LOWER_CUBLAS_SOURCE,
        functions=None,
        extra_cuda_cflags=["-O3"],
        extra_ldflags=["-lcublas"],
        with_cuda=True,
        verbose=False,
    )


def _cuda_cholesky32(data: torch.Tensor) -> torch.Tensor:
    try:
        return _cuda32_extension().cholesky32(data)
    except Exception:
        return _triton_cholesky32(data)


def _cuda_cholesky64(data: torch.Tensor) -> torch.Tensor:
    try:
        return _cuda64_extension().cholesky64(data)
    except Exception:
        return torch.linalg.cholesky_ex(data, check_errors=False).L


def _cuda_cholesky64_inverse(
    data: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    try:
        factor, inverse = _cuda64_extension().cholesky64_inverse(data)
        return factor, inverse
    except Exception:
        factor = torch.linalg.cholesky_ex(
            data,
            check_errors=False,
        ).L
        identity = torch.eye(
            64,
            device=data.device,
            dtype=data.dtype,
        ).expand(data.shape[0], -1, -1)
        inverse = torch.linalg.solve_triangular(
            factor,
            identity,
            upper=False,
            left=True,
        )
        return factor, inverse


def _cuda_cholesky128(data: torch.Tensor) -> torch.Tensor:
    try:
        return _cuda128_extension().cholesky128(data)
    except Exception:
        return _blocked_cholesky128(data)


_cuda128_inverse_failed = False


def _cuda_cholesky128_inverse(
    data: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    global _cuda128_inverse_failed
    if not _cuda128_inverse_failed:
        try:
            factor, inverse = (
                _cuda128_inverse_extension()
                .cholesky128_inverse(data)
            )
            return factor, inverse
        except Exception:
            _cuda128_inverse_failed = True

    factor = torch.linalg.cholesky_ex(
        data,
        check_errors=False,
    ).L
    identity = torch.eye(
        128,
        dtype=data.dtype,
        device=data.device,
    ).expand(data.shape[0], 128, 128)
    inverse = torch.linalg.solve_triangular(
        factor,
        identity,
        upper=False,
        left=True,
    )
    return factor, inverse


def _cuda_cholesky256_packed(data: torch.Tensor) -> torch.Tensor:
    try:
        return _cuda256_packed_extension().cholesky256_packed(data)
    except Exception:
        return torch.linalg.cholesky_ex(
            data,
            check_errors=False,
        ).L


def _persistent_wmma_cholesky(data: torch.Tensor) -> torch.Tensor:
    return (
        _persistent_wmma_extension()
        .persistent_wmma_cholesky(data)
    )


def _syrk_lower_update(
    output: torch.Tensor,
    panel: torch.Tensor,
) -> torch.Tensor:
    return _lower_cublas_extension().syrk_lower_in_place(
        output,
        panel,
    )


def _gemm_lower_update(
    output: torch.Tensor,
    panel: torch.Tensor,
    row_block: int,
) -> torch.Tensor:
    return _lower_cublas_extension().gemm_lower_blocks_in_place(
        output,
        panel,
        row_block,
    )


def _bf16_update(
    output: torch.Tensor,
    left: torch.Tensor,
    right: torch.Tensor,
) -> torch.Tensor:
    return _fast_update_extension().bf16_update(
        left,
        right,
        output,
    )


@triton.jit
def _bf16_update_kernel(
    output_ptr,
    left_ptr,
    right_ptr,
    m_size: tl.constexpr,
    n_size: tl.constexpr,
    k_size: tl.constexpr,
    output_row_stride: tl.constexpr,
    left_row_stride: tl.constexpr,
    right_row_stride: tl.constexpr,
    output_batch_stride: tl.constexpr,
    left_batch_stride: tl.constexpr,
    right_batch_stride: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr,
):
    program = tl.program_id(0)
    batch = tl.program_id(1)
    output_ptr += batch * output_batch_stride
    left_ptr += batch * left_batch_stride
    right_ptr += batch * right_batch_stride
    programs_m = tl.cdiv(m_size, BLOCK_M)
    programs_n = tl.cdiv(n_size, BLOCK_N)
    programs_per_group = GROUP_M * programs_n
    group = program // programs_per_group
    first_m = group * GROUP_M
    group_m = tl.minimum(programs_m - first_m, GROUP_M)
    local = program % programs_per_group
    program_m = first_m + (local % group_m)
    program_n = local // group_m

    rows = program_m * BLOCK_M + tl.arange(0, BLOCK_M)
    cols = program_n * BLOCK_N + tl.arange(0, BLOCK_N)
    inner = tl.arange(0, BLOCK_K)
    accumulator = tl.zeros((BLOCK_M, BLOCK_N), tl.float32)

    for start in range(0, k_size, BLOCK_K):
        inner_offsets = start + inner
        left_values = tl.load(
            left_ptr
            + rows[:, None] * left_row_stride
            + inner_offsets[None, :],
            mask=(rows[:, None] < m_size)
            & (inner_offsets[None, :] < k_size),
            other=0.0,
        ).to(tl.bfloat16)
        right_values = tl.load(
            right_ptr
            + cols[:, None] * right_row_stride
            + inner_offsets[None, :],
            mask=(cols[:, None] < n_size)
            & (inner_offsets[None, :] < k_size),
            other=0.0,
        ).to(tl.bfloat16)
        accumulator += tl.dot(
            left_values,
            tl.trans(right_values),
            out_dtype=tl.float32,
        )

    output_offsets = (
        rows[:, None] * output_row_stride + cols[None, :]
    )
    output_mask = (rows[:, None] < m_size) & (cols[None, :] < n_size)
    previous = tl.load(
        output_ptr + output_offsets,
        mask=output_mask,
        other=0.0,
    )
    tl.store(
        output_ptr + output_offsets,
        previous - accumulator,
        mask=output_mask,
    )


def _triton_bf16_update(
    output: torch.Tensor,
    left: torch.Tensor,
    right: torch.Tensor,
) -> torch.Tensor:
    m_size = output.shape[-2]
    n_size = output.shape[-1]
    k_size = left.shape[-1]
    block = 128
    grid = (
        triton.cdiv(m_size, block) * triton.cdiv(n_size, block),
        output.shape[0] if output.dim() == 3 else 1,
    )
    output_batch_stride = output.stride(0) if output.dim() == 3 else 0
    left_batch_stride = left.stride(0) if left.dim() == 3 else 0
    right_batch_stride = right.stride(0) if right.dim() == 3 else 0
    _bf16_update_kernel[grid](
        output,
        left,
        right,
        m_size,
        n_size,
        k_size,
        output.stride(-2),
        left.stride(-2),
        right.stride(-2),
        output_batch_stride,
        left_batch_stride,
        right_batch_stride,
        BLOCK_M=block,
        BLOCK_N=block,
        BLOCK_K=32,
        GROUP_M=8,
        num_warps=8,
        num_stages=4,
    )
    return output


@triton.jit
def _cholesky32_kernel(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    matrix = 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 * matrix_stride + rows * 32 + cols

    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))

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

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

    tl.store(output_ptr + offsets, values)


def _triton_cholesky32(data: torch.Tensor) -> torch.Tensor:
    output = torch.empty_like(data)
    _cholesky32_kernel[(data.shape[0],)](
        data,
        output,
        32 * 32,
        num_warps=1,
    )
    return output


def _individual_cholesky(data: torch.Tensor) -> torch.Tensor:
    """Avoid cuSOLVER's slow batched path for a few large matrices."""
    output = torch.empty_like(data)
    info = torch.empty((data.shape[0],), device=data.device, dtype=torch.int32)
    for matrix in range(data.shape[0]):
        torch.linalg.cholesky_ex(
            data[matrix],
            check_errors=False,
            out=(output[matrix], info[matrix]),
        )
    return output


def _blocked_cholesky128(data: torch.Tensor) -> torch.Tensor:
    """Two 64-wide panels, using the register kernel on both diagonals."""
    output = torch.empty_like(data)
    output[:, :64, 64:].zero_()

    diagonal0 = _cuda_cholesky64(data[:, :64, :64].contiguous())
    output[:, :64, :64].copy_(diagonal0)

    right_hand_side = data[:, 64:, :64].transpose(-1, -2)
    solved = torch.linalg.solve_triangular(
        diagonal0,
        right_hand_side,
        upper=False,
        left=True,
    )
    panel = solved.transpose(-1, -2)
    output[:, 64:, :64].copy_(panel)

    trailing = data[:, 64:, 64:].clone()
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        torch.baddbmm(
            trailing,
            panel,
            panel.transpose(-1, -2),
            beta=1.0,
            alpha=-1.0,
            out=trailing,
        )
    finally:
        torch.set_float32_matmul_precision(previous_precision)

    diagonal1 = _cuda_cholesky64(trailing)
    output[:, 64:, 64:].copy_(diagonal1)
    return output


def _blocked_custom64(data: torch.Tensor) -> torch.Tensor:
    """Blocked factorization with register-resident 64x64 diagonals."""
    work = data.clone()
    n = data.shape[-1]
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, 64):
            end = start + 64
            diagonal = _cuda_cholesky64(
                work[:, start:end, start:end].contiguous()
            )
            work[:, start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            work[:, start:end, end:].zero_()
            right_hand_side = work[:, end:, start:end].transpose(-1, -2)
            solved = torch.linalg.solve_triangular(
                diagonal,
                right_hand_side,
                upper=False,
                left=True,
            )
            panel = solved.transpose(-1, -2)
            work[:, end:, start:end].copy_(panel)

            trailing = work[:, end:, end:]
            torch.baddbmm(
                trailing,
                panel,
                panel.transpose(-1, -2),
                beta=1.0,
                alpha=-1.0,
                out=trailing,
            )
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return work


def _blocked_batched_cholesky(
    data: torch.Tensor,
    block_size: int,
) -> torch.Tensor:
    """Use exact panels and tensor-core trailing updates for medium matrices."""
    work = data.clone()
    n = data.shape[-1]
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, block_size):
            end = min(start + block_size, n)
            diagonal = torch.linalg.cholesky_ex(
                work[:, start:end, start:end],
                check_errors=False,
            ).L
            work[:, start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            right_hand_side = work[:, end:, start:end].transpose(-1, -2)
            solved = torch.linalg.solve_triangular(
                diagonal,
                right_hand_side,
                upper=False,
                left=True,
            )
            panel = solved.transpose(-1, -2)
            work[:, end:, start:end].copy_(panel)

            trailing = work[:, end:, end:]
            torch.baddbmm(
                trailing,
                panel,
                panel.transpose(-1, -2),
                beta=1.0,
                alpha=-1.0,
                out=trailing,
            )
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return torch.tril(work)


def _left_blocked_batched_cholesky(
    data: torch.Tensor,
    block_size: int,
    custom64: bool = False,
    fast_updates: bool = False,
    inverse_panels: bool = False,
) -> torch.Tensor:
    """Batched left-looking factorization with lower-panel updates only."""
    factor = torch.zeros_like(data)
    n = data.shape[-1]
    identity = (
        torch.eye(
            block_size,
            device=data.device,
            dtype=data.dtype,
        ).expand(data.shape[0], -1, -1)
        if inverse_panels
        else None
    )
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, block_size):
            end = min(start + block_size, n)
            diagonal_input = data[:, start:end, start:end].clone()
            if start:
                if fast_updates:
                    _bf16_update(
                        diagonal_input,
                        factor[:, start:end, :start],
                        factor[:, start:end, :start],
                    )
                else:
                    torch.baddbmm(
                        diagonal_input,
                        factor[:, start:end, :start],
                        factor[:, start:end, :start].transpose(-1, -2),
                        beta=1.0,
                        alpha=-1.0,
                        out=diagonal_input,
                    )
            if custom64:
                diagonal = _cuda_cholesky64(diagonal_input)
            else:
                diagonal = torch.linalg.cholesky_ex(
                    diagonal_input,
                    check_errors=False,
                ).L
            factor[:, start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            panel_input = data[:, end:, start:end].clone()
            if start:
                if fast_updates:
                    _bf16_update(
                        panel_input,
                        factor[:, end:, :start],
                        factor[:, start:end, :start],
                    )
                else:
                    torch.baddbmm(
                        panel_input,
                        factor[:, end:, :start],
                        factor[:, start:end, :start].transpose(-1, -2),
                        beta=1.0,
                        alpha=-1.0,
                        out=panel_input,
                    )
            if inverse_panels:
                diagonal_inverse = torch.linalg.solve_triangular(
                    diagonal,
                    identity,
                    upper=False,
                    left=True,
                )
                panel = torch.bmm(
                    panel_input,
                    diagonal_inverse.transpose(-1, -2),
                )
            else:
                solved = torch.linalg.solve_triangular(
                    diagonal,
                    panel_input.transpose(-1, -2),
                    upper=False,
                    left=True,
                )
                panel = solved.transpose(-1, -2)
            factor[:, end:, start:end].copy_(panel)
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return factor


def _left_blocked_batched128(data: torch.Tensor) -> torch.Tensor:
    """Two or four 128-wide panels with a fused factor-and-inverse base."""
    factor = torch.zeros_like(data)
    n = data.shape[-1]
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, 128):
            end = start + 128
            diagonal_input = data[:, start:end, start:end].clone()
            if start:
                previous_row = factor[:, start:end, :start]
                torch.baddbmm(
                    diagonal_input,
                    previous_row,
                    previous_row.transpose(-1, -2),
                    beta=1.0,
                    alpha=-1.0,
                    out=diagonal_input,
                )

            diagonal, diagonal_inverse = _cuda_cholesky128_inverse(
                diagonal_input
            )
            factor[:, start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            panel_input = data[:, end:, start:end].clone()
            if start:
                torch.baddbmm(
                    panel_input,
                    factor[:, end:, :start],
                    factor[:, start:end, :start].transpose(-1, -2),
                    beta=1.0,
                    alpha=-1.0,
                    out=panel_input,
                )
            torch.bmm(
                panel_input,
                diagonal_inverse.transpose(-1, -2),
                out=factor[:, end:, start:end],
            )
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return factor


def _left_blocked_batched64_inverse(data: torch.Tensor) -> torch.Tensor:
    """Left-looking 64-wide panels with fused factor and inverse."""
    factor = torch.zeros_like(data)
    n = data.shape[-1]
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, 64):
            end = start + 64
            diagonal_input = data[:, start:end, start:end].clone()
            if start:
                previous_row = factor[:, start:end, :start]
                torch.baddbmm(
                    diagonal_input,
                    previous_row,
                    previous_row.transpose(-1, -2),
                    beta=1.0,
                    alpha=-1.0,
                    out=diagonal_input,
                )

            diagonal, diagonal_inverse = _cuda_cholesky64_inverse(
                diagonal_input
            )
            factor[:, start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            panel_input = data[:, end:, start:end].clone()
            if start:
                torch.baddbmm(
                    panel_input,
                    factor[:, end:, :start],
                    factor[:, start:end, :start].transpose(-1, -2),
                    beta=1.0,
                    alpha=-1.0,
                    out=panel_input,
                )
            torch.bmm(
                panel_input,
                diagonal_inverse.transpose(-1, -2),
                out=factor[:, end:, start:end],
            )
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return factor


def _column_blocked_batched_cholesky(
    data: torch.Tensor,
    block_size: int,
    fast_updates: bool = False,
    custom64_inverse: bool = False,
) -> torch.Tensor:
    """Left-looking factorization updating each complete block column once."""
    factor = torch.zeros_like(data)
    batch = data.shape[0]
    n = data.shape[-1]
    identity = (
        torch.eye(
            block_size,
            device=data.device,
            dtype=data.dtype,
        ).expand(batch, -1, -1)
        if not custom64_inverse
        else None
    )
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, block_size):
            end = start + block_size
            column = data[:, start:, start:end].clone()
            if start:
                if fast_updates:
                    _bf16_update(
                        column,
                        factor[:, start:, :start],
                        factor[:, start:end, :start],
                    )
                else:
                    torch.baddbmm(
                        column,
                        factor[:, start:, :start],
                        factor[:, start:end, :start].transpose(-1, -2),
                        beta=1.0,
                        alpha=-1.0,
                        out=column,
                    )

            diagonal_input = column[:, :block_size, :]
            if custom64_inverse:
                diagonal, diagonal_inverse = _cuda_cholesky64_inverse(
                    diagonal_input.contiguous()
                )
            else:
                diagonal = torch.linalg.cholesky_ex(
                    diagonal_input,
                    check_errors=False,
                ).L
                diagonal_inverse = torch.linalg.solve_triangular(
                    diagonal,
                    identity,
                    upper=False,
                    left=True,
                )
            factor[:, start:end, start:end].copy_(diagonal)

            if end == n:
                continue
            torch.bmm(
                column[:, block_size:, :],
                diagonal_inverse.transpose(-1, -2),
                out=factor[:, end:, start:end],
            )
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return factor


def _blocked_cholesky(
    data: torch.Tensor,
    block_size: int,
    inverse_panels: bool = False,
) -> torch.Tensor:
    # All routed shapes have batch one. Squeezing the batch dimension makes
    # PyTorch select the ordinary cuSOLVER/cuBLAS paths instead of their
    # strided-batched variants.
    work = data[0].clone()
    n = data.shape[-1]
    identity = (
        torch.eye(
            block_size,
            device=data.device,
            dtype=data.dtype,
        )
        if inverse_panels
        else None
    )
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, block_size):
            end = min(start + block_size, n)
            diagonal = torch.linalg.cholesky_ex(
                work[start:end, start:end],
                check_errors=False,
            ).L
            work[start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            if inverse_panels:
                diagonal_inverse = torch.linalg.solve_triangular(
                    diagonal,
                    identity,
                    upper=False,
                    left=True,
                )
                panel = torch.mm(
                    work[end:, start:end],
                    diagonal_inverse.transpose(-1, -2),
                )
            else:
                right_hand_side = work[end:, start:end].transpose(-1, -2)
                solved = torch.linalg.solve_triangular(
                    diagonal,
                    right_hand_side,
                    upper=False,
                    left=True,
                )
                panel = solved.transpose(-1, -2)
            work[end:, start:end].copy_(panel)

            trailing = work[end:, end:]
            torch.addmm(
                trailing,
                panel,
                panel.transpose(-1, -2),
                beta=1.0,
                alpha=-1.0,
                out=trailing,
            )
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return torch.tril(work).unsqueeze(0)


def _left_blocked_cholesky(
    data: torch.Tensor,
    block_size: int,
    inverse_panels: bool = False,
    fast_updates: bool = False,
) -> torch.Tensor:
    """Left-looking factorization that never updates the unused upper half."""
    source = data[0]
    factor = torch.zeros_like(source)
    n = source.shape[-1]
    identity = (
        torch.eye(
            block_size,
            device=data.device,
            dtype=data.dtype,
        )
        if inverse_panels
        else None
    )
    previous_precision = torch.get_float32_matmul_precision()
    torch.set_float32_matmul_precision("high")
    try:
        for start in range(0, n, block_size):
            end = min(start + block_size, n)
            diagonal_input = source[start:end, start:end].clone()
            if start:
                if fast_updates:
                    _bf16_update(
                        diagonal_input,
                        factor[start:end, :start],
                        factor[start:end, :start],
                    )
                else:
                    torch.addmm(
                        diagonal_input,
                        factor[start:end, :start],
                        factor[start:end, :start].transpose(-1, -2),
                        beta=1.0,
                        alpha=-1.0,
                        out=diagonal_input,
                    )
            diagonal = torch.linalg.cholesky_ex(
                diagonal_input,
                check_errors=False,
            ).L
            factor[start:end, start:end].copy_(diagonal)

            if end == n:
                continue

            panel_input = source[end:, start:end].clone()
            if start:
                if fast_updates:
                    _bf16_update(
                        panel_input,
                        factor[end:, :start],
                        factor[start:end, :start],
                    )
                else:
                    torch.addmm(
                        panel_input,
                        factor[end:, :start],
                        factor[start:end, :start].transpose(-1, -2),
                        beta=1.0,
                        alpha=-1.0,
                        out=panel_input,
                    )

            if inverse_panels:
                diagonal_inverse = torch.linalg.solve_triangular(
                    diagonal,
                    identity,
                    upper=False,
                    left=True,
                )
                panel = torch.mm(
                    panel_input,
                    diagonal_inverse.transpose(-1, -2),
                )
            else:
                solved = torch.linalg.solve_triangular(
                    diagonal,
                    panel_input.transpose(-1, -2),
                    upper=False,
                    left=True,
                )
                panel = solved.transpose(-1, -2)
            factor[end:, start:end].copy_(panel)
    finally:
        torch.set_float32_matmul_precision(previous_precision)
    return factor.unsqueeze(0)


def _dispatch_eager(data: torch.Tensor) -> torch.Tensor:
    shape = tuple(data.shape)
    if shape == (4, 512, 512):
        return _persistent_wmma_cholesky(data)
    if shape == (4096, 32, 32):
        return _cuda_cholesky32(data)
    if shape == (1024, 64, 64):
        return _cuda_cholesky64(data)
    if shape == (256, 128, 128):
        return _cuda_cholesky128(data)
    if shape == (64, 256, 256):
        return _cuda_cholesky256_packed(data)
    if shape == (16, 512, 512):
        return _blocked_custom64(data)
    if shape == (640, 512, 512):
        return _persistent_wmma_cholesky(data)
    if shape == (60, 1024, 1024):
        return _left_blocked_batched_cholesky(
            data,
            256,
            fast_updates=True,
            inverse_panels=True,
        )
    if shape in {
        (2, 2048, 2048),
        (2, 4096, 4096),
    }:
        return _individual_cholesky(data)
    if shape == (1, 8192, 8192):
        return _left_blocked_cholesky(
            data,
            2048,
            fast_updates=True,
        )
    if shape == (1, 16384, 16384):
        return _left_blocked_cholesky(
            data,
            4096,
            inverse_panels=True,
            fast_updates=True,
        )
    if shape == (1, 32768, 32768):
        return _left_blocked_cholesky(
            data,
            4096,
            inverse_panels=True,
            fast_updates=True,
        )
    return torch.linalg.cholesky_ex(data, check_errors=False).L


_graph_shape = None
_graph_position = 0
_graph_slots = []
_graph_disabled = set()


def _capture_graph_slot(data: torch.Tensor):
    fixed = data.clone()
    warm = _dispatch_eager(fixed)
    torch.cuda.synchronize()
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        answer = _dispatch_eager(fixed)
    warm = None
    graph.replay()
    return fixed, answer, graph


def _run_graphed(data: torch.Tensor, width: int) -> torch.Tensor:
    global _graph_shape, _graph_position, _graph_slots
    shape = tuple(data.shape)
    if shape in _graph_disabled:
        return _dispatch_eager(data)
    if _graph_shape != shape:
        _graph_shape = shape
        _graph_position = 0
        _graph_slots = []

    position = _graph_position
    if position == len(_graph_slots):
        try:
            slot = _capture_graph_slot(data)
        except Exception:
            _graph_disabled.add(shape)
            return _dispatch_eager(data)
        _graph_slots.append(slot)
        answer = slot[1]
    else:
        fixed, answer, graph = _graph_slots[position]
        fixed.copy_(data)
        graph.replay()
    _graph_position = (position + 1) % width
    return answer


def custom_kernel(data: input_t) -> output_t:
    shape = tuple(data.shape)
    graph_widths = {
        (60, 1024, 1024): 1,
    }
    width = graph_widths.get(shape)
    if width is not None:
        return _run_graphed(data, width)
    return _dispatch_eager(data)
scrolls · 2756 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