Skip to content
KernelIndex
Search⌘K

submission 898138

Praneeth Veligeti · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b555a464c5181cd52fca9e588461061a6265b32b4c33f7e64d52bfe40dcf4d22
license declaredunknown
license concludedunknown
authorsPraneeth Veligeti
imported2026-08-26

Techniques

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

autotune@triton.autotune(
shared-memory__shared__ float pivot[128];

Kernel source

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

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

from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = True


def _blocked_tf32(matrix: torch.Tensor) -> torch.Tensor:
    n = matrix.shape[-1]
    if n <= 256:
        return torch.linalg.cholesky_ex(matrix, check_errors=False).L

    half = n // 2
    a11 = matrix[..., :half, :half]
    a21 = matrix[..., half:, :half]
    a22 = matrix[..., half:, half:]
    l11 = _blocked_tf32(a11)
    l21 = torch.linalg.solve_triangular(l11.mT, a21, upper=True, left=False)
    schur = torch.baddbmm(a22, l21, l21.mT, beta=1.0, alpha=-1.0)
    l22 = _blocked_tf32(schur)

    output = torch.zeros_like(matrix)
    output[..., :half, :half] = l11
    output[..., half:, :half] = l21
    output[..., half:, half:] = l22
    return output


def _blocked_large(matrix: torch.Tensor, base: int) -> torch.Tensor:
    """Recursive tensor-core Schur updates with cuSOLVER base panels."""
    if matrix.shape[-1] <= base:
        return torch.linalg.cholesky_ex(matrix, check_errors=False).L
    half = matrix.shape[-1] // 2
    a11 = matrix[..., :half, :half]
    a21 = matrix[..., half:, :half]
    a22 = matrix[..., half:, half:]
    l11 = _blocked_large(a11, base)
    l21 = torch.linalg.solve_triangular(l11.mT, a21, upper=True, left=False)
    schur = torch.baddbmm(a22, l21, l21.mT, beta=1.0, alpha=-1.0)
    l22 = _blocked_large(schur, base)
    output = torch.empty_like(matrix)
    _assemble_large[(triton.cdiv(output.numel(), 512),)](
        l11, l21, l22, output,
        l11.stride(0), l11.stride(1), l11.stride(2),
        l21.stride(0), l21.stride(1), l21.stride(2),
        l22.stride(0), l22.stride(1), l22.stride(2),
        N=matrix.shape[-1], H=half, TOTAL=output.numel(), BLOCK=512)
    return output


@triton.jit
def _assemble_large(l11, l21, l22, output,
                    l11_s0, l11_s1, l11_s2,
                    l21_s0, l21_s1, l21_s2,
                    l22_s0, l22_s1, l22_s2,
                    N: tl.constexpr, H: tl.constexpr,
                    TOTAL: tl.constexpr, BLOCK: tl.constexpr):
    index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    valid = index < TOTAL
    matrix_elements = N * N
    matrix = index // matrix_elements
    local = index - matrix * matrix_elements
    row = local // N
    column = local - row * N
    top_left = (row < H) & (column < H) & valid
    bottom_left = (row >= H) & (column < H) & valid
    bottom_right = (row >= H) & (column >= H) & valid
    value = tl.zeros((BLOCK,), tl.float32)
    value += tl.load(
        l11 + matrix * l11_s0 + row * l11_s1 + column * l11_s2,
        mask=top_left, other=0.0)
    value += tl.load(
        l21 + matrix * l21_s0 + (row - H) * l21_s1 + column * l21_s2,
        mask=bottom_left, other=0.0)
    value += tl.load(
        l22 + matrix * l22_s0 + (row - H) * l22_s1
        + (column - H) * l22_s2,
        mask=bottom_right, other=0.0)
    tl.store(output + index, value, mask=valid)


@triton.jit
def _assemble_1024(l11, l21, l22, output,
                   l11_s0, l11_s1, l11_s2,
                   l21_s0, l21_s1, l21_s2,
                   l22_s0, l22_s1, l22_s2,
                   TOTAL: tl.constexpr, BLOCK: tl.constexpr):
    index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
    valid = index < TOTAL
    matrix = index // (1024 * 1024)
    local = index - matrix * (1024 * 1024)
    row = local // 1024
    column = local - row * 1024
    top_left = (row < 512) & (column < 512) & valid
    bottom_left = (row >= 512) & (column < 512) & valid
    bottom_right = (row >= 512) & (column >= 512) & valid

    value = tl.zeros((BLOCK,), tl.float32)
    value += tl.load(
        l11 + matrix * l11_s0 + row * l11_s1 + column * l11_s2,
        mask=top_left, other=0.0)
    value += tl.load(
        l21 + matrix * l21_s0 + (row - 512) * l21_s1 + column * l21_s2,
        mask=bottom_left, other=0.0)
    value += tl.load(
        l22 + matrix * l22_s0 + (row - 512) * l22_s1
        + (column - 512) * l22_s2,
        mask=bottom_right, other=0.0)
    tl.store(output + index, value, mask=valid)


def _blocked_1024(data: torch.Tensor) -> torch.Tensor:
    a11 = data[..., :512, :512]
    a21 = data[..., 512:, :512]
    a22 = data[..., 512:, 512:]
    l11 = torch.linalg.cholesky_ex(a11, check_errors=False).L
    l21 = torch.linalg.solve_triangular(
        l11.mT, a21, upper=True, left=False
    )
    schur = a22 - l21 @ l21.mT
    l22 = torch.linalg.cholesky_ex(schur, check_errors=False).L
    output = torch.empty_like(data)
    _assemble_1024[(triton.cdiv(output.numel(), 256),)](
        l11, l21, l22, output,
        l11.stride(0), l11.stride(1), l11.stride(2),
        l21.stride(0), l21.stride(1), l21.stride(2),
        l22.stride(0), l22.stride(1), l22.stride(2),
        TOTAL=output.numel(), BLOCK=256)
    return output


CPP_SRC = r"""
torch::Tensor emulated_potrf(torch::Tensor input);
torch::Tensor blocked_potrf(torch::Tensor input, int64_t block);
torch::Tensor register_row_potrf64(torch::Tensor input);
torch::Tensor register_row_potrf128(torch::Tensor input);
torch::Tensor blocked256_register(torch::Tensor input);
torch::Tensor blocked512_register(torch::Tensor input);
torch::Tensor blocked1024_register(torch::Tensor input);
"""


CUDA_SRC = r"""
#include <torch/extension.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <unordered_map>
#include <vector>

namespace {

cusolverDnHandle_t solver_handle() {
    static cusolverDnHandle_t handle = [] {
        cusolverDnHandle_t value;
        TORCH_CHECK(cusolverDnCreate(&value) == CUSOLVER_STATUS_SUCCESS,
                    "cusolverDnCreate failed");
        TORCH_CHECK(
            cusolverDnSetMathMode(value, CUSOLVER_FP32_EMULATED_BF16X9_MATH)
                == CUSOLVER_STATUS_SUCCESS,
            "cusolverDnSetMathMode failed");
        return value;
    }();
    return handle;
}

cusolverDnParams_t solver_params() {
    static cusolverDnParams_t params = [] {
        cusolverDnParams_t value;
        TORCH_CHECK(cusolverDnCreateParams(&value) == CUSOLVER_STATUS_SUCCESS,
                    "cusolverDnCreateParams failed");
        return value;
    }();
    return params;
}

cusolverDnHandle_t blocked_solver_handle() {
    static cusolverDnHandle_t handle = [] {
        cusolverDnHandle_t value;
        TORCH_CHECK(cusolverDnCreate(&value) == CUSOLVER_STATUS_SUCCESS,
                    "blocked cusolverDnCreate failed");
        return value;
    }();
    return handle;
}

cublasHandle_t blocked_blas_handle() {
    static cublasHandle_t handle = [] {
        cublasHandle_t value;
        TORCH_CHECK(cublasCreate(&value) == CUBLAS_STATUS_SUCCESS,
                    "cublasCreate failed");
        TORCH_CHECK(cublasSetMathMode(value, CUBLAS_TF32_TENSOR_OP_MATH)
                        == CUBLAS_STATUS_SUCCESS,
                    "cublasSetMathMode failed");
        return value;
    }();
    return handle;
}

cublasHandle_t precise_blas_handle() {
    static cublasHandle_t handle = [] {
        cublasHandle_t value;
        TORCH_CHECK(cublasCreate(&value) == CUBLAS_STATUS_SUCCESS,
                    "precise cublasCreate failed");
        TORCH_CHECK(cublasSetMathMode(value, CUBLAS_DEFAULT_MATH)
                        == CUBLAS_STATUS_SUCCESS,
                    "precise cublasSetMathMode failed");
        return value;
    }();
    return handle;
}

__global__ void make_block_pointers(
    float* matrices, float** diagonal, float** panel,
    int batch, int64_t matrix_stride, int64_t diagonal_offset,
    int64_t panel_offset) {
    int matrix = blockIdx.x * blockDim.x + threadIdx.x;
    if (matrix < batch) {
        float* base = matrices + int64_t(matrix) * matrix_stride;
        diagonal[matrix] = base + diagonal_offset;
        panel[matrix] = base + panel_offset;
    }
}

__global__ void clear_batched_row_major_upper(
    float* matrices, int64_t total, int n, int64_t matrix_stride) {
    int64_t index = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
    if (index < total) {
        int64_t local = index % matrix_stride;
        int row = local / n;
        int column = local - int64_t(row) * n;
        if (column > row) matrices[index] = 0.0f;
    }
}

__global__ void clear_batched_row_major_upper_2d(
    float* matrices, int n, int64_t matrix_stride) {
    int column = blockIdx.x * blockDim.x + threadIdx.x;
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    int matrix = blockIdx.z;
    if (row < n && column < n && column > row) {
        matrices[int64_t(matrix) * matrix_stride + int64_t(row) * n + column]
            = 0.0f;
    }
}

__global__ void clear_upper(float* matrix, int64_t n) {
    int64_t index = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
    int64_t count = n * n;
    if (index < count) {
        int64_t row = index / n;
        int64_t col = index - row * n;
        if (col > row) matrix[index] = 0.0f;
    }
}

template <int K>
struct RegisterRowStep64 {
    __device__ __forceinline__ static void run(
        float (&row0)[64], float (&row1)[64], int lane) {
        constexpr unsigned mask = 0xffffffffu;
        constexpr int owner = K & 31;
        float diagonal = 0.0f;
        if (lane == owner) {
            diagonal = K < 32 ? row0[K] : row1[K];
#pragma unroll
            for (int j = 0; j < K; ++j) {
                float value = K < 32 ? row0[j] : row1[j];
                diagonal -= value * value;
            }
        }
        diagonal = __shfl_sync(mask, diagonal, owner);
        float inverse = rsqrtf(diagonal);
        float sum0 = row0[K];
        float sum1 = row1[K];
#pragma unroll
        for (int j = 0; j < K; ++j) {
            float owned = K < 32 ? row0[j] : row1[j];
            float pivot = __shfl_sync(mask, owned, owner);
            sum0 -= row0[j] * pivot;
            sum1 -= row1[j] * pivot;
        }
        if (lane == K) row0[K] = sqrtf(diagonal);
        else if (lane > K) row0[K] = sum0 * inverse;
        if (lane + 32 == K) row1[K] = sqrtf(diagonal);
        else if (lane + 32 > K) row1[K] = sum1 * inverse;
        RegisterRowStep64<K + 1>::run(row0, row1, lane);
    }
};

template <>
struct RegisterRowStep64<64> {
    __device__ __forceinline__ static void run(
        float (&)[64], float (&)[64], int) {}
};

__global__ __launch_bounds__(32) void register_row_kernel64(
    const float* __restrict__ input, float* __restrict__ output) {
    int matrix = blockIdx.x;
    int lane = threadIdx.x;
    constexpr int64_t stride = 64 * 64;
    const float* source = input + int64_t(matrix) * stride;
    float* destination = output + int64_t(matrix) * stride;
    float row0[64];
    float row1[64];
#pragma unroll
    for (int column = 0; column < 64; ++column) {
        row0[column] = column <= lane ? source[lane * 64 + column] : 0.0f;
        row1[column] = column <= lane + 32
            ? source[(lane + 32) * 64 + column] : 0.0f;
    }
    RegisterRowStep64<0>::run(row0, row1, lane);
#pragma unroll
    for (int column = 0; column < 64; ++column) {
        destination[lane * 64 + column] =
            column <= lane ? row0[column] : 0.0f;
        destination[(lane + 32) * 64 + column] =
            column <= lane + 32 ? row1[column] : 0.0f;
    }
}

template <int K>
struct RegisterRowStep128 {
    __device__ __forceinline__ static void run(
        float (&row)[128], float* pivot, int matrix_row) {
        if (matrix_row == K) {
            float diagonal = row[K];
#pragma unroll
            for (int j = 0; j < K; ++j) {
                diagonal -= row[j] * row[j];
                pivot[j] = row[j];
            }
            row[K] = sqrtf(diagonal);
            pivot[K] = row[K];
        }
        __syncthreads();
        if (matrix_row > K) {
            float value = row[K];
#pragma unroll
            for (int j = 0; j < K; ++j) {
                value -= row[j] * pivot[j];
            }
            row[K] = value / pivot[K];
        }
        __syncthreads();
        RegisterRowStep128<K + 1>::run(row, pivot, matrix_row);
    }
};

template <>
struct RegisterRowStep128<128> {
    __device__ __forceinline__ static void run(
        float (&)[128], float*, int) {}
};

__global__ __launch_bounds__(128, 1) void register_row_kernel128(
    const float* __restrict__ input, float* __restrict__ output) {
    int matrix = blockIdx.x;
    int matrix_row = threadIdx.x;
    constexpr int64_t stride = 128 * 128;
    const float* source = input + int64_t(matrix) * stride;
    float* destination = output + int64_t(matrix) * stride;
    __shared__ float pivot[128];
    float row[128];
#pragma unroll
    for (int column = 0; column < 128; ++column) {
        row[column] = column <= matrix_row
            ? source[matrix_row * 128 + column] : 0.0f;
    }
    RegisterRowStep128<0>::run(row, pivot, matrix_row);
#pragma unroll
    for (int column = 0; column < 128; ++column) {
        destination[matrix_row * 128 + column] =
            column <= matrix_row ? row[column] : 0.0f;
    }
}

__global__ __launch_bounds__(128, 1) void factor128_inplace(
    float* matrices, int leading, int64_t stride, int64_t offset) {
    int matrix_row = threadIdx.x;
    float* tile = matrices + int64_t(blockIdx.x) * stride + offset;
    __shared__ float pivot[128];
    float row[128];
#pragma unroll
    for (int column = 0; column < 128; ++column) {
        row[column] = column <= matrix_row
            ? tile[matrix_row * leading + column] : 0.0f;
    }
    RegisterRowStep128<0>::run(row, pivot, matrix_row);
#pragma unroll
    for (int column = 0; column < 128; ++column) {
        if (column <= matrix_row) {
            tile[matrix_row * leading + column] = row[column];
        }
    }
}

} // namespace

torch::Tensor emulated_potrf(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda(), "input must be CUDA");
    TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(0) == 1, "batch must be one");

    auto output = input.clone();
    const int64_t n = output.size(1);
    auto handle = solver_handle();
    auto params = solver_params();

    size_t device_bytes = 0;
    size_t host_bytes = 0;
    auto status = cusolverDnXpotrf_bufferSize(
        handle, params, CUBLAS_FILL_MODE_UPPER, n,
        CUDA_R_32F, output.data_ptr<float>(), n, CUDA_R_32F,
        &device_bytes, &host_bytes);
    TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, "bufferSize failed: ", int(status));

    static std::unordered_map<int64_t, torch::Tensor> workspaces;
    static std::unordered_map<int64_t, torch::Tensor> infos;
    static std::unordered_map<int64_t, std::vector<unsigned char>> host_workspaces;
    auto byte_options = torch::TensorOptions().device(input.device()).dtype(torch::kUInt8);
    if (!workspaces.count(n) || workspaces.at(n).numel() < static_cast<int64_t>(device_bytes)) {
        workspaces[n] = torch::empty({static_cast<int64_t>(device_bytes)}, byte_options);
        infos[n] = torch::empty(
            {1}, torch::TensorOptions().device(input.device()).dtype(torch::kInt32));
        host_workspaces[n] = std::vector<unsigned char>(host_bytes);
    }
    auto& workspace = workspaces.at(n);
    auto& info = infos.at(n);
    auto& host_workspace = host_workspaces.at(n);

    status = cusolverDnXpotrf(
        handle, params, CUBLAS_FILL_MODE_UPPER, n,
        CUDA_R_32F, output.data_ptr<float>(), n, CUDA_R_32F,
        workspace.data_ptr(), device_bytes,
        host_workspace.data(), host_bytes,
        info.data_ptr<int>());
    TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, "Xpotrf failed: ", int(status));

    int threads = 256;
    int64_t count = n * n;
    clear_upper<<<(count + threads - 1) / threads, threads>>>(output.data_ptr<float>(), n);
    return output;
}

torch::Tensor blocked_potrf(torch::Tensor input, int64_t block_value) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
                "input must be CUDA FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == input.size(2),
                "input must be batched square matrices");
    const int batch = input.size(0);
    const int n = input.size(1);
    const int block = static_cast<int>(block_value);
    TORCH_CHECK(n % block == 0, "block must divide n");

    auto output = input.clone();
    auto pointers = torch::empty(
        {2, batch}, torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
    auto info = torch::empty(
        {batch}, torch::TensorOptions().device(input.device()).dtype(torch::kInt32));
    auto diagonal = reinterpret_cast<float**>(pointers[0].data_ptr<int64_t>());
    auto panel = reinterpret_cast<float**>(pointers[1].data_ptr<int64_t>());
    const int64_t stride = int64_t(n) * n;
    const float one = 1.0f;
    const float minus_one = -1.0f;

    for (int k = 0; k < n; k += block) {
        const int trailing = n - k - block;
        const int64_t diagonal_offset = int64_t(k) * n + k;
        const int64_t panel_offset = trailing > 0
            ? int64_t(k + block) * n + k : diagonal_offset;
        make_block_pointers<<<(batch + 127) / 128, 128>>>(
            output.data_ptr<float>(), diagonal, panel, batch, stride,
            diagonal_offset, panel_offset);
        auto solver_status = cusolverDnSpotrfBatched(
            blocked_solver_handle(), CUBLAS_FILL_MODE_UPPER, block,
            diagonal, n, info.data_ptr<int>(), batch);
        TORCH_CHECK(solver_status == CUSOLVER_STATUS_SUCCESS,
                    "batched panel potrf failed: ", int(solver_status));

        if (trailing > 0) {
            auto trsm_status = cublasStrsmBatched(
                blocked_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
                CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
                block, trailing, &one,
                const_cast<const float**>(diagonal), n, panel, n, batch);
            TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
                        "batched TRSM failed: ", int(trsm_status));

            float* panel_base = output.data_ptr<float>() + panel_offset;
            float* trailing_base = output.data_ptr<float>()
                + int64_t(k + block) * n + (k + block);
            auto gemm_status = cublasSgemmStridedBatched(
                blocked_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
                trailing, trailing, block, &minus_one,
                panel_base, n, stride, panel_base, n, stride,
                &one, trailing_base, n, stride, batch);
            TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
                        "batched update GEMM failed: ", int(gemm_status));
        }
    }

    const int64_t total = int64_t(batch) * stride;
    dim3 clear_threads(32, 8);
    dim3 clear_grid((n + clear_threads.x - 1) / clear_threads.x,
                    (n + clear_threads.y - 1) / clear_threads.y,
                    batch);
    clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
        output.data_ptr<float>(), n, stride);
    auto error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess,
                "blocked launch failed: ", cudaGetErrorString(error));
    return output;
}

torch::Tensor register_row_potrf64(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
                "input must be CUDA FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 64 && input.size(2) == 64,
                "register-row kernel requires 64x64 matrices");
    auto output = torch::empty_like(input);
    register_row_kernel64<<<input.size(0), 32>>>(
        input.data_ptr<float>(), output.data_ptr<float>());
    auto error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "register-row launch failed: ",
                cudaGetErrorString(error));
    return output;
}

torch::Tensor register_row_potrf128(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
                "input must be CUDA FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 128
                    && input.size(2) == 128,
                "register-row kernel requires 128x128 matrices");
    auto output = torch::empty_like(input);
    register_row_kernel128<<<input.size(0), 128>>>(
        input.data_ptr<float>(), output.data_ptr<float>());
    auto error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "register-row launch failed: ",
                cudaGetErrorString(error));
    return output;
}

torch::Tensor blocked256_register(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
                "input must be CUDA FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 256
                    && input.size(2) == 256,
                "blocked256 requires 256x256 matrices");
    const int batch = input.size(0);
    constexpr int n = 256;
    constexpr int64_t stride = int64_t(n) * n;
    auto output = input.clone();
    auto pointers = torch::empty(
        {2, batch},
        torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
    auto diagonal = reinterpret_cast<float**>(
        pointers[0].data_ptr<int64_t>());
    auto panel = reinterpret_cast<float**>(
        pointers[1].data_ptr<int64_t>());

    factor128_inplace<<<batch, 128>>>(
        output.data_ptr<float>(), n, stride, 0);
    make_block_pointers<<<(batch + 127) / 128, 128>>>(
        output.data_ptr<float>(), diagonal, panel, batch, stride,
        0, 128 * n);

    const float one = 1.0f;
    const float minus_one = -1.0f;
    auto trsm_status = cublasStrsmBatched(
        precise_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
        CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
        128, 128, &one,
        const_cast<const float**>(diagonal), n,
        panel, n, batch);
    TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
                "precise batched TRSM failed: ", int(trsm_status));

    float* panel_base = output.data_ptr<float>() + 128 * n;
    float* trailing_base = panel_base + 128;
    auto gemm_status = cublasSgemmStridedBatched(
        precise_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
        128, 128, 128, &minus_one,
        panel_base, n, stride,
        panel_base, n, stride,
        &one, trailing_base, n, stride, batch);
    TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
                "precise batched GEMM failed: ", int(gemm_status));

    factor128_inplace<<<batch, 128>>>(
        output.data_ptr<float>(), n, stride, 128 * n + 128);
    dim3 clear_threads(32, 8);
    dim3 clear_grid(8, 32, batch);
    clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
        output.data_ptr<float>(), n, stride);
    auto error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "blocked256 launch failed: ",
                cudaGetErrorString(error));
    return output;
}

torch::Tensor blocked512_register(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
                "input must be CUDA FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 512
                    && input.size(2) == 512,
                "blocked512 requires 512x512 matrices");
    const int batch = input.size(0);
    constexpr int n = 512;
    constexpr int block = 128;
    constexpr int64_t stride = int64_t(n) * n;
    auto output = input.clone();
    auto pointers = torch::empty(
        {2, batch},
        torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
    auto diagonal = reinterpret_cast<float**>(
        pointers[0].data_ptr<int64_t>());
    auto panel = reinterpret_cast<float**>(
        pointers[1].data_ptr<int64_t>());
    const float one = 1.0f;
    const float minus_one = -1.0f;

    for (int k = 0; k < n; k += block) {
        factor128_inplace<<<batch, 128>>>(
            output.data_ptr<float>(), n, stride, int64_t(k) * n + k);
        const int trailing = n - k - block;
        if (trailing == 0) continue;
        const int64_t diagonal_offset = int64_t(k) * n + k;
        const int64_t panel_offset = int64_t(k + block) * n + k;
        make_block_pointers<<<(batch + 127) / 128, 128>>>(
            output.data_ptr<float>(), diagonal, panel, batch, stride,
            diagonal_offset, panel_offset);
        auto trsm_status = cublasStrsmBatched(
            precise_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
            block, trailing, &one,
            const_cast<const float**>(diagonal), n,
            panel, n, batch);
        TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
                    "blocked512 TRSM failed: ", int(trsm_status));
        float* panel_base = output.data_ptr<float>() + panel_offset;
        float* trailing_base = output.data_ptr<float>()
            + int64_t(k + block) * n + (k + block);
        auto gemm_status = cublasSgemmStridedBatched(
            precise_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
            trailing, trailing, block, &minus_one,
            panel_base, n, stride,
            panel_base, n, stride,
            &one, trailing_base, n, stride, batch);
        TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
                    "blocked512 GEMM failed: ", int(gemm_status));
    }
    dim3 clear_threads(32, 8);
    dim3 clear_grid(16, 64, batch);
    clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
        output.data_ptr<float>(), n, stride);
    auto error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "blocked512 launch failed: ",
                cudaGetErrorString(error));
    return output;
}

torch::Tensor blocked1024_register(torch::Tensor input) {
    TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
                "input must be CUDA FP32");
    TORCH_CHECK(input.dim() == 3 && input.size(1) == 1024
                    && input.size(2) == 1024,
                "blocked1024 requires 1024x1024 matrices");
    const int batch = input.size(0);
    constexpr int n = 1024;
    constexpr int block = 128;
    constexpr int64_t stride = int64_t(n) * n;
    auto output = input.clone();
    auto pointers = torch::empty(
        {2, batch},
        torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
    auto diagonal = reinterpret_cast<float**>(
        pointers[0].data_ptr<int64_t>());
    auto panel = reinterpret_cast<float**>(
        pointers[1].data_ptr<int64_t>());
    const float one = 1.0f;
    const float minus_one = -1.0f;

    for (int k = 0; k < n; k += block) {
        factor128_inplace<<<batch, 128>>>(
            output.data_ptr<float>(), n, stride, int64_t(k) * n + k);
        const int trailing = n - k - block;
        if (trailing == 0) continue;
        const int64_t diagonal_offset = int64_t(k) * n + k;
        const int64_t panel_offset = int64_t(k + block) * n + k;
        make_block_pointers<<<(batch + 127) / 128, 128>>>(
            output.data_ptr<float>(), diagonal, panel, batch, stride,
            diagonal_offset, panel_offset);
        auto trsm_status = cublasStrsmBatched(
            precise_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
            block, trailing, &one,
            const_cast<const float**>(diagonal), n,
            panel, n, batch);
        TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
                    "blocked1024 TRSM failed: ", int(trsm_status));
        float* panel_base = output.data_ptr<float>() + panel_offset;
        float* trailing_base = output.data_ptr<float>()
            + int64_t(k + block) * n + (k + block);
        auto gemm_status = cublasSgemmStridedBatched(
            precise_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
            trailing, trailing, block, &minus_one,
            panel_base, n, stride,
            panel_base, n, stride,
            &one, trailing_base, n, stride, batch);
        TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
                    "blocked1024 GEMM failed: ", int(gemm_status));
    }
    dim3 clear_threads(32, 8);
    dim3 clear_grid(32, 128, batch);
    clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
        output.data_ptr<float>(), n, stride);
    auto error = cudaGetLastError();
    TORCH_CHECK(error == cudaSuccess, "blocked1024 launch failed: ",
                cudaGetErrorString(error));
    return output;
}
"""


native = load_inline(
    name="chol_ranked_hybrid_v3",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=[
        "emulated_potrf",
        "blocked_potrf",
        "register_row_potrf64",
        "register_row_potrf128",
        "blocked256_register",
        "blocked512_register",
        "blocked1024_register",
    ],
    extra_ldflags=["-lcusolver"],
    extra_cuda_cflags=["-O3"],
    verbose=False,
)


@triton.autotune(
    configs=[triton.Config({}, num_warps=w) for w in (1, 2, 4, 8)],
    key=["N"],
)
@triton.jit
def _chol_small_kernel(A_ptr, L_ptr, N: tl.constexpr, BS: tl.constexpr):
    pid = tl.program_id(0)
    r = tl.arange(0, BS)
    row = r[:, None]
    col = r[None, :]
    bounds = (row < N) & (col < N)
    base = pid * N * N
    a = tl.load(A_ptr + base + row * N + col, mask=bounds, other=0.0)

    for k in tl.static_range(N):  # compile-time unroll: N is a constexpr
        dk = tl.sum(tl.where((row == k) & (col == k), a, 0.0))
        inv = 1.0 / tl.sqrt(dk)
        v = tl.sum(tl.where(col == k, a, 0.0), axis=1) * inv
        v = tl.where(r >= k, v, 0.0)
        a = a - v[:, None] * v[None, :]        # update; also zeroes row/col k
        a = tl.where(col == k, v[:, None], a)  # write finished L column in place

    out = tl.where(row >= col, a, 0.0)
    tl.store(L_ptr + base + row * N + col, out, mask=bounds)


def _triton_cholesky(data):
    x = data.contiguous()
    out = torch.empty_like(x)
    n = x.shape[-1]
    _chol_small_kernel[(x.shape[0],)](x, out, N=n, BS=triton.next_power_of_2(n))
    return out


def custom_kernel(data: input_t) -> output_t:
    batch = data.shape[0]
    n = data.shape[-1]

    if n == 32:
        return _triton_cholesky(data)

    if n == 64:
        return native.register_row_potrf64(data)

    if n == 128:
        return native.register_row_potrf128(data)

    if n == 256:
        return native.blocked256_register(data)

    if batch == 16 and n == 512:
        return native.blocked512_register(data)

    if batch == 4 and n == 1024:
        return native.blocked1024_register(data)

    if batch == 640 and n == 512:
        return native.blocked_potrf(data, 128)

    if batch == 60 and n == 1024:
        return native.blocked_potrf(data, 256)

    if batch == 1 and n == 8192:
        return _blocked_large(data, 4096)

    if batch == 1 and n == 16384:
        return _blocked_large(data, 4096)

    if batch == 1 and n == 32768:
        return _blocked_large(data, 4096)

    if batch == 8 and n == 2048:
        return native.blocked_potrf(data, 128)

    if 2 <= batch <= 4 and n >= 2048:
        results = []
        for i in range(batch):
            L = torch.linalg.cholesky_ex(data[i], check_errors=False).L
            results.append(L)
        return torch.stack(results)

    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 833 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