Skip to content
KernelIndex
Search⌘K

submission 910279

Clark Kitchen · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e511aa464729b81db6729e168d42cbf9222f3d19bb4e0d3cff8476ba369a2a72
license declaredunknown
license concludedunknown
authorsClark Kitchen
imported2026-08-26

Techniques

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

mmaupdate = tl.dot(left, tl.trans(right), input_precision="tf32")
num-warps = 8num_warps=8,
shared-memoryextern __shared__ float packed_storage[];
stages = 1num_stages=1,
tile-m = 16BLOCK_M=16,
tile-n = 256BLOCK_N=256,
vector-width = float4const float4* a4 = reinterpret_cast<const float4*>(a);

Kernel source

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

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

from task import input_t, output_t


CPP_SRC = r"""
torch::Tensor raw_cholesky(torch::Tensor input);
torch::Tensor copy_lower(torch::Tensor input);
void factor_panel(torch::Tensor output, int64_t panel_start);
void factor_panel_128(torch::Tensor output, int64_t panel_start);
void factor_panel_inverse(torch::Tensor output, int64_t panel_start);
// B640_FUSED_V1_BEGIN
void factor_panel_cutlass_inverse_b640(
    torch::Tensor output,
    torch::Tensor inverse,
    int64_t panel_start);
// B640_FUSED_V1_END
"""


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

#include <cfloat>
#include <cstdint>

namespace {

constexpr int kSmallMax = 128;
constexpr int kPanel = 64;
constexpr int kLargePanel = 128;
constexpr int kTile = 64;
constexpr int kThreads = 256;
constexpr int kLargeSolveThreads = 512;

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

// Small benchmark batches contain thousands of matrices. A warp owns one
// matrix so independent factorizations share a launch without sharing any
// synchronization. Vectorized full-row movement is safe because every
// benchmark size is a multiple of four and each matrix starts aligned.
template <int N, int MatricesPerCta>
__global__ void warp_packed_cholesky_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch,
    int64_t matrix_stride) {
    constexpr int kPacked = N * (N + 1) / 2;
    extern __shared__ float packed_storage[];

    const int lane = static_cast<int>(threadIdx.x) & 31;
    const int local_matrix = static_cast<int>(threadIdx.x) >> 5;
    const int matrix = static_cast<int>(blockIdx.x) * MatricesPerCta + local_matrix;
    if (matrix >= batch) {
        return;
    }

    float* lower = packed_storage + local_matrix * kPacked;
    const float* a = input + static_cast<int64_t>(matrix) * matrix_stride;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
    const float4* a4 = reinterpret_cast<const float4*>(a);
    float4* l4 = reinterpret_cast<float4*>(l);

    constexpr int kVectors = N * N / 4;
    for (int vector = lane; vector < kVectors; vector += 32) {
        const float4 values = a4[vector];
        const int linear = vector * 4;
        const int row = linear / N;
        const int col = linear - row * N;
        if (col <= row) {
            lower[packed_index(row, col)] = values.x;
        }
        if (col + 1 <= row) {
            lower[packed_index(row, col + 1)] = values.y;
        }
        if (col + 2 <= row) {
            lower[packed_index(row, col + 2)] = values.z;
        }
        if (col + 3 <= row) {
            lower[packed_index(row, col + 3)] = values.w;
        }
    }
    __syncwarp();

    for (int pivot = 0; pivot < N; ++pivot) {
        if (lane == 0) {
            const int diagonal_index = packed_index(pivot, pivot);
            lower[diagonal_index] = sqrtf(fmaxf(lower[diagonal_index], FLT_MIN));
        }
        __syncwarp();

        const float diagonal = lower[packed_index(pivot, pivot)];
        for (int row = pivot + 1 + lane; row < N; row += 32) {
            lower[packed_index(row, pivot)] /= diagonal;
        }
        __syncwarp();

        for (int row = pivot + 1 + lane; row < N; row += 32) {
            const float row_value = lower[packed_index(row, pivot)];
            for (int col = pivot + 1; col <= row; ++col) {
                const int index = packed_index(row, col);
                lower[index] = fmaf(
                    -row_value,
                    lower[packed_index(col, pivot)],
                    lower[index]);
            }
        }
        __syncwarp();
    }

    for (int vector = lane; vector < kVectors; vector += 32) {
        const int linear = vector * 4;
        const int row = linear / N;
        const int col = linear - row * N;
        float4 values;
        values.x = col <= row ? lower[packed_index(row, col)] : 0.0f;
        values.y = col + 1 <= row ? lower[packed_index(row, col + 1)] : 0.0f;
        values.z = col + 2 <= row ? lower[packed_index(row, col + 2)] : 0.0f;
        values.w = col + 3 <= row ? lower[packed_index(row, col + 3)] : 0.0f;
        l4[vector] = values;
    }
}

// One block owns one matrix. Keeping the lower triangle packed makes n=128
// fit comfortably below the default per-block shared-memory limit.
template <int N>
__global__ void small_cholesky_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int64_t matrix_stride,
    int runtime_n) {
    extern __shared__ float lower[];
    const int matrix = static_cast<int>(blockIdx.x);
    const int tid = static_cast<int>(threadIdx.x);
    const int n = N == 0 ? runtime_n : N;
    const float* a = input + static_cast<int64_t>(matrix) * matrix_stride;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    const int elements = n * n;
    for (int linear = tid; linear < elements; linear += blockDim.x) {
        const int row = linear / n;
        const int col = linear - row * n;
        if (col <= row) {
            lower[packed_index(row, col)] = a[linear];
        }
    }
    __syncthreads();

    // Right-looking Cholesky. A thread owns a complete trailing row during
    // each rank-1 update, so no two threads write the same shared element.
    for (int pivot = 0; pivot < n; ++pivot) {
        if (tid == 0) {
            float diagonal = lower[packed_index(pivot, pivot)];
            lower[packed_index(pivot, pivot)] = sqrtf(fmaxf(diagonal, FLT_MIN));
        }
        __syncthreads();

        const float diagonal = lower[packed_index(pivot, pivot)];
        for (int row = pivot + 1 + tid; row < n; row += blockDim.x) {
            lower[packed_index(row, pivot)] /= diagonal;
        }
        __syncthreads();

        for (int row = pivot + 1 + tid; row < n; row += blockDim.x) {
            const float row_value = lower[packed_index(row, pivot)];
            for (int col = pivot + 1; col <= row; ++col) {
                const int index = packed_index(row, col);
                lower[index] = fmaf(
                    -row_value,
                    lower[packed_index(col, pivot)],
                    lower[index]);
            }
        }
        // Thread 0 owns the next diagonal row; the next post-sqrt barrier
        // safely completes every independent later-row update.
    }

    for (int linear = tid; linear < elements; linear += blockDim.x) {
        const int row = linear / n;
        const int col = linear - row * n;
        l[linear] = col <= row ? lower[packed_index(row, col)] : 0.0f;
    }
}

__global__ void copy_lower_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int64_t total_elements,
    int n) {
    for (int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
         linear < total_elements;
         linear += static_cast<int64_t>(gridDim.x) * blockDim.x) {
        const int within_matrix = static_cast<int>(linear % (static_cast<int64_t>(n) * n));
        const int row = within_matrix / n;
        const int col = within_matrix - row * n;
        output[linear] = col <= row ? input[linear] : 0.0f;
    }
}

template <int PanelWidth>
__global__ void diagonal_factor_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_start,
    int runtime_panel_width) {
    __shared__ float tile[kPanel][kPanel + 1];
    const int matrix = static_cast<int>(blockIdx.x);
    const int tid = static_cast<int>(threadIdx.x);
    const int panel_width = PanelWidth == 0 ? runtime_panel_width : PanelWidth;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    const int tile_elements = panel_width * panel_width;
    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / panel_width;
        const int col = linear - row * panel_width;
        tile[row][col] = col <= row
            ? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
            : 0.0f;
    }
    __syncthreads();

    for (int pivot = 0; pivot < panel_width; ++pivot) {
        if (tid == 0) {
            tile[pivot][pivot] = sqrtf(fmaxf(tile[pivot][pivot], FLT_MIN));
        }
        __syncthreads();

        const float diagonal = tile[pivot][pivot];
        for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
            tile[row][pivot] /= diagonal;
        }
        __syncthreads();

        for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
            const float row_value = tile[row][pivot];
            for (int col = pivot + 1; col <= row; ++col) {
                tile[row][col] = fmaf(-row_value, tile[col][pivot], tile[row][col]);
            }
        }
        // Thread 0 exclusively updates the next diagonal row. The following
        // post-sqrt barrier completes all later rows before they are consumed.
    }

    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / panel_width;
        const int col = linear - row * panel_width;
        if (col <= row) {
            l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
                tile[row][col];
        }
    }
}

// The large-matrix route factors a full 128-column diagonal block in FP32.
// Dynamic shared memory is required because the padded tile exceeds the
// legacy 48 KiB static-shared limit.
__global__ void diagonal_factor_128_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_start) {
    extern __shared__ float tile[];
    const int matrix = static_cast<int>(blockIdx.x);
    const int tid = static_cast<int>(threadIdx.x);
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    constexpr int tile_stride = kLargePanel + 1;
    constexpr int tile_elements = kLargePanel * kLargePanel;
    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / kLargePanel;
        const int col = linear - row * kLargePanel;
        tile[row * tile_stride + col] = col <= row
            ? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
            : 0.0f;
    }
    __syncthreads();

    for (int pivot = 0; pivot < kLargePanel; ++pivot) {
        if (tid == 0) {
            float& diagonal = tile[pivot * tile_stride + pivot];
            diagonal = sqrtf(fmaxf(diagonal, FLT_MIN));
        }
        __syncthreads();

        const float diagonal = tile[pivot * tile_stride + pivot];
        for (int row = pivot + 1 + tid;
             row < kLargePanel;
             row += blockDim.x) {
            tile[row * tile_stride + pivot] /= diagonal;
        }
        __syncthreads();

        for (int row = pivot + 1 + tid;
             row < kLargePanel;
             row += blockDim.x) {
            const float row_value = tile[row * tile_stride + pivot];
            for (int col = pivot + 1; col <= row; ++col) {
                float& value = tile[row * tile_stride + col];
                value = fmaf(
                    -row_value,
                    tile[col * tile_stride + pivot],
                    value);
            }
        }
    }

    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / kLargePanel;
        const int col = linear - row * kLargePanel;
        if (col <= row) {
            l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
                tile[row * tile_stride + col];
        }
    }
}

template <int PanelWidth>
__global__ void diagonal_factor_inverse_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_start,
    int runtime_panel_width) {
    __shared__ float tile[kPanel][kPanel + 1];
    const int matrix = static_cast<int>(blockIdx.x);
    const int tid = static_cast<int>(threadIdx.x);
    const int panel_width = PanelWidth == 0 ? runtime_panel_width : PanelWidth;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    const int tile_elements = panel_width * panel_width;
    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / panel_width;
        const int col = linear - row * panel_width;
        tile[row][col] = col <= row
            ? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
            : 0.0f;
    }
    __syncthreads();

    for (int pivot = 0; pivot < panel_width; ++pivot) {
        if (tid == 0) {
            tile[pivot][pivot] = sqrtf(fmaxf(tile[pivot][pivot], FLT_MIN));
        }
        __syncthreads();

        const float diagonal = tile[pivot][pivot];
        for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
            tile[row][pivot] /= diagonal;
        }
        __syncthreads();

        for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
            const float row_value = tile[row][pivot];
            for (int col = pivot + 1; col <= row; ++col) {
                tile[row][col] = fmaf(-row_value, tile[col][pivot], tile[row][col]);
            }
        }
    }

    // Store U=L^{-T} in the unused strict upper triangle. The diagonal
    // reciprocal is synthesized from L when U is staged by the solve.
    __syncthreads();
    for (int inverse_row = 1; inverse_row < kPanel; ++inverse_row) {
        for (int inverse_col = tid;
             inverse_col < inverse_row;
             inverse_col += blockDim.x) {
            float value = 0.0f;
            for (int k = inverse_col; k < inverse_row; ++k) {
                const float inverse_value = k == inverse_col
                    ? 1.0f / tile[k][k]
                    : tile[inverse_col][k];
                value = fmaf(tile[inverse_row][k], inverse_value, value);
            }
            tile[inverse_col][inverse_row] =
                -value / tile[inverse_row][inverse_row];
        }
        __syncthreads();
    }

    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / panel_width;
        const int col = linear - row * panel_width;
        l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
            tile[row][col];
    }
}

// CUTLASS panel path: preserve L in the output diagonal block and write the
// full row-major U=L^{-T} into a reusable per-matrix workspace.
__global__ void diagonal_factor_cutlass_inverse_kernel(
    float* __restrict__ output,
    float* __restrict__ inverse,
    int64_t matrix_stride,
    int64_t inverse_stride,
    int n,
    int panel_start) {
    __shared__ float tile[kPanel][kPanel + 1];
    __shared__ float inverse_tile[kPanel][kPanel + 1];
    const int matrix = static_cast<int>(blockIdx.x);
    const int tid = static_cast<int>(threadIdx.x);
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    constexpr int tile_elements = kPanel * kPanel;
    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / kPanel;
        const int col = linear - row * kPanel;
        tile[row][col] = col <= row
            ? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
            : 0.0f;
    }
    __syncthreads();

    for (int pivot = 0; pivot < kPanel; ++pivot) {
        if (tid == 0) {
            tile[pivot][pivot] = sqrtf(fmaxf(tile[pivot][pivot], FLT_MIN));
        }
        __syncthreads();

        const float diagonal = tile[pivot][pivot];
        for (int row = pivot + 1 + tid; row < kPanel; row += blockDim.x) {
            tile[row][pivot] /= diagonal;
        }
        __syncthreads();

        for (int row = pivot + 1 + tid; row < kPanel; row += blockDim.x) {
            const float row_value = tile[row][pivot];
            for (int col = pivot + 1; col <= row; ++col) {
                tile[row][col] = fmaf(-row_value, tile[col][pivot], tile[row][col]);
            }
        }
    }

    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        inverse_tile[linear / kPanel][linear % kPanel] = 0.0f;
    }
    __syncthreads();

    if (tid < kPanel) {
        const int inverse_row = tid;
        for (int row = inverse_row; row < kPanel; ++row) {
            float sum = 0.0f;
            for (int p = inverse_row; p < row; ++p) {
                sum = fmaf(
                    tile[row][p],
                    inverse_tile[inverse_row][p],
                    sum);
            }
            inverse_tile[inverse_row][row] =
                ((row == inverse_row ? 1.0f : 0.0f) - sum) /
                tile[row][row];
        }
    }
    __syncthreads();

    float* u = inverse + static_cast<int64_t>(matrix) * inverse_stride;
    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        u[linear] = inverse_tile[linear / kPanel][linear % kPanel];
    }

    for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
        const int row = linear / kPanel;
        const int col = linear - row * kPanel;
        if (col <= row) {
            l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
                tile[row][col];
        }
    }
}

// A CTA stages U=L^{-T} and 16 right-hand-side rows, then computes every
// solved panel value independently in FP32.
__global__ void panel_solve_inverse_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_start) {
    extern __shared__ float shared[];
    float* upper = shared;
    float* rhs = upper + kPanel * (kPanel + 1);
    const int tid = static_cast<int>(threadIdx.x);
    const int matrix = static_cast<int>(blockIdx.y);
    const int first_local_row = static_cast<int>(blockIdx.x) * 16;
    const int trailing_start = panel_start + kPanel;
    const int trailing_rows = n - trailing_start;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    for (int linear = tid; linear < kPanel * kPanel; linear += blockDim.x) {
        const int row = linear / kPanel;
        const int col = linear - row * kPanel;
        const float value = l[
            static_cast<int64_t>(panel_start + row) * n + panel_start + col];
        upper[row * (kPanel + 1) + col] = col < row
            ? 0.0f
            : (col == row ? 1.0f / value : value);
    }

    for (int linear = tid; linear < 16 * kPanel; linear += blockDim.x) {
        const int local_row = linear / kPanel;
        const int col = linear - local_row * kPanel;
        const int trailing_row = first_local_row + local_row;
        rhs[local_row * (kPanel + 1) + col] = trailing_row < trailing_rows
            ? l[static_cast<int64_t>(trailing_start + trailing_row) * n +
                panel_start + col]
            : 0.0f;
    }
    __syncthreads();

    for (int linear = tid; linear < 16 * kPanel; linear += blockDim.x) {
        const int local_row = linear / kPanel;
        const int col = linear - local_row * kPanel;
        const int trailing_row = first_local_row + local_row;
        if (trailing_row < trailing_rows) {
            float solved = 0.0f;
            #pragma unroll
            for (int k = 0; k < kPanel; ++k) {
                if (k <= col) {
                    solved = fmaf(
                        rhs[local_row * (kPanel + 1) + k],
                        upper[k * (kPanel + 1) + col],
                        solved);
                }
            }
            l[static_cast<int64_t>(trailing_start + trailing_row) * n +
                panel_start + col] = solved;
        }
    }
}

__global__ void clear_diagonal_upper_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_count) {
    const int panel = static_cast<int>(blockIdx.x);
    const int matrix = static_cast<int>(blockIdx.y);
    const int panel_start = panel * kPanel;
    const int panel_width = n - panel_start < kPanel ? n - panel_start : kPanel;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
    for (int linear = static_cast<int>(threadIdx.x);
         linear < panel_width * panel_width;
         linear += blockDim.x) {
        const int row = linear / panel_width;
        const int col = linear - row * panel_width;
        if (col > row) {
            l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
                0.0f;
        }
    }
}

__device__ __forceinline__ float warp_sum(float value) {
    constexpr unsigned mask = 0xffffffffu;
    value += __shfl_down_sync(mask, value, 16);
    value += __shfl_down_sync(mask, value, 8);
    value += __shfl_down_sync(mask, value, 4);
    value += __shfl_down_sync(mask, value, 2);
    value += __shfl_down_sync(mask, value, 1);
    return value;
}

// One warp solves one row while all warps share the factored diagonal panel.
// A 2-D launch keeps each CTA inside a single matrix and removes the flat
// launch's per-warp division and modulo.
template <int PanelWidth>
__global__ void panel_solve_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_start,
    int runtime_panel_width) {
    extern __shared__ float diagonal_panel[];
    const int tid = static_cast<int>(threadIdx.x);
    const int lane = static_cast<int>(threadIdx.x) & 31;
    const int local_warp = static_cast<int>(threadIdx.x) >> 5;
    const int panel_width = PanelWidth == 0 ? runtime_panel_width : PanelWidth;
    const int matrix = static_cast<int>(blockIdx.y);
    const int trailing_rows = n - panel_start - panel_width;
    constexpr int rows_per_cta =
        PanelWidth == kLargePanel ? kLargeSolveThreads / 32 : 8;
    const int local_row =
        static_cast<int>(blockIdx.x) * rows_per_cta + local_warp;
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    const int diagonal_elements = panel_width * panel_width;
    for (int linear = tid; linear < diagonal_elements; linear += blockDim.x) {
        const int row = linear / panel_width;
        const int col = linear - row * panel_width;
        diagonal_panel[linear] = col <= row
            ? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
            : 0.0f;
    }
    __syncthreads();

    if (local_row >= trailing_rows) {
        return;
    }

    const int row = panel_start + panel_width + local_row;
    float* row_ptr = l + static_cast<int64_t>(row) * n;

    if (PanelWidth == kLargePanel) {
        constexpr unsigned mask = 0xffffffffu;
        float row_values[4] = {
            row_ptr[panel_start + lane],
            row_ptr[panel_start + lane + 32],
            row_ptr[panel_start + lane + 64],
            row_ptr[panel_start + lane + 96],
        };
        #pragma unroll
        for (int col = 0; col < kLargePanel; ++col) {
            float partial = 0.0f;
            #pragma unroll
            for (int segment = 0; segment < 4; ++segment) {
                const int p = lane + segment * 32;
                if (p < col) {
                    partial = fmaf(
                        row_values[segment],
                        diagonal_panel[col * kLargePanel + p],
                        partial);
                }
            }
            partial = warp_sum(partial);

            const int segment = col >> 5;
            const float rhs = __shfl_sync(
                mask, row_values[segment], col & 31);
            float solved = 0.0f;
            if (lane == 0) {
                solved = (rhs - partial) /
                    diagonal_panel[col * kLargePanel + col];
                row_ptr[panel_start + col] = solved;
            }
            solved = __shfl_sync(mask, solved, 0);
            if (lane == (col & 31)) {
                row_values[segment] = solved;
            }
        }
        return;
    }

    if (PanelWidth == kPanel) {
        constexpr unsigned mask = 0xffffffffu;
        float row_lo = row_ptr[panel_start + lane];
        float row_hi = row_ptr[panel_start + lane + 32];
        #pragma unroll
        for (int col = 0; col < kPanel; ++col) {
            float partial = 0.0f;
            if (lane < col) {
                partial = row_lo * diagonal_panel[col * kPanel + lane];
            }
            if (lane + 32 < col) {
                partial = fmaf(
                    row_hi,
                    diagonal_panel[col * kPanel + lane + 32],
                    partial);
            }
            partial = warp_sum(partial);

            const float rhs = col < 32
                ? __shfl_sync(mask, row_lo, col)
                : __shfl_sync(mask, row_hi, col - 32);
            float solved = 0.0f;
            if (lane == 0) {
                solved = (rhs - partial) /
                    diagonal_panel[col * kPanel + col];
                row_ptr[panel_start + col] = solved;
            }
            solved = __shfl_sync(mask, solved, 0);
            if (col < 32) {
                if (lane == col) {
                    row_lo = solved;
                }
            } else if (lane == col - 32) {
                row_hi = solved;
            }
        }
        return;
    }

    for (int col = 0; col < panel_width; ++col) {
        float partial = 0.0f;
        for (int p = lane; p < col; p += 32) {
            partial = fmaf(
                row_ptr[panel_start + p],
                diagonal_panel[col * panel_width + p],
                partial);
        }
        partial = warp_sum(partial);
        if (lane == 0) {
            row_ptr[panel_start + col] =
                (row_ptr[panel_start + col] - partial) /
                diagonal_panel[col * panel_width + col];
        }
        __syncwarp();
    }
}

// Each 16x16 CTA computes a 64x64 output tile with a 4x4 register tile per
// thread. Padding the panel dimension removes same-bank row strides.
__global__ void trailing_update_kernel(
    float* __restrict__ output,
    int64_t matrix_stride,
    int n,
    int panel_start,
    int panel_width,
    int trailing_start,
    int tile_count) {
    __shared__ float panel_rows[2][kTile][kPanel + 1];
    __shared__ int tile_row_shared;
    __shared__ int tile_col_shared;

    const int tid = static_cast<int>(threadIdx.y) * blockDim.x + threadIdx.x;
    if (tid == 0) {
        const int64_t linear_tile = blockIdx.x;
        int tile_row = static_cast<int>(
            floor((sqrt(8.0 * static_cast<double>(linear_tile) + 1.0) - 1.0) * 0.5));
        while (static_cast<int64_t>(tile_row + 1) * (tile_row + 2) / 2 <= linear_tile) {
            ++tile_row;
        }
        while (static_cast<int64_t>(tile_row) * (tile_row + 1) / 2 > linear_tile) {
            --tile_row;
        }
        tile_row_shared = tile_row;
        tile_col_shared = static_cast<int>(
            linear_tile - static_cast<int64_t>(tile_row) * (tile_row + 1) / 2);
    }
    __syncthreads();

    const int tile_row = tile_row_shared;
    const int tile_col = tile_col_shared;
    if (tile_row >= tile_count || tile_col > tile_row) {
        return;
    }

    const int matrix = static_cast<int>(blockIdx.y);
    float* l = output + static_cast<int64_t>(matrix) * matrix_stride;

    const int staged_elements = 2 * kTile * panel_width;
    for (int linear = tid; linear < staged_elements; linear += blockDim.x * blockDim.y) {
        const int side = linear / (kTile * panel_width);
        const int remaining = linear - side * kTile * panel_width;
        const int local_row = remaining / panel_width;
        const int p = remaining - local_row * panel_width;
        const int global_row = trailing_start +
            (side == 0 ? tile_row : tile_col) * kTile + local_row;
        panel_rows[side][local_row][p] = global_row < n
            ? l[static_cast<int64_t>(global_row) * n + panel_start + p]
            : 0.0f;
    }
    __syncthreads();

    float accum[4][4] = {};
    #pragma unroll 1
    for (int p = 0; p < panel_width; ++p) {
        float left[4];
        float right[4];
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            left[i] = panel_rows[0][threadIdx.y + i * 16][p];
            right[i] = panel_rows[1][threadIdx.x + i * 16][p];
        }
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            #pragma unroll
            for (int j = 0; j < 4; ++j) {
                accum[i][j] = fmaf(-left[i], right[j], accum[i][j]);
            }
        }
    }

    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        const int row = trailing_start + tile_row * kTile + threadIdx.y + i * 16;
        #pragma unroll
        for (int j = 0; j < 4; ++j) {
            const int col = trailing_start + tile_col * kTile + threadIdx.x + j * 16;
            if (row < n && col < n && col <= row) {
                const int64_t index = static_cast<int64_t>(row) * n + col;
                l[index] += accum[i][j];
            }
        }
    }
}

}  // namespace

torch::Tensor raw_cholesky(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be torch.float32");
    TORCH_CHECK(input.dim() == 3, "input must have shape [batch, n, n]");
    TORCH_CHECK(input.size(1) == input.size(2), "matrices must be square");
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");

    c10::cuda::CUDAGuard device_guard(input.device());
    auto output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int n = static_cast<int>(input.size(1));
    const int64_t matrix_stride = static_cast<int64_t>(n) * n;
    const int64_t total_elements = static_cast<int64_t>(batch) * matrix_stride;
    if (batch == 0 || n == 0) {
        return output;
    }

    const float* input_ptr = input.data_ptr<float>();
    float* output_ptr = output.data_ptr<float>();

    if (n == 32) {
        constexpr int matrices_per_cta = 16;
        constexpr size_t shared_bytes =
            matrices_per_cta * 32 * 33 / 2 * sizeof(float);
        const int blocks = (batch + matrices_per_cta - 1) / matrices_per_cta;
        warp_packed_cholesky_kernel<32, matrices_per_cta>
            <<<blocks, matrices_per_cta * 32, shared_bytes>>>(
                input_ptr, output_ptr, batch, matrix_stride);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return output;
    }

    if (n <= kSmallMax) {
        const size_t shared_bytes =
            static_cast<size_t>(n) * (n + 1) / 2 * sizeof(float);
        if (n == 64) {
            small_cholesky_kernel<64><<<batch, kSmallMax, shared_bytes>>>(
                input_ptr, output_ptr, matrix_stride, n);
        } else if (n == 128) {
            small_cholesky_kernel<128><<<batch, kSmallMax, shared_bytes>>>(
                input_ptr, output_ptr, matrix_stride, n);
        } else {
            small_cholesky_kernel<0><<<batch, kSmallMax, shared_bytes>>>(
                input_ptr, output_ptr, matrix_stride, n);
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return output;
    }

    const int64_t copy_blocks_64 = (total_elements + kThreads - 1) / kThreads;
    const int copy_blocks = static_cast<int>(copy_blocks_64 > 2147483647LL
        ? 2147483647LL
        : copy_blocks_64);
    copy_lower_kernel<<<copy_blocks, kThreads>>>(
        input_ptr, output_ptr, total_elements, n);

    for (int panel_start = 0; panel_start < n; panel_start += kPanel) {
        const int remaining = n - panel_start;
        const int panel_width = remaining < kPanel ? remaining : kPanel;
        if (panel_width == kPanel) {
            diagonal_factor_kernel<kPanel><<<batch, 128>>>(
                output_ptr, matrix_stride, n, panel_start, panel_width);
        } else {
            diagonal_factor_kernel<0><<<batch, 128>>>(
                output_ptr, matrix_stride, n, panel_start, panel_width);
        }

        const int trailing_start = panel_start + panel_width;
        const int trailing_rows = n - trailing_start;
        if (trailing_rows == 0) {
            continue;
        }

        const dim3 solve_grid(
            static_cast<unsigned>((trailing_rows + 7) / 8),
            static_cast<unsigned>(batch));
        const size_t solve_shared_bytes =
            static_cast<size_t>(panel_width) * panel_width * sizeof(float);
        if (panel_width == kPanel) {
            panel_solve_kernel<kPanel>
                <<<solve_grid, kThreads, solve_shared_bytes>>>(
                output_ptr,
                matrix_stride,
                n,
                panel_start,
                panel_width);
        } else {
            panel_solve_kernel<0>
                <<<solve_grid, kThreads, solve_shared_bytes>>>(
                output_ptr,
                matrix_stride,
                n,
                panel_start,
                panel_width);
        }

        const int tile_count = (trailing_rows + kTile - 1) / kTile;
        const int64_t triangular_tiles =
            static_cast<int64_t>(tile_count) * (tile_count + 1) / 2;
        const dim3 grid(
            static_cast<unsigned>(triangular_tiles),
            static_cast<unsigned>(batch));
        const dim3 block(16, 16);
        trailing_update_kernel<<<grid, block>>>(
            output_ptr,
            matrix_stride,
            n,
            panel_start,
            panel_width,
            trailing_start,
            tile_count);
    }

    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return output;
}

torch::Tensor copy_lower(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
    TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be torch.float32");
    TORCH_CHECK(input.dim() == 3, "input must have shape [batch, n, n]");
    TORCH_CHECK(input.size(1) == input.size(2), "matrices must be square");
    TORCH_CHECK(input.is_contiguous(), "input must be contiguous");

    c10::cuda::CUDAGuard device_guard(input.device());
    auto output = torch::empty_like(input);
    const int batch = static_cast<int>(input.size(0));
    const int n = static_cast<int>(input.size(1));
    const int64_t total_elements = static_cast<int64_t>(batch) * n * n;
    if (total_elements == 0) {
        return output;
    }

    const int64_t copy_blocks_64 = (total_elements + kThreads - 1) / kThreads;
    const int copy_blocks = static_cast<int>(copy_blocks_64 > 2147483647LL
        ? 2147483647LL
        : copy_blocks_64);
    copy_lower_kernel<<<copy_blocks, kThreads>>>(
        input.data_ptr<float>(), output.data_ptr<float>(), total_elements, n);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return output;
}

void factor_panel(torch::Tensor output, int64_t panel_start_64) {
    TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
    TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
    TORCH_CHECK(output.dim() == 3, "output must have shape [batch, n, n]");
    TORCH_CHECK(output.size(1) == output.size(2), "matrices must be square");
    TORCH_CHECK(output.is_contiguous(), "output must be contiguous");

    c10::cuda::CUDAGuard device_guard(output.device());
    const int batch = static_cast<int>(output.size(0));
    const int n = static_cast<int>(output.size(1));
    const int panel_start = static_cast<int>(panel_start_64);
    TORCH_CHECK(panel_start >= 0 && panel_start < n, "invalid panel start");
    const int remaining = n - panel_start;
    const int panel_width = remaining < kPanel ? remaining : kPanel;
    const int64_t matrix_stride = static_cast<int64_t>(n) * n;
    float* output_ptr = output.data_ptr<float>();

    if (panel_width == kPanel) {
        diagonal_factor_kernel<kPanel><<<batch, 128>>>(
            output_ptr, matrix_stride, n, panel_start, panel_width);
    } else {
        diagonal_factor_kernel<0><<<batch, 128>>>(
            output_ptr, matrix_stride, n, panel_start, panel_width);
    }
    const int trailing_rows = n - panel_start - panel_width;
    if (trailing_rows > 0) {
        const dim3 solve_grid(
            static_cast<unsigned>((trailing_rows + 7) / 8),
            static_cast<unsigned>(batch));
        const size_t solve_shared_bytes =
            static_cast<size_t>(panel_width) * panel_width * sizeof(float);
        if (panel_width == kPanel) {
            panel_solve_kernel<kPanel>
                <<<solve_grid, kThreads, solve_shared_bytes>>>(
                output_ptr,
                matrix_stride,
                n,
                panel_start,
                panel_width);
        } else {
            panel_solve_kernel<0>
                <<<solve_grid, kThreads, solve_shared_bytes>>>(
                output_ptr,
                matrix_stride,
                n,
                panel_start,
                panel_width);
        }
    }
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void factor_panel_128(torch::Tensor output, int64_t panel_start_64) {
    TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
    TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
    TORCH_CHECK(output.dim() == 3, "output must have shape [batch, n, n]");
    TORCH_CHECK(output.size(1) == output.size(2), "matrices must be square");
    TORCH_CHECK(output.is_contiguous(), "output must be contiguous");

    c10::cuda::CUDAGuard device_guard(output.device());
    const int batch = static_cast<int>(output.size(0));
    const int n = static_cast<int>(output.size(1));
    const int panel_start = static_cast<int>(panel_start_64);
    TORCH_CHECK(
        n >= 2048 && n % kLargePanel == 0,
        "128-column panel tier requires n>=2048 divisible by 128");
    TORCH_CHECK(
        panel_start >= 0 && panel_start + kLargePanel <= n &&
            panel_start % kLargePanel == 0,
        "invalid 128-column panel start");
    if (batch == 0) {
        return;
    }

    constexpr int diagonal_shared_bytes =
        kLargePanel * (kLargePanel + 1) * sizeof(float);
    constexpr int solve_shared_bytes =
        kLargePanel * kLargePanel * sizeof(float);
    static thread_local int configured_device = -1;
    int current_device = -1;
    C10_CUDA_CHECK(cudaGetDevice(&current_device));
    if (configured_device != current_device) {
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            diagonal_factor_128_kernel,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            diagonal_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            panel_solve_kernel<kLargePanel>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            solve_shared_bytes));
        configured_device = current_device;
    }

    const int64_t matrix_stride = static_cast<int64_t>(n) * n;
    float* output_ptr = output.data_ptr<float>();
    diagonal_factor_128_kernel<<<batch, kThreads, diagonal_shared_bytes>>>(
        output_ptr, matrix_stride, n, panel_start);

    const int trailing_rows = n - panel_start - kLargePanel;
    if (trailing_rows > 0) {
        constexpr int rows_per_cta = kLargeSolveThreads / 32;
        const dim3 solve_grid(
            static_cast<unsigned>(
                (trailing_rows + rows_per_cta - 1) / rows_per_cta),
            static_cast<unsigned>(batch));
        panel_solve_kernel<kLargePanel>
            <<<solve_grid, kLargeSolveThreads, solve_shared_bytes>>>(
                output_ptr,
                matrix_stride,
                n,
                panel_start,
                kLargePanel);
    }
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void factor_panel_inverse(torch::Tensor output, int64_t panel_start_64) {
    TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
    TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
    TORCH_CHECK(output.dim() == 3, "output must have shape [batch, n, n]");
    TORCH_CHECK(output.size(1) == output.size(2), "matrices must be square");
    TORCH_CHECK(output.is_contiguous(), "output must be contiguous");

    c10::cuda::CUDAGuard device_guard(output.device());
    const int batch = static_cast<int>(output.size(0));
    const int n = static_cast<int>(output.size(1));
    const int panel_start = static_cast<int>(panel_start_64);
    TORCH_CHECK(n == 512 && batch >= 128, "inverse panel tier requires n=512, batch>=128");
    TORCH_CHECK(panel_start >= 0 && panel_start < n, "invalid panel start");
    const int panel_width = n - panel_start < kPanel ? n - panel_start : kPanel;
    TORCH_CHECK(panel_width == kPanel, "inverse panel tier requires full panels");
    const int64_t matrix_stride = static_cast<int64_t>(n) * n;
    float* output_ptr = output.data_ptr<float>();

    diagonal_factor_inverse_kernel<kPanel><<<batch, 128>>>(
        output_ptr, matrix_stride, n, panel_start, panel_width);

    const int trailing_rows = n - panel_start - panel_width;
    if (trailing_rows > 0) {
        const dim3 solve_grid(
            static_cast<unsigned>((trailing_rows + 15) / 16),
            static_cast<unsigned>(batch));
        constexpr size_t inverse_shared_bytes =
            (kPanel * (kPanel + 1) + 16 * (kPanel + 1)) * sizeof(float);
        panel_solve_inverse_kernel<<<solve_grid, kThreads, inverse_shared_bytes>>>(
            output_ptr, matrix_stride, n, panel_start);
    }

    if (panel_start + panel_width == n) {
        constexpr int panel_count = 512 / kPanel;
        const dim3 clear_grid(
            static_cast<unsigned>(panel_count),
            static_cast<unsigned>(batch));
        clear_diagonal_upper_kernel<<<clear_grid, kThreads>>>(
            output_ptr, matrix_stride, n, panel_count);
    }
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

// B640_FUSED_V1_BEGIN
void factor_panel_cutlass_inverse_b640(
    torch::Tensor output,
    torch::Tensor inverse,
    int64_t panel_start_64) {
    TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
    TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
    TORCH_CHECK(
        output.dim() == 3 && output.size(0) == 640 &&
            output.size(1) == 512 && output.size(2) == 512,
        "output must have exact shape [640, 512, 512]");
    TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
    TORCH_CHECK(inverse.is_cuda(), "inverse must be a CUDA tensor");
    TORCH_CHECK(inverse.scalar_type() == at::kFloat, "inverse must be torch.float32");
    TORCH_CHECK(
        inverse.dim() == 3 && inverse.size(0) == 640 &&
            inverse.size(1) == kPanel && inverse.size(2) == kPanel,
        "inverse must have exact shape [640, 64, 64]");
    TORCH_CHECK(inverse.is_contiguous(), "inverse must be contiguous");
    TORCH_CHECK(
        inverse.device() == output.device(),
        "output and inverse must be on the same CUDA device");

    c10::cuda::CUDAGuard device_guard(output.device());
    const int panel_start = static_cast<int>(panel_start_64);
    TORCH_CHECK(
        panel_start >= 0 && panel_start + kPanel <= 512 &&
            panel_start % kPanel == 0,
        "panel start must identify a full aligned 64-column panel");

    constexpr int batch = 640;
    constexpr int n = 512;
    constexpr int64_t matrix_stride = static_cast<int64_t>(n) * n;
    constexpr int64_t inverse_stride = kPanel * kPanel;
    diagonal_factor_cutlass_inverse_kernel<<<batch, 128>>>(
        output.data_ptr<float>(),
        inverse.data_ptr<float>(),
        matrix_stride,
        inverse_stride,
        n,
        panel_start);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}
// B640_FUSED_V1_END

"""


_extension = load_inline(
    # B640_FUSED_V1_BEGIN
    name="gpumode_raw_cholesky_combo_k32_cutlass64_b640_fused_v1",
    # B640_FUSED_V1_END
    cpp_sources=CPP_SRC,
    cuda_sources=CUDA_SRC,
    functions=[
        "raw_cholesky",
        "copy_lower",
        "factor_panel",
        "factor_panel_128",
        "factor_panel_inverse",
        # B640_FUSED_V1_BEGIN
        "factor_panel_cutlass_inverse_b640",
        # B640_FUSED_V1_END
    ],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3"],
    with_cuda=True,
    verbose=False,
)


@triton.jit
def _copy_lower_2d(
    input,
    output,
    n,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    row_block = tl.program_id(0)
    col_block = tl.program_id(1)
    matrix = tl.program_id(2)

    rows = row_block * BLOCK_M + tl.arange(0, BLOCK_M)
    cols = col_block * BLOCK_N + tl.arange(0, BLOCK_N)
    offsets = matrix * n * n + rows[:, None] * n + cols[None, :]
    valid = (rows[:, None] < n) & (cols[None, :] < n)
    lower = cols[None, :] <= rows[:, None]
    values = tl.load(input + offsets, mask=valid & lower, other=0.0)
    tl.store(output + offsets, values, mask=valid)


@triton.jit
def _update_factor_final64_n256(
    output,
    N: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    matrix = tl.program_id(0)
    rows = tl.arange(0, BLOCK_SIZE)
    cols = tl.arange(0, BLOCK_SIZE)
    matrix_base = matrix * N * N
    panel_offsets = (
        matrix_base
        + (192 + rows[:, None]) * N
        + 128
        + cols[None, :]
    )
    left = tl.load(output + panel_offsets)
    right = tl.load(output + panel_offsets)
    left_hi = left.to(tl.bfloat16)
    right_hi = right.to(tl.bfloat16)
    left_lo = (left - left_hi.to(tl.float32)).to(tl.bfloat16)
    right_lo = (right - right_hi.to(tl.float32)).to(tl.bfloat16)
    update = tl.dot(
        left_hi,
        tl.trans(right_hi),
        out_dtype=tl.float32,
    )
    update += tl.dot(
        left_hi,
        tl.trans(right_lo),
        out_dtype=tl.float32,
    )
    update += tl.dot(
        left_lo,
        tl.trans(right_hi),
        out_dtype=tl.float32,
    )

    final_offsets = (
        matrix_base
        + (192 + rows[:, None]) * N
        + 192
        + cols[None, :]
    )
    lower = cols[None, :] <= rows[:, None]
    factor = tl.load(
        output + final_offsets,
        mask=lower,
        other=0.0,
    ).to(tl.float32)
    factor -= update

    for pivot in tl.range(0, BLOCK_SIZE):
        diagonal_rows = tl.sum(
            tl.where(
                (rows[:, None] == pivot) & (cols[None, :] == pivot),
                factor,
                0.0,
            ),
            axis=1,
        )
        diagonal = tl.sum(diagonal_rows, axis=0)
        diagonal = tl.sqrt(
            tl.maximum(diagonal, 1.1754943508222875e-38)
        )
        column = tl.sum(
            tl.where(cols[None, :] == pivot, factor, 0.0),
            axis=1,
        )
        column = column / diagonal
        column = tl.where(rows < pivot, 0.0, column)
        factor = tl.where(
            (cols[None, :] == pivot) & (rows[:, None] >= pivot),
            column[:, None],
            factor,
        )
        trailing_mask = (
            (rows[:, None] > pivot)
            & (cols[None, :] > pivot)
            & (cols[None, :] <= rows[:, None])
        )
        factor = tl.where(
            trailing_mask,
            factor - column[:, None] * column[None, :],
            factor,
        )

    tl.store(output + final_offsets, factor, mask=lower)


@triton.jit
def _factor_solve_panel_128_256(output, n):
    matrix = tl.program_id(0)
    rr = tl.arange(0, 64)[:, None]
    cc = tl.arange(0, 64)[None, :]
    matrix_base = matrix * n * n

    diagonal_offsets = matrix_base + (128 + rr) * n + 128 + cc
    solve_offsets = matrix_base + (192 + rr) * n + 128 + cc
    lower = cc <= rr
    factor = tl.load(
        output + diagonal_offsets,
        mask=lower,
        other=0.0,
    ).to(tl.float32)
    solved = tl.load(output + solve_offsets).to(tl.float32)

    for p in tl.range(0, 64):
        diagonal = tl.sum(
            tl.sum(
                tl.where((rr == p) & (cc == p), factor, 0.0),
                axis=1,
            ),
            axis=0,
        )
        diagonal = tl.sqrt(
            tl.maximum(diagonal, 1.1754943508222875e-38)
        )
        column = tl.sum(
            tl.where(cc == p, factor, 0.0),
            axis=1,
        )[:, None]
        column = tl.where(
            rr > p,
            column / diagonal,
            tl.where(rr == p, diagonal, 0.0),
        )
        factor = tl.where(
            cc == p,
            column,
            factor,
        )
        factor = tl.where(
            (rr > p) & (cc > p) & (cc <= rr),
            factor - column * tl.trans(column),
            factor,
        )

        solved = tl.where(cc == p, solved / diagonal, solved)
        pivot = tl.sum(
            tl.where(cc == p, solved, 0.0),
            axis=1,
        )[:, None]
        coefficients = tl.sum(
            tl.where(cc == p, factor, 0.0),
            axis=1,
        )[None, :]
        solved = tl.where(
            cc > p,
            solved - pivot * coefficients,
            solved,
        )

    tl.store(output + diagonal_offsets, factor, mask=lower)
    tl.store(output + solve_offsets, solved)


@triton.jit
def _trailing_update_tf32(
    output,
    n,
    panel_start,
    BLOCK: tl.constexpr,
    COMPENSATED: tl.constexpr,
):
    triangular_tile = tl.program_id(0)
    matrix = tl.program_id(1)

    # Invert row*(row+1)/2 to map one-dimensional program IDs onto only the
    # lower-triangular output tiles. This avoids launching/computing the unused
    # upper half of the trailing matrix.
    triangular_float = triangular_tile.to(tl.float32)
    tile_row = tl.floor(
        (tl.sqrt(8.0 * triangular_float + 1.0) - 1.0) * 0.5
    ).to(tl.int32)
    tile_col = triangular_tile - tile_row * (tile_row + 1) // 2

    trailing_start = panel_start + BLOCK
    rows = trailing_start + tile_row * BLOCK + tl.arange(0, BLOCK)
    cols = trailing_start + tile_col * BLOCK + tl.arange(0, BLOCK)
    reduction = panel_start + tl.arange(0, BLOCK)

    matrix_base = matrix * n * n
    left = tl.load(
        output + matrix_base + rows[:, None] * n + reduction[None, :]
    )
    right = tl.load(
        output + matrix_base + cols[:, None] * n + reduction[None, :]
    )
    if COMPENSATED:
        left_hi = left.to(tl.bfloat16)
        right_hi = right.to(tl.bfloat16)
        left_lo = (left - left_hi.to(tl.float32)).to(tl.bfloat16)
        right_lo = (right - right_hi.to(tl.float32)).to(tl.bfloat16)
        update = tl.dot(
            left_hi, tl.trans(right_hi), out_dtype=tl.float32
        )
        update += tl.dot(
            left_hi, tl.trans(right_lo), out_dtype=tl.float32
        )
        update += tl.dot(
            left_lo, tl.trans(right_hi), out_dtype=tl.float32
        )
    else:
        update = tl.dot(left, tl.trans(right), input_precision="tf32")

    output_offsets = matrix_base + rows[:, None] * n + cols[None, :]
    old = tl.load(output + output_offsets)
    lower_mask = cols[None, :] <= rows[:, None]
    tl.store(output + output_offsets, old - update, mask=lower_mask)


@triton.jit
def _trailing_update_tf32_panel128(
    output,
    n,
    panel_start,
    BLOCK: tl.constexpr,
    PANEL: tl.constexpr,
):
    triangular_tile = tl.program_id(0)
    matrix = tl.program_id(1)

    triangular_float = triangular_tile.to(tl.float32)
    tile_row = tl.floor(
        (tl.sqrt(8.0 * triangular_float + 1.0) - 1.0) * 0.5
    ).to(tl.int32)
    tile_col = triangular_tile - tile_row * (tile_row + 1) // 2

    trailing_start = panel_start + PANEL
    rows = trailing_start + tile_row * BLOCK + tl.arange(0, BLOCK)
    cols = trailing_start + tile_col * BLOCK + tl.arange(0, BLOCK)
    reduction = panel_start + tl.arange(0, BLOCK // 2)

    matrix_base = matrix * n * n
    left = tl.load(
        output + matrix_base + rows[:, None] * n + reduction[None, :]
    )
    right = tl.load(
        output + matrix_base + cols[:, None] * n + reduction[None, :]
    )
    update = tl.dot(left, tl.trans(right), input_precision="tf32")

    reduction += BLOCK // 2
    left = tl.load(
        output + matrix_base + rows[:, None] * n + reduction[None, :]
    )
    right = tl.load(
        output + matrix_base + cols[:, None] * n + reduction[None, :]
    )
    update += tl.dot(left, tl.trans(right), input_precision="tf32")

    reduction += BLOCK // 2
    left = tl.load(
        output + matrix_base + rows[:, None] * n + reduction[None, :]
    )
    right = tl.load(
        output + matrix_base + cols[:, None] * n + reduction[None, :]
    )
    update += tl.dot(left, tl.trans(right), input_precision="tf32")

    reduction += BLOCK // 2
    left = tl.load(
        output + matrix_base + rows[:, None] * n + reduction[None, :]
    )
    right = tl.load(
        output + matrix_base + cols[:, None] * n + reduction[None, :]
    )
    update += tl.dot(left, tl.trans(right), input_precision="tf32")

    output_offsets = matrix_base + rows[:, None] * n + cols[None, :]
    old = tl.load(output + output_offsets)
    lower_mask = cols[None, :] <= rows[:, None]
    tl.store(output + output_offsets, old - update, mask=lower_mask)


# B640_FUSED_V1_BEGIN
@triton.jit
def _b640_fused_panel_schur_inverse_apply_v1(
    data,
    output,
    inverse,
    PANEL_START,
    BLOCK_M: tl.constexpr,
):
    row_block = tl.program_id(0)
    matrix = tl.program_id(1)

    rows = PANEL_START + 64 + row_block * BLOCK_M + tl.arange(0, BLOCK_M)
    panel_cols = PANEL_START + tl.arange(0, 64)
    matrix_base = matrix * 512 * 512
    panel = tl.load(
        data + matrix_base + rows[:, None] * 512 + panel_cols[None, :]
    ).to(tl.float32)

    for previous_start in tl.static_range(0, 512, 64):
        if previous_start < PANEL_START:
            previous_cols = previous_start + tl.arange(0, 64)
            left = tl.load(
                output + matrix_base + rows[:, None] * 512 + previous_cols[None, :]
            )
            right = tl.load(
                output
                + matrix_base
                + panel_cols[:, None] * 512
                + previous_cols[None, :]
            )
            panel -= tl.dot(left, tl.trans(right), input_precision="tf32")

    inverse_rows = tl.arange(0, 64)
    inverse_cols = tl.arange(0, 64)
    upper_inverse = tl.load(
        inverse
        + matrix * 64 * 64
        + inverse_rows[:, None] * 64
        + inverse_cols[None, :]
    )
    solved = tl.dot(panel, upper_inverse, input_precision="tf32")
    tl.store(
        output + matrix_base + rows[:, None] * 512 + panel_cols[None, :],
        solved,
    )
# B640_FUSED_V1_END


def _old_custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    if n < 256 or n % 64 != 0:
        return _extension.raw_cholesky(data)

    output = torch.empty_like(data)
    batch = data.shape[0]
    _copy_lower_2d[
        (triton.cdiv(n, 16), triton.cdiv(n, 256), batch)
    ](
        data,
        output,
        n,
        BLOCK_M=16,
        BLOCK_N=256,
        num_warps=8,
        num_stages=1,
    )
    if n == 256 and batch == 64:
        for panel_start in (0, 64):
            _extension.factor_panel(output, panel_start)
            trailing_rows = n - panel_start - 64
            tile_count = triton.cdiv(trailing_rows, 64)
            triangular_tiles = tile_count * (tile_count + 1) // 2
            _trailing_update_tf32[(triangular_tiles, batch)](
                output,
                n,
                panel_start,
                BLOCK=64,
                COMPENSATED=True,
                num_warps=8,
                num_stages=2,
            )
        _factor_solve_panel_128_256[(batch,)](
            output,
            n,
            num_warps=8,
            num_stages=1,
        )
        _update_factor_final64_n256[(batch,)](
            output,
            N=256,
            BLOCK_SIZE=64,
            num_warps=8,
            num_stages=2,
        )
        return output

    if n == 32768:
        for panel_start in range(0, n, 128):
            _extension.factor_panel_128(output, panel_start)
            trailing_rows = n - panel_start - 128
            if trailing_rows > 0:
                tile_count = triton.cdiv(trailing_rows, 64)
                triangular_tiles = tile_count * (tile_count + 1) // 2
                _trailing_update_tf32_panel128[(triangular_tiles, batch)](
                    output,
                    n,
                    panel_start,
                    BLOCK=64,
                    PANEL=128,
                    num_warps=8,
                    num_stages=2,
                )
        return output

    use_inverse = n == 512 and batch >= 128
    for panel_start in range(0, n, 64):
        if use_inverse and panel_start + 64 < n:
            _extension.factor_panel_inverse(output, panel_start)
        else:
            _extension.factor_panel(output, panel_start)
        trailing_rows = n - panel_start - 64
        if trailing_rows > 0:
            tile_count = triton.cdiv(trailing_rows, 64)
            triangular_tiles = tile_count * (tile_count + 1) // 2
            _trailing_update_tf32[(triangular_tiles, batch)](
                output,
                n,
                panel_start,
                BLOCK=64,
                COMPENSATED=n < 2048,
                num_warps=8,
                num_stages=1 if n >= 2048 else 2,
            )
    return output


import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = True

CPP_SRC = r"""torch::Tensor chol_small(torch::Tensor input);"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cmath>

__global__ void register_cholesky32(const float* input, float* output) {
    const int lane = threadIdx.x;
    const size_t base = static_cast<size_t>(blockIdx.x) * 1024;
    float row[32];
    // A matrix is 4 KiB aligned, so each lane can move its complete row as
    // eight naturally aligned 16-byte transactions rather than scalar loads.
    const float4* input4 = reinterpret_cast<const float4*>(input + base + lane * 32);
    float4* output4 = reinterpret_cast<float4*>(output + base + lane * 32);
#pragma unroll
    for (int vector = 0; vector < 8; ++vector) {
        const float4 values = input4[vector];
        const int col = vector * 4;
        row[col + 0] = col + 0 <= lane ? values.x : 0.0f;
        row[col + 1] = col + 1 <= lane ? values.y : 0.0f;
        row[col + 2] = col + 2 <= lane ? values.z : 0.0f;
        row[col + 3] = col + 3 <= lane ? values.w : 0.0f;
    }
#pragma unroll
    for (int k = 0; k < 32; ++k) {
        if (lane == k) row[k] = sqrtf(fmaxf(row[k], 0.0f));
        const float diagonal = __shfl_sync(0xffffffffu, row[k], k);
        if (lane > k) row[k] /= diagonal;
        const float own = row[k];
#pragma unroll
        for (int col = k + 1; col < 32; ++col) {
            const float other = __shfl_sync(0xffffffffu, row[k], col);
            if (lane >= col) row[col] = fmaf(-own, other, row[col]);
        }
    }
#pragma unroll
    for (int vector = 0; vector < 8; ++vector) {
        const int col = vector * 4;
        output4[vector] = make_float4(
            row[col + 0], row[col + 1], row[col + 2], row[col + 3]);
    }
}

__global__ void shared_cholesky64(const float* input, float* output) {
    extern __shared__ float matrix[];
    const int tid = threadIdx.x;
    const size_t base = static_cast<size_t>(blockIdx.x) * 4096;
    for (int index = tid; index < 4096; index += blockDim.x) {
        const int row = index / 64, col = index - row * 64;
        matrix[index] = row >= col ? input[base + index] : 0.0f;
    }
    __syncthreads();
#pragma unroll
    for (int k = 0; k < 64; ++k) {
        if (tid == 0) {
            float value = matrix[k * 64 + k];
#pragma unroll 4
            for (int j = 0; j < k; ++j) {
                const float x = matrix[k * 64 + j];
                value = fmaf(-x, x, value);
            }
            matrix[k * 64 + k] = sqrtf(fmaxf(value, 0.0f));
        }
        __syncthreads();
        if (tid > k && tid < 64) {
            float value = matrix[tid * 64 + k];
#pragma unroll 4
            for (int j = 0; j < k; ++j)
                value = fmaf(-matrix[tid * 64 + j], matrix[k * 64 + j], value);
            matrix[tid * 64 + k] = value / matrix[k * 64 + k];
        }
        __syncthreads();
    }
    for (int index = tid; index < 4096; index += blockDim.x)
        output[base + index] = matrix[index];
}

torch::Tensor chol_small(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == at::kFloat && input.is_contiguous(),
                "expected contiguous CUDA float32");
    auto output = torch::empty_like(input);
    const int64_t n = input.size(1), batch = input.size(0);
    if (n == 32) register_cholesky32<<<batch, 32>>>(input.data_ptr<float>(), output.data_ptr<float>());
    else if (n == 64) shared_cholesky64<<<batch, 64, 16384>>>(input.data_ptr<float>(), output.data_ptr<float>());
    else TORCH_CHECK(false, "unsupported n");
    cudaError_t error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
    return output;
}
"""

_native = load_inline(
    name="hopper_chol_merged_v5",
    cpp_sources=[CPP_SRC], cuda_sources=[CUDA_SRC], functions=["chol_small"],
    extra_cuda_cflags=["-O3"], verbose=False,
)


def _blocked_inverse(data: torch.Tensor, block_size: int) -> torch.Tensor:
    batch, n, _ = data.shape
    lower = torch.zeros_like(data)
    identity = torch.eye(block_size, dtype=data.dtype, device=data.device).expand(
        batch, block_size, block_size
    )
    for start in range(0, n, block_size):
        end = start + block_size
        diagonal = data[:, start:end, start:end]
        if start:
            previous = lower[:, start:end, :start]
            diagonal = diagonal - previous @ previous.transpose(-1, -2)
        factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
        lower[:, start:end, start:end] = factor
        if end < n:
            panel = data[:, end:, start:end]
            if start:
                panel = panel - lower[:, end:, :start] @ lower[:, start:end, :start].transpose(-1, -2)
            inverse = torch.linalg.solve_triangular(factor, identity, upper=False)
            lower[:, end:, start:end] = panel @ inverse.transpose(-1, -2)
    return lower


# B640_FUSED_V1_BEGIN
def _b640_fused_cholesky_v1(data: torch.Tensor) -> torch.Tensor:
    output = torch.zeros_like(data)
    inverse = torch.empty((640, 64, 64), dtype=torch.float32, device=data.device)

    for panel_start in range(0, 512, 64):
        panel_end = panel_start + 64
        diagonal = data[:, panel_start:panel_end, panel_start:panel_end]
        if panel_start:
            previous = output[:, panel_start:panel_end, :panel_start]
            diagonal = diagonal - previous @ previous.transpose(-1, -2)

        diagonal_output = output[
            :, panel_start:panel_end, panel_start:panel_end
        ]
        diagonal_output.copy_(diagonal)
        diagonal_output.tril_()
        _extension.factor_panel_cutlass_inverse_b640(
            output, inverse, panel_start
        )

        if panel_end < 512:
            _b640_fused_panel_schur_inverse_apply_v1[
                ((512 - panel_end) // 64, 640)
            ](
                data,
                output,
                inverse,
                PANEL_START=panel_start,
                BLOCK_M=64,
                num_warps=8,
                num_stages=2,
            )

    return output
# B640_FUSED_V1_END


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n in (32, 64):
        return _native.chol_small(data)
    if batch == 640 and n == 512:
        # B640_FUSED_V1_BEGIN
        return _b640_fused_cholesky_v1(data)
        # B640_FUSED_V1_END
    if batch == 60 and n == 1024:
        return _blocked_inverse(data, 128)
    if batch == 8 and n == 2048:
        return _old_custom_kernel(data)
    if batch == 1 and n == 8192:
        return _blocked_inverse(data, 2048)
    if batch == 1 and n == 16384:
        return _blocked_inverse(data, 2048)
    if batch == 1 and n == 32768:
        return _blocked_inverse(data, 1024)
    if batch == 2 and n in (2048, 4096):
        return torch.cat([
            torch.linalg.cholesky_ex(
                data[i:i+1].transpose(-1, -2), check_errors=False
            ).L
            for i in range(batch)
        ], dim=0)
    if (batch, n) in (
        (256, 128),
        (64, 256),
        (16, 512),
        (4, 1024),
        (1, 4096),
    ):
        return torch.linalg.cholesky_ex(
            data.transpose(-1, -2), check_errors=False
        ).L
    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 1809 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