Skip to content
KernelIndex
Search⌘K

submission 888810

serverinspector · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-888810?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.28ms
#151 of 337
2026-07-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:23de69f5484ea0c3047dfc2b1bed0377fff6f3d35695b072fd691139718286bf
license declaredunknown
license concludedunknown
authorsserverinspector
imported2026-08-26

Techniques

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

shared-memory__shared__ float transfer[kTransferWidth][kTransferWidth + 1];
vector-width = float4apart in global memory is moved with consecutive-address float4

Kernel source

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

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

from task import input_t, output_t

_CPP_SOURCE = r"""
torch::Tensor cholesky512(torch::Tensor input);
torch::Tensor cholesky1024(torch::Tensor input);
torch::Tensor cholesky1024_d64(torch::Tensor input);
torch::Tensor cholesky2048(torch::Tensor input);
torch::Tensor cholesky32(torch::Tensor input);
torch::Tensor cholesky64(torch::Tensor input);
torch::Tensor cholesky128(torch::Tensor input);

void schur_update(
    torch::Tensor target,
    torch::Tensor row_factor,
    torch::Tensor col_factor,
    float alpha,
    bool use_tf32
);

void split_tf32_scaled(
    torch::Tensor input,
    torch::Tensor high,
    torch::Tensor expanded
);

void split_tf32_packed65(
    torch::Tensor input,
    torch::Tensor packed
);

void schur_update_tf32x2(
    torch::Tensor target,
    torch::Tensor row_high,
    torch::Tensor row_expanded,
    torch::Tensor col_high,
    torch::Tensor col_expanded
);

void triangular_solve_in_place(
    torch::Tensor factor,
    torch::Tensor right_hand_side
);

void split_tf32_packed_residual(
    torch::Tensor input,
    torch::Tensor packed
);

void triangular_solve_inverse_in_place(
    torch::Tensor factor,
    torch::Tensor right_hand_side
);

void split_packed65_fp16(
    torch::Tensor input,
    torch::Tensor packed
);

void schur_update_fp16(
    torch::Tensor target,
    torch::Tensor row_factor,
    torch::Tensor col_factor,
    float alpha
);
"""

_CUDA_SOURCE = r"""
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <limits>
#include <c10/cuda/CUDAException.h>
#include <cuda_fp16.h>

/*
  The left-looking eight-column panel structure below is adapted from MAGMA's
  batched POTF2 implementation, retrieved 2026-07-17:
  https://github.com/icl-utk-edu/magma/blob/master/magmablas/zpotf2_devicesfunc.cuh
  https://github.com/icl-utk-edu/magma/blob/master/magmablas/zpotf2_kernels.cu
  https://github.com/icl-utk-edu/magma/blob/master/src/zpotrf_batched.cpp
  https://github.com/icl-utk-edu/magma/blob/master/src/zpotrf_panel_batched.cpp

  Copyright (c) 2009-2023, The University of Tennessee
  All rights reserved.

  Redistribution and use in source and binary forms, with or without
  modification, are permitted provided that the following conditions are met:
  * Redistributions of source code must retain the above copyright notice,
    this list of conditions and the following disclaimer.
  * Redistributions in binary form must reproduce the above copyright notice,
    this list of conditions and the following disclaimer in the documentation
    and/or other materials provided with the distribution.
  * Neither the name of the University of Tennessee, Knoxville nor the names
    of its contributors may be used to endorse or promote products derived
    from this software without specific prior written permission.

  THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
  AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
  IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
  ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE
  LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
  CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
  SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
  INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN
  CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
  ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE
  POSSIBILITY OF SUCH DAMAGE.
*/

namespace {
constexpr int kCholeskyOrder = 512;
constexpr int kInnerPanelWidth = 8;
constexpr int kOuterPanelWidth = 64;
constexpr int kTransferWidth = 32;
constexpr int kPanelRowSplit = 256;

__global__ __launch_bounds__(512, 2) void transpose_input512_kernel(
    const float* __restrict__ input,
    float* __restrict__ output
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder * kCholeskyOrder;
    const float* source = input + matrix_offset;
    float* factor = output + matrix_offset;
    __shared__ float transfer[kTransferWidth][kTransferWidth + 1];

    const int tile_row = static_cast<int>(blockIdx.y);
    for (int tile_col = 0; tile_col <= tile_row; ++tile_col) {
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                transfer[local_row][local_col] = source[
                    (tile_row * kTransferWidth + local_row) * kCholeskyOrder +
                    tile_col * kTransferWidth + local_col
                ];
            }
            __syncthreads();
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                factor[
                    (tile_col * kTransferWidth + local_row) * kCholeskyOrder +
                    tile_row * kTransferWidth + local_col
                ] = transfer[local_col][local_row];
            }
            __syncthreads();
        }
}

__global__ __launch_bounds__(256, 4) void factor_panel128_kernel(
    float* __restrict__ output,
    int outer_start
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder * kCholeskyOrder;
    float* factor = output + matrix_offset;
    __shared__ float basis[kInnerPanelWidth][kInnerPanelWidth];
    __shared__ float panel_diagonal[kInnerPanelWidth];
    __shared__ float pivot_column[kInnerPanelWidth];

    for (int local_start = 0; local_start < kOuterPanelWidth;
         local_start += kInnerPanelWidth) {
        const int panel_start = outer_start + local_start;
        const int first_row = panel_start + lane;
        const int second_row = first_row + kPanelRowSplit;
        const bool first_active = first_row < kCholeskyOrder;
        const bool second_active = second_row < kCholeskyOrder;
        float first_products[kInnerPanelWidth];
        float second_products[kInnerPanelWidth];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth; ++column) {
            first_products[column] = 0.0f;
            second_products[column] = 0.0f;
        }

        for (int prior_start = 0; prior_start < local_start;
             prior_start += kInnerPanelWidth) {
            float first_values[kInnerPanelWidth];
            float second_values[kInnerPanelWidth];
            #pragma unroll
            for (int depth = 0; depth < kInnerPanelWidth; ++depth) {
                const int prior_column =
                    outer_start + prior_start + depth;
                first_values[depth] = first_active
                    ? factor[first_row + prior_column * kCholeskyOrder]
                    : 0.0f;
                second_values[depth] = second_active
                    ? factor[second_row + prior_column * kCholeskyOrder]
                    : 0.0f;
            }
            if (lane < kInnerPanelWidth) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth; ++depth) {
                    basis[depth][lane] = factor[
                        panel_start + lane +
                        (outer_start + prior_start + depth) *
                            kCholeskyOrder
                    ];
                }
            }
            __syncthreads();
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth; ++column) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth; ++depth) {
                    first_products[column] = fmaf(
                        first_values[depth],
                        basis[depth][column],
                        first_products[column]
                    );
                    second_products[column] = fmaf(
                        second_values[depth],
                        basis[depth][column],
                        second_products[column]
                    );
                }
            }
            __syncthreads();
        }

        float first_panel[kInnerPanelWidth];
        float second_panel[kInnerPanelWidth];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth; ++column) {
            first_panel[column] =
                first_active && lane >= column
                ? factor[first_row +
                         (panel_start + column) * kCholeskyOrder] -
                      first_products[column]
                : 0.0f;
            second_panel[column] = second_active
                ? factor[second_row +
                         (panel_start + column) * kCholeskyOrder] -
                      second_products[column]
                : 0.0f;
        }
        __syncthreads();

        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth; ++column) {
            if (lane == column) {
                panel_diagonal[column] = sqrtf(first_panel[column]);
            }
            __syncthreads();
            if (first_active && lane >= column) {
                first_panel[column] /= panel_diagonal[column];
            }
            if (second_active) {
                second_panel[column] /= panel_diagonal[column];
            }
            __syncthreads();
            if (lane < kInnerPanelWidth) {
                pivot_column[lane] = first_panel[column];
            }
            __syncthreads();
            if (first_active) {
                const float current = first_panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth;
                     ++later) {
                    if (lane >= later) {
                        first_panel[later] = fmaf(
                            -current,
                            pivot_column[later],
                            first_panel[later]
                        );
                    }
                }
            }
            if (second_active) {
                const float current = second_panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth;
                     ++later) {
                    second_panel[later] = fmaf(
                        -current,
                        pivot_column[later],
                        second_panel[later]
                    );
                }
            }
            __syncthreads();
        }

        if (first_active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth; ++column) {
                if (lane >= column) {
                    factor[first_row +
                           (panel_start + column) * kCholeskyOrder] =
                        first_panel[column];
                }
            }
        }
        if (second_active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth; ++column) {
                factor[second_row +
                       (panel_start + column) * kCholeskyOrder] =
                    second_panel[column];
            }
        }
        __syncthreads();
    }
}

__global__ __launch_bounds__(512, 2) void transpose_output512_kernel(
    float* __restrict__ output
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder * kCholeskyOrder;
    float* factor = output + matrix_offset;
    __shared__ float transfer[kTransferWidth][kTransferWidth + 1];

    const int tile_row = static_cast<int>(blockIdx.y);
        #pragma unroll
        for (int item = lane; item < 1024; item += 512) {
            const int local_row = item >> 5;
            const int local_col = item & 31;
            transfer[local_row][local_col] = factor[
                (tile_row * kTransferWidth + local_row) * kCholeskyOrder +
                tile_row * kTransferWidth + local_col
            ];
        }
        __syncthreads();
        #pragma unroll
        for (int item = lane; item < 1024; item += 512) {
            const int local_row = item >> 5;
            const int local_col = item & 31;
            factor[
                (tile_row * kTransferWidth + local_row) * kCholeskyOrder +
                tile_row * kTransferWidth + local_col
            ] = local_row >= local_col
                ? transfer[local_col][local_row]
                : 0.0f;
        }
        __syncthreads();

        for (int tile_col = 0; tile_col < tile_row; ++tile_col) {
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                transfer[local_row][local_col] = factor[
                    (tile_col * kTransferWidth + local_row) * kCholeskyOrder +
                    tile_row * kTransferWidth + local_col
                ];
            }
            __syncthreads();
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                factor[
                    (tile_row * kTransferWidth + local_row) * kCholeskyOrder +
                    tile_col * kTransferWidth + local_col
                ] = transfer[local_col][local_row];
                factor[
                    (tile_col * kTransferWidth + local_row) * kCholeskyOrder +
                    tile_row * kTransferWidth + local_col
                ] = 0.0f;
            }
            __syncthreads();
        }
}
}  // namespace

torch::Tensor cholesky512(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int64_t matrix_stride =
        static_cast<int64_t>(kCholeskyOrder) * kCholeskyOrder;
    float* factor = output.data_ptr<float>();
    const float alpha = -1.0f;
    const float beta = 1.0f;

    cudaError_t cuda_status = cudaDeviceSynchronize();
    static thread_local cublasHandle_t handle = nullptr;
    cublasStatus_t operation_status = CUBLAS_STATUS_SUCCESS;
    if (handle == nullptr) {
        operation_status = cublasCreate(&handle);
        if (operation_status == CUBLAS_STATUS_SUCCESS) {
            operation_status =
                cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
        }
        if (operation_status != CUBLAS_STATUS_SUCCESS &&
            handle != nullptr) {
            cublasDestroy(handle);
            handle = nullptr;
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_input512_kernel<<<dim3(batch, 16), 512>>>(
            input.data_ptr<float>(), factor
        );
        cuda_status = cudaGetLastError();
    }

    for (int outer_start = 0;
         outer_start < kCholeskyOrder &&
             operation_status == CUBLAS_STATUS_SUCCESS &&
             cuda_status == cudaSuccess;
         outer_start += kOuterPanelWidth) {
        if (outer_start > 0) {
            const int remaining = kCholeskyOrder - outer_start;
            operation_status = cublasGemmStridedBatchedEx(
                handle,
                CUBLAS_OP_N,
                CUBLAS_OP_T,
                remaining,
                kOuterPanelWidth,
                outer_start,
                &alpha,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder,
                matrix_stride,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder,
                matrix_stride,
                &beta,
                factor + outer_start +
                    outer_start * kCholeskyOrder,
                CUDA_R_32F,
                kCholeskyOrder,
                matrix_stride,
                batch,
                CUBLAS_COMPUTE_32F_PEDANTIC,
                CUBLAS_GEMM_DEFAULT_TENSOR_OP
            );
        }
        if (operation_status == CUBLAS_STATUS_SUCCESS &&
            cuda_status == cudaSuccess) {
            factor_panel128_kernel<<<batch, 256>>>(factor, outer_start);
            cuda_status = cudaGetLastError();
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_output512_kernel<<<dim3(batch, 16), 512>>>(factor);
        cuda_status = cudaGetLastError();
    }
    if (cuda_status == cudaSuccess) {
        cuda_status = cudaDeviceSynchronize();
    }

    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGemmStridedBatchedEx failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n512 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}
__global__ void split_tf32_scaled_kernel(
    const float* __restrict__ input,
    float* __restrict__ high,
    float* __restrict__ expanded,
    int64_t count,
    int64_t rows,
    int64_t depth,
    int64_t input_batch_stride,
    int64_t input_row_stride
) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index >= count) {
        return;
    }
    const int64_t column = index % depth;
    const int64_t matrix_row = (index / depth) % rows;
    const int64_t matrix = index / (rows * depth);
    const float value = input[
        matrix * input_batch_stride +
        matrix_row * input_row_stride +
        column
    ];
    unsigned int high_bits;
    asm volatile(
        "cvt.rna.tf32.f32 %0, %1;"
        : "=r"(high_bits)
        : "f"(value)
    );
    const float rounded = __uint_as_float(high_bits);
    const float amplified = isfinite(value)
        ? fmaf(64.0f, value - rounded, rounded)
        : rounded;
    unsigned int expanded_bits;
    asm volatile(
        "cvt.rna.tf32.f32 %0, %1;"
        : "=r"(expanded_bits)
        : "f"(amplified)
    );
    high[index] = rounded;
    expanded[index] = __uint_as_float(expanded_bits);
}

void split_tf32_scaled(
    torch::Tensor input,
    torch::Tensor high,
    torch::Tensor expanded
) {
    TORCH_CHECK(
        input.is_cuda() && high.is_cuda() && expanded.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        input.device() == high.device() &&
        input.device() == expanded.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        input.scalar_type() == at::kFloat &&
        high.scalar_type() == at::kFloat &&
        expanded.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        input.dim() == 3 && high.dim() == 3 && expanded.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        input.sizes() == high.sizes() &&
        input.sizes() == expanded.sizes(),
        "argument shapes must match"
    );
    TORCH_CHECK(
        input.stride(2) == 1 &&
        input.stride(0) > 0 &&
        input.stride(1) > 0,
        "input layout is unsupported"
    );
    TORCH_CHECK(
        high.is_contiguous() && expanded.is_contiguous(),
        "outputs must be contiguous"
    );

    const int64_t count = input.numel();
    TORCH_CHECK(count > 0, "input must be non-empty");
    constexpr int threads = 256;
    const int64_t blocks64 = (count + threads - 1) / threads;
    TORCH_CHECK(
        blocks64 <= std::numeric_limits<int>::max(),
        "input is too large"
    );
    c10::cuda::CUDAGuard device_guard(input.device());
    split_tf32_scaled_kernel<<<static_cast<int>(blocks64), threads>>>(
        input.data_ptr<float>(),
        high.data_ptr<float>(),
        expanded.data_ptr<float>(),
        count,
        input.size(1),
        input.size(2),
        input.stride(0),
        input.stride(1)
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "scaled TF32 split failed: ",
        cudaGetErrorString(cuda_status)
    );
}

__global__ void split_tf32_packed65_kernel(
    const float* __restrict__ input,
    float* __restrict__ packed,
    int64_t count,
    int64_t rows,
    int64_t depth,
    int64_t input_batch_stride,
    int64_t input_row_stride,
    int64_t packed_batch_stride,
    int64_t packed_row_stride
) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index >= count) {
        return;
    }
    const int64_t column = index % depth;
    const int64_t matrix_row = (index / depth) % rows;
    const int64_t matrix = index / (rows * depth);
    const float value = input[
        matrix * input_batch_stride +
        matrix_row * input_row_stride +
        column
    ];
    unsigned int high_bits;
    asm volatile(
        "cvt.rna.tf32.f32 %0, %1;"
        : "=r"(high_bits)
        : "f"(value)
    );
    const float high = __uint_as_float(high_bits);
    const float amplified = isfinite(value)
        ? fmaf(65.0f, value - high, high)
        : high;
    unsigned int expanded_bits;
    asm volatile(
        "cvt.rna.tf32.f32 %0, %1;"
        : "=r"(expanded_bits)
        : "f"(amplified)
    );
    const int64_t packed_offset =
        matrix * packed_batch_stride +
        matrix_row * packed_row_stride +
        column;
    packed[packed_offset] = high;
    packed[packed_offset + depth] =
        __uint_as_float(expanded_bits) * 0.125f;
}

void split_tf32_packed65(
    torch::Tensor input,
    torch::Tensor packed
) {
    TORCH_CHECK(
        input.is_cuda() && packed.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        input.device() == packed.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        input.scalar_type() == at::kFloat &&
        packed.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        input.dim() == 3 && packed.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        input.size(2) <= std::numeric_limits<int64_t>::max() / 2,
        "input depth is too large"
    );
    TORCH_CHECK(
        input.size(0) == packed.size(0) &&
        input.size(1) == packed.size(1) &&
        packed.size(2) == 2 * input.size(2),
        "packed shape must double the input depth"
    );
    TORCH_CHECK(
        input.stride(2) == 1 &&
        input.stride(0) > 0 &&
        input.stride(1) > 0,
        "input layout is unsupported"
    );
    TORCH_CHECK(
        packed.stride(2) == 1 &&
        packed.stride(0) > 0 &&
        packed.stride(1) > 0 &&
        packed.stride(1) >= packed.size(2),
        "packed output layout is unsupported"
    );

    const int64_t count = input.numel();
    TORCH_CHECK(count > 0, "input must be non-empty");
    constexpr int threads = 256;
    const int64_t blocks64 = (count + threads - 1) / threads;
    TORCH_CHECK(
        blocks64 <= std::numeric_limits<int>::max(),
        "input is too large"
    );
    c10::cuda::CUDAGuard device_guard(input.device());
    split_tf32_packed65_kernel<<<static_cast<int>(blocks64), threads>>>(
        input.data_ptr<float>(),
        packed.data_ptr<float>(),
        count,
        input.size(1),
        input.size(2),
        input.stride(0),
        input.stride(1),
        packed.stride(0),
        packed.stride(1)
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "packed scale65 TF32 split failed: ",
        cudaGetErrorString(cuda_status)
    );
}

__global__ void split_packed65_fp16_kernel(
    const float* __restrict__ input,
    __half* __restrict__ packed,
    int64_t count,
    int64_t rows,
    int64_t depth,
    int64_t input_batch_stride,
    int64_t input_row_stride,
    int64_t packed_batch_stride,
    int64_t packed_row_stride
) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index >= count) {
        return;
    }
    const int64_t column = index % depth;
    const int64_t matrix_row = (index / depth) % rows;
    const int64_t matrix = index / (rows * depth);
    const float value = input[
        matrix * input_batch_stride +
        matrix_row * input_row_stride +
        column
    ];
    const __half high_half = __float2half(value);
    const float high = __half2float(high_half);
    const float amplified = isfinite(value)
        ? fmaf(65.0f, value - high, high)
        : high;
    const __half expanded_half = __float2half(amplified);
    const float expanded_scaled = __half2float(expanded_half) * 0.125f;
    const int64_t packed_offset =
        matrix * packed_batch_stride +
        matrix_row * packed_row_stride +
        column;
    packed[packed_offset] = high_half;
    packed[packed_offset + depth] = __float2half(expanded_scaled);
}

void split_packed65_fp16(
    torch::Tensor input,
    torch::Tensor packed
) {
    TORCH_CHECK(
        input.is_cuda() && packed.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        input.device() == packed.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        input.scalar_type() == at::kFloat &&
        packed.scalar_type() == at::kHalf,
        "input must be float32 and packed must be float16"
    );
    TORCH_CHECK(
        input.dim() == 3 && packed.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        input.size(2) <= std::numeric_limits<int64_t>::max() / 2,
        "input depth is too large"
    );
    TORCH_CHECK(
        input.size(0) == packed.size(0) &&
        input.size(1) == packed.size(1) &&
        packed.size(2) == 2 * input.size(2),
        "packed shape must double the input depth"
    );
    TORCH_CHECK(
        input.stride(2) == 1 &&
        input.stride(0) > 0 &&
        input.stride(1) > 0,
        "input layout is unsupported"
    );
    TORCH_CHECK(
        packed.stride(2) == 1 &&
        packed.stride(0) > 0 &&
        packed.stride(1) > 0 &&
        packed.stride(1) >= packed.size(2),
        "packed output layout is unsupported"
    );

    const int64_t count = input.numel();
    TORCH_CHECK(count > 0, "input must be non-empty");
    constexpr int threads = 256;
    const int64_t blocks64 = (count + threads - 1) / threads;
    TORCH_CHECK(
        blocks64 <= std::numeric_limits<int>::max(),
        "input is too large"
    );
    c10::cuda::CUDAGuard device_guard(input.device());
    split_packed65_fp16_kernel<<<static_cast<int>(blocks64), threads>>>(
        input.data_ptr<float>(),
        reinterpret_cast<__half*>(packed.data_ptr<at::Half>()),
        count,
        input.size(1),
        input.size(2),
        input.stride(0),
        input.stride(1),
        packed.stride(0),
        packed.stride(1)
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "packed scale65 fp16 split failed: ",
        cudaGetErrorString(cuda_status)
    );
}

__global__ void split_tf32_packed_residual_kernel(
    const float* __restrict__ input,
    float* __restrict__ packed,
    int64_t count,
    int64_t rows,
    int64_t depth,
    int64_t input_batch_stride,
    int64_t input_row_stride,
    int64_t packed_batch_stride,
    int64_t packed_row_stride
) {
    const int64_t index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index >= count) {
        return;
    }
    const int64_t column = index % depth;
    const int64_t matrix_row = (index / depth) % rows;
    const int64_t matrix = index / (rows * depth);
    const float value = input[
        matrix * input_batch_stride +
        matrix_row * input_row_stride +
        column
    ];
    unsigned int high_bits;
    asm volatile(
        "cvt.rna.tf32.f32 %0, %1;"
        : "=r"(high_bits)
        : "f"(value)
    );
    const float high = __uint_as_float(high_bits);
    const float residual = __fsub_rn(value, high);
    unsigned int residual_bits;
    asm volatile(
        "cvt.rna.tf32.f32 %0, %1;"
        : "=r"(residual_bits)
        : "f"(residual)
    );
    const int64_t packed_offset =
        matrix * packed_batch_stride +
        matrix_row * packed_row_stride +
        column;
    packed[packed_offset] = high;
    packed[packed_offset + depth] =
        __uint_as_float(residual_bits);
}

void split_tf32_packed_residual(
    torch::Tensor input,
    torch::Tensor packed
) {
    TORCH_CHECK(
        input.is_cuda() && packed.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        input.device() == packed.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        input.scalar_type() == at::kFloat &&
        packed.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        input.dim() == 3 && packed.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        input.size(2) <= std::numeric_limits<int64_t>::max() / 2,
        "input depth is too large"
    );
    TORCH_CHECK(
        input.size(0) == packed.size(0) &&
        input.size(1) == packed.size(1) &&
        packed.size(2) == 2 * input.size(2),
        "packed shape must double the input depth"
    );
    TORCH_CHECK(
        input.stride(2) == 1 &&
        input.stride(0) > 0 &&
        input.stride(1) > 0,
        "input layout is unsupported"
    );
    TORCH_CHECK(
        packed.stride(2) == 1 &&
        packed.stride(0) > 0 &&
        packed.stride(1) > 0 &&
        packed.stride(1) >= packed.size(2),
        "packed output layout is unsupported"
    );

    const int64_t count = input.numel();
    TORCH_CHECK(count > 0, "input must be non-empty");
    constexpr int threads = 256;
    const int64_t blocks64 = (count + threads - 1) / threads;
    TORCH_CHECK(
        blocks64 <= std::numeric_limits<int>::max(),
        "input is too large"
    );
    c10::cuda::CUDAGuard device_guard(input.device());
    split_tf32_packed_residual_kernel<<<
        static_cast<int>(blocks64), threads
    >>>(
        input.data_ptr<float>(),
        packed.data_ptr<float>(),
        count,
        input.size(1),
        input.size(2),
        input.stride(0),
        input.stride(1),
        packed.stride(0),
        packed.stride(1)
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "packed TF32 residual split failed: ",
        cudaGetErrorString(cuda_status)
    );
}

void schur_update(
    torch::Tensor target,
    torch::Tensor row_factor,
    torch::Tensor col_factor,
    float alpha,
    bool use_tf32
) {
    TORCH_CHECK(
        target.is_cuda() && row_factor.is_cuda() && col_factor.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        target.device() == row_factor.device() &&
        target.device() == col_factor.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        target.scalar_type() == at::kFloat &&
        row_factor.scalar_type() == at::kFloat &&
        col_factor.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        target.dim() == 3 && row_factor.dim() == 3 && col_factor.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        target.size(0) > 0 &&
        target.size(0) == row_factor.size(0) &&
        target.size(0) == col_factor.size(0),
        "batch dimensions must match and be positive"
    );
    TORCH_CHECK(
        target.size(1) == row_factor.size(1),
        "target rows must match row factor"
    );
    TORCH_CHECK(
        target.size(2) == col_factor.size(1),
        "target columns must match column factor"
    );
    TORCH_CHECK(
        row_factor.size(2) == col_factor.size(2),
        "factor depths must match"
    );
    TORCH_CHECK(
        target.size(1) > 0 &&
        target.size(2) > 0 &&
        row_factor.size(2) > 0,
        "matrix dimensions must be positive"
    );
    TORCH_CHECK(
        target.stride(2) == 1 &&
        row_factor.stride(2) == 1 &&
        col_factor.stride(2) == 1,
        "last dimensions must be contiguous"
    );
    TORCH_CHECK(
        target.stride(1) >= target.size(2) &&
        row_factor.stride(1) >= row_factor.size(2) &&
        col_factor.stride(1) >= col_factor.size(2),
        "matrix rows must be non-overlapping"
    );
    TORCH_CHECK(
        target.stride(0) > 0 &&
        row_factor.stride(0) > 0 &&
        col_factor.stride(0) > 0,
        "batch strides must be positive"
    );

    const int64_t batch64 = target.size(0);
    const int64_t rows64 = target.size(1);
    const int64_t cols64 = target.size(2);
    const int64_t depth64 = row_factor.size(2);
    const int64_t target_leading64 = target.stride(1);
    const int64_t row_leading64 = row_factor.stride(1);
    const int64_t col_leading64 = col_factor.stride(1);
    TORCH_CHECK(
        batch64 <= std::numeric_limits<int>::max() &&
        rows64 <= std::numeric_limits<int>::max() &&
        cols64 <= std::numeric_limits<int>::max() &&
        depth64 <= std::numeric_limits<int>::max() &&
        target_leading64 <= std::numeric_limits<int>::max() &&
        row_leading64 <= std::numeric_limits<int>::max() &&
        col_leading64 <= std::numeric_limits<int>::max(),
        "matrix dimension or leading dimension exceeds cuBLAS integer range"
    );

    const int batch = static_cast<int>(batch64);
    const int rows = static_cast<int>(rows64);
    const int cols = static_cast<int>(cols64);
    const int depth = static_cast<int>(depth64);
    const int target_leading = static_cast<int>(target_leading64);
    const int row_leading = static_cast<int>(row_leading64);
    const int col_leading = static_cast<int>(col_leading64);
    const float beta = 1.0f;

    c10::cuda::CUDAGuard device_guard(target.device());
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    cublasPointerMode_t saved_pointer_mode;
    cublasStatus_t operation_status =
        cublasGetPointerMode(handle, &saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGetPointerMode failed: ",
        static_cast<int>(operation_status)
    );

    operation_status =
        cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
    if (operation_status == CUBLAS_STATUS_SUCCESS) {
        operation_status = cublasGemmStridedBatchedEx(
            handle,
            CUBLAS_OP_T,
            CUBLAS_OP_N,
            cols,
            rows,
            depth,
            &alpha,
            col_factor.data_ptr<float>(),
            CUDA_R_32F,
            col_leading,
            col_factor.stride(0),
            row_factor.data_ptr<float>(),
            CUDA_R_32F,
            row_leading,
            row_factor.stride(0),
            &beta,
            target.data_ptr<float>(),
            CUDA_R_32F,
            target_leading,
            target.stride(0),
            batch,
            use_tf32
                ? CUBLAS_COMPUTE_32F_FAST_TF32
                : CUBLAS_COMPUTE_32F_PEDANTIC,
            CUBLAS_GEMM_DEFAULT_TENSOR_OP
        );
    }
    const cublasStatus_t pointer_restore_status =
        cublasSetPointerMode(handle, saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGemmStridedBatchedEx failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        pointer_restore_status == CUBLAS_STATUS_SUCCESS,
        "cuBLAS pointer mode restoration failed: ",
        static_cast<int>(pointer_restore_status)
    );
}

void schur_update_fp16(
    torch::Tensor target,
    torch::Tensor row_factor,
    torch::Tensor col_factor,
    float alpha
) {
    TORCH_CHECK(
        target.is_cuda() && row_factor.is_cuda() && col_factor.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        target.device() == row_factor.device() &&
        target.device() == col_factor.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        target.scalar_type() == at::kFloat &&
        row_factor.scalar_type() == at::kHalf &&
        col_factor.scalar_type() == at::kHalf,
        "target must be float32 and factors must be float16"
    );
    TORCH_CHECK(
        target.dim() == 3 && row_factor.dim() == 3 && col_factor.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        target.size(0) > 0 &&
        target.size(0) == row_factor.size(0) &&
        target.size(0) == col_factor.size(0),
        "batch dimensions must match and be positive"
    );
    TORCH_CHECK(
        target.size(1) == row_factor.size(1),
        "target rows must match row factor"
    );
    TORCH_CHECK(
        target.size(2) == col_factor.size(1),
        "target columns must match column factor"
    );
    TORCH_CHECK(
        row_factor.size(2) == col_factor.size(2),
        "factor depths must match"
    );
    TORCH_CHECK(
        target.size(1) > 0 &&
        target.size(2) > 0 &&
        row_factor.size(2) > 0,
        "matrix dimensions must be positive"
    );
    TORCH_CHECK(
        target.stride(2) == 1 &&
        row_factor.stride(2) == 1 &&
        col_factor.stride(2) == 1,
        "last dimensions must be contiguous"
    );
    TORCH_CHECK(
        target.stride(1) >= target.size(2) &&
        row_factor.stride(1) >= row_factor.size(2) &&
        col_factor.stride(1) >= col_factor.size(2),
        "matrix rows must be non-overlapping"
    );
    TORCH_CHECK(
        target.stride(0) > 0 &&
        row_factor.stride(0) > 0 &&
        col_factor.stride(0) > 0,
        "batch strides must be positive"
    );

    const int64_t batch64 = target.size(0);
    const int64_t rows64 = target.size(1);
    const int64_t cols64 = target.size(2);
    const int64_t depth64 = row_factor.size(2);
    const int64_t target_leading64 = target.stride(1);
    const int64_t row_leading64 = row_factor.stride(1);
    const int64_t col_leading64 = col_factor.stride(1);
    TORCH_CHECK(
        batch64 <= std::numeric_limits<int>::max() &&
        rows64 <= std::numeric_limits<int>::max() &&
        cols64 <= std::numeric_limits<int>::max() &&
        depth64 <= std::numeric_limits<int>::max() &&
        target_leading64 <= std::numeric_limits<int>::max() &&
        row_leading64 <= std::numeric_limits<int>::max() &&
        col_leading64 <= std::numeric_limits<int>::max(),
        "matrix dimension or leading dimension exceeds cuBLAS integer range"
    );

    const int batch = static_cast<int>(batch64);
    const int rows = static_cast<int>(rows64);
    const int cols = static_cast<int>(cols64);
    const int depth = static_cast<int>(depth64);
    const int target_leading = static_cast<int>(target_leading64);
    const int row_leading = static_cast<int>(row_leading64);
    const int col_leading = static_cast<int>(col_leading64);
    const float beta = 1.0f;

    c10::cuda::CUDAGuard device_guard(target.device());
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    cublasPointerMode_t saved_pointer_mode;
    cublasStatus_t operation_status =
        cublasGetPointerMode(handle, &saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGetPointerMode failed: ",
        static_cast<int>(operation_status)
    );

    operation_status =
        cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
    if (operation_status == CUBLAS_STATUS_SUCCESS) {
        operation_status = cublasGemmStridedBatchedEx(
            handle,
            CUBLAS_OP_T,
            CUBLAS_OP_N,
            cols,
            rows,
            depth,
            &alpha,
            col_factor.data_ptr<at::Half>(),
            CUDA_R_16F,
            col_leading,
            col_factor.stride(0),
            row_factor.data_ptr<at::Half>(),
            CUDA_R_16F,
            row_leading,
            row_factor.stride(0),
            &beta,
            target.data_ptr<float>(),
            CUDA_R_32F,
            target_leading,
            target.stride(0),
            batch,
            CUBLAS_COMPUTE_32F,
            CUBLAS_GEMM_DEFAULT_TENSOR_OP
        );
    }
    const cublasStatus_t pointer_restore_status =
        cublasSetPointerMode(handle, saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "fp16 cublasGemmStridedBatchedEx failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        pointer_restore_status == CUBLAS_STATUS_SUCCESS,
        "cuBLAS pointer mode restoration failed: ",
        static_cast<int>(pointer_restore_status)
    );
}

void triangular_solve_in_place(
    torch::Tensor factor,
    torch::Tensor right_hand_side
) {
    TORCH_CHECK(
        factor.is_cuda() && right_hand_side.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        factor.device() == right_hand_side.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        factor.scalar_type() == at::kFloat &&
        right_hand_side.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        factor.dim() == 3 && right_hand_side.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        factor.size(0) == 1 &&
        right_hand_side.size(0) == 1,
        "in-place triangular solve requires batch one"
    );
    TORCH_CHECK(
        factor.size(1) > 0 &&
        factor.size(1) == factor.size(2) &&
        right_hand_side.size(1) > 0 &&
        right_hand_side.size(2) == factor.size(1),
        "factor and right-hand-side dimensions are incompatible"
    );
    TORCH_CHECK(
        factor.stride(2) == 1 &&
        right_hand_side.stride(2) == 1 &&
        factor.stride(1) >= factor.size(2) &&
        right_hand_side.stride(1) >= right_hand_side.size(2),
        "matrix rows must be non-overlapping and contiguous"
    );

    const int64_t order64 = factor.size(1);
    const int64_t columns64 = right_hand_side.size(1);
    const int64_t factor_leading64 = factor.stride(1);
    const int64_t right_hand_side_leading64 =
        right_hand_side.stride(1);
    TORCH_CHECK(
        order64 <= std::numeric_limits<int>::max() &&
        columns64 <= std::numeric_limits<int>::max() &&
        factor_leading64 <= std::numeric_limits<int>::max() &&
        right_hand_side_leading64 <=
            std::numeric_limits<int>::max(),
        "matrix dimension or leading dimension exceeds cuBLAS range"
    );

    const int order = static_cast<int>(order64);
    const int columns = static_cast<int>(columns64);
    const int factor_leading = static_cast<int>(factor_leading64);
    const int right_hand_side_leading =
        static_cast<int>(right_hand_side_leading64);
    const float alpha = 1.0f;

    c10::cuda::CUDAGuard device_guard(factor.device());
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    cublasPointerMode_t saved_pointer_mode;
    cublasStatus_t operation_status =
        cublasGetPointerMode(handle, &saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGetPointerMode failed: ",
        static_cast<int>(operation_status)
    );

    operation_status =
        cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
    if (operation_status == CUBLAS_STATUS_SUCCESS) {
        operation_status = cublasStrsm(
            handle,
            CUBLAS_SIDE_LEFT,
            CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T,
            CUBLAS_DIAG_NON_UNIT,
            order,
            columns,
            &alpha,
            factor.data_ptr<float>(),
            factor_leading,
            right_hand_side.data_ptr<float>(),
            right_hand_side_leading
        );
    }
    const cublasStatus_t pointer_restore_status =
        cublasSetPointerMode(handle, saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasStrsm failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        pointer_restore_status == CUBLAS_STATUS_SUCCESS,
        "cuBLAS pointer mode restoration failed: ",
        static_cast<int>(pointer_restore_status)
    );
}

void triangular_solve_inverse_in_place(
    torch::Tensor factor,
    torch::Tensor right_hand_side
) {
    TORCH_CHECK(
        factor.is_cuda() && right_hand_side.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        factor.device() == right_hand_side.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        factor.scalar_type() == at::kFloat &&
        right_hand_side.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        factor.dim() == 3 && right_hand_side.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        factor.size(0) == 1 &&
        right_hand_side.size(0) == 1,
        "in-place triangular solve requires batch one"
    );
    TORCH_CHECK(
        factor.size(1) > 0 &&
        factor.size(1) == factor.size(2) &&
        right_hand_side.size(1) > 0 &&
        right_hand_side.size(2) == factor.size(1),
        "factor and right-hand-side dimensions are incompatible"
    );
    TORCH_CHECK(
        factor.stride(2) == 1 &&
        right_hand_side.stride(2) == 1 &&
        factor.stride(1) >= factor.size(2) &&
        right_hand_side.stride(1) >= right_hand_side.size(2),
        "matrix rows must be non-overlapping and contiguous"
    );

    const int64_t order64 = factor.size(1);
    const int64_t columns64 = right_hand_side.size(1);
    const int64_t factor_leading64 = factor.stride(1);
    const int64_t right_hand_side_leading64 =
        right_hand_side.stride(1);
    TORCH_CHECK(
        order64 <= std::numeric_limits<int>::max() &&
        columns64 <= std::numeric_limits<int>::max() &&
        factor_leading64 <= std::numeric_limits<int>::max() &&
        right_hand_side_leading64 <=
            std::numeric_limits<int>::max(),
        "matrix dimension or leading dimension exceeds cuBLAS range"
    );

    const int order = static_cast<int>(order64);
    const int columns = static_cast<int>(columns64);
    const int factor_leading = static_cast<int>(factor_leading64);
    const int right_hand_side_leading =
        static_cast<int>(right_hand_side_leading64);
    const float alpha = 1.0f;

    c10::cuda::CUDAGuard device_guard(factor.device());
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    cublasPointerMode_t saved_pointer_mode;
    cublasStatus_t operation_status =
        cublasGetPointerMode(handle, &saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGetPointerMode failed: ",
        static_cast<int>(operation_status)
    );

    operation_status =
        cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
    if (operation_status == CUBLAS_STATUS_SUCCESS) {
        operation_status = cublasStrsm(
            handle,
            CUBLAS_SIDE_LEFT,
            CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_N,
            CUBLAS_DIAG_NON_UNIT,
            order,
            columns,
            &alpha,
            factor.data_ptr<float>(),
            factor_leading,
            right_hand_side.data_ptr<float>(),
            right_hand_side_leading
        );
    }
    const cublasStatus_t pointer_restore_status =
        cublasSetPointerMode(handle, saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "inverse-orientation cublasStrsm failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        pointer_restore_status == CUBLAS_STATUS_SUCCESS,
        "cuBLAS pointer mode restoration failed: ",
        static_cast<int>(pointer_restore_status)
    );
}

void schur_update_tf32x2(
    torch::Tensor target,
    torch::Tensor row_high,
    torch::Tensor row_expanded,
    torch::Tensor col_high,
    torch::Tensor col_expanded
) {
    TORCH_CHECK(
        target.is_cuda() &&
        row_high.is_cuda() &&
        row_expanded.is_cuda() &&
        col_high.is_cuda() &&
        col_expanded.is_cuda(),
        "arguments must be CUDA tensors"
    );
    TORCH_CHECK(
        target.device() == row_high.device() &&
        target.device() == row_expanded.device() &&
        target.device() == col_high.device() &&
        target.device() == col_expanded.device(),
        "arguments must be on the same CUDA device"
    );
    TORCH_CHECK(
        target.scalar_type() == at::kFloat &&
        row_high.scalar_type() == at::kFloat &&
        row_expanded.scalar_type() == at::kFloat &&
        col_high.scalar_type() == at::kFloat &&
        col_expanded.scalar_type() == at::kFloat,
        "arguments must be float32"
    );
    TORCH_CHECK(
        target.dim() == 3 &&
        row_high.dim() == 3 &&
        row_expanded.dim() == 3 &&
        col_high.dim() == 3 &&
        col_expanded.dim() == 3,
        "arguments must have three dimensions"
    );
    TORCH_CHECK(
        target.size(0) > 0 &&
        target.size(0) == row_high.size(0) &&
        target.size(0) == col_high.size(0),
        "batch dimensions must match and be positive"
    );
    TORCH_CHECK(
        row_expanded.sizes() == row_high.sizes() &&
        col_expanded.sizes() == col_high.sizes(),
        "high and expanded factor shapes must match"
    );
    TORCH_CHECK(
        target.size(1) == row_high.size(1) &&
        target.size(2) == col_high.size(1) &&
        row_high.size(2) == col_high.size(2),
        "target and factor dimensions must match"
    );
    TORCH_CHECK(
        target.size(1) > 0 &&
        target.size(2) > 0 &&
        row_high.size(2) > 0,
        "matrix dimensions must be positive"
    );
    TORCH_CHECK(
        target.stride(2) == 1 &&
        row_high.stride(2) == 1 &&
        row_expanded.stride(2) == 1 &&
        col_high.stride(2) == 1 &&
        col_expanded.stride(2) == 1,
        "last dimensions must be contiguous"
    );
    TORCH_CHECK(
        target.stride(1) >= target.size(2) &&
        row_high.stride(1) >= row_high.size(2) &&
        row_expanded.stride(1) >= row_expanded.size(2) &&
        col_high.stride(1) >= col_high.size(2) &&
        col_expanded.stride(1) >= col_expanded.size(2),
        "matrix rows must be non-overlapping"
    );
    TORCH_CHECK(
        target.stride(0) > 0 &&
        row_high.stride(0) > 0 &&
        row_expanded.stride(0) > 0 &&
        col_high.stride(0) > 0 &&
        col_expanded.stride(0) > 0,
        "batch strides must be positive"
    );

    const int64_t batch64 = target.size(0);
    const int64_t rows64 = target.size(1);
    const int64_t cols64 = target.size(2);
    const int64_t depth64 = row_high.size(2);
    const int64_t target_leading64 = target.stride(1);
    const int64_t row_high_leading64 = row_high.stride(1);
    const int64_t row_expanded_leading64 = row_expanded.stride(1);
    const int64_t col_high_leading64 = col_high.stride(1);
    const int64_t col_expanded_leading64 = col_expanded.stride(1);
    TORCH_CHECK(
        batch64 <= std::numeric_limits<int>::max() &&
        rows64 <= std::numeric_limits<int>::max() &&
        cols64 <= std::numeric_limits<int>::max() &&
        depth64 <= std::numeric_limits<int>::max() &&
        target_leading64 <= std::numeric_limits<int>::max() &&
        row_high_leading64 <= std::numeric_limits<int>::max() &&
        row_expanded_leading64 <= std::numeric_limits<int>::max() &&
        col_high_leading64 <= std::numeric_limits<int>::max() &&
        col_expanded_leading64 <= std::numeric_limits<int>::max(),
        "matrix dimension or leading dimension exceeds cuBLAS integer range"
    );

    const int batch = static_cast<int>(batch64);
    const int rows = static_cast<int>(rows64);
    const int cols = static_cast<int>(cols64);
    const int depth = static_cast<int>(depth64);
    const int target_leading = static_cast<int>(target_leading64);
    const float alpha_high = -0.984375f;
    const float alpha_expanded = -0.015625f;
    const float beta = 1.0f;

    c10::cuda::CUDAGuard device_guard(target.device());
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    cublasPointerMode_t saved_pointer_mode;
    cublasStatus_t operation_status =
        cublasGetPointerMode(handle, &saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGetPointerMode failed: ",
        static_cast<int>(operation_status)
    );

    operation_status =
        cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
    auto run_product = [&](
        const torch::Tensor& row_factor,
        const torch::Tensor& col_factor,
        const float* product_alpha
    ) {
        return cublasGemmStridedBatchedEx(
            handle,
            CUBLAS_OP_T,
            CUBLAS_OP_N,
            cols,
            rows,
            depth,
            product_alpha,
            col_factor.data_ptr<float>(),
            CUDA_R_32F,
            static_cast<int>(col_factor.stride(1)),
            col_factor.stride(0),
            row_factor.data_ptr<float>(),
            CUDA_R_32F,
            static_cast<int>(row_factor.stride(1)),
            row_factor.stride(0),
            &beta,
            target.data_ptr<float>(),
            CUDA_R_32F,
            target_leading,
            target.stride(0),
            batch,
            CUBLAS_COMPUTE_32F_FAST_TF32,
            CUBLAS_GEMM_DEFAULT_TENSOR_OP
        );
    };
    if (operation_status == CUBLAS_STATUS_SUCCESS) {
        operation_status =
            run_product(row_high, col_high, &alpha_high);
    }
    if (operation_status == CUBLAS_STATUS_SUCCESS) {
        operation_status = run_product(
            row_expanded,
            col_expanded,
            &alpha_expanded
        );
    }
    const cublasStatus_t pointer_restore_status =
        cublasSetPointerMode(handle, saved_pointer_mode);
    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "scaled TF32x2 cublasGemmStridedBatchedEx failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        pointer_restore_status == CUBLAS_STATUS_SUCCESS,
        "cuBLAS pointer mode restoration failed: ",
        static_cast<int>(pointer_restore_status)
    );
}
"""

_CUDA_1024_SOURCE = r"""
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <limits>
#include <c10/cuda/CUDAException.h>

// Uses the MAGMA-derived panel structure and retained license above.
namespace {
constexpr int kCholeskyOrder1024 = 1024;
constexpr int kInnerPanelWidth1024 = 8;
constexpr int kOuterPanelWidth1024 = 64;
constexpr int kTransferWidth1024 = 32;
constexpr int kPanelRowSplit1024 = 512;

__global__ __launch_bounds__(512, 2) void transpose_input1024_kernel(
    const float* __restrict__ input,
    float* __restrict__ output
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder1024 * kCholeskyOrder1024;
    const float* source = input + matrix_offset;
    float* factor = output + matrix_offset;
    __shared__ float transfer[kTransferWidth1024][kTransferWidth1024 + 1];

    const int tile_row = static_cast<int>(blockIdx.y);
    for (int tile_col = 0; tile_col <= tile_row; ++tile_col) {
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                transfer[local_row][local_col] = source[
                    (tile_row * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                    tile_col * kTransferWidth1024 + local_col
                ];
            }
            __syncthreads();
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                factor[
                    (tile_col * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                    tile_row * kTransferWidth1024 + local_col
                ] = transfer[local_col][local_row];
            }
            __syncthreads();
        }
}

__global__ __launch_bounds__(512, 2) void factor_panel1024_kernel(
    float* __restrict__ output,
    int outer_start
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder1024 * kCholeskyOrder1024;
    float* factor = output + matrix_offset;
    __shared__ float basis[kInnerPanelWidth1024][kInnerPanelWidth1024];
    __shared__ float panel_diagonal[kInnerPanelWidth1024];
    __shared__ float pivot_column[kInnerPanelWidth1024];

    for (int local_start = 0; local_start < kOuterPanelWidth1024;
         local_start += kInnerPanelWidth1024) {
        const int panel_start = outer_start + local_start;
        const int first_row = panel_start + lane;
        const int second_row = first_row + kPanelRowSplit1024;
        const bool first_active = first_row < kCholeskyOrder1024;
        const bool second_active = second_row < kCholeskyOrder1024;
        float first_products[kInnerPanelWidth1024];
        float second_products[kInnerPanelWidth1024];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth1024; ++column) {
            first_products[column] = 0.0f;
            second_products[column] = 0.0f;
        }

        for (int prior_start = 0; prior_start < local_start;
             prior_start += kInnerPanelWidth1024) {
            float first_values[kInnerPanelWidth1024];
            float second_values[kInnerPanelWidth1024];
            #pragma unroll
            for (int depth = 0; depth < kInnerPanelWidth1024; ++depth) {
                const int prior_column =
                    outer_start + prior_start + depth;
                first_values[depth] = first_active
                    ? factor[first_row + prior_column * kCholeskyOrder1024]
                    : 0.0f;
                second_values[depth] = second_active
                    ? factor[second_row + prior_column * kCholeskyOrder1024]
                    : 0.0f;
            }
            if (lane < kInnerPanelWidth1024) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth1024; ++depth) {
                    basis[depth][lane] = factor[
                        panel_start + lane +
                        (outer_start + prior_start + depth) *
                            kCholeskyOrder1024
                    ];
                }
            }
            __syncthreads();
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth1024; ++column) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth1024; ++depth) {
                    first_products[column] = fmaf(
                        first_values[depth],
                        basis[depth][column],
                        first_products[column]
                    );
                    second_products[column] = fmaf(
                        second_values[depth],
                        basis[depth][column],
                        second_products[column]
                    );
                }
            }
            __syncthreads();
        }

        float first_panel[kInnerPanelWidth1024];
        float second_panel[kInnerPanelWidth1024];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth1024; ++column) {
            first_panel[column] =
                first_active && lane >= column
                ? factor[first_row +
                         (panel_start + column) * kCholeskyOrder1024] -
                      first_products[column]
                : 0.0f;
            second_panel[column] = second_active
                ? factor[second_row +
                         (panel_start + column) * kCholeskyOrder1024] -
                      second_products[column]
                : 0.0f;
        }
        __syncthreads();

        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth1024; ++column) {
            if (lane == column) {
                panel_diagonal[column] = sqrtf(first_panel[column]);
            }
            __syncthreads();
            if (first_active && lane >= column) {
                first_panel[column] /= panel_diagonal[column];
            }
            if (second_active) {
                second_panel[column] /= panel_diagonal[column];
            }
            __syncthreads();
            if (lane < kInnerPanelWidth1024) {
                pivot_column[lane] = first_panel[column];
            }
            __syncthreads();
            if (first_active) {
                const float current = first_panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth1024;
                     ++later) {
                    if (lane >= later) {
                        first_panel[later] = fmaf(
                            -current,
                            pivot_column[later],
                            first_panel[later]
                        );
                    }
                }
            }
            if (second_active) {
                const float current = second_panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth1024;
                     ++later) {
                    second_panel[later] = fmaf(
                        -current,
                        pivot_column[later],
                        second_panel[later]
                    );
                }
            }
            __syncthreads();
        }

        if (first_active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth1024; ++column) {
                if (lane >= column) {
                    factor[first_row +
                           (panel_start + column) * kCholeskyOrder1024] =
                        first_panel[column];
                }
            }
        }
        if (second_active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth1024; ++column) {
                factor[second_row +
                       (panel_start + column) * kCholeskyOrder1024] =
                    second_panel[column];
            }
        }
        __syncthreads();
    }
}

__global__ __launch_bounds__(512, 2) void transpose_output1024_kernel(
    float* __restrict__ output
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder1024 * kCholeskyOrder1024;
    float* factor = output + matrix_offset;
    __shared__ float transfer[kTransferWidth1024][kTransferWidth1024 + 1];

    const int tile_row = static_cast<int>(blockIdx.y);
        #pragma unroll
        for (int item = lane; item < 1024; item += 512) {
            const int local_row = item >> 5;
            const int local_col = item & 31;
            transfer[local_row][local_col] = factor[
                (tile_row * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                tile_row * kTransferWidth1024 + local_col
            ];
        }
        __syncthreads();
        #pragma unroll
        for (int item = lane; item < 1024; item += 512) {
            const int local_row = item >> 5;
            const int local_col = item & 31;
            factor[
                (tile_row * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                tile_row * kTransferWidth1024 + local_col
            ] = local_row >= local_col
                ? transfer[local_col][local_row]
                : 0.0f;
        }
        __syncthreads();

        for (int tile_col = 0; tile_col < tile_row; ++tile_col) {
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                transfer[local_row][local_col] = factor[
                    (tile_col * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                    tile_row * kTransferWidth1024 + local_col
                ];
            }
            __syncthreads();
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                factor[
                    (tile_row * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                    tile_col * kTransferWidth1024 + local_col
                ] = transfer[local_col][local_row];
                factor[
                    (tile_col * kTransferWidth1024 + local_row) * kCholeskyOrder1024 +
                    tile_row * kTransferWidth1024 + local_col
                ] = 0.0f;
            }
            __syncthreads();
        }
}
}  // namespace

torch::Tensor cholesky1024(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int64_t matrix_stride =
        static_cast<int64_t>(kCholeskyOrder1024) * kCholeskyOrder1024;
    float* factor = output.data_ptr<float>();
    const float alpha = -1.0f;
    const float beta = 1.0f;

    cudaError_t cuda_status = cudaDeviceSynchronize();
    static thread_local cublasHandle_t handle = nullptr;
    cublasStatus_t operation_status = CUBLAS_STATUS_SUCCESS;
    if (handle == nullptr) {
        operation_status = cublasCreate(&handle);
        if (operation_status == CUBLAS_STATUS_SUCCESS) {
            operation_status =
                cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
        }
        if (operation_status != CUBLAS_STATUS_SUCCESS &&
            handle != nullptr) {
            cublasDestroy(handle);
            handle = nullptr;
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_input1024_kernel<<<dim3(batch, 32), 512>>>(
            input.data_ptr<float>(), factor
        );
        cuda_status = cudaGetLastError();
    }

    for (int outer_start = 0;
         outer_start < kCholeskyOrder1024 &&
             operation_status == CUBLAS_STATUS_SUCCESS &&
             cuda_status == cudaSuccess;
         outer_start += kOuterPanelWidth1024) {
        if (outer_start > 0) {
            const int remaining = kCholeskyOrder1024 - outer_start;
            operation_status = cublasGemmStridedBatchedEx(
                handle,
                CUBLAS_OP_N,
                CUBLAS_OP_T,
                remaining,
                kOuterPanelWidth1024,
                outer_start,
                &alpha,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder1024,
                matrix_stride,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder1024,
                matrix_stride,
                &beta,
                factor + outer_start +
                    outer_start * kCholeskyOrder1024,
                CUDA_R_32F,
                kCholeskyOrder1024,
                matrix_stride,
                batch,
                CUBLAS_COMPUTE_32F_PEDANTIC,
                CUBLAS_GEMM_DEFAULT_TENSOR_OP
            );
        }
        if (operation_status == CUBLAS_STATUS_SUCCESS &&
            cuda_status == cudaSuccess) {
            factor_panel1024_kernel<<<batch, 512>>>(factor, outer_start);
            cuda_status = cudaGetLastError();
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_output1024_kernel<<<dim3(batch, 32), 512>>>(factor);
        cuda_status = cudaGetLastError();
    }
    if (cuda_status == cudaSuccess) {
        cuda_status = cudaDeviceSynchronize();
    }

    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "cublasGemmStridedBatchedEx failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n1024 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}
"""

_CUDA_1024_D64_SOURCE = r"""
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <limits>
#include <c10/cuda/CUDAException.h>

// Small-batch n1024 driver: reuses the column-major transpose kernels,
// panel constants, and PEDANTIC prefix GEMM pattern already defined in
// this translation unit, factors only the strict 64x64 diagonal block
// per stage with the ordered inner8 recurrence, and solves every
// external row through one exact FP32 cublasStrsmBatched call per
// stage. The Python dispatch gates this driver to small batches.
namespace {
constexpr int kSolveStageCount1024 =
    kCholeskyOrder1024 / kOuterPanelWidth1024 - 1;

__global__ __launch_bounds__(64) void factor_diagonal1024_64_inner8_kernel(
    float* __restrict__ output,
    int outer_start
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) *
        kCholeskyOrder1024 * kCholeskyOrder1024;
    float* factor = output + matrix_offset;
    __shared__ float basis[kInnerPanelWidth1024][kInnerPanelWidth1024];
    __shared__ float panel_diagonal[kInnerPanelWidth1024];
    __shared__ float pivot_column[kInnerPanelWidth1024];

    for (int local_start = 0; local_start < kOuterPanelWidth1024;
         local_start += kInnerPanelWidth1024) {
        const int panel_start = outer_start + local_start;
        const int matrix_row = panel_start + lane;
        const bool active =
            local_start + lane < kOuterPanelWidth1024;
        float products[kInnerPanelWidth1024];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth1024; ++column) {
            products[column] = 0.0f;
        }

        for (int prior_start = 0; prior_start < local_start;
             prior_start += kInnerPanelWidth1024) {
            float row_values[kInnerPanelWidth1024];
            #pragma unroll
            for (int depth = 0; depth < kInnerPanelWidth1024; ++depth) {
                const int prior_column =
                    outer_start + prior_start + depth;
                row_values[depth] = active
                    ? factor[
                        matrix_row +
                        prior_column * kCholeskyOrder1024
                    ]
                    : 0.0f;
            }
            if (lane < kInnerPanelWidth1024) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth1024; ++depth) {
                    basis[depth][lane] = factor[
                        panel_start + lane +
                        (outer_start + prior_start + depth) *
                            kCholeskyOrder1024
                    ];
                }
            }
            __syncthreads();
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth1024; ++column) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth1024; ++depth) {
                    products[column] = fmaf(
                        row_values[depth],
                        basis[depth][column],
                        products[column]
                    );
                }
            }
            __syncthreads();
        }

        float panel[kInnerPanelWidth1024];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth1024; ++column) {
            panel[column] =
                active && lane >= column
                ? factor[
                    matrix_row +
                    (panel_start + column) * kCholeskyOrder1024
                ] - products[column]
                : 0.0f;
        }
        __syncthreads();

        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth1024; ++column) {
            if (lane == column) {
                panel_diagonal[column] = sqrtf(panel[column]);
            }
            __syncthreads();
            if (active && lane >= column) {
                panel[column] /= panel_diagonal[column];
            }
            __syncthreads();
            if (lane < kInnerPanelWidth1024) {
                pivot_column[lane] = panel[column];
            }
            __syncthreads();
            if (active) {
                const float current = panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth1024;
                     ++later) {
                    if (lane >= later) {
                        panel[later] = fmaf(
                            -current,
                            pivot_column[later],
                            panel[later]
                        );
                    }
                }
            }
            __syncthreads();
        }

        if (active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth1024; ++column) {
                if (lane >= column) {
                    factor[
                        matrix_row +
                        (panel_start + column) * kCholeskyOrder1024
                    ] = panel[column];
                }
            }
        }
        __syncthreads();
    }
}

__global__ void build_trsm_pointer_arrays1024_kernel(
    float* __restrict__ output,
    int batch,
    const float** __restrict__ diagonal_pointers,
    float** __restrict__ rhs_pointers
) {
    const int matrix = static_cast<int>(blockIdx.x);
    const int stage = static_cast<int>(blockIdx.y);
    const int outer_start = stage * kOuterPanelWidth1024;
    const size_t pointer_index =
        static_cast<size_t>(stage) * batch + matrix;
    float* factor =
        output +
        static_cast<size_t>(matrix) *
            kCholeskyOrder1024 * kCholeskyOrder1024;
    diagonal_pointers[pointer_index] =
        factor +
        outer_start +
        outer_start * kCholeskyOrder1024;
    rhs_pointers[pointer_index] =
        factor +
        outer_start + kOuterPanelWidth1024 +
        outer_start * kCholeskyOrder1024;
}
}  // namespace

torch::Tensor cholesky1024_d64(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int64_t matrix_stride =
        static_cast<int64_t>(kCholeskyOrder1024) * kCholeskyOrder1024;
    torch::Tensor diagonal_pointer_storage = torch::empty(
        {kSolveStageCount1024, batch},
        input.options().dtype(at::kLong)
    );
    torch::Tensor rhs_pointer_storage = torch::empty(
        {kSolveStageCount1024, batch},
        input.options().dtype(at::kLong)
    );
    float* factor = output.data_ptr<float>();
    const float** diagonal_pointers = reinterpret_cast<const float**>(
        diagonal_pointer_storage.data_ptr<int64_t>()
    );
    float** rhs_pointers = reinterpret_cast<float**>(
        rhs_pointer_storage.data_ptr<int64_t>()
    );
    const float alpha = -1.0f;
    const float beta = 1.0f;
    const float solve_alpha = 1.0f;

    cudaError_t cuda_status = cudaDeviceSynchronize();
    static thread_local cublasHandle_t handle = nullptr;
    cublasStatus_t operation_status = CUBLAS_STATUS_SUCCESS;
    if (handle == nullptr) {
        operation_status = cublasCreate(&handle);
        if (operation_status == CUBLAS_STATUS_SUCCESS) {
            operation_status =
                cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
        }
        if (operation_status == CUBLAS_STATUS_SUCCESS) {
            operation_status =
                cublasSetMathMode(handle, CUBLAS_PEDANTIC_MATH);
        }
        if (operation_status != CUBLAS_STATUS_SUCCESS &&
            handle != nullptr) {
            cublasDestroy(handle);
            handle = nullptr;
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_input1024_kernel<<<dim3(batch, 32), 512>>>(
            input.data_ptr<float>(), factor
        );
        cuda_status = cudaGetLastError();
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        build_trsm_pointer_arrays1024_kernel<<<
            dim3(batch, kSolveStageCount1024), 1
        >>>(
            factor,
            batch,
            diagonal_pointers,
            rhs_pointers
        );
        cuda_status = cudaGetLastError();
    }

    for (int outer_start = 0;
         outer_start < kCholeskyOrder1024 &&
             operation_status == CUBLAS_STATUS_SUCCESS &&
             cuda_status == cudaSuccess;
         outer_start += kOuterPanelWidth1024) {
        const int stage = outer_start / kOuterPanelWidth1024;
        const int remaining = kCholeskyOrder1024 - outer_start;
        const int trailing_rows =
            remaining - kOuterPanelWidth1024;
        if (outer_start > 0) {
            operation_status = cublasGemmStridedBatchedEx(
                handle,
                CUBLAS_OP_N,
                CUBLAS_OP_T,
                remaining,
                kOuterPanelWidth1024,
                outer_start,
                &alpha,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder1024,
                matrix_stride,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder1024,
                matrix_stride,
                &beta,
                factor + outer_start +
                    outer_start * kCholeskyOrder1024,
                CUDA_R_32F,
                kCholeskyOrder1024,
                matrix_stride,
                batch,
                CUBLAS_COMPUTE_32F_PEDANTIC,
                CUBLAS_GEMM_DEFAULT_TENSOR_OP
            );
        }
        if (operation_status == CUBLAS_STATUS_SUCCESS &&
            cuda_status == cudaSuccess) {
            factor_diagonal1024_64_inner8_kernel<<<
                batch, kOuterPanelWidth1024
            >>>(factor, outer_start);
            cuda_status = cudaGetLastError();
        }
        if (trailing_rows > 0 &&
            operation_status == CUBLAS_STATUS_SUCCESS &&
            cuda_status == cudaSuccess) {
            operation_status = cublasStrsmBatched(
                handle,
                CUBLAS_SIDE_RIGHT,
                CUBLAS_FILL_MODE_LOWER,
                CUBLAS_OP_T,
                CUBLAS_DIAG_NON_UNIT,
                trailing_rows,
                kOuterPanelWidth1024,
                &solve_alpha,
                diagonal_pointers +
                    static_cast<size_t>(stage) * batch,
                kCholeskyOrder1024,
                rhs_pointers +
                    static_cast<size_t>(stage) * batch,
                kCholeskyOrder1024,
                batch
            );
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_output1024_kernel<<<dim3(batch, 32), 512>>>(factor);
        cuda_status = cudaGetLastError();
    }
    if (cuda_status == cudaSuccess) {
        cuda_status = cudaDeviceSynchronize();
    }

    TORCH_CHECK(
        operation_status == CUBLAS_STATUS_SUCCESS,
        "n1024 d64 cuBLAS operation failed: ",
        static_cast<int>(operation_status)
    );
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n1024 d64 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}
"""

_CUDA_2048_SOURCE = r"""
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <limits>
#include <c10/cuda/CUDAException.h>

// Uses the MAGMA-derived panel structure and retained license above.
namespace {
constexpr int kCholeskyOrder2048 = 2048;
constexpr int kInnerPanelWidth2048 = 8;
constexpr int kOuterPanelWidth2048 = 64;
constexpr int kTransferWidth2048 = 32;
constexpr int kPanelRowSplit2048 = 1024;

__global__ __launch_bounds__(512, 2) void transpose_input2048_kernel(
    const float* __restrict__ input,
    float* __restrict__ output
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder2048 * kCholeskyOrder2048;
    const float* source = input + matrix_offset;
    float* factor = output + matrix_offset;
    __shared__ float transfer[kTransferWidth2048][kTransferWidth2048 + 1];

    const int tile_row = static_cast<int>(blockIdx.y);
    for (int tile_col = 0; tile_col <= tile_row; ++tile_col) {
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                transfer[local_row][local_col] = source[
                    (tile_row * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                    tile_col * kTransferWidth2048 + local_col
                ];
            }
            __syncthreads();
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                factor[
                    (tile_col * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                    tile_row * kTransferWidth2048 + local_col
                ] = transfer[local_col][local_row];
            }
            __syncthreads();
        }
}

__global__ __launch_bounds__(1024, 1) void factor_panel2048_kernel(
    float* __restrict__ output,
    int outer_start
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder2048 * kCholeskyOrder2048;
    float* factor = output + matrix_offset;
    __shared__ float basis[kInnerPanelWidth2048][kInnerPanelWidth2048];
    __shared__ float panel_diagonal[kInnerPanelWidth2048];
    __shared__ float pivot_column[kInnerPanelWidth2048];

    for (int local_start = 0; local_start < kOuterPanelWidth2048;
         local_start += kInnerPanelWidth2048) {
        const int panel_start = outer_start + local_start;
        const int first_row = panel_start + lane;
        const int second_row = first_row + kPanelRowSplit2048;
        const bool first_active = first_row < kCholeskyOrder2048;
        const bool second_active = second_row < kCholeskyOrder2048;
        float first_products[kInnerPanelWidth2048];
        float second_products[kInnerPanelWidth2048];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth2048; ++column) {
            first_products[column] = 0.0f;
            second_products[column] = 0.0f;
        }

        for (int prior_start = 0; prior_start < local_start;
             prior_start += kInnerPanelWidth2048) {
            float first_values[kInnerPanelWidth2048];
            float second_values[kInnerPanelWidth2048];
            #pragma unroll
            for (int depth = 0; depth < kInnerPanelWidth2048; ++depth) {
                const int prior_column =
                    outer_start + prior_start + depth;
                first_values[depth] = first_active
                    ? factor[first_row + prior_column * kCholeskyOrder2048]
                    : 0.0f;
                second_values[depth] = second_active
                    ? factor[second_row + prior_column * kCholeskyOrder2048]
                    : 0.0f;
            }
            if (lane < kInnerPanelWidth2048) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth2048; ++depth) {
                    basis[depth][lane] = factor[
                        panel_start + lane +
                        (outer_start + prior_start + depth) *
                            kCholeskyOrder2048
                    ];
                }
            }
            __syncthreads();
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth2048; ++column) {
                #pragma unroll
                for (int depth = 0; depth < kInnerPanelWidth2048; ++depth) {
                    first_products[column] = fmaf(
                        first_values[depth],
                        basis[depth][column],
                        first_products[column]
                    );
                    second_products[column] = fmaf(
                        second_values[depth],
                        basis[depth][column],
                        second_products[column]
                    );
                }
            }
            __syncthreads();
        }

        float first_panel[kInnerPanelWidth2048];
        float second_panel[kInnerPanelWidth2048];
        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth2048; ++column) {
            first_panel[column] =
                first_active && lane >= column
                ? factor[first_row +
                         (panel_start + column) * kCholeskyOrder2048] -
                      first_products[column]
                : 0.0f;
            second_panel[column] = second_active
                ? factor[second_row +
                         (panel_start + column) * kCholeskyOrder2048] -
                      second_products[column]
                : 0.0f;
        }
        __syncthreads();

        #pragma unroll
        for (int column = 0; column < kInnerPanelWidth2048; ++column) {
            if (lane == column) {
                panel_diagonal[column] = sqrtf(first_panel[column]);
            }
            __syncthreads();
            if (first_active && lane >= column) {
                first_panel[column] /= panel_diagonal[column];
            }
            if (second_active) {
                second_panel[column] /= panel_diagonal[column];
            }
            __syncthreads();
            if (lane < kInnerPanelWidth2048) {
                pivot_column[lane] = first_panel[column];
            }
            __syncthreads();
            if (first_active) {
                const float current = first_panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth2048;
                     ++later) {
                    if (lane >= later) {
                        first_panel[later] = fmaf(
                            -current,
                            pivot_column[later],
                            first_panel[later]
                        );
                    }
                }
            }
            if (second_active) {
                const float current = second_panel[column];
                #pragma unroll
                for (int later = column + 1;
                     later < kInnerPanelWidth2048;
                     ++later) {
                    second_panel[later] = fmaf(
                        -current,
                        pivot_column[later],
                        second_panel[later]
                    );
                }
            }
            __syncthreads();
        }

        if (first_active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth2048; ++column) {
                if (lane >= column) {
                    factor[first_row +
                           (panel_start + column) * kCholeskyOrder2048] =
                        first_panel[column];
                }
            }
        }
        if (second_active) {
            #pragma unroll
            for (int column = 0; column < kInnerPanelWidth2048; ++column) {
                factor[second_row +
                       (panel_start + column) * kCholeskyOrder2048] =
                    second_panel[column];
            }
        }
        __syncthreads();
    }
}

__global__ __launch_bounds__(512, 2) void transpose_output2048_kernel(
    float* __restrict__ output
) {
    const int lane = threadIdx.x;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kCholeskyOrder2048 * kCholeskyOrder2048;
    float* factor = output + matrix_offset;
    __shared__ float transfer[kTransferWidth2048][kTransferWidth2048 + 1];

    const int tile_row = static_cast<int>(blockIdx.y);
        #pragma unroll
        for (int item = lane; item < 1024; item += 512) {
            const int local_row = item >> 5;
            const int local_col = item & 31;
            transfer[local_row][local_col] = factor[
                (tile_row * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                tile_row * kTransferWidth2048 + local_col
            ];
        }
        __syncthreads();
        #pragma unroll
        for (int item = lane; item < 1024; item += 512) {
            const int local_row = item >> 5;
            const int local_col = item & 31;
            factor[
                (tile_row * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                tile_row * kTransferWidth2048 + local_col
            ] = local_row >= local_col
                ? transfer[local_col][local_row]
                : 0.0f;
        }
        __syncthreads();

        for (int tile_col = 0; tile_col < tile_row; ++tile_col) {
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                transfer[local_row][local_col] = factor[
                    (tile_col * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                    tile_row * kTransferWidth2048 + local_col
                ];
            }
            __syncthreads();
            #pragma unroll
            for (int item = lane; item < 1024; item += 512) {
                const int local_row = item >> 5;
                const int local_col = item & 31;
                factor[
                    (tile_row * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                    tile_col * kTransferWidth2048 + local_col
                ] = transfer[local_col][local_row];
                factor[
                    (tile_col * kTransferWidth2048 + local_row) * kCholeskyOrder2048 +
                    tile_row * kTransferWidth2048 + local_col
                ] = 0.0f;
            }
            __syncthreads();
        }
}
}  // namespace

torch::Tensor cholesky2048(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int64_t matrix_stride =
        static_cast<int64_t>(kCholeskyOrder2048) * kCholeskyOrder2048;
    float* factor = output.data_ptr<float>();
    const float alpha = -1.0f;
    const float beta = 1.0f;

    cudaError_t cuda_status = cudaDeviceSynchronize();
    static thread_local cublasHandle_t handle = nullptr;
    cublasStatus_t operation_status = CUBLAS_STATUS_SUCCESS;
    if (handle == nullptr) {
        operation_status = cublasCreate(&handle);
        if (operation_status == CUBLAS_STATUS_SUCCESS) {
            operation_status =
                cublasSetPointerMode(handle, CUBLAS_POINTER_MODE_HOST);
        }
        if (operation_status != CUBLAS_STATUS_SUCCESS &&
            handle != nullptr) {
            cublasDestroy(handle);
            handle = nullptr;
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_input2048_kernel<<<dim3(batch, 64), 512>>>(
            input.data_ptr<float>(), factor
        );
        cuda_status = cudaGetLastError();
    }

    for (int outer_start = 0;
         outer_start < kCholeskyOrder2048 &&
             operation_status == CUBLAS_STATUS_SUCCESS &&
             cuda_status == cudaSuccess;
         outer_start += kOuterPanelWidth2048) {
        if (outer_start > 0) {
            const int remaining = kCholeskyOrder2048 - outer_start;
            operation_status = cublasGemmStridedBatchedEx(
                handle,
                CUBLAS_OP_N,
                CUBLAS_OP_T,
                remaining,
                kOuterPanelWidth2048,
                outer_start,
                &alpha,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder2048,
                matrix_stride,
                factor + outer_start,
                CUDA_R_32F,
                kCholeskyOrder2048,
                matrix_stride,
                &beta,
                factor + outer_start +
                    outer_start * kCholeskyOrder2048,
                CUDA_R_32F,
                kCholeskyOrder2048,
                matrix_stride,
                batch,
                CUBLAS_COMPUTE_32F_PEDANTIC,
                CUBLAS_GEMM_DEFAULT_TENSOR_OP
            );
        }
        if (operation_status == CUBLAS_STATUS_SUCCESS &&
            cuda_status == cudaSuccess) {
            factor_panel2048_kernel<<<batch, 1024>>>(factor, outer_start);
            cuda_status = cudaGetLastError();
        }
    }

    if (operation_status == CUBLAS_STATUS_SUCCESS &&
        cuda_status == cudaSuccess) {
        transpose_output2048_kernel<<<dim3(batch, 64), 512>>>(factor);
        cuda_status = cudaGetLastError();
    }
    if (cuda_status == cudaSuccess) {
        cuda_status = cudaDeviceSynchronize();
    }

    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n2048 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}
"""



_CUDA_SMALL_SOURCE = r"""
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <limits>
#include <c10/cuda/CUDAException.h>

namespace {
constexpr unsigned int kFullLaneMask = 0xffffffffu;
constexpr int kLaneCount = 32;
constexpr int kSmallOrder32 = 32;
constexpr int kSmallOrder64 = 64;
constexpr int kMatrixWarps32 = 4;
constexpr int kRowQuads32 = kSmallOrder32 / 4;
constexpr int kRowQuads64 = kSmallOrder64 / 4;
constexpr int kHalfQuads64 = kRowQuads64 / 2;
constexpr int kSmallOrder128 = 128;
constexpr int kBlocks128 = kSmallOrder128 / kLaneCount;
constexpr int kBlockQuads32 = kLaneCount / 4;
constexpr int kTilePad = kLaneCount + 1;

/*
  Warp-cooperative coalesced staging between global memory and a padded
  32x33 shared tile. A 32x32 float block whose rows sit stride4 float4s
  apart in global memory is moved with consecutive-address float4
  transactions: on request chunk c, lane l touches row 4*c + (l >> 3),
  quad l & 7, so every warp request covers four fully utilized 128-byte
  lines instead of the up-to-32 distinct lines of the row-per-lane
  pattern these helpers replace. The +1 padding keeps both the quad-wise
  staging accesses and the per-lane row copies free of shared-memory
  bank conflicts. Pure data movement: the values, their register
  destinations, and their global locations are exactly those of the
  direct row-per-lane transfers.
*/
__device__ __forceinline__ void stage_tile_from_global(
    float (*tile)[kTilePad],
    const float4* __restrict__ source4,
    int stride4,
    int lane
) {
    #pragma unroll
    for (int chunk = 0; chunk < 8; ++chunk) {
        const int row = 4 * chunk + (lane >> 3);
        const int quad = lane & 7;
        const float4 value = source4[row * stride4 + quad];
        tile[row][4 * quad + 0] = value.x;
        tile[row][4 * quad + 1] = value.y;
        tile[row][4 * quad + 2] = value.z;
        tile[row][4 * quad + 3] = value.w;
    }
}

__device__ __forceinline__ void stage_tile_to_global(
    float4* __restrict__ destination4,
    const float (*tile)[kTilePad],
    int stride4,
    int lane
) {
    #pragma unroll
    for (int chunk = 0; chunk < 8; ++chunk) {
        const int row = 4 * chunk + (lane >> 3);
        const int quad = lane & 7;
        float4 value;
        value.x = tile[row][4 * quad + 0];
        value.y = tile[row][4 * quad + 1];
        value.z = tile[row][4 * quad + 2];
        value.w = tile[row][4 * quad + 3];
        destination4[row * stride4 + quad] = value;
    }
}

__device__ __forceinline__ void stage_zero_to_global(
    float4* __restrict__ destination4,
    int stride4,
    int lane
) {
    const float4 zero_value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
    #pragma unroll
    for (int chunk = 0; chunk < 8; ++chunk) {
        destination4[(4 * chunk + (lane >> 3)) * stride4 + (lane & 7)] =
            zero_value;
    }
}

__device__ __forceinline__ void rows_from_stage(
    float (&row)[kLaneCount],
    const float (*tile)[kTilePad],
    int lane
) {
    #pragma unroll
    for (int c = 0; c < kLaneCount; ++c) {
        row[c] = tile[lane][c];
    }
}

__device__ __forceinline__ void rows_to_stage(
    float (*tile)[kTilePad],
    const float (&row)[kLaneCount],
    int lane
) {
    #pragma unroll
    for (int c = 0; c < kLaneCount; ++c) {
        tile[lane][c] = row[c];
    }
}

/* Same zero predicate, element for element, as the direct masked stores
   this replaces: column c of row `lane` survives iff c <= lane. */
__device__ __forceinline__ void rows_to_stage_masked(
    float (*tile)[kTilePad],
    const float (&row)[kLaneCount],
    int lane
) {
    #pragma unroll
    for (int c = 0; c < kLaneCount; ++c) {
        tile[lane][c] = c <= lane ? row[c] : 0.0f;
    }
}

/*
  Right-looking register-resident 32x32 factor recurrence. Lane i owns row i
  of the tile. Per column k: broadcast the pivot from lane k, take
  sqrt(max(pivot, 0)), divide the Schur column by the pivot root, then apply
  one rank-1 trailing update. This mirrors, column for column, the audited
  Triton recurrence used by the accepted experiment 127/130 kernels; entries
  above the diagonal are never read and are masked to zero at store time.
*/
__device__ __forceinline__ void factor_tile32_rows(
    float (&row)[kLaneCount],
    int lane
) {
    #pragma unroll
    for (int k = 0; k < kLaneCount; ++k) {
        const float pivot = __shfl_sync(kFullLaneMask, row[k], k);
        const float diagonal = sqrtf(fmaxf(pivot, 0.0f));
        const float column = lane >= k ? row[k] / diagonal : 0.0f;
        row[k] = column;
        #pragma unroll
        for (int j = k + 1; j < kLaneCount; ++j) {
            const float mirror = __shfl_sync(kFullLaneMask, column, j);
            row[j] = fmaf(-column, mirror, row[j]);
        }
    }
}

__global__ __launch_bounds__(kMatrixWarps32 * kLaneCount)
void cholesky32_warp_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch
) {
    const int lane = static_cast<int>(threadIdx.x) & (kLaneCount - 1);
    const int warp = static_cast<int>(threadIdx.x) >> 5;
    const int matrix = static_cast<int>(blockIdx.x) * kMatrixWarps32 + warp;

    __shared__ float stage[kMatrixWarps32][kLaneCount][kTilePad];

    if (matrix >= batch) {
        return;
    }

    const size_t matrix_offset =
        static_cast<size_t>(matrix) * kSmallOrder32 * kSmallOrder32;
    const float4* source =
        reinterpret_cast<const float4*>(input + matrix_offset);
    float4* destination = reinterpret_cast<float4*>(output + matrix_offset);
    float (*tile)[kTilePad] = stage[warp];

    stage_tile_from_global(tile, source, kRowQuads32, lane);
    __syncwarp();

    float row[kLaneCount];
    rows_from_stage(row, tile, lane);

    factor_tile32_rows(row, lane);

    rows_to_stage_masked(tile, row, lane);
    __syncwarp();
    stage_tile_to_global(destination, tile, kRowQuads32, lane);
}

/*
  Blocked 64x64 factorization with one warp per matrix, mirroring the audited
  experiment 130 recurrence order: factor A11, solve L21 against the L11
  columns (plain division by the L11 diagonal), subtract the rank-1 Schur
  updates of L21 into A22 with ascending k, then factor the updated A22.
  Lane i owns row i of each 32x32 block; the upper-right block is written as
  explicit zeros exactly like the Triton kernel it replaces.
*/
__global__ __launch_bounds__(kLaneCount)
void cholesky64_warp_kernel(
    const float* __restrict__ input,
    float* __restrict__ output
) {
    const int lane = static_cast<int>(threadIdx.x);
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kSmallOrder64 * kSmallOrder64;
    const float4* source =
        reinterpret_cast<const float4*>(input + matrix_offset);
    float4* destination = reinterpret_cast<float4*>(output + matrix_offset);

    __shared__ float stage[kLaneCount][kTilePad];

    // Block bases in float4 units (row stride kRowQuads64): leading (0,0)
    // at 0, strict-upper (0,1) at kHalfQuads64, solved (1,0) at
    // kLaneCount * kRowQuads64, trailing (1,1) at that plus kHalfQuads64.
    constexpr int kSolvedBase4 = kLaneCount * kRowQuads64;

    stage_tile_from_global(stage, source, kRowQuads64, lane);
    __syncwarp();
    float leading[kLaneCount];
    rows_from_stage(leading, stage, lane);

    factor_tile32_rows(leading, lane);

    __syncwarp();
    stage_tile_from_global(stage, source + kSolvedBase4, kRowQuads64, lane);
    __syncwarp();
    float solved[kLaneCount];
    rows_from_stage(solved, stage, lane);

    #pragma unroll
    for (int k = 0; k < kLaneCount; ++k) {
        const float diagonal = __shfl_sync(kFullLaneMask, leading[k], k);
        const float column = solved[k] / diagonal;
        solved[k] = column;
        #pragma unroll
        for (int j = k + 1; j < kLaneCount; ++j) {
            const float mirror = __shfl_sync(kFullLaneMask, leading[k], j);
            solved[j] = fmaf(-column, mirror, solved[j]);
        }
    }

    rows_to_stage_masked(stage, leading, lane);
    __syncwarp();
    stage_tile_to_global(destination, stage, kRowQuads64, lane);
    stage_zero_to_global(destination + kHalfQuads64, kRowQuads64, lane);

    __syncwarp();
    stage_tile_from_global(
        stage, source + kSolvedBase4 + kHalfQuads64, kRowQuads64, lane
    );
    __syncwarp();
    float trailing[kLaneCount];
    rows_from_stage(trailing, stage, lane);

    #pragma unroll
    for (int k = 0; k < kLaneCount; ++k) {
        const float column = solved[k];
        #pragma unroll
        for (int j = 0; j < kLaneCount; ++j) {
            const float mirror = __shfl_sync(kFullLaneMask, column, j);
            trailing[j] = fmaf(-column, mirror, trailing[j]);
        }
    }

    rows_to_stage(stage, solved, lane);
    __syncwarp();
    stage_tile_to_global(destination + kSolvedBase4, stage, kRowQuads64, lane);

    factor_tile32_rows(trailing, lane);

    __syncwarp();
    rows_to_stage_masked(stage, trailing, lane);
    __syncwarp();
    stage_tile_to_global(
        destination + kSolvedBase4 + kHalfQuads64, stage, kRowQuads64, lane
    );
}

/*
  Triangular-solve body for one 32x32 block row set: identical, operation
  for operation, to the L21 solve inside the n64 kernel above, except that
  the factor rows arrive through a register array (staged from a shared
  tile) instead of the same warp's leading registers. Lane i owns row i of
  the solved block; factor_rows holds lane i's row i of the already
  factored diagonal block.
*/
__device__ __forceinline__ void solve_tile32_rows(
    float (&row)[kLaneCount],
    const float (&factor_rows)[kLaneCount]
) {
    #pragma unroll
    for (int k = 0; k < kLaneCount; ++k) {
        const float diagonal = __shfl_sync(kFullLaneMask, factor_rows[k], k);
        const float column = row[k] / diagonal;
        row[k] = column;
        #pragma unroll
        for (int j = k + 1; j < kLaneCount; ++j) {
            const float mirror =
                __shfl_sync(kFullLaneMask, factor_rows[k], j);
            row[j] = fmaf(-column, mirror, row[j]);
        }
    }
}

/*
  Rank-32 trailing update of one 32x32 block: row[j] accumulates
  -row_factor[k] * col_factor[k][j] with ascending k, exactly the Schur
  update loop of the n64 kernel above (there row_factor == col_factor ==
  the solved panel rows). Single-rounding fmaf throughout.
*/
__device__ __forceinline__ void update_tile32_rows(
    float (&row)[kLaneCount],
    const float (&row_factor)[kLaneCount],
    const float (&col_factor)[kLaneCount]
) {
    #pragma unroll
    for (int k = 0; k < kLaneCount; ++k) {
        const float left = row_factor[k];
        #pragma unroll
        for (int j = 0; j < kLaneCount; ++j) {
            const float mirror =
                __shfl_sync(kFullLaneMask, col_factor[k], j);
            row[j] = fmaf(-left, mirror, row[j]);
        }
    }
}

__device__ __forceinline__ void tile_store_rows(
    float (*tile)[kTilePad],
    const float (&row)[kLaneCount],
    int lane
) {
    #pragma unroll
    for (int c = 0; c < kLaneCount; ++c) {
        tile[lane][c] = row[c];
    }
}

__device__ __forceinline__ void tile_load_rows(
    float (&row)[kLaneCount],
    const float (*tile)[kTilePad],
    int lane
) {
    #pragma unroll
    for (int c = 0; c < kLaneCount; ++c) {
        row[c] = tile[lane][c];
    }
}

__device__ __forceinline__ void load_block128(
    float (&row)[kLaneCount],
    const float* __restrict__ matrix,
    int block_row,
    int block_col,
    int lane
) {
    const float4* source = reinterpret_cast<const float4*>(
        matrix
        + (block_row * kLaneCount + lane) * kSmallOrder128
        + block_col * kLaneCount
    );
    #pragma unroll
    for (int quad = 0; quad < kBlockQuads32; ++quad) {
        const float4 value = source[quad];
        row[4 * quad + 0] = value.x;
        row[4 * quad + 1] = value.y;
        row[4 * quad + 2] = value.z;
        row[4 * quad + 3] = value.w;
    }
}

__device__ __forceinline__ void store_block128(
    float* __restrict__ matrix,
    const float (&row)[kLaneCount],
    int block_row,
    int block_col,
    int lane
) {
    float4* destination = reinterpret_cast<float4*>(
        matrix
        + (block_row * kLaneCount + lane) * kSmallOrder128
        + block_col * kLaneCount
    );
    #pragma unroll
    for (int quad = 0; quad < kBlockQuads32; ++quad) {
        float4 value;
        value.x = row[4 * quad + 0];
        value.y = row[4 * quad + 1];
        value.z = row[4 * quad + 2];
        value.w = row[4 * quad + 3];
        destination[quad] = value;
    }
}

__device__ __forceinline__ void store_diag_block128(
    float* __restrict__ matrix,
    const float (&row)[kLaneCount],
    int block_row,
    int block_col,
    int lane
) {
    float4* destination = reinterpret_cast<float4*>(
        matrix
        + (block_row * kLaneCount + lane) * kSmallOrder128
        + block_col * kLaneCount
    );
    #pragma unroll
    for (int quad = 0; quad < kBlockQuads32; ++quad) {
        float4 value;
        value.x = 4 * quad + 0 <= lane ? row[4 * quad + 0] : 0.0f;
        value.y = 4 * quad + 1 <= lane ? row[4 * quad + 1] : 0.0f;
        value.z = 4 * quad + 2 <= lane ? row[4 * quad + 2] : 0.0f;
        value.w = 4 * quad + 3 <= lane ? row[4 * quad + 3] : 0.0f;
        destination[quad] = value;
    }
}

__device__ __forceinline__ void store_zero_block128(
    float* __restrict__ matrix,
    int block_row,
    int block_col,
    int lane
) {
    float4* destination = reinterpret_cast<float4*>(
        matrix
        + (block_row * kLaneCount + lane) * kSmallOrder128
        + block_col * kLaneCount
    );
    const float4 zero_value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
    #pragma unroll
    for (int quad = 0; quad < kBlockQuads32; ++quad) {
        destination[quad] = zero_value;
    }
}

/*
  Blocked right-looking 128x128 factorization, one matrix per four-warp
  CTA, extending the register-resident n32/n64 family above. The 4x4 grid
  of 32x32 blocks is distributed over the warps so that no warp ever holds
  more than two matrix blocks live at once:
    warp 0: (0,0) then (2,1), and (3,3)
    warp 1: (1,0), and (2,2)
    warp 2: (2,0) then (3,1), and (3,2)
    warp 3: (3,0), and (1,1)
  Lane i owns row i of each owned block, exactly as in the n32/n64
  kernels; panel results move between warps only through padded 32x33
  shared tiles (tile[i] holds block (i,k) of the current panel step k).
  Per step k: the diagonal owner factors block (k,k) in registers with the
  audited n32 recurrence and publishes it; owners of blocks (i,k) run the
  audited n64 triangular-solve body against the published factor rows and
  publish L(i,k); owners of trailing blocks (i,j) apply register-resident
  rank-32 updates with ascending k. Every element therefore accumulates
  its updates in exactly the ascending-k order of the unblocked
  right-looking recurrence: strict FP32 with IEEE sqrtf, true division,
  single-rounding fmaf, the same max(pivot, 0) guard, and an explicit
  all-zero strict upper triangle at store time.
*/
__global__ __launch_bounds__(kBlocks128 * kLaneCount)
void cholesky128_warp_kernel(
    const float* __restrict__ input,
    float* __restrict__ output
) {
    const int lane = static_cast<int>(threadIdx.x) & (kLaneCount - 1);
    const int warp = static_cast<int>(threadIdx.x) >> 5;
    const size_t matrix_offset =
        static_cast<size_t>(blockIdx.x) * kSmallOrder128 * kSmallOrder128;
    const float* source = input + matrix_offset;
    float* destination = output + matrix_offset;

    __shared__ float tiles[kBlocks128][kLaneCount][kTilePad];

    float first[kLaneCount];
    float second[kLaneCount];
    float row_stage[kLaneCount];
    float col_stage[kLaneCount];

    // Phase 0: step-0 panel loads; warp 0 factors the leading diagonal.
    if (warp == 0) {
        load_block128(first, source, 0, 0, lane);
        factor_tile32_rows(first, lane);
        store_diag_block128(destination, first, 0, 0, lane);
        tile_store_rows(tiles[0], first, lane);
    } else if (warp == 1) {
        load_block128(first, source, 1, 0, lane);
    } else if (warp == 2) {
        load_block128(first, source, 2, 0, lane);
    } else {
        load_block128(first, source, 3, 0, lane);
    }
    __syncthreads();

    // Phase 1: step-0 panel solves; warp 0 loads its trailing blocks and
    // stores the explicit-zero blocks of the strict upper triangle.
    if (warp == 0) {
        load_block128(first, source, 2, 1, lane);
        load_block128(second, source, 3, 3, lane);
        store_zero_block128(destination, 1, 2, lane);
        store_zero_block128(destination, 1, 3, lane);
        store_zero_block128(destination, 2, 3, lane);
    } else if (warp == 1) {
        load_block128(second, source, 2, 2, lane);
        tile_load_rows(col_stage, tiles[0], lane);
        solve_tile32_rows(first, col_stage);
        store_block128(destination, first, 1, 0, lane);
        tile_store_rows(tiles[1], first, lane);
        store_zero_block128(destination, 0, 1, lane);
    } else if (warp == 2) {
        load_block128(second, source, 3, 2, lane);
        tile_load_rows(col_stage, tiles[0], lane);
        solve_tile32_rows(first, col_stage);
        store_block128(destination, first, 2, 0, lane);
        tile_store_rows(tiles[2], first, lane);
        load_block128(first, source, 3, 1, lane);
        store_zero_block128(destination, 0, 2, lane);
    } else {
        load_block128(second, source, 1, 1, lane);
        tile_load_rows(col_stage, tiles[0], lane);
        solve_tile32_rows(first, col_stage);
        store_block128(destination, first, 3, 0, lane);
        tile_store_rows(tiles[3], first, lane);
        store_zero_block128(destination, 0, 3, lane);
    }
    __syncthreads();

    // Phase 2: step-0 rank-32 trailing updates, ascending k throughout.
    if (warp == 0) {
        tile_load_rows(row_stage, tiles[2], lane);
        tile_load_rows(col_stage, tiles[1], lane);
        update_tile32_rows(first, row_stage, col_stage);
        tile_load_rows(row_stage, tiles[3], lane);
        update_tile32_rows(second, row_stage, row_stage);
    } else if (warp == 1) {
        tile_load_rows(row_stage, tiles[2], lane);
        update_tile32_rows(second, row_stage, row_stage);
    } else if (warp == 2) {
        tile_load_rows(row_stage, tiles[3], lane);
        tile_load_rows(col_stage, tiles[1], lane);
        update_tile32_rows(first, row_stage, col_stage);
        tile_load_rows(col_stage, tiles[2], lane);
        update_tile32_rows(second, row_stage, col_stage);
    } else {
        tile_load_rows(row_stage, tiles[1], lane);
        update_tile32_rows(second, row_stage, row_stage);
    }
    __syncthreads();

    // Phase 3: step-1 diagonal factor.
    if (warp == 3) {
        factor_tile32_rows(second, lane);
        store_diag_block128(destination, second, 1, 1, lane);
        tile_store_rows(tiles[1], second, lane);
    }
    __syncthreads();

    // Phase 4: step-1 panel solves.
    if (warp == 0) {
        tile_load_rows(col_stage, tiles[1], lane);
        solve_tile32_rows(first, col_stage);
        store_block128(destination, first, 2, 1, lane);
        tile_store_rows(tiles[2], first, lane);
    } else if (warp == 2) {
        tile_load_rows(col_stage, tiles[1], lane);
        solve_tile32_rows(first, col_stage);
        store_block128(destination, first, 3, 1, lane);
        tile_store_rows(tiles[3], first, lane);
    }
    __syncthreads();

    // Phase 5: step-1 rank-32 trailing updates; warp 2 reuses its own
    // register-resident L(3,1) rows as the row factor.
    if (warp == 0) {
        tile_load_rows(row_stage, tiles[3], lane);
        update_tile32_rows(second, row_stage, row_stage);
    } else if (warp == 1) {
        tile_load_rows(row_stage, tiles[2], lane);
        update_tile32_rows(second, row_stage, row_stage);
    } else if (warp == 2) {
        tile_load_rows(col_stage, tiles[2], lane);
        update_tile32_rows(second, first, col_stage);
    }
    __syncthreads();

    // Phase 6: step-2 diagonal factor.
    if (warp == 1) {
        factor_tile32_rows(second, lane);
        store_diag_block128(destination, second, 2, 2, lane);
        tile_store_rows(tiles[2], second, lane);
    }
    __syncthreads();

    // Phase 7: step-2 panel solve.
    if (warp == 2) {
        tile_load_rows(col_stage, tiles[2], lane);
        solve_tile32_rows(second, col_stage);
        store_block128(destination, second, 3, 2, lane);
        tile_store_rows(tiles[3], second, lane);
    }
    __syncthreads();

    // Phase 8: final trailing update and diagonal factor.
    if (warp == 0) {
        tile_load_rows(row_stage, tiles[3], lane);
        update_tile32_rows(second, row_stage, row_stage);
        factor_tile32_rows(second, lane);
        store_diag_block128(destination, second, 3, 3, lane);
    }
}
}

torch::Tensor cholesky32(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(
        input.size(1) == kSmallOrder32 && input.size(2) == kSmallOrder32,
        "input matrices must be 32x32"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int blocks = (batch + kMatrixWarps32 - 1) / kMatrixWarps32;

    cholesky32_warp_kernel<<<blocks, kMatrixWarps32 * kLaneCount>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>(),
        batch
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n32 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}

torch::Tensor cholesky64(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(
        input.size(1) == kSmallOrder64 && input.size(2) == kSmallOrder64,
        "input matrices must be 64x64"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));

    cholesky64_warp_kernel<<<batch, kLaneCount>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>()
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n64 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}

torch::Tensor cholesky128(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
    TORCH_CHECK(
        input.dim() == 3 && input.size(0) > 0,
        "input must be a non-empty batch of matrices"
    );
    TORCH_CHECK(
        input.size(1) == kSmallOrder128 && input.size(2) == kSmallOrder128,
        "input matrices must be 128x128"
    );
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
    TORCH_CHECK(
        input.size(0) <= std::numeric_limits<int>::max(),
        "batch dimension exceeds CUDA grid range"
    );

    c10::cuda::CUDAGuard device_guard(input.device());
    torch::Tensor output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));

    cholesky128_warp_kernel<<<batch, kBlocks128 * kLaneCount>>>(
        input.data_ptr<float>(),
        output.data_ptr<float>()
    );
    const cudaError_t cuda_status = cudaGetLastError();
    TORCH_CHECK(
        cuda_status == cudaSuccess,
        "n128 CUDA operation failed: ",
        cudaGetErrorString(cuda_status)
    );
    return output;
}
"""


_schur_extension = load_inline(
    name="cholesky_native_exp242b_coalesced_n32n64",
    cpp_sources=[_CPP_SOURCE],
    cuda_sources=[
        _CUDA_SOURCE,
        _CUDA_1024_SOURCE,
        _CUDA_1024_D64_SOURCE,
        _CUDA_2048_SOURCE,
        _CUDA_SMALL_SOURCE,
    ],
    functions=[
        "cholesky512",
        "cholesky1024",
        "cholesky1024_d64",
        "cholesky2048",
        "cholesky32",
        "cholesky64",
        "cholesky128",
        "schur_update",
        "schur_update_tf32x2",
        "triangular_solve_in_place",
        "split_tf32_scaled",
        "split_tf32_packed65",
        "split_tf32_packed_residual",
        "split_packed65_fp16",
        "schur_update_fp16",
        "triangular_solve_inverse_in_place",
    ],
    extra_cuda_cflags=["-O3"],
    extra_ldflags=["-lcublas"],
    verbose=False,
)

_NATIVE_MATRIX_KERNELS = {
    128: _schur_extension.cholesky128,
    512: _schur_extension.cholesky512,
    1024: _schur_extension.cholesky1024,
    2048: _schur_extension.cholesky2048,
}
_NATIVE_BATCH_THRESHOLDS = {
    128: 1,
    512: 16,
    1024: 4,
    2048: 8,
}
_NATIVE_SMALL_BATCH_KERNELS = {
    1024: _schur_extension.cholesky1024_d64,
}
_NATIVE_SMALL_BATCH_LIMITS = {
    1024: 4,
}



@triton.jit
def _cholesky32_left_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):
        column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
        diagonal = tl.sum(
            tl.where(row_ids == k, column, 0.0),
            axis=0,
        )
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
        column = tl.where(row_ids >= k, column / diagonal, 0.0)
        values = tl.where((rows >= k) & (cols == k), column[:, None], values)
        update = column[:, None] * column[None, :]
        trailing_lower = (rows > k) & (cols > k) & (rows >= cols)
        values = tl.where(trailing_lower, values - update, values)

    tl.store(output_ptr + offsets, values)


@triton.jit
def _cholesky64_blocked32_kernel(
    input_ptr,
    output_ptr,
    matrix_stride: tl.constexpr,
):
    matrix = tl.program_id(0)
    ids = tl.arange(0, 32)
    rows = ids[:, None]
    cols = ids[None, :]
    base = matrix * matrix_stride

    offsets11 = base + rows * 64 + cols
    l11 = tl.where(rows >= cols, tl.load(input_ptr + offsets11), 0.0)
    for k in range(32):
        column = tl.sum(tl.where(cols == k, l11, 0.0), axis=1)
        diagonal = tl.sum(tl.where(ids == k, column, 0.0), axis=0)
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
        column = tl.where(ids >= k, column / diagonal, 0.0)
        l11 = tl.where((rows >= k) & (cols == k), column[:, None], l11)
        update = column[:, None] * column[None, :]
        trailing_lower = (rows > k) & (cols > k) & (rows >= cols)
        l11 = tl.where(trailing_lower, l11 - update, l11)

    offsets21 = base + (rows + 32) * 64 + cols
    l21 = tl.load(input_ptr + offsets21)
    for k in range(32):
        factor_column = tl.sum(tl.where(cols == k, l11, 0.0), axis=1)
        diagonal = tl.sum(tl.where(ids == k, factor_column, 0.0), axis=0)
        column = tl.sum(tl.where(cols == k, l21, 0.0), axis=1) / diagonal
        l21 = tl.where(cols == k, column[:, None], l21)
        update = column[:, None] * factor_column[None, :]
        l21 = tl.where(cols > k, l21 - update, l21)

    offsets22 = base + (rows + 32) * 64 + (cols + 32)
    l22 = tl.where(rows >= cols, tl.load(input_ptr + offsets22), 0.0)
    for k in range(32):
        column = tl.sum(tl.where(cols == k, l21, 0.0), axis=1)
        update = column[:, None] * column[None, :]
        l22 = tl.where(rows >= cols, l22 - update, l22)

    for k in range(32):
        column = tl.sum(tl.where(cols == k, l22, 0.0), axis=1)
        diagonal = tl.sum(tl.where(ids == k, column, 0.0), axis=0)
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
        column = tl.where(ids >= k, column / diagonal, 0.0)
        l22 = tl.where((rows >= k) & (cols == k), column[:, None], l22)
        update = column[:, None] * column[None, :]
        trailing_lower = (rows > k) & (cols > k) & (rows >= cols)
        l22 = tl.where(trailing_lower, l22 - update, l22)

    offsets12 = base + rows * 64 + (cols + 32)
    tl.store(output_ptr + offsets11, l11)
    tl.store(output_ptr + offsets12, 0.0)
    tl.store(output_ptr + offsets21, l21)
    tl.store(output_ptr + offsets22, l22)


def _blocked_cholesky(data: torch.Tensor, block_size: int) -> torch.Tensor:
    work = data.clone()
    n = data.shape[-1]
    for start in range(0, n, block_size):
        stop = min(start + block_size, n)
        diagonal = work[:, start:stop, start:stop]
        factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
        diagonal.copy_(factor)

        if stop == n:
            continue

        right_hand_side = work[:, stop:, start:stop].transpose(-1, -2)
        solved = torch.linalg.solve_triangular(
            factor,
            right_hand_side,
            upper=False,
            left=True,
        )
        column = solved.transpose(-1, -2)
        work[:, stop:, start:stop].copy_(column)
        column = work[:, stop:, start:stop]
        trailing = work[:, stop:, stop:]
        _schur_extension.schur_update(
            trailing,
            column,
            column,
            -1.0,
            False,
        )

    return work.tril_()


def _blocked_cholesky_lower_tiles(data: torch.Tensor, block_size: int) -> torch.Tensor:
    work = data.clone()
    n = data.shape[-1]
    for start in range(0, n, block_size):
        stop = min(start + block_size, n)
        diagonal = work[:, start:stop, start:stop]
        factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
        diagonal.copy_(factor)

        if stop == n:
            continue

        right_hand_side = work[:, stop:, start:stop].transpose(-1, -2)
        solved = torch.linalg.solve_triangular(
            factor,
            right_hand_side,
            upper=False,
            left=True,
        )
        column = solved.transpose(-1, -2)
        work[:, stop:, start:stop].copy_(column)
        column = work[:, stop:, start:stop]
        trailing = work[:, stop:, stop:]
        trailing_size = n - stop
        for row_tile_start in range(0, trailing_size, 1024):
            row_tile_stop = min(row_tile_start + 1024, trailing_size)
            row_factor = column[:, row_tile_start:row_tile_stop, :]
            for col_tile_start in range(0, row_tile_start + 1, 1024):
                col_tile_stop = min(col_tile_start + 1024, trailing_size)
                target = trailing[
                    :, row_tile_start:row_tile_stop, col_tile_start:col_tile_stop
                ]
                col_factor = column[:, col_tile_start:col_tile_stop, :]
                _schur_extension.schur_update(
                    target,
                    row_factor,
                    col_factor,
                    -1.0,
                    False,
                )

    return work.tril_()


def _factor_solve_pack_left512_s65(
    work: torch.Tensor,
    start: int,
    use_external_inverse: bool,
    factor_info: torch.Tensor,
) -> torch.Tensor:
    n = work.shape[-1]
    panel_width = 2048
    leaf_size = 512
    packed_alpha = -64.0 / 65.0
    packed = torch.empty(
        (1, n - start, 2 * panel_width),
        dtype=work.dtype,
        device=work.device,
    )
    identity = None
    inverse = None
    inverse_packed = None
    external_solved = None
    if use_external_inverse:
        identity = torch.eye(
            leaf_size,
            dtype=work.dtype,
            device=work.device,
        ).unsqueeze(0)
        inverse = torch.empty_like(identity)
        if start + panel_width < n:
            inverse_packed = torch.empty(
                (1, leaf_size, 2 * leaf_size),
                dtype=work.dtype,
                device=work.device,
            )
            external_solved = torch.empty(
                (1, n - start - panel_width, leaf_size),
                dtype=work.dtype,
                device=work.device,
            )
    for p in range(0, panel_width, leaf_size):
        q = p + leaf_size
        if p > 0:
            panel_block_column = work[
                :,
                start + p : start + panel_width,
                start + p : start + q,
            ]
            panel_prior_row = work[
                :,
                start + p : start + panel_width,
                start : start + p,
            ]
            panel_prior_col = work[
                :, start + p : start + q, start : start + p
            ]
            _schur_extension.schur_update(
                panel_block_column,
                panel_prior_row,
                panel_prior_col,
                -1.0,
                False,
            )

            if start + panel_width < n:
                external_block_column = work[
                    :,
                    start + panel_width :,
                    start + p : start + q,
                ]
                external_prior_row_packed = packed[
                    :, panel_width:, : 2 * p
                ]
                external_prior_col_packed = packed[:, p:q, : 2 * p]
                for prior_chunk_start in range(
                    0,
                    2 * p,
                    _GCERT_A_UPDATE_K_CHUNK,
                ):
                    prior_chunk_stop = (
                        prior_chunk_start + _GCERT_A_UPDATE_K_CHUNK
                    )
                    _schur_extension.schur_update(
                        external_block_column,
                        external_prior_row_packed[
                            :, :, prior_chunk_start:prior_chunk_stop
                        ],
                        external_prior_col_packed[
                            :, :, prior_chunk_start:prior_chunk_stop
                        ],
                        packed_alpha,
                        True,
                    )

        diagonal = work[
            :,
            start + p : start + q,
            start + p : start + q,
        ]
        leaf_result = torch.linalg.cholesky_ex(
            diagonal,
            check_errors=False,
        )
        factor_info.add_(leaf_result.info.abs().sum())
        diagonal.copy_(leaf_result.L)

        if not use_external_inverse:
            if start + q == n:
                continue
            below = work[:, start + q :, start + p : start + q]
            _schur_extension.triangular_solve_in_place(
                diagonal,
                below,
            )
            packed_destination = packed[:, q:, 2 * p : 2 * q]
            _schur_extension.split_tf32_packed65(
                below,
                packed_destination,
            )
            continue

        if q < panel_width:
            strict_below = work[
                :,
                start + q : start + panel_width,
                start + p : start + q,
            ]
            _schur_extension.triangular_solve_in_place(
                diagonal,
                strict_below,
            )
            strict_packed_destination = packed[
                :,
                q:panel_width,
                2 * p : 2 * q,
            ]
            _schur_extension.split_tf32_packed65(
                strict_below,
                strict_packed_destination,
            )

        if start + panel_width < n:
            external_rhs = work[
                :,
                start + panel_width :,
                start + p : start + q,
            ]
            inverse.copy_(identity)
            _schur_extension.triangular_solve_inverse_in_place(
                diagonal,
                inverse,
            )
            external_packed_slot = packed[
                :,
                panel_width:,
                2 * p : 2 * q,
            ]
            _schur_extension.split_tf32_packed_residual(
                external_rhs,
                external_packed_slot,
            )
            _schur_extension.split_tf32_packed_residual(
                inverse,
                inverse_packed,
            )
            external_rhs_high = external_packed_slot[
                :, :, :leaf_size
            ]
            external_rhs_residual = external_packed_slot[
                :, :, leaf_size:
            ]
            inverse_high = inverse_packed[:, :, :leaf_size]
            inverse_residual = inverse_packed[:, :, leaf_size:]
            external_solved.zero_()
            _schur_extension.schur_update(
                external_solved,
                external_rhs_high,
                inverse_high,
                1.0,
                True,
            )
            _schur_extension.schur_update(
                external_solved,
                external_rhs_high,
                inverse_residual,
                1.0,
                True,
            )
            _schur_extension.schur_update(
                external_solved,
                external_rhs_residual,
                inverse_high,
                1.0,
                True,
            )
            external_rhs.copy_(external_solved)
            _schur_extension.split_tf32_packed65(
                external_solved,
                external_packed_slot,
            )

    return packed[:, panel_width:, :]


def _blocked_cholesky_cublas_packed_s65_lower_tiles(
    data: torch.Tensor,
    block_size: int,
    use_external_inverse: bool,
) -> tuple:
    work = data.clone()
    factor_info = torch.zeros((), dtype=torch.int64, device=data.device)
    n = data.shape[-1]
    update_tile_size = 6144
    diagonal_tile_size = 3072
    packed_alpha = -64.0 / 65.0
    for start in range(0, n, block_size):
        stop = min(start + block_size, n)
        packed = _factor_solve_pack_left512_s65(
            work,
            start,
            use_external_inverse,
            factor_info,
        )

        if stop == n:
            continue

        trailing = work[:, stop:, stop:]
        trailing_size = n - stop
        for row_tile_start in range(0, trailing_size, update_tile_size):
            row_tile_stop = min(
                row_tile_start + update_tile_size,
                trailing_size,
            )
            row_packed = packed[:, row_tile_start:row_tile_stop, :]
            for col_tile_start in range(
                0,
                row_tile_start,
                update_tile_size,
            ):
                col_tile_stop = min(
                    col_tile_start + update_tile_size,
                    trailing_size,
                )
                target = trailing[
                    :,
                    row_tile_start:row_tile_stop,
                    col_tile_start:col_tile_stop,
                ]
                col_packed = packed[:, col_tile_start:col_tile_stop, :]
                for chunk_start in range(
                    0,
                    2 * block_size,
                    _GCERT_A_UPDATE_K_CHUNK,
                ):
                    chunk_stop = chunk_start + _GCERT_A_UPDATE_K_CHUNK
                    _schur_extension.schur_update(
                        target,
                        row_packed[:, :, chunk_start:chunk_stop],
                        col_packed[:, :, chunk_start:chunk_stop],
                        packed_alpha,
                        True,
                    )

            for diagonal_row_start in range(
                row_tile_start,
                row_tile_stop,
                diagonal_tile_size,
            ):
                diagonal_row_stop = min(
                    diagonal_row_start + diagonal_tile_size,
                    row_tile_stop,
                )
                diagonal_row_packed = packed[
                    :,
                    diagonal_row_start:diagonal_row_stop,
                    :,
                ]
                for diagonal_col_start in range(
                    row_tile_start,
                    diagonal_row_start + 1,
                    diagonal_tile_size,
                ):
                    diagonal_col_stop = min(
                        diagonal_col_start + diagonal_tile_size,
                        row_tile_stop,
                    )
                    diagonal_target = trailing[
                        :,
                        diagonal_row_start:diagonal_row_stop,
                        diagonal_col_start:diagonal_col_stop,
                    ]
                    diagonal_col_packed = packed[
                        :,
                        diagonal_col_start:diagonal_col_stop,
                        :,
                    ]
                    for chunk_start in range(
                        0,
                        2 * block_size,
                        _GCERT_A_UPDATE_K_CHUNK,
                    ):
                        chunk_stop = (
                            chunk_start + _GCERT_A_UPDATE_K_CHUNK
                        )
                        _schur_extension.schur_update(
                            diagonal_target,
                            diagonal_row_packed[
                                :, :, chunk_start:chunk_stop
                            ],
                            diagonal_col_packed[
                                :, :, chunk_start:chunk_stop
                            ],
                            packed_alpha,
                            True,
                        )

    return work.tril_(), factor_info


def _factor_solve_pack_left512_s65_fp16(
    work: torch.Tensor,
    start: int,
    factor_info: torch.Tensor,
) -> torch.Tensor:
    n = work.shape[-1]
    panel_width = 2048
    leaf_size = 512
    packed_alpha = -64.0 / 65.0
    packed = torch.empty(
        (1, n - start, 2 * panel_width),
        dtype=torch.float16,
        device=work.device,
    )
    for p in range(0, panel_width, leaf_size):
        q = p + leaf_size
        if p > 0:
            panel_block_column = work[
                :,
                start + p : start + panel_width,
                start + p : start + q,
            ]
            panel_prior_row = work[
                :,
                start + p : start + panel_width,
                start : start + p,
            ]
            panel_prior_col = work[
                :, start + p : start + q, start : start + p
            ]
            _schur_extension.schur_update(
                panel_block_column,
                panel_prior_row,
                panel_prior_col,
                -1.0,
                False,
            )

            if start + panel_width < n:
                external_block_column = work[
                    :,
                    start + panel_width :,
                    start + p : start + q,
                ]
                external_prior_row_packed = packed[
                    :, panel_width:, : 2 * p
                ]
                external_prior_col_packed = packed[:, p:q, : 2 * p]
                for prior_chunk_start in range(
                    0,
                    2 * p,
                    _GCERT_A_UPDATE_K_CHUNK,
                ):
                    prior_chunk_stop = (
                        prior_chunk_start + _GCERT_A_UPDATE_K_CHUNK
                    )
                    _schur_extension.schur_update_fp16(
                        external_block_column,
                        external_prior_row_packed[
                            :, :, prior_chunk_start:prior_chunk_stop
                        ],
                        external_prior_col_packed[
                            :, :, prior_chunk_start:prior_chunk_stop
                        ],
                        packed_alpha,
                    )

        diagonal = work[
            :,
            start + p : start + q,
            start + p : start + q,
        ]
        leaf_result = torch.linalg.cholesky_ex(
            diagonal,
            check_errors=False,
        )
        factor_info.add_(leaf_result.info.abs().sum())
        diagonal.copy_(leaf_result.L)

        if start + q == n:
            continue
        below = work[:, start + q :, start + p : start + q]
        _schur_extension.triangular_solve_in_place(
            diagonal,
            below,
        )
        packed_destination = packed[:, q:, 2 * p : 2 * q]
        _schur_extension.split_packed65_fp16(
            below,
            packed_destination,
        )

    return packed[:, panel_width:, :]


def _blocked_cholesky_cublas_packed_s65_fp16_lower_tiles(
    data: torch.Tensor,
    block_size: int,
) -> tuple:
    work = data.clone()
    factor_info = torch.zeros((), dtype=torch.int64, device=data.device)
    n = data.shape[-1]
    update_tile_size = 6144
    diagonal_tile_size = 3072
    packed_alpha = -64.0 / 65.0
    for start in range(0, n, block_size):
        stop = min(start + block_size, n)
        packed = _factor_solve_pack_left512_s65_fp16(
            work,
            start,
            factor_info,
        )

        if stop == n:
            continue

        trailing = work[:, stop:, stop:]
        trailing_size = n - stop
        for row_tile_start in range(0, trailing_size, update_tile_size):
            row_tile_stop = min(
                row_tile_start + update_tile_size,
                trailing_size,
            )
            row_packed = packed[:, row_tile_start:row_tile_stop, :]
            for col_tile_start in range(
                0,
                row_tile_start,
                update_tile_size,
            ):
                col_tile_stop = min(
                    col_tile_start + update_tile_size,
                    trailing_size,
                )
                target = trailing[
                    :,
                    row_tile_start:row_tile_stop,
                    col_tile_start:col_tile_stop,
                ]
                col_packed = packed[:, col_tile_start:col_tile_stop, :]
                for chunk_start in range(
                    0,
                    2 * block_size,
                    _GCERT_A_UPDATE_K_CHUNK,
                ):
                    chunk_stop = chunk_start + _GCERT_A_UPDATE_K_CHUNK
                    _schur_extension.schur_update_fp16(
                        target,
                        row_packed[:, :, chunk_start:chunk_stop],
                        col_packed[:, :, chunk_start:chunk_stop],
                        packed_alpha,
                    )

            for diagonal_row_start in range(
                row_tile_start,
                row_tile_stop,
                diagonal_tile_size,
            ):
                diagonal_row_stop = min(
                    diagonal_row_start + diagonal_tile_size,
                    row_tile_stop,
                )
                diagonal_row_packed = packed[
                    :,
                    diagonal_row_start:diagonal_row_stop,
                    :,
                ]
                for diagonal_col_start in range(
                    row_tile_start,
                    diagonal_row_start + 1,
                    diagonal_tile_size,
                ):
                    diagonal_col_stop = min(
                        diagonal_col_start + diagonal_tile_size,
                        row_tile_stop,
                    )
                    diagonal_target = trailing[
                        :,
                        diagonal_row_start:diagonal_row_stop,
                        diagonal_col_start:diagonal_col_stop,
                    ]
                    diagonal_col_packed = packed[
                        :,
                        diagonal_col_start:diagonal_col_stop,
                        :,
                    ]
                    for chunk_start in range(
                        0,
                        2 * block_size,
                        _GCERT_A_UPDATE_K_CHUNK,
                    ):
                        chunk_stop = (
                            chunk_start + _GCERT_A_UPDATE_K_CHUNK
                        )
                        _schur_extension.schur_update_fp16(
                            diagonal_target,
                            diagonal_row_packed[
                                :, :, chunk_start:chunk_stop
                            ],
                            diagonal_col_packed[
                                :, :, chunk_start:chunk_stop
                            ],
                            packed_alpha,
                        )

    return work.tril_(), factor_info


def _blocked_cholesky_cublas_tf32x2_lower_tiles(
    data: torch.Tensor,
    block_size: int,
) -> tuple:
    work = data.clone()
    factor_info = torch.zeros((), dtype=torch.int64, device=data.device)
    n = data.shape[-1]
    update_tile_size = 6144
    diagonal_tile_size = 3072
    for start in range(0, n, block_size):
        stop = min(start + block_size, n)
        diagonal = work[:, start:stop, start:stop]
        panel_result = torch.linalg.cholesky_ex(diagonal, check_errors=False)
        factor_info.add_(panel_result.info.abs().sum())
        factor = panel_result.L
        diagonal.copy_(factor)

        if stop == n:
            continue

        right_hand_side = work[:, stop:, start:stop].transpose(-1, -2)
        solved = torch.linalg.solve_triangular(
            factor,
            right_hand_side,
            upper=False,
            left=True,
        )
        column = solved.transpose(-1, -2)
        work[:, stop:, start:stop].copy_(column)
        column = work[:, stop:, start:stop]
        column_high = torch.empty(
            column.shape,
            dtype=column.dtype,
            device=column.device,
        )
        column_expanded = torch.empty(
            column.shape,
            dtype=column.dtype,
            device=column.device,
        )
        _schur_extension.split_tf32_scaled(
            column,
            column_high,
            column_expanded,
        )

        trailing = work[:, stop:, stop:]
        trailing_size = n - stop
        for row_tile_start in range(0, trailing_size, update_tile_size):
            row_tile_stop = min(
                row_tile_start + update_tile_size,
                trailing_size,
            )
            row_high = column_high[:, row_tile_start:row_tile_stop, :]
            row_expanded = column_expanded[:, row_tile_start:row_tile_stop, :]
            for col_tile_start in range(
                0,
                row_tile_start,
                update_tile_size,
            ):
                col_tile_stop = min(
                    col_tile_start + update_tile_size,
                    trailing_size,
                )
                target = trailing[
                    :,
                    row_tile_start:row_tile_stop,
                    col_tile_start:col_tile_stop,
                ]
                col_high = column_high[:, col_tile_start:col_tile_stop, :]
                col_expanded = column_expanded[
                    :,
                    col_tile_start:col_tile_stop,
                    :,
                ]
                for chunk_start in range(
                    0,
                    stop - start,
                    _GCERT_B_UPDATE_K_CHUNK,
                ):
                    chunk_stop = chunk_start + _GCERT_B_UPDATE_K_CHUNK
                    _schur_extension.schur_update_tf32x2(
                        target,
                        row_high[:, :, chunk_start:chunk_stop],
                        row_expanded[:, :, chunk_start:chunk_stop],
                        col_high[:, :, chunk_start:chunk_stop],
                        col_expanded[:, :, chunk_start:chunk_stop],
                    )

            for diagonal_row_start in range(
                row_tile_start,
                row_tile_stop,
                diagonal_tile_size,
            ):
                diagonal_row_stop = min(
                    diagonal_row_start + diagonal_tile_size,
                    row_tile_stop,
                )
                diagonal_row_high = column_high[
                    :,
                    diagonal_row_start:diagonal_row_stop,
                    :,
                ]
                diagonal_row_expanded = column_expanded[
                    :,
                    diagonal_row_start:diagonal_row_stop,
                    :,
                ]
                for diagonal_col_start in range(
                    row_tile_start,
                    diagonal_row_start + 1,
                    diagonal_tile_size,
                ):
                    diagonal_col_stop = min(
                        diagonal_col_start + diagonal_tile_size,
                        row_tile_stop,
                    )
                    diagonal_target = trailing[
                        :,
                        diagonal_row_start:diagonal_row_stop,
                        diagonal_col_start:diagonal_col_stop,
                    ]
                    diagonal_col_high = column_high[
                        :,
                        diagonal_col_start:diagonal_col_stop,
                        :,
                    ]
                    diagonal_col_expanded = column_expanded[
                        :,
                        diagonal_col_start:diagonal_col_stop,
                        :,
                    ]
                    for chunk_start in range(
                        0,
                        stop - start,
                        _GCERT_B_UPDATE_K_CHUNK,
                    ):
                        chunk_stop = (
                            chunk_start + _GCERT_B_UPDATE_K_CHUNK
                        )
                        _schur_extension.schur_update_tf32x2(
                            diagonal_target,
                            diagonal_row_high[
                                :, :, chunk_start:chunk_stop
                            ],
                            diagonal_row_expanded[
                                :, :, chunk_start:chunk_stop
                            ],
                            diagonal_col_high[
                                :, :, chunk_start:chunk_stop
                            ],
                            diagonal_col_expanded[
                                :, :, chunk_start:chunk_stop
                            ],
                        )

    return work.tril_(), factor_info


_GCERT_SCAN_TILE = 2048
_GCERT_UNIT = 2.0 ** -24
_GCERT_EPS32 = 2.0 ** -23
_GCERT_RTOL_FACTOR = 20.0
_GCERT_TOL_HAIRCUT = 1.0 - 4e-6
_GCERT_GROWTH_INFLATE = 1.0 + 2e-3

_GCERT_A_PANEL_WIDTH = 512
_GCERT_A_UPDATE_K_CHUNK = 1024
_GCERT_A_TERM_BIAS = 3.15e-5
_GCERT_A_ACCUM_COEFF = 2.0
_GCERT_A_ACCUM_CHAIN = 1024.0
_GCERT_A_ACCUM_EXCESS = 1.02
_GCERT_A_MERGE_COEFF = 2.0
_GCERT_A_MERGE_COUNT = 63.0
_GCERT_A_MERGE_EXCESS = 2.05
_GCERT_A_SOLVE_TERM = 0.0
_GCERT_A_CHOL_TERM = 1.221448347e-4
_GCERT_A_HARDENING = 1.05

_GCERT_FP16_Q_UNIT = 7.5e-8
_GCERT_FP16_CLAMP_Q0 = 2.0 ** -25
_GCERT_FP16_CLAMP_COEFF = 64.0
_GCERT_FP16_FLOOR_HARDENING = 1.05
_GCERT_FP16_MAX_ABS = 2.0 ** 14

_GCERT_B_PANEL_WIDTH = 2048
_GCERT_B_UPDATE_K_CHUNK = 512
_GCERT_B_TERM_BIAS = 3.15e-5
_GCERT_B_ACCUM_COEFF = 2.0
_GCERT_B_ACCUM_CHAIN = 512.0
_GCERT_B_ACCUM_EXCESS = 1.02
_GCERT_B_MERGE_COEFF = 2.0
_GCERT_B_MERGE_COUNT = 56.0
_GCERT_B_MERGE_EXCESS = 2.05
_GCERT_B_SOLVE_TERM = 0.0
_GCERT_B_CHOL_TERM = 1.221448347e-4
_GCERT_B_HARDENING = 1.05


def _growth_certificate_coefficient(
    term_bias: float,
    accum_coeff: float,
    accum_chain: float,
    accum_excess: float,
    merge_coeff: float,
    merge_count: float,
    merge_excess: float,
    solve_term: float,
    chol_term: float,
    hardening: float,
) -> float:
    return (
        term_bias
        + accum_coeff * _GCERT_UNIT * accum_chain * accum_excess
        + merge_coeff * _GCERT_UNIT * merge_count * merge_excess
        + solve_term
        + chol_term
    ) * hardening


_GCERT_A_COEFFICIENT = _growth_certificate_coefficient(
    _GCERT_A_TERM_BIAS,
    _GCERT_A_ACCUM_COEFF,
    _GCERT_A_ACCUM_CHAIN,
    _GCERT_A_ACCUM_EXCESS,
    _GCERT_A_MERGE_COEFF,
    _GCERT_A_MERGE_COUNT,
    _GCERT_A_MERGE_EXCESS,
    _GCERT_A_SOLVE_TERM,
    _GCERT_A_CHOL_TERM,
    _GCERT_A_HARDENING,
)
_GCERT_B_COEFFICIENT = _growth_certificate_coefficient(
    _GCERT_B_TERM_BIAS,
    _GCERT_B_ACCUM_COEFF,
    _GCERT_B_ACCUM_CHAIN,
    _GCERT_B_ACCUM_EXCESS,
    _GCERT_B_MERGE_COEFF,
    _GCERT_B_MERGE_COUNT,
    _GCERT_B_MERGE_EXCESS,
    _GCERT_B_SOLVE_TERM,
    _GCERT_B_CHOL_TERM,
    _GCERT_B_HARDENING,
)


def _growth_certificate_accepts(
    data: torch.Tensor,
    factor: torch.Tensor,
    panel_width: int,
    coefficient: float,
) -> bool:
    n = data.shape[-1]
    source = data[0]
    matrix = factor[0]

    input_column_sums = torch.zeros(
        n, dtype=torch.float64, device=data.device
    )
    for start in range(0, n, _GCERT_SCAN_TILE):
        stop = min(start + _GCERT_SCAN_TILE, n)
        input_column_sums[start:stop] = source[:, start:stop].abs().sum(
            dim=0, dtype=torch.float64
        )
    allowance = (
        _GCERT_RTOL_FACTOR
        * n
        * _GCERT_EPS32
        * float(input_column_sums.max().item())
        * _GCERT_TOL_HAIRCUT
    )

    panel_count = n // panel_width
    panel_terms = torch.zeros(
        panel_count, dtype=torch.float64, device=data.device
    )
    for index in range(panel_count):
        start = index * panel_width
        stop = start + panel_width
        panel_block = matrix[start:, start:stop].abs()
        panel_terms[index] = panel_block.sum(
            dim=0, dtype=torch.float64
        ).max() * panel_block.sum(dim=1, dtype=torch.float64).max()
    growth = float(panel_terms.sum().item()) * _GCERT_GROWTH_INFLATE
    certified = coefficient * growth
    return bool(certified <= allowance)


def _growth_certificate_accepts_fp16(
    data: torch.Tensor,
    factor: torch.Tensor,
    panel_width: int,
    coefficient: float,
) -> bool:
    n = data.shape[-1]
    source = data[0]
    matrix = factor[0]

    if bool((matrix.abs().max() > _GCERT_FP16_MAX_ABS).item()):
        return False

    input_column_sums = torch.zeros(
        n, dtype=torch.float64, device=data.device
    )
    for start in range(0, n, _GCERT_SCAN_TILE):
        stop = min(start + _GCERT_SCAN_TILE, n)
        input_column_sums[start:stop] = source[:, start:stop].abs().sum(
            dim=0, dtype=torch.float64
        )
    allowance = (
        _GCERT_RTOL_FACTOR
        * n
        * _GCERT_EPS32
        * float(input_column_sums.max().item())
        * _GCERT_TOL_HAIRCUT
    )

    panel_count = n // panel_width
    panel_terms = torch.zeros(
        panel_count, dtype=torch.float64, device=data.device
    )
    q_terms = torch.zeros(
        panel_count, dtype=torch.float64, device=data.device
    )
    clamp_terms = torch.zeros(
        panel_count, dtype=torch.float64, device=data.device
    )
    for index in range(panel_count):
        start = index * panel_width
        stop = start + panel_width
        panel_block = matrix[start:, start:stop].abs()
        column_sums = panel_block.sum(dim=0, dtype=torch.float64)
        row_sums = panel_block.sum(dim=1, dtype=torch.float64)
        panel_terms[index] = column_sums.max() * row_sums.max()
        q_terms[index] = column_sums.sum() + (
            panel_block.shape[0] * row_sums.max()
        )
        clamped = panel_block.clamp(max=_GCERT_FP16_CLAMP_Q0)
        clamp_terms[index] = clamped.sum(
            dim=0, dtype=torch.float64
        ).max() * clamped.sum(dim=1, dtype=torch.float64).max()
    growth = float(panel_terms.sum().item()) * _GCERT_GROWTH_INFLATE
    q_sum = float(q_terms.sum().item())
    gclamp_sum = float(clamp_terms.sum().item())
    floor = (
        _GCERT_FP16_FLOOR_HARDENING
        * (
            _GCERT_FP16_Q_UNIT * q_sum
            + _GCERT_FP16_CLAMP_COEFF * gclamp_sum
        )
        * _GCERT_GROWTH_INFLATE
    )
    certified = coefficient * growth + floor
    return bool(certified <= allowance)


def _growth_certified_leaf512_singleton(
    data: torch.Tensor,
) -> torch.Tensor:
    factor, factor_info = _blocked_cholesky_cublas_packed_s65_lower_tiles(
        data, 2048, False
    )
    diagonal = factor.diagonal(dim1=-2, dim2=-1)
    factor_valid = (
        bool(factor_info.item() == 0)
        and bool(torch.isfinite(diagonal).all().item())
        and bool((diagonal > 0).all().item())
    )
    if factor_valid and _growth_certificate_accepts(
        data,
        factor,
        _GCERT_A_PANEL_WIDTH,
        _GCERT_A_COEFFICIENT,
    ):
        return factor
    return torch.linalg.cholesky_ex(data, check_errors=False).L


def _growth_certified_leaf512_fp16_singleton(
    data: torch.Tensor,
) -> torch.Tensor:
    factor, factor_info = (
        _blocked_cholesky_cublas_packed_s65_fp16_lower_tiles(data, 2048)
    )
    diagonal = factor.diagonal(dim1=-2, dim2=-1)
    factor_valid = (
        bool(factor_info.item() == 0)
        and bool(torch.isfinite(diagonal).all().item())
        and bool((diagonal > 0).all().item())
    )
    if factor_valid and _growth_certificate_accepts_fp16(
        data,
        factor,
        _GCERT_A_PANEL_WIDTH,
        _GCERT_A_COEFFICIENT,
    ):
        return factor
    return torch.linalg.cholesky_ex(data, check_errors=False).L


def _growth_certified_panel2048_singleton(
    data: torch.Tensor,
) -> torch.Tensor:
    factor, factor_info = _blocked_cholesky_cublas_tf32x2_lower_tiles(
        data, 2048
    )
    diagonal = factor.diagonal(dim1=-2, dim2=-1)
    factor_valid = (
        bool(factor_info.item() == 0)
        and bool(torch.isfinite(diagonal).all().item())
        and bool((diagonal > 0).all().item())
    )
    if factor_valid and _growth_certificate_accepts(
        data,
        factor,
        _GCERT_B_PANEL_WIDTH,
        _GCERT_B_COEFFICIENT,
    ):
        return factor
    return torch.linalg.cholesky_ex(data, check_errors=False).L


_CERTIFIED_SINGLETON_DRIVERS = {
    (32768, 1): _growth_certified_leaf512_fp16_singleton,
    (16384, 1): _growth_certified_panel2048_singleton,
}


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n == 32:
        return _schur_extension.cholesky32(data)
    if n == 64:
        return _schur_extension.cholesky64(data)
    native_matrix_kernel = _NATIVE_MATRIX_KERNELS.get(n)
    native_batch_threshold = _NATIVE_BATCH_THRESHOLDS.get(n)
    if (
        native_matrix_kernel is not None
        and native_batch_threshold is not None
        and batch >= native_batch_threshold
    ):
        small_batch_kernel = _NATIVE_SMALL_BATCH_KERNELS.get(n)
        small_batch_limit = _NATIVE_SMALL_BATCH_LIMITS.get(n)
        if (
            small_batch_kernel is not None
            and small_batch_limit is not None
            and batch <= small_batch_limit
        ):
            return small_batch_kernel(data)
        return native_matrix_kernel(data)
    if n == 2048 and batch <= 2:
        return _blocked_cholesky(data, 512)
    if n == 4096 and batch > 1:
        return _blocked_cholesky_lower_tiles(data, 512)
    certified_driver = _CERTIFIED_SINGLETON_DRIVERS.get((n, batch))
    if certified_driver is not None:
        return certified_driver(data)
    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 4657 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