Skip to content
KernelIndex
Search⌘K

submission 892007

coder-2011 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

candidate137_structured_compact.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-892007?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
2.00ms
#260 of 337
2026-07-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:db93fc40fafd6b68e6bcd817c5138a6308c1b7d4f634e3300bbac6e34e2bf2cc
license declaredunknown
license concludedunknown
authorscoder-2011
imported2026-08-26

Techniques

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

mbarrier"mbarrier.init.shared::cta.b64 [%0], 1;\n\t"
shared-memory__shared__ int warp_structures[kInitializeRowsPerBlock];
tcgen05"tcgen05.mma.cta_group::1.kind::f16 "
vector-width = float4float4 values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);

Kernel source

candidate137_structured_compact.py1923 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

from pathlib import Path

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

# Host-only tests use PyTorch because macOS cannot compile the CUDA extension.
if torch.version.cuda is None:
    blocked_cholesky_cuda = torch.linalg.cholesky
    blocked_cholesky_cuda_medium = torch.linalg.cholesky
    blocked_cholesky_cuda_large = torch.linalg.cholesky
else:
    cuda_source = r"""
#include <ATen/ATen.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>
#include <torch/types.h>

#include <algorithm>
#include <climits>
#include <cstdint>

namespace {

#ifndef CHOLESKY_PANEL
#define CHOLESKY_PANEL 64
#endif

#ifndef CHOLESKY_ENTRYPOINT
#define CHOLESKY_ENTRYPOINT blocked_cholesky_cuda
#endif

constexpr int kPanel = CHOLESKY_PANEL;
constexpr int kScalarUpdateTile = 16;
constexpr int kTensorUpdateTile = 128;
constexpr int kTensorInner = 16;
constexpr int kTensorChunkHalfs = kTensorUpdateTile * kTensorInner;
constexpr int kTensorChunkWords = kTensorChunkHalfs / 2;
constexpr int kTensorOperandHalfs = kTensorUpdateTile * kPanel;
constexpr int kTensorMemoryColumns = 128;
constexpr int kTensorThreads = 128;
constexpr int kTensorSharedBytes = 4 * kTensorOperandHalfs * sizeof(uint16_t) + 16;
constexpr int kWarpThreads = 32;
constexpr int kInitializeThreads = 256;
constexpr int kInitializeRowsPerBlock = kInitializeThreads / kWarpThreads;
constexpr int kDiagonalMicroblock = 8;
constexpr int kPanelThreads = 256;
constexpr int kWarpCholeskySize = 32;
constexpr int kWarpMatricesPerBlock = 2;
constexpr int kFullShared64Size = 64;
constexpr int kFullShared64Threads = 256;
constexpr int kFullShared64Bytes =
    kFullShared64Size * (kFullShared64Size + 1) * sizeof(float);
constexpr int kFullShared128Size = 128;
constexpr int kFullShared128Threads = 256;
constexpr int kFullShared128Bytes =
    kFullShared128Size * (kFullShared128Size + 1) * sizeof(float);
constexpr int kSolvePanelSharedBytes =
    (kPanel * kPanel + kPanelThreads * kPanel) * sizeof(float);
constexpr int kTensorDispatchThreshold = 512;
constexpr int kStructuredSize = 512;
constexpr int kStructureDiagonal = 0;
constexpr int kStructureTridiagonal = 1;
constexpr int kStructureGeneral = 2;

// Copy the lower input into independent work storage and make the upper triangle exact zero.
__global__ void initialize_lower_kernel(
    const float* input,
    float* output,
    int64_t n) {
    const int lane = threadIdx.x & (kWarpThreads - 1);
    const int warp = threadIdx.x >> 5;
    const int64_t row =
        static_cast<int64_t>(blockIdx.x) * kInitializeRowsPerBlock + warp;
    if (row >= n) {
        return;
    }

    const int64_t matrix = blockIdx.y;
    const int64_t row_offset = (matrix * n + row) * n;
    for (int64_t column = lane; column < n; column += kWarpThreads) {
        const int64_t index = row_offset + column;
        output[index] = row >= column ? input[index] : 0.0f;
    }
}

// Copy and classify low-batch n=512 matrices in the same memory pass.
__global__ void initialize_lower_structured_kernel(
    const float* input,
    float* output) {
    constexpr int n = kStructuredSize;
    const int lane = threadIdx.x & (kWarpThreads - 1);
    const int warp = threadIdx.x >> 5;
    const int64_t row =
        static_cast<int64_t>(blockIdx.x) * kInitializeRowsPerBlock + warp;
    const int64_t matrix = blockIdx.y;
    const int64_t row_offset = (matrix * n + row) * n;
    const int64_t block_row =
        static_cast<int64_t>(blockIdx.x) * kInitializeRowsPerBlock;
    const int64_t probe_row = block_row == 0 ? 2 : block_row;
    const bool inspect_structure =
        input[(matrix * n + probe_row) * n + probe_row - 2] == 0.0f;
    bool has_off_diagonal = false;
    bool has_off_band = false;

    for (int64_t column = lane; column < n; column += kWarpThreads) {
        const int64_t index = row_offset + column;
        const float value = input[index];
        const bool structure_word = row == block_row && column == n - 1;
        if (!structure_word) {
            output[index] = row >= column ? value : 0.0f;
        }
        if (inspect_structure && row > column && value != 0.0f) {
            has_off_diagonal = true;
            has_off_band |= row > column + 1;
        }
    }

    __shared__ int warp_structures[kInitializeRowsPerBlock];
    int block_structure = kStructureGeneral;
    if (inspect_structure) {
        const bool warp_has_off_diagonal =
            __any_sync(0xffffffffu, has_off_diagonal);
        const bool warp_has_off_band = __any_sync(0xffffffffu, has_off_band);
        if (lane == 0) {
            warp_structures[warp] = warp_has_off_band
                ? kStructureGeneral
                : (warp_has_off_diagonal
                    ? kStructureTridiagonal
                    : kStructureDiagonal);
        }
        __syncthreads();
        if (threadIdx.x == 0) {
            block_structure = kStructureDiagonal;
            #pragma unroll
            for (int item = 0; item < kInitializeRowsPerBlock; ++item) {
                block_structure = max(block_structure, warp_structures[item]);
            }
        }
    }

    if (threadIdx.x == 0) {
        const int64_t matrix_offset = matrix * n * n;
        const int64_t scratch = block_row * n + n - 1;
        output[matrix_offset + scratch] = static_cast<float>(block_structure);
    }
}

// Factor one diagonal tile per batch matrix with 8-wide scalar FP32 microblocks.
template <bool RecognizeStructure>
__global__ void factor_diagonal_kernel(
    float* output,
    int64_t n,
    int64_t matrix_elements,
    int panel_start,
    int panel_width) {
    __shared__ float tile[kPanel][kPanel + 1];

    const int thread = threadIdx.x;
    const int lane = thread & (kWarpThreads - 1);
    const int warp = thread >> 5;
    const int block_warps = blockDim.x >> 5;
    float* matrix = output + static_cast<int64_t>(blockIdx.x) * matrix_elements;

    if constexpr (RecognizeStructure) {
        if (thread < kWarpThreads) {
            if (panel_start == 0) {
                int structure = kStructureDiagonal;
                const int row_blocks = static_cast<int>(n) / kInitializeRowsPerBlock;
                for (int index = lane; index < row_blocks; index += kWarpThreads) {
                    const int64_t scratch =
                        static_cast<int64_t>(index) * kInitializeRowsPerBlock * n + n - 1;
                    structure = max(
                        structure,
                        static_cast<int>(matrix[scratch]));
                    matrix[scratch] = 0.0f;
                }
                for (int offset = 16; offset > 0; offset >>= 1) {
                    structure = max(
                        structure,
                        __shfl_down_sync(0xffffffffu, structure, offset));
                }
                if (lane == 0) {
                    tile[0][kPanel] = structure < kStructureGeneral
                        ? static_cast<float>(structure + 1)
                        : 0.0f;
                }
            } else if (lane == 0) {
                tile[0][kPanel] = matrix[n - 1];
            }
        }
    }

    for (int row = warp; row < panel_width; row += block_warps) {
        for (int column = lane; column < panel_width; column += kWarpThreads) {
            if (row >= column) {
                const int64_t offset =
                    static_cast<int64_t>(panel_start + row) * n + panel_start + column;
                tile[row][column] = matrix[offset];
            } else {
                tile[row][column] = 0.0f;
            }
        }
    }

    // The first warp factors a small diagonal block; the whole CTA handles its row/update work.
    __syncthreads();
    if constexpr (RecognizeStructure) {
        const int structure_state = static_cast<int>(tile[0][kPanel]);
        if (structure_state != 0) {
            if (panel_start == 0) {
                if (structure_state == kStructureDiagonal + 1) {
                    for (int diagonal = thread; diagonal < n; diagonal += blockDim.x) {
                        const int64_t index = static_cast<int64_t>(diagonal) * n + diagonal;
                        matrix[index] = sqrtf(matrix[index]);
                    }
                } else if (thread == 0) {
                    float previous_diagonal = sqrtf(matrix[0]);
                    matrix[0] = previous_diagonal;
                    for (int row = 1; row < n; ++row) {
                        const int64_t diagonal = static_cast<int64_t>(row) * n + row;
                        const int64_t subdiagonal = diagonal - 1;
                        const float factor = matrix[subdiagonal] / previous_diagonal;
                        matrix[subdiagonal] = factor;
                        previous_diagonal = sqrtf(fmaf(
                            -factor,
                            factor,
                            matrix[diagonal]));
                        matrix[diagonal] = previous_diagonal;
                    }
                }
                __syncthreads();
                if (thread == 0) {
                    matrix[n - 1] = static_cast<float>(structure_state);
                }
            }
            if (panel_start + panel_width == n && thread == 0) {
                matrix[n - 1] = 0.0f;
            }
            return;
        }
    }

    for (
        int microblock_start = 0;
        microblock_start < panel_width;
        microblock_start += kDiagonalMicroblock) {
        const int microblock_end = min(
            microblock_start + kDiagonalMicroblock,
            panel_width);

        if (thread < warpSize) {
            const int lane = thread;
            for (int pivot = microblock_start; pivot < microblock_end; ++pivot) {
                const int pivot_owner = pivot - microblock_start;
                float diagonal = 0.0f;
                if (lane == pivot_owner) {
                    diagonal = sqrtf(tile[pivot][pivot]);
                    tile[pivot][pivot] = diagonal;
                }
                diagonal = __shfl_sync(0xffffffffu, diagonal, pivot_owner);

                if (pivot + 1 < microblock_end) {
                    for (
                        int row = pivot + 1 + lane;
                        row < microblock_end;
                        row += warpSize) {
                        tile[row][pivot] /= diagonal;
                    }

                    // The rank-1 update consumes pivot values produced by other lanes.
                    __syncwarp();
                    for (
                        int column = pivot + 1 + lane;
                        column < microblock_end;
                        column += warpSize) {
                        const float column_pivot = tile[column][pivot];
                        for (int row = column; row < microblock_end; ++row) {
                            tile[row][column] = fmaf(
                                -tile[row][pivot],
                                column_pivot,
                                tile[row][column]);
                        }
                    }

                    // The next pivot reads the diagonal updated by another lane.
                    __syncwarp();
                }
            }
        }

        // Trailing rows cannot consume the diagonal block until its warp publishes it.
        __syncthreads();
        if (microblock_end == panel_width) {
            break;
        }

        const int solve_row = microblock_end + thread;
        if (solve_row < panel_width) {
            for (int column = microblock_start; column < microblock_end; ++column) {
                float value = tile[solve_row][column];
                for (int previous = microblock_start; previous < column; ++previous) {
                    value = fmaf(
                        -tile[solve_row][previous],
                        tile[column][previous],
                        value);
                }
                value /= tile[column][column];
                tile[solve_row][column] = value;
            }
        }

        // Every lower trailing element reads rows solved by potentially different warps.
        __syncthreads();
        for (
            int column = microblock_end + warp;
            column < panel_width;
            column += block_warps) {
            for (
                int row = column + lane;
                row < panel_width;
                row += kWarpThreads) {
                float value = tile[row][column];
                for (int inner = microblock_start; inner < microblock_end; ++inner) {
                    value = fmaf(
                        -tile[row][inner],
                        tile[column][inner],
                        value);
                }
                tile[row][column] = value;
            }
        }

        // The next diagonal microblock consumes the completed rank-8 update.
        __syncthreads();
    }

    // Writing only the lower tile preserves the exact-zero upper-triangle invariant.
    for (int row = warp; row < panel_width; row += block_warps) {
        for (int column = lane; column <= row; column += kWarpThreads) {
            const int64_t offset =
                static_cast<int64_t>(panel_start + row) * n + panel_start + column;
            matrix[offset] = tile[row][column];
        }
    }
}

// Solve independent trailing rows against one complete compile-time-width diagonal factor.
template <bool RecognizeStructure>
__global__ void solve_panel_kernel(
    float* output,
    int64_t n,
    int64_t matrix_elements,
    int panel_start) {
    // A single dynamic buffer opts the exact compile-time layout into shared memory.
    extern __shared__ float solve_panel_shared[];
    float (*diagonal_tile)[kPanel] =
        reinterpret_cast<float (*)[kPanel]>(solve_panel_shared);
    float (*panel_rows)[kPanel] = reinterpret_cast<float (*)[kPanel]>(
        solve_panel_shared + kPanel * kPanel);

    const int thread = threadIdx.x;
    float* matrix = output + static_cast<int64_t>(blockIdx.y) * matrix_elements;
    if constexpr (RecognizeStructure) {
        if (matrix[n - 1] != 0.0f) {
            return;
        }
    }

    for (int index = thread; index < kPanel * kPanel; index += blockDim.x) {
        const int row = index / kPanel;
        const int column = index - row * kPanel;
        if (row >= column) {
            const int64_t offset =
                static_cast<int64_t>(panel_start + row) * n + panel_start + column;
            diagonal_tile[row][column] = matrix[offset];
        } else {
            diagonal_tile[row][column] = 0.0f;
        }
    }

    const int trailing_row_start =
        panel_start + kPanel + static_cast<int>(blockIdx.x) * kPanelThreads;
    if ((n & 3) == 0) {
        constexpr int vectors_per_row = kPanel / 4;
        constexpr int vector_count = kPanelThreads * vectors_per_row;
        for (
            int vector_index = thread;
            vector_index < vector_count;
            vector_index += blockDim.x) {
            const int tile_row = vector_index / vectors_per_row;
            const int column = vector_index % vectors_per_row * 4;
            const int64_t global_row = static_cast<int64_t>(trailing_row_start + tile_row);
            float4 values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
            if (global_row < n) {
                const float* source = matrix + global_row * n + panel_start + column;
                values = *reinterpret_cast<const float4*>(source);
            }

            // XOR keeps same-column reads by adjacent row owners on distinct banks.
            const int swizzle = tile_row & 31;
            panel_rows[tile_row][(column + 0) ^ swizzle] = values.x;
            panel_rows[tile_row][(column + 1) ^ swizzle] = values.y;
            panel_rows[tile_row][(column + 2) ^ swizzle] = values.z;
            panel_rows[tile_row][(column + 3) ^ swizzle] = values.w;
        }
    } else {
        constexpr int tile_elements = kPanelThreads * kPanel;
        for (int index = thread; index < tile_elements; index += blockDim.x) {
            const int tile_row = index / kPanel;
            const int column = index % kPanel;
            const int64_t global_row = static_cast<int64_t>(trailing_row_start + tile_row);
            const float value = global_row < n
                ? matrix[global_row * n + panel_start + column]
                : 0.0f;
            panel_rows[tile_row][column ^ (tile_row & 31)] = value;
        }
    }

    // Invalid edge-row threads still participate in both shared-memory barriers.
    __syncthreads();
    const int64_t row = static_cast<int64_t>(trailing_row_start + thread);

    float row_values[kPanel];
    if (row < n) {
        #pragma unroll
        for (int column = 0; column < kPanel; ++column) {
            float value = panel_rows[thread][column ^ (thread & 31)];
            #pragma unroll
            for (int previous = 0; previous < column; ++previous) {
                value = fmaf(
                    -row_values[previous],
                    diagonal_tile[column][previous],
                    value);
            }
            value /= diagonal_tile[column][column];
            row_values[column] = value;
            panel_rows[thread][column ^ (thread & 31)] = value;
        }
    }
    __syncthreads();

    if ((n & 3) == 0) {
        constexpr int vectors_per_row = kPanel / 4;
        constexpr int vector_count = kPanelThreads * vectors_per_row;
        for (
            int vector_index = thread;
            vector_index < vector_count;
            vector_index += blockDim.x) {
            const int tile_row = vector_index / vectors_per_row;
            const int column = vector_index % vectors_per_row * 4;
            const int64_t global_row = static_cast<int64_t>(trailing_row_start + tile_row);
            if (global_row < n) {
                const int swizzle = tile_row & 31;
                const float4 values = make_float4(
                    panel_rows[tile_row][(column + 0) ^ swizzle],
                    panel_rows[tile_row][(column + 1) ^ swizzle],
                    panel_rows[tile_row][(column + 2) ^ swizzle],
                    panel_rows[tile_row][(column + 3) ^ swizzle]);
                float* destination = matrix + global_row * n + panel_start + column;
                *reinterpret_cast<float4*>(destination) = values;
            }
        }
    } else {
        constexpr int tile_elements = kPanelThreads * kPanel;
        for (int index = thread; index < tile_elements; index += blockDim.x) {
            const int tile_row = index / kPanel;
            const int column = index % kPanel;
            const int64_t global_row = static_cast<int64_t>(trailing_row_start + tile_row);
            if (global_row < n) {
                matrix[global_row * n + panel_start + column] =
                    panel_rows[tile_row][column ^ (tile_row & 31)];
            }
        }
    }
}

// Apply one compile-time-width panel update to a unique 16x16 lower-triangular tile.
template <bool RecognizeStructure>
__global__ void update_trailing_scalar_kernel(
    float* output,
    int64_t n,
    int64_t matrix_elements,
    int panel_start) {
    __shared__ float panel_rows[kScalarUpdateTile][kPanel + 1];
    __shared__ float panel_columns[kScalarUpdateTile][kPanel + 1];

    // A tile strictly above the diagonal can exit uniformly before the block barrier.
    if (blockIdx.x > blockIdx.y) {
        return;
    }

    float* matrix = output + static_cast<int64_t>(blockIdx.z) * matrix_elements;
    if constexpr (RecognizeStructure) {
        if (matrix[n - 1] != 0.0f) {
            return;
        }
    }

    const int local_column = threadIdx.x;
    const int local_row = threadIdx.y;
    const int thread = local_row * kScalarUpdateTile + local_column;
    const int trailing_start = panel_start + kPanel;
    const int row_start = trailing_start + blockIdx.y * kScalarUpdateTile;
    const int column_start = trailing_start + blockIdx.x * kScalarUpdateTile;

    if ((n & 3) == 0) {
        constexpr int vectors_per_row = kPanel / 4;
        constexpr int vector_count = kScalarUpdateTile * vectors_per_row;
        // The fixed CTA stride clamps narrow panels and adds one partial wave for wide panels.
        for (
            int vector_index = thread;
            vector_index < vector_count;
            vector_index += kScalarUpdateTile * kScalarUpdateTile) {
            const int tile_row = vector_index / vectors_per_row;
            const int panel_column = vector_index % vectors_per_row * 4;
            const int64_t global_row = static_cast<int64_t>(row_start + tile_row);
            const int64_t global_column = static_cast<int64_t>(column_start + tile_row);
            float4 row_values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
            float4 column_values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
            if (global_row < n) {
                const float* source = matrix + global_row * n + panel_start + panel_column;
                row_values = *reinterpret_cast<const float4*>(source);
            }
            if (global_column < n) {
                const float* source = matrix + global_column * n + panel_start + panel_column;
                column_values = *reinterpret_cast<const float4*>(source);
            }

            // Scalar scatters preserve the padded shared rows used by the dot-product lanes.
            panel_rows[tile_row][panel_column + 0] = row_values.x;
            panel_rows[tile_row][panel_column + 1] = row_values.y;
            panel_rows[tile_row][panel_column + 2] = row_values.z;
            panel_rows[tile_row][panel_column + 3] = row_values.w;
            panel_columns[tile_row][panel_column + 0] = column_values.x;
            panel_columns[tile_row][panel_column + 1] = column_values.y;
            panel_columns[tile_row][panel_column + 2] = column_values.z;
            panel_columns[tile_row][panel_column + 3] = column_values.w;
        }
    } else {
        // Linear cooperative loads preserve arbitrary contiguous matrix row strides.
        for (
            int index = thread;
            index < kScalarUpdateTile * kPanel;
            index += blockDim.x * blockDim.y) {
            const int tile_row = index / kPanel;
            const int panel_column = index - tile_row * kPanel;
            const int64_t global_row = static_cast<int64_t>(row_start + tile_row);
            const int64_t global_column = static_cast<int64_t>(column_start + tile_row);

            panel_rows[tile_row][panel_column] = global_row < n
                ? matrix[global_row * n + panel_start + panel_column]
                : 0.0f;
            panel_columns[tile_row][panel_column] = global_column < n
                ? matrix[global_column * n + panel_start + panel_column]
                : 0.0f;
        }
    }

    // Padding each shared row by one word keeps column-varying consumers off one bank.
    __syncthreads();
    const int64_t row = static_cast<int64_t>(row_start + local_row);
    const int64_t column = static_cast<int64_t>(column_start + local_column);
    if (row < n && column < n && row >= column) {
        float value = matrix[row * n + column];
        #pragma unroll
        for (int inner = 0; inner < kPanel; ++inner) {
            value = fmaf(
                -panel_rows[local_row][inner],
                panel_columns[local_column][inner],
                value);
        }
        matrix[row * n + column] = value;
    }
}

// Split eight FP32 values and write each BF16 plane with one aligned 128-bit store.
__device__ __forceinline__ void split_bf16_eight(
    const float (&values)[8],
    uint16_t* high_destination,
    uint16_t* low_destination) {
    uint32_t packed_high[4];
    uint32_t packed_low[4];
    #pragma unroll
    for (int pair = 0; pair < 4; ++pair) {
        // PTX puts its first source in bits 31:16, so reverse the sources to
        // preserve increasing shared-memory element order on little-endian SM100.
        asm volatile(
            "cvt.rn.bf16x2.f32 %0, %1, %2;"
            : "=r"(packed_high[pair])
            : "f"(values[pair * 2 + 1]), "f"(values[pair * 2]));
        const float residual0 =
            values[pair * 2] - __uint_as_float(packed_high[pair] << 16);
        const float residual1 =
            values[pair * 2 + 1] -
            __uint_as_float(packed_high[pair] & 0xffff0000u);
        asm volatile(
            "cvt.rn.bf16x2.f32 %0, %1, %2;"
            : "=r"(packed_low[pair])
            : "f"(residual1), "f"(residual0));
    }
    *reinterpret_cast<uint4*>(high_destination) = make_uint4(
        packed_high[0], packed_high[1], packed_high[2], packed_high[3]);
    *reinterpret_cast<uint4*>(low_destination) = make_uint4(
        packed_low[0], packed_low[1], packed_low[2], packed_low[3]);
}

// Map one BF16 row/inner pair into a 32-byte-swizzled K-major shared tile.
__device__ __forceinline__ int tensor_shared_half(int row, int inner) {
    const int chunk = inner / kTensorInner;
    const int atom_inner = inner % kTensorInner;
    const int atom_word_inner = atom_inner / 2;
    const int atom_row = row % 8;
    const int row_half = atom_row / 4;
    const int inner_group = atom_word_inner / 4 ^ row_half;
    const int atom_word =
        row_half * 32 +
        ((atom_row % 4) * 2 + inner_group) * 4 +
        atom_word_inner % 4;
    return
        chunk * kTensorChunkHalfs +
        row / 8 * 128 +
        atom_word * 2 +
        atom_inner % 2;
}

// Describe one aligned 128x16 BF16 operand in Blackwell's shared-memory format.
__device__ __forceinline__ uint64_t tensor_shared_descriptor(const uint16_t* tile) {
    const uint32_t address = static_cast<uint32_t>(__cvta_generic_to_shared(tile));
    const uint64_t start = static_cast<uint64_t>(address >> 4) & 0x3fffull;
    const uint64_t leading = 1ull << 16;   // 16 bytes between K-major vector elements.
    const uint64_t stride = 16ull << 32;   // 256 bytes between eight-row atoms.
    const uint64_t version = 1ull << 46;
    const uint64_t layout = 6ull << 61;    // 32-byte swizzle.
    return start | leading | stride | version | layout;
}

// Issue one complete m128n128k16 BF16 tensor-core operation from a single thread.
__device__ __forceinline__ void issue_tensor_mma(
    uint32_t destination,
    uint64_t a_descriptor,
    uint64_t b_descriptor,
    bool accumulate) {
    constexpr uint32_t instruction =
        (1u << 4) | (1u << 7) | (1u << 10) | (16u << 17) | (8u << 24);
    const uint32_t zero = 0;
    asm volatile(
        "{\n\t"
        ".reg .pred use_input;\n\t"
        "setp.ne.b32 use_input, %4, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::f16 "
        "[%0], %1, %2, %3, {%5, %6, %7, %8}, use_input;\n\t"
        "}\n"
        :
        : "r"(destination),
          "l"(a_descriptor),
          "l"(b_descriptor),
          "r"(instruction),
          "r"(static_cast<uint32_t>(accumulate)),
          "r"(zero),
          "r"(zero),
          "r"(zero),
          "r"(zero)
        : "memory");
}

// Start one warp-collective load of eight adjacent FP32 tensor-memory columns.
__device__ __forceinline__ void load_tensor_eight(
    uint32_t address,
    uint32_t (&values)[8]) {
    asm volatile(
        "tcgen05.ld.sync.aligned.32x32b.x8.b32 "
        "{%0, %1, %2, %3, %4, %5, %6, %7}, [%8];\n"
        : "=r"(values[0]),
          "=r"(values[1]),
          "=r"(values[2]),
          "=r"(values[3]),
          "=r"(values[4]),
          "=r"(values[5]),
          "=r"(values[6]),
          "=r"(values[7])
        : "r"(address));
}

// Load one naturally aligned 32-byte global segment with SM100's native vector form.
__device__ __forceinline__ void load_global_eight(
    const float* source,
    float (&values)[8]) {
    asm volatile(
        "ld.global.v8.f32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];\n"
        : "=f"(values[0]),
          "=f"(values[1]),
          "=f"(values[2]),
          "=f"(values[3]),
          "=f"(values[4]),
          "=f"(values[5]),
          "=f"(values[6]),
          "=f"(values[7])
        : "l"(source)
        : "memory");
}

// Store one naturally aligned 32-byte output segment with SM100's native vector form.
__device__ __forceinline__ void store_global_eight(
    float* destination,
    const float (&values)[8]) {
    asm volatile(
        "st.global.v8.f32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8};\n"
        :
        : "l"(destination),
          "f"(values[0]),
          "f"(values[1]),
          "f"(values[2]),
          "f"(values[3]),
          "f"(values[4]),
          "f"(values[5]),
          "f"(values[6]),
          "f"(values[7])
        : "memory");
}

// Factor independent 32x32 matrices with one warp per matrix and one row per lane.
template <bool RecognizeDiagonal>
__global__ void factor_warp_32_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch) {
    const int lane = threadIdx.x % kWarpCholeskySize;
    const int warp = threadIdx.x / kWarpCholeskySize;
    const int matrix_index =
        static_cast<int>(blockIdx.x) * kWarpMatricesPerBlock + warp;
    if (matrix_index >= batch) {
        return;
    }

    constexpr int matrix_elements = kWarpCholeskySize * kWarpCholeskySize;
    const int64_t matrix_offset = static_cast<int64_t>(matrix_index) * matrix_elements;
    const int64_t row_offset = matrix_offset + lane * kWarpCholeskySize;
    const float* input_row = input + row_offset;
    float* output_row = output + row_offset;
    float row_values[kWarpCholeskySize];
    const bool aligned_segments =
        (reinterpret_cast<uintptr_t>(input) & (8 * sizeof(float) - 1)) == 0;

    // Row and matrix strides preserve the base alignment, so this branch is warp-uniform.
    if (aligned_segments) {
        #pragma unroll
        for (int segment = 0; segment < kWarpCholeskySize; segment += 8) {
            float values[8];
            load_global_eight(input_row + segment, values);
            #pragma unroll
            for (int item = 0; item < 8; ++item) {
                row_values[segment + item] = values[item];
            }
        }
    } else {
        #pragma unroll
        for (int column = 0; column < kWarpCholeskySize; ++column) {
            row_values[column] = input_row[column];
        }
    }

    if constexpr (RecognizeDiagonal) {
        bool diagonal_matrix = __all_sync(
            0xffffffffu,
            row_values[lane ^ 1] == 0.0f);
        if (diagonal_matrix) {
            bool row_is_diagonal = true;
            #pragma unroll
            for (int column = 0; column < kWarpCholeskySize; ++column) {
                row_is_diagonal &= column == lane || row_values[column] == 0.0f;
            }
            diagonal_matrix = __all_sync(0xffffffffu, row_is_diagonal);
        }

        if (diagonal_matrix) {
            const float diagonal = sqrtf(row_values[lane]);
            if (aligned_segments) {
                #pragma unroll
                for (int segment = 0; segment < kWarpCholeskySize; segment += 8) {
                    float values[8] = {};
                    #pragma unroll
                    for (int item = 0; item < 8; ++item) {
                        values[item] = segment + item == lane ? diagonal : 0.0f;
                    }
                    store_global_eight(output_row + segment, values);
                }
            } else {
                #pragma unroll
                for (int column = 0; column < kWarpCholeskySize; ++column) {
                    output_row[column] = column == lane ? diagonal : 0.0f;
                }
            }
            return;
        }
    }

    #pragma unroll
    for (int pivot = 0; pivot < kWarpCholeskySize; ++pivot) {
        float diagonal = 0.0f;
        if (lane == pivot) {
            diagonal = sqrtf(row_values[pivot]);
            row_values[pivot] = diagonal;
        }
        diagonal = __shfl_sync(0xffffffffu, diagonal, pivot);
        if (lane > pivot) {
            row_values[pivot] /= diagonal;
        }

        const float row_factor = row_values[pivot];
        #pragma unroll
        for (int column = pivot + 1; column < kWarpCholeskySize; ++column) {
            // Every lane executes the shuffle; divergence under a full mask is invalid.
            const float column_factor =
                __shfl_sync(0xffffffffu, row_factor, column);
            if (lane >= column) {
                row_values[column] = fmaf(
                    -row_factor,
                    column_factor,
                    row_values[column]);
            }
        }
    }

    // The store path matches the load legality and always makes the upper triangle zero.
    if (aligned_segments) {
        #pragma unroll
        for (int segment = 0; segment < kWarpCholeskySize; segment += 8) {
            float values[8];
            #pragma unroll
            for (int item = 0; item < 8; ++item) {
                const int column = segment + item;
                values[item] = lane >= column ? row_values[column] : 0.0f;
            }
            store_global_eight(output_row + segment, values);
        }
    } else {
        #pragma unroll
        for (int column = 0; column < kWarpCholeskySize; ++column) {
            output_row[column] = lane >= column ? row_values[column] : 0.0f;
        }
    }
}

// Factor one fixed-size matrix per CTA entirely inside padded shared memory.
template <
    int Size,
    int Microblock = kDiagonalMicroblock,
    bool RecognizeDiagonal = false>
__global__ void factor_full_shared_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
    static_assert(Size % Microblock == 0);
    static_assert(Size <= kFullShared64Threads);
    extern __shared__ float tile[];
    constexpr int row_stride = Size + 1;
    constexpr int matrix_elements = Size * Size;
    constexpr int vector_elements = 4;
    constexpr int vectors_per_row = Size / vector_elements;
    constexpr int matrix_vectors = Size * vectors_per_row;

    const int thread = threadIdx.x;
    const int64_t matrix_offset =
        static_cast<int64_t>(blockIdx.x) * matrix_elements;
    const float* input_matrix = input + matrix_offset;
    float* output_matrix = output + matrix_offset;
    const bool aligned_vectors =
        ((reinterpret_cast<uintptr_t>(input_matrix) |
          reinterpret_cast<uintptr_t>(output_matrix)) &
         (alignof(float4) - 1)) == 0;

    // The vector path reduces global instructions; shared rows retain scalar padding.
    if (aligned_vectors) {
        for (
            int vector_index = thread;
            vector_index < matrix_vectors;
            vector_index += blockDim.x) {
            const int row = vector_index / vectors_per_row;
            const int column =
                (vector_index - row * vectors_per_row) * vector_elements;
            const int global_index = row * Size + column;
            float4 values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);

            if (column <= row) {
                values =
                    *reinterpret_cast<const float4*>(input_matrix + global_index);
                values.y = row >= column + 1 ? values.y : 0.0f;
                values.z = row >= column + 2 ? values.z : 0.0f;
                values.w = row >= column + 3 ? values.w : 0.0f;
            }

            float* shared_row = tile + row * row_stride + column;
            shared_row[0] = values.x;
            shared_row[1] = values.y;
            shared_row[2] = values.z;
            shared_row[3] = values.w;
        }
    } else {
        // Contiguous views with unusual storage offsets retain the exact scalar path.
        for (int index = thread; index < matrix_elements; index += blockDim.x) {
            const int row = index / Size;
            const int column = index - row * Size;
            tile[row * row_stride + column] =
                row >= column ? input_matrix[index] : 0.0f;
        }
    }

    // Warp zero factors each microblock; the full CTA solves and updates the remainder.
    __syncthreads();
    if constexpr (RecognizeDiagonal) {
        bool diagonal_matrix = false;
        if (tile[row_stride] == 0.0f) {
            bool has_off_diagonal = false;
            for (int index = thread; index < matrix_elements; index += blockDim.x) {
                const int row = index / Size;
                const int column = index - row * Size;
                has_off_diagonal |=
                    row > column && tile[row * row_stride + column] != 0.0f;
            }
            diagonal_matrix = __syncthreads_count(has_off_diagonal) == 0;
        }
        if (diagonal_matrix) {
            for (int diagonal = thread; diagonal < Size; diagonal += blockDim.x) {
                tile[diagonal * row_stride + diagonal] =
                    sqrtf(tile[diagonal * row_stride + diagonal]);
            }
            __syncthreads();

            if (aligned_vectors) {
                for (
                    int vector_index = thread;
                    vector_index < matrix_vectors;
                    vector_index += blockDim.x) {
                    const int row = vector_index / vectors_per_row;
                    const int column =
                        (vector_index - row * vectors_per_row) * vector_elements;
                    const float* shared_row = tile + row * row_stride + column;
                    const float4 values = make_float4(
                        shared_row[0],
                        shared_row[1],
                        shared_row[2],
                        shared_row[3]);
                    *reinterpret_cast<float4*>(
                        output_matrix + row * Size + column) = values;
                }
            } else {
                for (int index = thread; index < matrix_elements; index += blockDim.x) {
                    const int row = index / Size;
                    const int column = index - row * Size;
                    output_matrix[index] = tile[row * row_stride + column];
                }
            }
            return;
        }
    }

    for (
        int microblock_start = 0;
        microblock_start < Size;
        microblock_start += Microblock) {
        const int microblock_end = microblock_start + Microblock;

        if (thread < warpSize) {
            const int lane = thread;
            for (int pivot = microblock_start; pivot < microblock_end; ++pivot) {
                const int pivot_owner = pivot - microblock_start;
                float diagonal = 0.0f;
                if (lane == pivot_owner) {
                    diagonal = sqrtf(tile[pivot * row_stride + pivot]);
                    tile[pivot * row_stride + pivot] = diagonal;
                }
                diagonal = __shfl_sync(0xffffffffu, diagonal, pivot_owner);

                if (pivot + 1 < microblock_end) {
                    for (
                        int row = pivot + 1 + lane;
                        row < microblock_end;
                        row += warpSize) {
                        tile[row * row_stride + pivot] /= diagonal;
                    }

                    // The rank-1 update consumes pivot values produced by other lanes.
                    __syncwarp();
                    for (
                        int column = pivot + 1 + lane;
                        column < microblock_end;
                        column += warpSize) {
                        const float column_pivot =
                            tile[column * row_stride + pivot];
                        for (int row = column; row < microblock_end; ++row) {
                            const int element = row * row_stride + column;
                            tile[element] = fmaf(
                                -tile[row * row_stride + pivot],
                                column_pivot,
                                tile[element]);
                        }
                    }

                    // The next pivot reads a diagonal updated by another lane.
                    __syncwarp();
                }
            }
        }

        // Trailing rows cannot consume the diagonal block until warp zero publishes it.
        __syncthreads();
        if (microblock_end == Size) {
            break;
        }

        const int solve_row = microblock_end + thread;
        if (solve_row < Size) {
            float solved_fragment[Microblock];
            #pragma unroll
            for (int local_column = 0; local_column < Microblock; ++local_column) {
                const int column = microblock_start + local_column;
                const int element = solve_row * row_stride + column;
                float value = tile[element];
                #pragma unroll
                for (int local_previous = 0; local_previous < local_column; ++local_previous) {
                    const int previous = microblock_start + local_previous;
                    value = fmaf(
                        -solved_fragment[local_previous],
                        tile[column * row_stride + previous],
                        value);
                }
                value /= tile[column * row_stride + column];
                solved_fragment[local_column] = value;
                tile[element] = value;
            }
        }

        // Every lower trailing element reads rows solved by potentially different warps.
        __syncthreads();
        constexpr int warp_threads = 32;
        constexpr int update_warps = kFullShared64Threads / warp_threads;
        static_assert(kFullShared64Threads == kFullShared128Threads);
        const int update_warp = thread >> 5;
        const int update_lane = thread & (warp_threads - 1);

        // Each warp owns columns; padded rows make lanes' downward walk bank-consecutive.
        for (
            int column = microblock_end + update_warp;
            column < Size;
            column += update_warps) {
            for (
                int row = column + update_lane;
                row < Size;
                row += warp_threads) {
                const int element = row * row_stride + column;
                float value = tile[element];
                #pragma unroll
                for (
                    int local_inner = 0;
                    local_inner < Microblock;
                    ++local_inner) {
                    const int inner = microblock_start + local_inner;
                    value = fmaf(
                        -tile[row * row_stride + inner],
                        tile[column * row_stride + inner],
                        value);
                }
                tile[element] = value;
            }
        }

        // The next diagonal microblock consumes the complete rank update.
        __syncthreads();
    }

    // Every output word is written once, including exact zeros in the upper triangle.
    if (aligned_vectors) {
        for (
            int vector_index = thread;
            vector_index < matrix_vectors;
            vector_index += blockDim.x) {
            const int row = vector_index / vectors_per_row;
            const int column =
                (vector_index - row * vectors_per_row) * vector_elements;
            const float* shared_row = tile + row * row_stride + column;
            const float4 values = make_float4(
                shared_row[0],
                shared_row[1],
                shared_row[2],
                shared_row[3]);
            *reinterpret_cast<float4*>(
                output_matrix + row * Size + column) = values;
        }
    } else {
        for (int index = thread; index < matrix_elements; index += blockDim.x) {
            const int row = index / Size;
            const int column = index - row * Size;
            output_matrix[index] = tile[row * row_stride + column];
        }
    }
}

// Apply one first-order compensated BF16 product into a single 128x128 accumulator.
__global__ void update_trailing_tensor_kernel(
    float* output,
    int64_t n,
    int64_t matrix_elements,
    int panel_start) {
    // A tile strictly above the diagonal can exit before any collective instruction.
    if (blockIdx.x > blockIdx.y) {
        return;
    }

    extern __shared__ __align__(256) uint8_t shared_storage[];
    uint16_t* a_high = reinterpret_cast<uint16_t*>(shared_storage);
    uint16_t* a_low = a_high + kTensorOperandHalfs;
    uint16_t* b_high = a_low + kTensorOperandHalfs;
    uint16_t* b_low = b_high + kTensorOperandHalfs;
    uint64_t* completion = reinterpret_cast<uint64_t*>(b_low + kTensorOperandHalfs);
    uint32_t* tensor_slot = reinterpret_cast<uint32_t*>(completion + 1);

    const int thread = threadIdx.x;
    const uint32_t tensor_slot_address =
        static_cast<uint32_t>(__cvta_generic_to_shared(tensor_slot));
    const uint32_t completion_address =
        static_cast<uint32_t>(__cvta_generic_to_shared(completion));

    // Exactly one full warp owns allocation and deallocation of this CTA's tensor memory.
    if (thread < 32) {
        asm volatile(
            "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
            :
            : "r"(tensor_slot_address), "r"(kTensorMemoryColumns)
            : "memory");
    }
    if (thread == 0) {
        asm volatile(
            "mbarrier.init.shared::cta.b64 [%0], 1;\n\t"
            "fence.mbarrier_init.release.cluster;"
            :
            : "r"(completion_address)
            : "memory");
    }
    __syncthreads();

    const int trailing_start = panel_start + kPanel;
    const int row_start = trailing_start + blockIdx.y * kTensorUpdateTile;
    const int column_start = trailing_start + blockIdx.x * kTensorUpdateTile;
    const bool diagonal_tile = blockIdx.x == blockIdx.y;
    float* matrix = output + static_cast<int64_t>(blockIdx.z) * matrix_elements;

    // Each owner converts 16 adjacent values after two native 32-byte loads.  In
    // rows 4..7 of each atom, the 32-byte swizzle reverses the two eight-value halves.
    constexpr int vectors_per_row = kPanel / 16;
    const bool aligned_vectors =
        (n & 7) == 0 &&
        (reinterpret_cast<uintptr_t>(matrix) & (8 * sizeof(float) - 1)) == 0;
    for (int vector = thread; vector < kTensorOperandHalfs / 16; vector += blockDim.x) {
        const int tile_row = vector / vectors_per_row;
        const int inner = (vector - tile_row * vectors_per_row) * 16;
        const int shared_half = tensor_shared_half(tile_row, inner);
        const int segment_stride = (tile_row & 4) == 0 ? 8 : -8;
        const int64_t a_row = static_cast<int64_t>(row_start + tile_row);
        const int64_t b_row = static_cast<int64_t>(column_start + tile_row);
        #pragma unroll
        for (int segment = 0; segment < 2; ++segment) {
            float a_values[8] = {};
            float b_values[8] = {};
            if (a_row < n) {
                const float* source =
                    matrix + a_row * n + panel_start + inner + segment * 8;
                if (aligned_vectors) {
                    load_global_eight(source, a_values);
                } else {
                    #pragma unroll
                    for (int item = 0; item < 8; ++item) {
                        a_values[item] = source[item];
                    }
                }
            }
            if (!diagonal_tile && b_row < n) {
                const float* source =
                    matrix + b_row * n + panel_start + inner + segment * 8;
                if (aligned_vectors) {
                    load_global_eight(source, b_values);
                } else {
                    #pragma unroll
                    for (int item = 0; item < 8; ++item) {
                        b_values[item] = source[item];
                    }
                }
            }
            const int destination = shared_half + segment * segment_stride;
            split_bf16_eight(
                a_values,
                a_high + destination,
                a_low + destination);
            if (!diagonal_tile) {
                split_bf16_eight(
                    b_values,
                    b_high + destination,
                    b_low + destination);
            }
        }
    }
    __syncthreads();

    const uint32_t tensor_base = *tensor_slot;
    if (thread == 0) {
        // The async tensor proxy must observe the generic shared-memory stores above.
        asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
        const uint64_t a_high_descriptor = tensor_shared_descriptor(a_high);
        const uint64_t a_low_descriptor = tensor_shared_descriptor(a_low);
        const uint64_t b_high_descriptor =
            diagonal_tile ? a_high_descriptor : tensor_shared_descriptor(b_high);
        const uint64_t b_low_descriptor =
            diagonal_tile ? a_low_descriptor : tensor_shared_descriptor(b_low);
        constexpr uint64_t descriptor_step =
            kTensorChunkWords * sizeof(uint32_t) / 16;

        // Accumulate the main product and both first-order cross terms in one TMEM tile.
        for (int chunk = 0; chunk < kPanel / kTensorInner; ++chunk) {
            const uint64_t offset = chunk * descriptor_step;
            issue_tensor_mma(
                tensor_base,
                a_high_descriptor + offset,
                b_high_descriptor + offset,
                chunk != 0);
        }
        for (int chunk = 0; chunk < kPanel / kTensorInner; ++chunk) {
            const uint64_t offset = chunk * descriptor_step;
            issue_tensor_mma(
                tensor_base,
                a_high_descriptor + offset,
                b_low_descriptor + offset,
                true);
        }
        for (int chunk = 0; chunk < kPanel / kTensorInner; ++chunk) {
            const uint64_t offset = chunk * descriptor_step;
            issue_tensor_mma(
                tensor_base,
                a_low_descriptor + offset,
                b_high_descriptor + offset,
                true);
        }
        asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 "
            "[%0];"
            :
            : "r"(completion_address)
            : "memory");
    }

    // Every consumer waits for the asynchronous MMAs before issuing its collective loads.
    asm volatile(
        "{\n\t"
        ".reg .pred done;\n\t"
        "WAIT_MMA:\n\t"
        "mbarrier.try_wait.parity.shared::cta.b64 done, [%0], 0;\n\t"
        "@!done bra WAIT_MMA;\n\t"
        "}\n"
        :
        : "r"(completion_address)
        : "memory");
    asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

    const int warp = thread / 32;
    const int lane = thread % 32;
    const int64_t row = static_cast<int64_t>(row_start + warp * 32 + lane);
    const uint32_t tensor_row = tensor_base + (static_cast<uint32_t>(warp * 32) << 16);
    for (int tile_column = 0; tile_column < kTensorUpdateTile; tile_column += 8) {
        uint32_t values[8];
        load_tensor_eight(tensor_row + tile_column, values);
        asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

        // The first-order dot product is subtracted once in ordinary FP32 arithmetic.
        const int64_t column = static_cast<int64_t>(column_start + tile_column);
        // The stride guard makes every full-eight row segment naturally 32-byte aligned.
        if ((n & 7) == 0 && row < n && column + 7 < n && row >= column + 7) {
            float* destination = matrix + row * n + column;
            float current[8];
            load_global_eight(destination, current);
            #pragma unroll
            for (int item = 0; item < 8; ++item) {
                current[item] -= __uint_as_float(values[item]);
            }
            store_global_eight(destination, current);
        } else {
            #pragma unroll
            for (int item = 0; item < 8; ++item) {
                const int64_t item_column = column + item;
                if (row < n && item_column < n && row >= item_column) {
                    matrix[row * n + item_column] -= __uint_as_float(values[item]);
                }
            }
        }
    }

    // All tensor-memory reads must finish before the owning warp returns the allocation.
    asm volatile("tcgen05.fence::before_thread_sync;" ::: "memory");
    __syncthreads();
    if (thread < 32) {
        asm volatile(
            "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n\t"
            "tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;"
            :
            : "r"(tensor_base), "r"(kTensorMemoryColumns));
    }
}

// Quantize each solved panel value once before any trailing-update CTA consumes it.
__global__ void quantize_panel_kernel(
    const float* output,
    uint16_t* quantized_high,
    uint16_t* quantized_low,
    int64_t n,
    int64_t matrix_elements,
    int64_t workspace_matrix_elements,
    int panel_start,
    int trailing) {
    constexpr int vectors_per_row = kPanel / 8;
    const int64_t vector_index =
        static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    const int64_t vector_count = static_cast<int64_t>(trailing) * vectors_per_row;
    if (vector_index >= vector_count) {
        return;
    }

    const int64_t trailing_row = vector_index / vectors_per_row;
    const int inner =
        static_cast<int>(vector_index - trailing_row * vectors_per_row) * 8;
    const int64_t global_row = panel_start + kPanel + trailing_row;
    const int64_t matrix = blockIdx.y;
    const float* source =
        output + matrix * matrix_elements + global_row * n + panel_start + inner;
    float values[8];
    if ((n & 3) == 0) {
        const float4 first = *reinterpret_cast<const float4*>(source);
        const float4 second = *reinterpret_cast<const float4*>(source + 4);
        values[0] = first.x;
        values[1] = first.y;
        values[2] = first.z;
        values[3] = first.w;
        values[4] = second.x;
        values[5] = second.y;
        values[6] = second.z;
        values[7] = second.w;
    } else {
        #pragma unroll
        for (int item = 0; item < 8; ++item) {
            values[item] = source[item];
        }
    }

    const int64_t destination =
        matrix * workspace_matrix_elements + global_row * kPanel + inner;
    split_bf16_eight(
        values,
        quantized_high + destination,
        quantized_low + destination);
}

// Apply one first-order compensated BF16 product into a single 128x128 accumulator.
__global__ void update_trailing_tensor_workspace_kernel(
    float* output,
    const uint16_t* quantized_high,
    const uint16_t* quantized_low,
    int64_t n,
    int64_t matrix_elements,
    int64_t workspace_matrix_elements,
    int panel_start) {
    // A tile strictly above the diagonal can exit before any collective instruction.
    if (blockIdx.x > blockIdx.y) {
        return;
    }

    extern __shared__ __align__(256) uint8_t shared_storage[];
    uint16_t* a_high = reinterpret_cast<uint16_t*>(shared_storage);
    uint16_t* a_low = a_high + kTensorOperandHalfs;
    uint16_t* b_high = a_low + kTensorOperandHalfs;
    uint16_t* b_low = b_high + kTensorOperandHalfs;
    uint64_t* completion = reinterpret_cast<uint64_t*>(b_low + kTensorOperandHalfs);
    uint32_t* tensor_slot = reinterpret_cast<uint32_t*>(completion + 1);

    const int thread = threadIdx.x;
    const uint32_t tensor_slot_address =
        static_cast<uint32_t>(__cvta_generic_to_shared(tensor_slot));
    const uint32_t completion_address =
        static_cast<uint32_t>(__cvta_generic_to_shared(completion));

    // Exactly one full warp owns allocation and deallocation of this CTA's tensor memory.
    if (thread < 32) {
        asm volatile(
            "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
            :
            : "r"(tensor_slot_address), "r"(kTensorMemoryColumns)
            : "memory");
    }
    if (thread == 0) {
        asm volatile(
            "mbarrier.init.shared::cta.b64 [%0], 1;\n\t"
            "fence.mbarrier_init.release.cluster;"
            :
            : "r"(completion_address)
            : "memory");
    }
    __syncthreads();

    const int trailing_start = panel_start + kPanel;
    const int row_start = trailing_start + blockIdx.y * kTensorUpdateTile;
    const int column_start = trailing_start + blockIdx.x * kTensorUpdateTile;
    const bool diagonal_tile = blockIdx.x == blockIdx.y;
    float* matrix = output + static_cast<int64_t>(blockIdx.z) * matrix_elements;
    const uint16_t* matrix_high =
        quantized_high + static_cast<int64_t>(blockIdx.z) * workspace_matrix_elements;
    const uint16_t* matrix_low =
        quantized_low + static_cast<int64_t>(blockIdx.z) * workspace_matrix_elements;

    // The solve kernel quantized this panel once. Each owner now copies two aligned
    // eight-value fragments into the swizzled tensor-core operand layout.
    constexpr int vectors_per_row = kPanel / 16;
    for (int vector = thread; vector < kTensorOperandHalfs / 16; vector += blockDim.x) {
        const int tile_row = vector / vectors_per_row;
        const int inner = (vector - tile_row * vectors_per_row) * 16;
        const int shared_half = tensor_shared_half(tile_row, inner);
        const int segment_stride = (tile_row & 4) == 0 ? 8 : -8;
        const int64_t a_row = static_cast<int64_t>(row_start + tile_row);
        const int64_t b_row = static_cast<int64_t>(column_start + tile_row);
        #pragma unroll
        for (int segment = 0; segment < 2; ++segment) {
            uint4 a_high_values = make_uint4(0, 0, 0, 0);
            uint4 a_low_values = make_uint4(0, 0, 0, 0);
            uint4 b_high_values = make_uint4(0, 0, 0, 0);
            uint4 b_low_values = make_uint4(0, 0, 0, 0);
            if (a_row < n) {
                const int64_t source = a_row * kPanel + inner + segment * 8;
                a_high_values = *reinterpret_cast<const uint4*>(matrix_high + source);
                a_low_values = *reinterpret_cast<const uint4*>(matrix_low + source);
            }
            if (!diagonal_tile && b_row < n) {
                const int64_t source = b_row * kPanel + inner + segment * 8;
                b_high_values = *reinterpret_cast<const uint4*>(matrix_high + source);
                b_low_values = *reinterpret_cast<const uint4*>(matrix_low + source);
            }
            const int destination = shared_half + segment * segment_stride;
            *reinterpret_cast<uint4*>(a_high + destination) = a_high_values;
            *reinterpret_cast<uint4*>(a_low + destination) = a_low_values;
            if (!diagonal_tile) {
                *reinterpret_cast<uint4*>(b_high + destination) = b_high_values;
                *reinterpret_cast<uint4*>(b_low + destination) = b_low_values;
            }
        }
    }
    __syncthreads();

    const uint32_t tensor_base = *tensor_slot;
    if (thread == 0) {
        // The async tensor proxy must observe the generic shared-memory stores above.
        asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
        const uint64_t a_high_descriptor = tensor_shared_descriptor(a_high);
        const uint64_t a_low_descriptor = tensor_shared_descriptor(a_low);
        const uint64_t b_high_descriptor =
            diagonal_tile ? a_high_descriptor : tensor_shared_descriptor(b_high);
        const uint64_t b_low_descriptor =
            diagonal_tile ? a_low_descriptor : tensor_shared_descriptor(b_low);
        constexpr uint64_t descriptor_step =
            kTensorChunkWords * sizeof(uint32_t) / 16;

        // Accumulate the main product and both first-order cross terms in one TMEM tile.
        for (int chunk = 0; chunk < kPanel / kTensorInner; ++chunk) {
            const uint64_t offset = chunk * descriptor_step;
            issue_tensor_mma(
                tensor_base,
                a_high_descriptor + offset,
                b_high_descriptor + offset,
                chunk != 0);
        }
        for (int chunk = 0; chunk < kPanel / kTensorInner; ++chunk) {
            const uint64_t offset = chunk * descriptor_step;
            issue_tensor_mma(
                tensor_base,
                a_high_descriptor + offset,
                b_low_descriptor + offset,
                true);
        }
        for (int chunk = 0; chunk < kPanel / kTensorInner; ++chunk) {
            const uint64_t offset = chunk * descriptor_step;
            issue_tensor_mma(
                tensor_base,
                a_low_descriptor + offset,
                b_high_descriptor + offset,
                true);
        }
        asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 "
            "[%0];"
            :
            : "r"(completion_address)
            : "memory");
    }

    // Every consumer waits for the asynchronous MMAs before issuing its collective loads.
    asm volatile(
        "{\n\t"
        ".reg .pred done;\n\t"
        "WAIT_MMA:\n\t"
        "mbarrier.try_wait.parity.shared::cta.b64 done, [%0], 0;\n\t"
        "@!done bra WAIT_MMA;\n\t"
        "}\n"
        :
        : "r"(completion_address)
        : "memory");
    asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

    const int warp = thread / 32;
    const int lane = thread % 32;
    const int64_t row = static_cast<int64_t>(row_start + warp * 32 + lane);
    const uint32_t tensor_row = tensor_base + (static_cast<uint32_t>(warp * 32) << 16);
    for (int tile_column = 0; tile_column < kTensorUpdateTile; tile_column += 8) {
        uint32_t values[8];
        load_tensor_eight(tensor_row + tile_column, values);
        asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

        // The first-order dot product is subtracted once in ordinary FP32 arithmetic.
        const int64_t column = static_cast<int64_t>(column_start + tile_column);
        // The stride guard makes every full-eight row segment naturally 32-byte aligned.
        if ((n & 7) == 0 && row < n && column + 7 < n && row >= column + 7) {
            float* destination = matrix + row * n + column;
            float current[8];
            load_global_eight(destination, current);
            #pragma unroll
            for (int item = 0; item < 8; ++item) {
                current[item] -= __uint_as_float(values[item]);
            }
            store_global_eight(destination, current);
        } else {
            #pragma unroll
            for (int item = 0; item < 8; ++item) {
                const int64_t item_column = column + item;
                if (row < n && item_column < n && row >= item_column) {
                    matrix[row * n + item_column] -= __uint_as_float(values[item]);
                }
            }
        }
    }

    // All tensor-memory reads must finish before the owning warp returns the allocation.
    asm volatile("tcgen05.fence::before_thread_sync;" ::: "memory");
    __syncthreads();
    if (thread < 32) {
        asm volatile(
            "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n\t"
            "tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;"
            :
            : "r"(tensor_base), "r"(kTensorMemoryColumns));
    }
}


}  // namespace

// Validate the fixed ABI, allocate independent output, and enqueue each dependent panel stage.
at::Tensor CHOLESKY_ENTRYPOINT(at::Tensor data) {
    TORCH_CHECK(data.is_cuda(), "data must be a CUDA tensor");
    TORCH_CHECK(data.scalar_type() == at::kFloat, "data must have dtype torch.float32");
    TORCH_CHECK(data.dim() == 3, "data must have shape [batch, n, n]");
    TORCH_CHECK(data.size(1) == data.size(2), "data matrices must be square");
    TORCH_CHECK(data.size(0) > 0 && data.size(1) > 0, "batch and n must be positive");
    TORCH_CHECK(data.is_contiguous(), "data must be contiguous");
    TORCH_CHECK(data.size(0) <= 65535, "batch exceeds the CUDA y/z grid limit");
    TORCH_CHECK(data.size(1) <= INT_MAX, "n exceeds the kernel index range");

    // The guard makes allocation and launches follow the input on non-default devices.
    const c10::cuda::CUDAGuard device_guard(data.device());
    at::Tensor output = at::empty_like(data);

    const int64_t batch64 = data.size(0);
    const int64_t n64 = data.size(1);
    const int batch = static_cast<int>(batch64);
    const int n = static_cast<int>(n64);

    if (n == kWarpCholeskySize) {
        const int blocks =
            (batch + kWarpMatricesPerBlock - 1) / kWarpMatricesPerBlock;
#if CHOLESKY_PANEL == 64
        if (batch <= 32) {
            factor_warp_32_kernel<true><<<
                blocks,
                kWarpMatricesPerBlock * kWarpCholeskySize>>>(
                data.data_ptr<float>(),
                output.data_ptr<float>(),
                batch);
        } else
#endif
        {
            factor_warp_32_kernel<false><<<
                blocks,
                kWarpMatricesPerBlock * kWarpCholeskySize>>>(
                data.data_ptr<float>(),
                output.data_ptr<float>(),
                batch);
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return output;
    }

    if (n == kFullShared64Size) {
#if CHOLESKY_PANEL == 64
        if (batch <= 16) {
            factor_full_shared_kernel<
                kFullShared64Size,
                kDiagonalMicroblock,
                true><<<
                batch,
                kFullShared64Threads,
                kFullShared64Bytes>>>(
                data.data_ptr<float>(),
                output.data_ptr<float>());
        } else
#endif
        {
            factor_full_shared_kernel<kFullShared64Size><<<
                batch,
                kFullShared64Threads,
                kFullShared64Bytes>>>(
                data.data_ptr<float>(),
                output.data_ptr<float>());
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return output;
    }

    if (n == kFullShared128Size) {
        // The padded 128-row tile exceeds the default 48 KiB dynamic shared limit.
#if CHOLESKY_PANEL == 64
        if (batch <= 8) {
            static const cudaError_t structured_shared_memory_status =
                cudaFuncSetAttribute(
                    factor_full_shared_kernel<kFullShared128Size, 4, true>,
                    cudaFuncAttributeMaxDynamicSharedMemorySize,
                    kFullShared128Bytes);
            TORCH_CHECK(
                structured_shared_memory_status == cudaSuccess,
                "could not configure structured n=128 shared memory: ",
                cudaGetErrorString(structured_shared_memory_status));
            factor_full_shared_kernel<kFullShared128Size, 4, true><<<
                batch,
                kFullShared128Threads,
                kFullShared128Bytes>>>(
                data.data_ptr<float>(),
                output.data_ptr<float>());
            C10_CUDA_KERNEL_LAUNCH_CHECK();
            return output;
        }
#endif
        static const cudaError_t shared_memory_status = cudaFuncSetAttribute(
            factor_full_shared_kernel<kFullShared128Size, 4>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            kFullShared128Bytes);
        TORCH_CHECK(
            shared_memory_status == cudaSuccess,
            "could not configure n=128 shared memory: ",
            cudaGetErrorString(shared_memory_status));
        factor_full_shared_kernel<kFullShared128Size, 4><<<
            batch,
            kFullShared128Threads,
            kFullShared128Bytes>>>(
            data.data_ptr<float>(),
            output.data_ptr<float>());
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return output;
    }

    const int64_t matrix_elements = n64 * n64;
    const int64_t workspace_matrix_elements = n64 * kPanel;
    const bool needs_quantized_workspace = n >= 8192;
    at::Tensor quantized_workspace;
    uint16_t* quantized_high = nullptr;
    uint16_t* quantized_low = nullptr;
    if (needs_quantized_workspace) {
        const int64_t workspace_plane_elements = batch64 * workspace_matrix_elements;
        quantized_workspace = at::empty({workspace_plane_elements}, data.options());
        quantized_high = reinterpret_cast<uint16_t*>(quantized_workspace.data_ptr<float>());
        quantized_low = quantized_high + workspace_plane_elements;
    }
    // Match diagonal parallelism to the measured B200 panel width and batch regimes.
    int diagonal_threads = 512;
    if (n == 64 || (n == 512 && batch >= 512)) {
        diagonal_threads = 256;
    }
    const int initialize_row_blocks = static_cast<int>(
        (n64 + kInitializeRowsPerBlock - 1) / kInitializeRowsPerBlock);
    const dim3 initialize_grid(initialize_row_blocks, batch);
#if CHOLESKY_PANEL == 64
    const bool recognize_structure = n == kStructuredSize && batch <= 4;
    if (recognize_structure) {
        initialize_lower_structured_kernel<<<initialize_grid, kInitializeThreads>>>(
            data.data_ptr<float>(),
            output.data_ptr<float>());
    } else
#endif
    {
        initialize_lower_kernel<<<initialize_grid, kInitializeThreads>>>(
            data.data_ptr<float>(),
            output.data_ptr<float>(),
            n64);
    }
    C10_CUDA_KERNEL_LAUNCH_CHECK();

    for (int panel_start = 0; panel_start < n; panel_start += kPanel) {
        const int panel_width = std::min(kPanel, n - panel_start);
#if CHOLESKY_PANEL == 64
        if (recognize_structure) {
            factor_diagonal_kernel<true><<<batch, diagonal_threads>>>(
                output.data_ptr<float>(),
                n64,
                matrix_elements,
                panel_start,
                panel_width);
        } else
#endif
        {
            factor_diagonal_kernel<false><<<batch, diagonal_threads>>>(
                output.data_ptr<float>(),
                n64,
                matrix_elements,
                panel_start,
                panel_width);
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        const int trailing = n - panel_start - panel_width;
        if (trailing == 0) {
            continue;
        }

        // Any non-final panel is complete, so both hot kernels use their fixed-width loop.
        TORCH_INTERNAL_ASSERT(panel_width == kPanel);
        const dim3 panel_grid(
            (trailing + kPanelThreads - 1) / kPanelThreads,
            batch);
        // CUDA requires explicit dynamic allocation and opt-in above 48 KiB per block.
        static const cudaError_t solve_shared_status = cudaFuncSetAttribute(
            solve_panel_kernel<false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            kSolvePanelSharedBytes);
#if CHOLESKY_PANEL == 64
        static const cudaError_t structured_solve_shared_status = cudaFuncSetAttribute(
            solve_panel_kernel<true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            kSolvePanelSharedBytes);
#endif
        TORCH_CHECK(
#if CHOLESKY_PANEL == 64
            (recognize_structure
                ? structured_solve_shared_status
                : solve_shared_status) == cudaSuccess,
#else
            solve_shared_status == cudaSuccess,
#endif
            "could not configure solve-panel shared memory: ",
#if CHOLESKY_PANEL == 64
            cudaGetErrorString(recognize_structure
                ? structured_solve_shared_status
                : solve_shared_status));
#else
            cudaGetErrorString(solve_shared_status));
#endif
#if CHOLESKY_PANEL == 64
        if (recognize_structure) {
            solve_panel_kernel<true><<<
                panel_grid,
                kPanelThreads,
                kSolvePanelSharedBytes>>>(
                output.data_ptr<float>(),
                n64,
                matrix_elements,
                panel_start);
        } else
#endif
        {
            solve_panel_kernel<false><<<
                panel_grid,
                kPanelThreads,
                kSolvePanelSharedBytes>>>(
                output.data_ptr<float>(),
                n64,
                matrix_elements,
                panel_start);
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        // Scalar FP32 wins below the measured tensor-core trailing-size crossover.
        if (trailing < kTensorDispatchThreshold) {
            const int update_tiles =
                (trailing + kScalarUpdateTile - 1) / kScalarUpdateTile;
            const dim3 update_grid(update_tiles, update_tiles, batch);
            const dim3 update_block(kScalarUpdateTile, kScalarUpdateTile);
#if CHOLESKY_PANEL == 64
            if (recognize_structure) {
                update_trailing_scalar_kernel<true><<<update_grid, update_block>>>(
                    output.data_ptr<float>(),
                    n64,
                    matrix_elements,
                    panel_start);
            } else
#endif
            {
                update_trailing_scalar_kernel<false><<<update_grid, update_block>>>(
                    output.data_ptr<float>(),
                    n64,
                    matrix_elements,
                    panel_start);
            }
        } else {
            if (!needs_quantized_workspace) {
                static const cudaError_t direct_shared_memory_status = cudaFuncSetAttribute(
                    update_trailing_tensor_kernel,
                    cudaFuncAttributeMaxDynamicSharedMemorySize,
                    kTensorSharedBytes);
                TORCH_CHECK(
                    direct_shared_memory_status == cudaSuccess,
                    "could not configure direct tcgen05 shared memory: ",
                    cudaGetErrorString(direct_shared_memory_status));
                const int direct_update_tiles =
                    (trailing + kTensorUpdateTile - 1) / kTensorUpdateTile;
                const dim3 direct_update_grid(
                    direct_update_tiles,
                    direct_update_tiles,
                    batch);
                update_trailing_tensor_kernel<<<
                    direct_update_grid,
                    kTensorThreads,
                    kTensorSharedBytes>>>(
                    output.data_ptr<float>(),
                    n64,
                    matrix_elements,
                    panel_start);
                C10_CUDA_KERNEL_LAUNCH_CHECK();
                continue;
            }

            constexpr int quantize_threads = 256;
            const int64_t quantize_vectors =
                static_cast<int64_t>(trailing) * (kPanel / 8);
            const dim3 quantize_grid(
                static_cast<unsigned int>(
                    (quantize_vectors + quantize_threads - 1) / quantize_threads),
                batch);
            quantize_panel_kernel<<<quantize_grid, quantize_threads>>>(
                output.data_ptr<float>(),
                quantized_high,
                quantized_low,
                n64,
                matrix_elements,
                workspace_matrix_elements,
                panel_start,
                trailing);
            C10_CUDA_KERNEL_LAUNCH_CHECK();

            // Opt in once to the compile-time operand buffer required by the tensor tile.
            static const cudaError_t shared_memory_status = cudaFuncSetAttribute(
                update_trailing_tensor_workspace_kernel,
                cudaFuncAttributeMaxDynamicSharedMemorySize,
                kTensorSharedBytes);
            TORCH_CHECK(
                shared_memory_status == cudaSuccess,
                "could not configure tcgen05 shared memory: ",
                cudaGetErrorString(shared_memory_status));
            const int update_tiles =
                (trailing + kTensorUpdateTile - 1) / kTensorUpdateTile;
            const dim3 update_grid(update_tiles, update_tiles, batch);
            update_trailing_tensor_workspace_kernel<<<
                update_grid,
                kTensorThreads,
                kTensorSharedBytes>>>(
                output.data_ptr<float>(),
                quantized_high,
                quantized_low,
                n64,
                matrix_elements,
                workspace_matrix_elements,
                panel_start);
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();
    }

    return output;
}

"""

    def _load_cholesky_variant(panel: int):
        """Compile one panel width under a unique extension and exported function name."""
        function_name = f"blocked_cholesky_cuda_panel_{panel}"
        extension = load_inline(
            name=function_name,
            cpp_sources=f"at::Tensor {function_name}(at::Tensor data);",
            cuda_sources=cuda_source,
            functions=[function_name],
            extra_cflags=["-O3"],
            extra_cuda_cflags=[
                "-O3",
                "-gencode=arch=compute_100a,code=sm_100a",
                f"-DCHOLESKY_PANEL={panel}",
                f"-DCHOLESKY_ENTRYPOINT={function_name}",
            ],
        )
        return getattr(extension, function_name)

    blocked_cholesky_cuda = _load_cholesky_variant(64)
    blocked_cholesky_cuda_medium = _load_cholesky_variant(96)
    blocked_cholesky_cuda_large = _load_cholesky_variant(128)


def custom_kernel(data: input_t) -> output_t:
    """Run the custom blocked CUDA factorization."""
    if data.shape[-1] == 16384:
        return blocked_cholesky_cuda_medium(data)
    if data.shape[-1] == 32768:
        return blocked_cholesky_cuda_large(data)
    return blocked_cholesky_cuda(data)
scrolls · 1923 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