submission 910279
Clark Kitchen · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1809 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-910279?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:e511aa464729b81db6729e168d42cbf9222f3d19bb4e0d3cff8476ba369a2a72
license declaredunknown
license concludedunknown
authorsClark Kitchen
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
update = tl.dot(left, tl.trans(right), input_precision="tf32")num-warps = 8
num_warps=8,shared-memory
extern __shared__ float packed_storage[];stages = 1
num_stages=1,tile-m = 16
BLOCK_M=16,tile-n = 256
BLOCK_N=256,vector-width = float4
const float4* a4 = reinterpret_cast<const float4*>(a);Kernel source
submission.py1809 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_SRC = r"""
torch::Tensor raw_cholesky(torch::Tensor input);
torch::Tensor copy_lower(torch::Tensor input);
void factor_panel(torch::Tensor output, int64_t panel_start);
void factor_panel_128(torch::Tensor output, int64_t panel_start);
void factor_panel_inverse(torch::Tensor output, int64_t panel_start);
// B640_FUSED_V1_BEGIN
void factor_panel_cutlass_inverse_b640(
torch::Tensor output,
torch::Tensor inverse,
int64_t panel_start);
// B640_FUSED_V1_END
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cfloat>
#include <cstdint>
namespace {
constexpr int kSmallMax = 128;
constexpr int kPanel = 64;
constexpr int kLargePanel = 128;
constexpr int kTile = 64;
constexpr int kThreads = 256;
constexpr int kLargeSolveThreads = 512;
__device__ __forceinline__ int packed_index(int row, int col) {
return (row * (row + 1)) / 2 + col;
}
// Small benchmark batches contain thousands of matrices. A warp owns one
// matrix so independent factorizations share a launch without sharing any
// synchronization. Vectorized full-row movement is safe because every
// benchmark size is a multiple of four and each matrix starts aligned.
template <int N, int MatricesPerCta>
__global__ void warp_packed_cholesky_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch,
int64_t matrix_stride) {
constexpr int kPacked = N * (N + 1) / 2;
extern __shared__ float packed_storage[];
const int lane = static_cast<int>(threadIdx.x) & 31;
const int local_matrix = static_cast<int>(threadIdx.x) >> 5;
const int matrix = static_cast<int>(blockIdx.x) * MatricesPerCta + local_matrix;
if (matrix >= batch) {
return;
}
float* lower = packed_storage + local_matrix * kPacked;
const float* a = input + static_cast<int64_t>(matrix) * matrix_stride;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
const float4* a4 = reinterpret_cast<const float4*>(a);
float4* l4 = reinterpret_cast<float4*>(l);
constexpr int kVectors = N * N / 4;
for (int vector = lane; vector < kVectors; vector += 32) {
const float4 values = a4[vector];
const int linear = vector * 4;
const int row = linear / N;
const int col = linear - row * N;
if (col <= row) {
lower[packed_index(row, col)] = values.x;
}
if (col + 1 <= row) {
lower[packed_index(row, col + 1)] = values.y;
}
if (col + 2 <= row) {
lower[packed_index(row, col + 2)] = values.z;
}
if (col + 3 <= row) {
lower[packed_index(row, col + 3)] = values.w;
}
}
__syncwarp();
for (int pivot = 0; pivot < N; ++pivot) {
if (lane == 0) {
const int diagonal_index = packed_index(pivot, pivot);
lower[diagonal_index] = sqrtf(fmaxf(lower[diagonal_index], FLT_MIN));
}
__syncwarp();
const float diagonal = lower[packed_index(pivot, pivot)];
for (int row = pivot + 1 + lane; row < N; row += 32) {
lower[packed_index(row, pivot)] /= diagonal;
}
__syncwarp();
for (int row = pivot + 1 + lane; row < N; row += 32) {
const float row_value = lower[packed_index(row, pivot)];
for (int col = pivot + 1; col <= row; ++col) {
const int index = packed_index(row, col);
lower[index] = fmaf(
-row_value,
lower[packed_index(col, pivot)],
lower[index]);
}
}
__syncwarp();
}
for (int vector = lane; vector < kVectors; vector += 32) {
const int linear = vector * 4;
const int row = linear / N;
const int col = linear - row * N;
float4 values;
values.x = col <= row ? lower[packed_index(row, col)] : 0.0f;
values.y = col + 1 <= row ? lower[packed_index(row, col + 1)] : 0.0f;
values.z = col + 2 <= row ? lower[packed_index(row, col + 2)] : 0.0f;
values.w = col + 3 <= row ? lower[packed_index(row, col + 3)] : 0.0f;
l4[vector] = values;
}
}
// One block owns one matrix. Keeping the lower triangle packed makes n=128
// fit comfortably below the default per-block shared-memory limit.
template <int N>
__global__ void small_cholesky_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int64_t matrix_stride,
int runtime_n) {
extern __shared__ float lower[];
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
const int n = N == 0 ? runtime_n : N;
const float* a = input + static_cast<int64_t>(matrix) * matrix_stride;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
const int elements = n * n;
for (int linear = tid; linear < elements; linear += blockDim.x) {
const int row = linear / n;
const int col = linear - row * n;
if (col <= row) {
lower[packed_index(row, col)] = a[linear];
}
}
__syncthreads();
// Right-looking Cholesky. A thread owns a complete trailing row during
// each rank-1 update, so no two threads write the same shared element.
for (int pivot = 0; pivot < n; ++pivot) {
if (tid == 0) {
float diagonal = lower[packed_index(pivot, pivot)];
lower[packed_index(pivot, pivot)] = sqrtf(fmaxf(diagonal, FLT_MIN));
}
__syncthreads();
const float diagonal = lower[packed_index(pivot, pivot)];
for (int row = pivot + 1 + tid; row < n; row += blockDim.x) {
lower[packed_index(row, pivot)] /= diagonal;
}
__syncthreads();
for (int row = pivot + 1 + tid; row < n; row += blockDim.x) {
const float row_value = lower[packed_index(row, pivot)];
for (int col = pivot + 1; col <= row; ++col) {
const int index = packed_index(row, col);
lower[index] = fmaf(
-row_value,
lower[packed_index(col, pivot)],
lower[index]);
}
}
// Thread 0 owns the next diagonal row; the next post-sqrt barrier
// safely completes every independent later-row update.
}
for (int linear = tid; linear < elements; linear += blockDim.x) {
const int row = linear / n;
const int col = linear - row * n;
l[linear] = col <= row ? lower[packed_index(row, col)] : 0.0f;
}
}
__global__ void copy_lower_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int64_t total_elements,
int n) {
for (int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
linear < total_elements;
linear += static_cast<int64_t>(gridDim.x) * blockDim.x) {
const int within_matrix = static_cast<int>(linear % (static_cast<int64_t>(n) * n));
const int row = within_matrix / n;
const int col = within_matrix - row * n;
output[linear] = col <= row ? input[linear] : 0.0f;
}
}
template <int PanelWidth>
__global__ void diagonal_factor_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_start,
int runtime_panel_width) {
__shared__ float tile[kPanel][kPanel + 1];
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
const int panel_width = PanelWidth == 0 ? runtime_panel_width : PanelWidth;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
const int tile_elements = panel_width * panel_width;
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / panel_width;
const int col = linear - row * panel_width;
tile[row][col] = col <= row
? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
: 0.0f;
}
__syncthreads();
for (int pivot = 0; pivot < panel_width; ++pivot) {
if (tid == 0) {
tile[pivot][pivot] = sqrtf(fmaxf(tile[pivot][pivot], FLT_MIN));
}
__syncthreads();
const float diagonal = tile[pivot][pivot];
for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
tile[row][pivot] /= diagonal;
}
__syncthreads();
for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
const float row_value = tile[row][pivot];
for (int col = pivot + 1; col <= row; ++col) {
tile[row][col] = fmaf(-row_value, tile[col][pivot], tile[row][col]);
}
}
// Thread 0 exclusively updates the next diagonal row. The following
// post-sqrt barrier completes all later rows before they are consumed.
}
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / panel_width;
const int col = linear - row * panel_width;
if (col <= row) {
l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
tile[row][col];
}
}
}
// The large-matrix route factors a full 128-column diagonal block in FP32.
// Dynamic shared memory is required because the padded tile exceeds the
// legacy 48 KiB static-shared limit.
__global__ void diagonal_factor_128_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_start) {
extern __shared__ float tile[];
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
constexpr int tile_stride = kLargePanel + 1;
constexpr int tile_elements = kLargePanel * kLargePanel;
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / kLargePanel;
const int col = linear - row * kLargePanel;
tile[row * tile_stride + col] = col <= row
? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
: 0.0f;
}
__syncthreads();
for (int pivot = 0; pivot < kLargePanel; ++pivot) {
if (tid == 0) {
float& diagonal = tile[pivot * tile_stride + pivot];
diagonal = sqrtf(fmaxf(diagonal, FLT_MIN));
}
__syncthreads();
const float diagonal = tile[pivot * tile_stride + pivot];
for (int row = pivot + 1 + tid;
row < kLargePanel;
row += blockDim.x) {
tile[row * tile_stride + pivot] /= diagonal;
}
__syncthreads();
for (int row = pivot + 1 + tid;
row < kLargePanel;
row += blockDim.x) {
const float row_value = tile[row * tile_stride + pivot];
for (int col = pivot + 1; col <= row; ++col) {
float& value = tile[row * tile_stride + col];
value = fmaf(
-row_value,
tile[col * tile_stride + pivot],
value);
}
}
}
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / kLargePanel;
const int col = linear - row * kLargePanel;
if (col <= row) {
l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
tile[row * tile_stride + col];
}
}
}
template <int PanelWidth>
__global__ void diagonal_factor_inverse_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_start,
int runtime_panel_width) {
__shared__ float tile[kPanel][kPanel + 1];
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
const int panel_width = PanelWidth == 0 ? runtime_panel_width : PanelWidth;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
const int tile_elements = panel_width * panel_width;
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / panel_width;
const int col = linear - row * panel_width;
tile[row][col] = col <= row
? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
: 0.0f;
}
__syncthreads();
for (int pivot = 0; pivot < panel_width; ++pivot) {
if (tid == 0) {
tile[pivot][pivot] = sqrtf(fmaxf(tile[pivot][pivot], FLT_MIN));
}
__syncthreads();
const float diagonal = tile[pivot][pivot];
for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
tile[row][pivot] /= diagonal;
}
__syncthreads();
for (int row = pivot + 1 + tid; row < panel_width; row += blockDim.x) {
const float row_value = tile[row][pivot];
for (int col = pivot + 1; col <= row; ++col) {
tile[row][col] = fmaf(-row_value, tile[col][pivot], tile[row][col]);
}
}
}
// Store U=L^{-T} in the unused strict upper triangle. The diagonal
// reciprocal is synthesized from L when U is staged by the solve.
__syncthreads();
for (int inverse_row = 1; inverse_row < kPanel; ++inverse_row) {
for (int inverse_col = tid;
inverse_col < inverse_row;
inverse_col += blockDim.x) {
float value = 0.0f;
for (int k = inverse_col; k < inverse_row; ++k) {
const float inverse_value = k == inverse_col
? 1.0f / tile[k][k]
: tile[inverse_col][k];
value = fmaf(tile[inverse_row][k], inverse_value, value);
}
tile[inverse_col][inverse_row] =
-value / tile[inverse_row][inverse_row];
}
__syncthreads();
}
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / panel_width;
const int col = linear - row * panel_width;
l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
tile[row][col];
}
}
// CUTLASS panel path: preserve L in the output diagonal block and write the
// full row-major U=L^{-T} into a reusable per-matrix workspace.
__global__ void diagonal_factor_cutlass_inverse_kernel(
float* __restrict__ output,
float* __restrict__ inverse,
int64_t matrix_stride,
int64_t inverse_stride,
int n,
int panel_start) {
__shared__ float tile[kPanel][kPanel + 1];
__shared__ float inverse_tile[kPanel][kPanel + 1];
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
constexpr int tile_elements = kPanel * kPanel;
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / kPanel;
const int col = linear - row * kPanel;
tile[row][col] = col <= row
? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
: 0.0f;
}
__syncthreads();
for (int pivot = 0; pivot < kPanel; ++pivot) {
if (tid == 0) {
tile[pivot][pivot] = sqrtf(fmaxf(tile[pivot][pivot], FLT_MIN));
}
__syncthreads();
const float diagonal = tile[pivot][pivot];
for (int row = pivot + 1 + tid; row < kPanel; row += blockDim.x) {
tile[row][pivot] /= diagonal;
}
__syncthreads();
for (int row = pivot + 1 + tid; row < kPanel; row += blockDim.x) {
const float row_value = tile[row][pivot];
for (int col = pivot + 1; col <= row; ++col) {
tile[row][col] = fmaf(-row_value, tile[col][pivot], tile[row][col]);
}
}
}
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
inverse_tile[linear / kPanel][linear % kPanel] = 0.0f;
}
__syncthreads();
if (tid < kPanel) {
const int inverse_row = tid;
for (int row = inverse_row; row < kPanel; ++row) {
float sum = 0.0f;
for (int p = inverse_row; p < row; ++p) {
sum = fmaf(
tile[row][p],
inverse_tile[inverse_row][p],
sum);
}
inverse_tile[inverse_row][row] =
((row == inverse_row ? 1.0f : 0.0f) - sum) /
tile[row][row];
}
}
__syncthreads();
float* u = inverse + static_cast<int64_t>(matrix) * inverse_stride;
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
u[linear] = inverse_tile[linear / kPanel][linear % kPanel];
}
for (int linear = tid; linear < tile_elements; linear += blockDim.x) {
const int row = linear / kPanel;
const int col = linear - row * kPanel;
if (col <= row) {
l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
tile[row][col];
}
}
}
// A CTA stages U=L^{-T} and 16 right-hand-side rows, then computes every
// solved panel value independently in FP32.
__global__ void panel_solve_inverse_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_start) {
extern __shared__ float shared[];
float* upper = shared;
float* rhs = upper + kPanel * (kPanel + 1);
const int tid = static_cast<int>(threadIdx.x);
const int matrix = static_cast<int>(blockIdx.y);
const int first_local_row = static_cast<int>(blockIdx.x) * 16;
const int trailing_start = panel_start + kPanel;
const int trailing_rows = n - trailing_start;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
for (int linear = tid; linear < kPanel * kPanel; linear += blockDim.x) {
const int row = linear / kPanel;
const int col = linear - row * kPanel;
const float value = l[
static_cast<int64_t>(panel_start + row) * n + panel_start + col];
upper[row * (kPanel + 1) + col] = col < row
? 0.0f
: (col == row ? 1.0f / value : value);
}
for (int linear = tid; linear < 16 * kPanel; linear += blockDim.x) {
const int local_row = linear / kPanel;
const int col = linear - local_row * kPanel;
const int trailing_row = first_local_row + local_row;
rhs[local_row * (kPanel + 1) + col] = trailing_row < trailing_rows
? l[static_cast<int64_t>(trailing_start + trailing_row) * n +
panel_start + col]
: 0.0f;
}
__syncthreads();
for (int linear = tid; linear < 16 * kPanel; linear += blockDim.x) {
const int local_row = linear / kPanel;
const int col = linear - local_row * kPanel;
const int trailing_row = first_local_row + local_row;
if (trailing_row < trailing_rows) {
float solved = 0.0f;
#pragma unroll
for (int k = 0; k < kPanel; ++k) {
if (k <= col) {
solved = fmaf(
rhs[local_row * (kPanel + 1) + k],
upper[k * (kPanel + 1) + col],
solved);
}
}
l[static_cast<int64_t>(trailing_start + trailing_row) * n +
panel_start + col] = solved;
}
}
}
__global__ void clear_diagonal_upper_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_count) {
const int panel = static_cast<int>(blockIdx.x);
const int matrix = static_cast<int>(blockIdx.y);
const int panel_start = panel * kPanel;
const int panel_width = n - panel_start < kPanel ? n - panel_start : kPanel;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
for (int linear = static_cast<int>(threadIdx.x);
linear < panel_width * panel_width;
linear += blockDim.x) {
const int row = linear / panel_width;
const int col = linear - row * panel_width;
if (col > row) {
l[static_cast<int64_t>(panel_start + row) * n + panel_start + col] =
0.0f;
}
}
}
__device__ __forceinline__ float warp_sum(float value) {
constexpr unsigned mask = 0xffffffffu;
value += __shfl_down_sync(mask, value, 16);
value += __shfl_down_sync(mask, value, 8);
value += __shfl_down_sync(mask, value, 4);
value += __shfl_down_sync(mask, value, 2);
value += __shfl_down_sync(mask, value, 1);
return value;
}
// One warp solves one row while all warps share the factored diagonal panel.
// A 2-D launch keeps each CTA inside a single matrix and removes the flat
// launch's per-warp division and modulo.
template <int PanelWidth>
__global__ void panel_solve_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_start,
int runtime_panel_width) {
extern __shared__ float diagonal_panel[];
const int tid = static_cast<int>(threadIdx.x);
const int lane = static_cast<int>(threadIdx.x) & 31;
const int local_warp = static_cast<int>(threadIdx.x) >> 5;
const int panel_width = PanelWidth == 0 ? runtime_panel_width : PanelWidth;
const int matrix = static_cast<int>(blockIdx.y);
const int trailing_rows = n - panel_start - panel_width;
constexpr int rows_per_cta =
PanelWidth == kLargePanel ? kLargeSolveThreads / 32 : 8;
const int local_row =
static_cast<int>(blockIdx.x) * rows_per_cta + local_warp;
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
const int diagonal_elements = panel_width * panel_width;
for (int linear = tid; linear < diagonal_elements; linear += blockDim.x) {
const int row = linear / panel_width;
const int col = linear - row * panel_width;
diagonal_panel[linear] = col <= row
? l[static_cast<int64_t>(panel_start + row) * n + panel_start + col]
: 0.0f;
}
__syncthreads();
if (local_row >= trailing_rows) {
return;
}
const int row = panel_start + panel_width + local_row;
float* row_ptr = l + static_cast<int64_t>(row) * n;
if (PanelWidth == kLargePanel) {
constexpr unsigned mask = 0xffffffffu;
float row_values[4] = {
row_ptr[panel_start + lane],
row_ptr[panel_start + lane + 32],
row_ptr[panel_start + lane + 64],
row_ptr[panel_start + lane + 96],
};
#pragma unroll
for (int col = 0; col < kLargePanel; ++col) {
float partial = 0.0f;
#pragma unroll
for (int segment = 0; segment < 4; ++segment) {
const int p = lane + segment * 32;
if (p < col) {
partial = fmaf(
row_values[segment],
diagonal_panel[col * kLargePanel + p],
partial);
}
}
partial = warp_sum(partial);
const int segment = col >> 5;
const float rhs = __shfl_sync(
mask, row_values[segment], col & 31);
float solved = 0.0f;
if (lane == 0) {
solved = (rhs - partial) /
diagonal_panel[col * kLargePanel + col];
row_ptr[panel_start + col] = solved;
}
solved = __shfl_sync(mask, solved, 0);
if (lane == (col & 31)) {
row_values[segment] = solved;
}
}
return;
}
if (PanelWidth == kPanel) {
constexpr unsigned mask = 0xffffffffu;
float row_lo = row_ptr[panel_start + lane];
float row_hi = row_ptr[panel_start + lane + 32];
#pragma unroll
for (int col = 0; col < kPanel; ++col) {
float partial = 0.0f;
if (lane < col) {
partial = row_lo * diagonal_panel[col * kPanel + lane];
}
if (lane + 32 < col) {
partial = fmaf(
row_hi,
diagonal_panel[col * kPanel + lane + 32],
partial);
}
partial = warp_sum(partial);
const float rhs = col < 32
? __shfl_sync(mask, row_lo, col)
: __shfl_sync(mask, row_hi, col - 32);
float solved = 0.0f;
if (lane == 0) {
solved = (rhs - partial) /
diagonal_panel[col * kPanel + col];
row_ptr[panel_start + col] = solved;
}
solved = __shfl_sync(mask, solved, 0);
if (col < 32) {
if (lane == col) {
row_lo = solved;
}
} else if (lane == col - 32) {
row_hi = solved;
}
}
return;
}
for (int col = 0; col < panel_width; ++col) {
float partial = 0.0f;
for (int p = lane; p < col; p += 32) {
partial = fmaf(
row_ptr[panel_start + p],
diagonal_panel[col * panel_width + p],
partial);
}
partial = warp_sum(partial);
if (lane == 0) {
row_ptr[panel_start + col] =
(row_ptr[panel_start + col] - partial) /
diagonal_panel[col * panel_width + col];
}
__syncwarp();
}
}
// Each 16x16 CTA computes a 64x64 output tile with a 4x4 register tile per
// thread. Padding the panel dimension removes same-bank row strides.
__global__ void trailing_update_kernel(
float* __restrict__ output,
int64_t matrix_stride,
int n,
int panel_start,
int panel_width,
int trailing_start,
int tile_count) {
__shared__ float panel_rows[2][kTile][kPanel + 1];
__shared__ int tile_row_shared;
__shared__ int tile_col_shared;
const int tid = static_cast<int>(threadIdx.y) * blockDim.x + threadIdx.x;
if (tid == 0) {
const int64_t linear_tile = blockIdx.x;
int tile_row = static_cast<int>(
floor((sqrt(8.0 * static_cast<double>(linear_tile) + 1.0) - 1.0) * 0.5));
while (static_cast<int64_t>(tile_row + 1) * (tile_row + 2) / 2 <= linear_tile) {
++tile_row;
}
while (static_cast<int64_t>(tile_row) * (tile_row + 1) / 2 > linear_tile) {
--tile_row;
}
tile_row_shared = tile_row;
tile_col_shared = static_cast<int>(
linear_tile - static_cast<int64_t>(tile_row) * (tile_row + 1) / 2);
}
__syncthreads();
const int tile_row = tile_row_shared;
const int tile_col = tile_col_shared;
if (tile_row >= tile_count || tile_col > tile_row) {
return;
}
const int matrix = static_cast<int>(blockIdx.y);
float* l = output + static_cast<int64_t>(matrix) * matrix_stride;
const int staged_elements = 2 * kTile * panel_width;
for (int linear = tid; linear < staged_elements; linear += blockDim.x * blockDim.y) {
const int side = linear / (kTile * panel_width);
const int remaining = linear - side * kTile * panel_width;
const int local_row = remaining / panel_width;
const int p = remaining - local_row * panel_width;
const int global_row = trailing_start +
(side == 0 ? tile_row : tile_col) * kTile + local_row;
panel_rows[side][local_row][p] = global_row < n
? l[static_cast<int64_t>(global_row) * n + panel_start + p]
: 0.0f;
}
__syncthreads();
float accum[4][4] = {};
#pragma unroll 1
for (int p = 0; p < panel_width; ++p) {
float left[4];
float right[4];
#pragma unroll
for (int i = 0; i < 4; ++i) {
left[i] = panel_rows[0][threadIdx.y + i * 16][p];
right[i] = panel_rows[1][threadIdx.x + i * 16][p];
}
#pragma unroll
for (int i = 0; i < 4; ++i) {
#pragma unroll
for (int j = 0; j < 4; ++j) {
accum[i][j] = fmaf(-left[i], right[j], accum[i][j]);
}
}
}
#pragma unroll
for (int i = 0; i < 4; ++i) {
const int row = trailing_start + tile_row * kTile + threadIdx.y + i * 16;
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int col = trailing_start + tile_col * kTile + threadIdx.x + j * 16;
if (row < n && col < n && col <= row) {
const int64_t index = static_cast<int64_t>(row) * n + col;
l[index] += accum[i][j];
}
}
}
}
} // namespace
torch::Tensor raw_cholesky(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be torch.float32");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch, n, n]");
TORCH_CHECK(input.size(1) == input.size(2), "matrices must be square");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
c10::cuda::CUDAGuard device_guard(input.device());
auto output = torch::empty_like(input);
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
const int64_t matrix_stride = static_cast<int64_t>(n) * n;
const int64_t total_elements = static_cast<int64_t>(batch) * matrix_stride;
if (batch == 0 || n == 0) {
return output;
}
const float* input_ptr = input.data_ptr<float>();
float* output_ptr = output.data_ptr<float>();
if (n == 32) {
constexpr int matrices_per_cta = 16;
constexpr size_t shared_bytes =
matrices_per_cta * 32 * 33 / 2 * sizeof(float);
const int blocks = (batch + matrices_per_cta - 1) / matrices_per_cta;
warp_packed_cholesky_kernel<32, matrices_per_cta>
<<<blocks, matrices_per_cta * 32, shared_bytes>>>(
input_ptr, output_ptr, batch, matrix_stride);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
if (n <= kSmallMax) {
const size_t shared_bytes =
static_cast<size_t>(n) * (n + 1) / 2 * sizeof(float);
if (n == 64) {
small_cholesky_kernel<64><<<batch, kSmallMax, shared_bytes>>>(
input_ptr, output_ptr, matrix_stride, n);
} else if (n == 128) {
small_cholesky_kernel<128><<<batch, kSmallMax, shared_bytes>>>(
input_ptr, output_ptr, matrix_stride, n);
} else {
small_cholesky_kernel<0><<<batch, kSmallMax, shared_bytes>>>(
input_ptr, output_ptr, matrix_stride, n);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
const int64_t copy_blocks_64 = (total_elements + kThreads - 1) / kThreads;
const int copy_blocks = static_cast<int>(copy_blocks_64 > 2147483647LL
? 2147483647LL
: copy_blocks_64);
copy_lower_kernel<<<copy_blocks, kThreads>>>(
input_ptr, output_ptr, total_elements, n);
for (int panel_start = 0; panel_start < n; panel_start += kPanel) {
const int remaining = n - panel_start;
const int panel_width = remaining < kPanel ? remaining : kPanel;
if (panel_width == kPanel) {
diagonal_factor_kernel<kPanel><<<batch, 128>>>(
output_ptr, matrix_stride, n, panel_start, panel_width);
} else {
diagonal_factor_kernel<0><<<batch, 128>>>(
output_ptr, matrix_stride, n, panel_start, panel_width);
}
const int trailing_start = panel_start + panel_width;
const int trailing_rows = n - trailing_start;
if (trailing_rows == 0) {
continue;
}
const dim3 solve_grid(
static_cast<unsigned>((trailing_rows + 7) / 8),
static_cast<unsigned>(batch));
const size_t solve_shared_bytes =
static_cast<size_t>(panel_width) * panel_width * sizeof(float);
if (panel_width == kPanel) {
panel_solve_kernel<kPanel>
<<<solve_grid, kThreads, solve_shared_bytes>>>(
output_ptr,
matrix_stride,
n,
panel_start,
panel_width);
} else {
panel_solve_kernel<0>
<<<solve_grid, kThreads, solve_shared_bytes>>>(
output_ptr,
matrix_stride,
n,
panel_start,
panel_width);
}
const int tile_count = (trailing_rows + kTile - 1) / kTile;
const int64_t triangular_tiles =
static_cast<int64_t>(tile_count) * (tile_count + 1) / 2;
const dim3 grid(
static_cast<unsigned>(triangular_tiles),
static_cast<unsigned>(batch));
const dim3 block(16, 16);
trailing_update_kernel<<<grid, block>>>(
output_ptr,
matrix_stride,
n,
panel_start,
panel_width,
trailing_start,
tile_count);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor copy_lower(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be a CUDA tensor");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be torch.float32");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch, n, n]");
TORCH_CHECK(input.size(1) == input.size(2), "matrices must be square");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
c10::cuda::CUDAGuard device_guard(input.device());
auto output = torch::empty_like(input);
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
const int64_t total_elements = static_cast<int64_t>(batch) * n * n;
if (total_elements == 0) {
return output;
}
const int64_t copy_blocks_64 = (total_elements + kThreads - 1) / kThreads;
const int copy_blocks = static_cast<int>(copy_blocks_64 > 2147483647LL
? 2147483647LL
: copy_blocks_64);
copy_lower_kernel<<<copy_blocks, kThreads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), total_elements, n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
void factor_panel(torch::Tensor output, int64_t panel_start_64) {
TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
TORCH_CHECK(output.dim() == 3, "output must have shape [batch, n, n]");
TORCH_CHECK(output.size(1) == output.size(2), "matrices must be square");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
c10::cuda::CUDAGuard device_guard(output.device());
const int batch = static_cast<int>(output.size(0));
const int n = static_cast<int>(output.size(1));
const int panel_start = static_cast<int>(panel_start_64);
TORCH_CHECK(panel_start >= 0 && panel_start < n, "invalid panel start");
const int remaining = n - panel_start;
const int panel_width = remaining < kPanel ? remaining : kPanel;
const int64_t matrix_stride = static_cast<int64_t>(n) * n;
float* output_ptr = output.data_ptr<float>();
if (panel_width == kPanel) {
diagonal_factor_kernel<kPanel><<<batch, 128>>>(
output_ptr, matrix_stride, n, panel_start, panel_width);
} else {
diagonal_factor_kernel<0><<<batch, 128>>>(
output_ptr, matrix_stride, n, panel_start, panel_width);
}
const int trailing_rows = n - panel_start - panel_width;
if (trailing_rows > 0) {
const dim3 solve_grid(
static_cast<unsigned>((trailing_rows + 7) / 8),
static_cast<unsigned>(batch));
const size_t solve_shared_bytes =
static_cast<size_t>(panel_width) * panel_width * sizeof(float);
if (panel_width == kPanel) {
panel_solve_kernel<kPanel>
<<<solve_grid, kThreads, solve_shared_bytes>>>(
output_ptr,
matrix_stride,
n,
panel_start,
panel_width);
} else {
panel_solve_kernel<0>
<<<solve_grid, kThreads, solve_shared_bytes>>>(
output_ptr,
matrix_stride,
n,
panel_start,
panel_width);
}
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void factor_panel_128(torch::Tensor output, int64_t panel_start_64) {
TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
TORCH_CHECK(output.dim() == 3, "output must have shape [batch, n, n]");
TORCH_CHECK(output.size(1) == output.size(2), "matrices must be square");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
c10::cuda::CUDAGuard device_guard(output.device());
const int batch = static_cast<int>(output.size(0));
const int n = static_cast<int>(output.size(1));
const int panel_start = static_cast<int>(panel_start_64);
TORCH_CHECK(
n >= 2048 && n % kLargePanel == 0,
"128-column panel tier requires n>=2048 divisible by 128");
TORCH_CHECK(
panel_start >= 0 && panel_start + kLargePanel <= n &&
panel_start % kLargePanel == 0,
"invalid 128-column panel start");
if (batch == 0) {
return;
}
constexpr int diagonal_shared_bytes =
kLargePanel * (kLargePanel + 1) * sizeof(float);
constexpr int solve_shared_bytes =
kLargePanel * kLargePanel * sizeof(float);
static thread_local int configured_device = -1;
int current_device = -1;
C10_CUDA_CHECK(cudaGetDevice(¤t_device));
if (configured_device != current_device) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
diagonal_factor_128_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
diagonal_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
panel_solve_kernel<kLargePanel>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
solve_shared_bytes));
configured_device = current_device;
}
const int64_t matrix_stride = static_cast<int64_t>(n) * n;
float* output_ptr = output.data_ptr<float>();
diagonal_factor_128_kernel<<<batch, kThreads, diagonal_shared_bytes>>>(
output_ptr, matrix_stride, n, panel_start);
const int trailing_rows = n - panel_start - kLargePanel;
if (trailing_rows > 0) {
constexpr int rows_per_cta = kLargeSolveThreads / 32;
const dim3 solve_grid(
static_cast<unsigned>(
(trailing_rows + rows_per_cta - 1) / rows_per_cta),
static_cast<unsigned>(batch));
panel_solve_kernel<kLargePanel>
<<<solve_grid, kLargeSolveThreads, solve_shared_bytes>>>(
output_ptr,
matrix_stride,
n,
panel_start,
kLargePanel);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void factor_panel_inverse(torch::Tensor output, int64_t panel_start_64) {
TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
TORCH_CHECK(output.dim() == 3, "output must have shape [batch, n, n]");
TORCH_CHECK(output.size(1) == output.size(2), "matrices must be square");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
c10::cuda::CUDAGuard device_guard(output.device());
const int batch = static_cast<int>(output.size(0));
const int n = static_cast<int>(output.size(1));
const int panel_start = static_cast<int>(panel_start_64);
TORCH_CHECK(n == 512 && batch >= 128, "inverse panel tier requires n=512, batch>=128");
TORCH_CHECK(panel_start >= 0 && panel_start < n, "invalid panel start");
const int panel_width = n - panel_start < kPanel ? n - panel_start : kPanel;
TORCH_CHECK(panel_width == kPanel, "inverse panel tier requires full panels");
const int64_t matrix_stride = static_cast<int64_t>(n) * n;
float* output_ptr = output.data_ptr<float>();
diagonal_factor_inverse_kernel<kPanel><<<batch, 128>>>(
output_ptr, matrix_stride, n, panel_start, panel_width);
const int trailing_rows = n - panel_start - panel_width;
if (trailing_rows > 0) {
const dim3 solve_grid(
static_cast<unsigned>((trailing_rows + 15) / 16),
static_cast<unsigned>(batch));
constexpr size_t inverse_shared_bytes =
(kPanel * (kPanel + 1) + 16 * (kPanel + 1)) * sizeof(float);
panel_solve_inverse_kernel<<<solve_grid, kThreads, inverse_shared_bytes>>>(
output_ptr, matrix_stride, n, panel_start);
}
if (panel_start + panel_width == n) {
constexpr int panel_count = 512 / kPanel;
const dim3 clear_grid(
static_cast<unsigned>(panel_count),
static_cast<unsigned>(batch));
clear_diagonal_upper_kernel<<<clear_grid, kThreads>>>(
output_ptr, matrix_stride, n, panel_count);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
// B640_FUSED_V1_BEGIN
void factor_panel_cutlass_inverse_b640(
torch::Tensor output,
torch::Tensor inverse,
int64_t panel_start_64) {
TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be torch.float32");
TORCH_CHECK(
output.dim() == 3 && output.size(0) == 640 &&
output.size(1) == 512 && output.size(2) == 512,
"output must have exact shape [640, 512, 512]");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
TORCH_CHECK(inverse.is_cuda(), "inverse must be a CUDA tensor");
TORCH_CHECK(inverse.scalar_type() == at::kFloat, "inverse must be torch.float32");
TORCH_CHECK(
inverse.dim() == 3 && inverse.size(0) == 640 &&
inverse.size(1) == kPanel && inverse.size(2) == kPanel,
"inverse must have exact shape [640, 64, 64]");
TORCH_CHECK(inverse.is_contiguous(), "inverse must be contiguous");
TORCH_CHECK(
inverse.device() == output.device(),
"output and inverse must be on the same CUDA device");
c10::cuda::CUDAGuard device_guard(output.device());
const int panel_start = static_cast<int>(panel_start_64);
TORCH_CHECK(
panel_start >= 0 && panel_start + kPanel <= 512 &&
panel_start % kPanel == 0,
"panel start must identify a full aligned 64-column panel");
constexpr int batch = 640;
constexpr int n = 512;
constexpr int64_t matrix_stride = static_cast<int64_t>(n) * n;
constexpr int64_t inverse_stride = kPanel * kPanel;
diagonal_factor_cutlass_inverse_kernel<<<batch, 128>>>(
output.data_ptr<float>(),
inverse.data_ptr<float>(),
matrix_stride,
inverse_stride,
n,
panel_start);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
// B640_FUSED_V1_END
"""
_extension = load_inline(
# B640_FUSED_V1_BEGIN
name="gpumode_raw_cholesky_combo_k32_cutlass64_b640_fused_v1",
# B640_FUSED_V1_END
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=[
"raw_cholesky",
"copy_lower",
"factor_panel",
"factor_panel_128",
"factor_panel_inverse",
# B640_FUSED_V1_BEGIN
"factor_panel_cutlass_inverse_b640",
# B640_FUSED_V1_END
],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3"],
with_cuda=True,
verbose=False,
)
@triton.jit
def _copy_lower_2d(
input,
output,
n,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
row_block = tl.program_id(0)
col_block = tl.program_id(1)
matrix = tl.program_id(2)
rows = row_block * BLOCK_M + tl.arange(0, BLOCK_M)
cols = col_block * BLOCK_N + tl.arange(0, BLOCK_N)
offsets = matrix * n * n + rows[:, None] * n + cols[None, :]
valid = (rows[:, None] < n) & (cols[None, :] < n)
lower = cols[None, :] <= rows[:, None]
values = tl.load(input + offsets, mask=valid & lower, other=0.0)
tl.store(output + offsets, values, mask=valid)
@triton.jit
def _update_factor_final64_n256(
output,
N: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
matrix = tl.program_id(0)
rows = tl.arange(0, BLOCK_SIZE)
cols = tl.arange(0, BLOCK_SIZE)
matrix_base = matrix * N * N
panel_offsets = (
matrix_base
+ (192 + rows[:, None]) * N
+ 128
+ cols[None, :]
)
left = tl.load(output + panel_offsets)
right = tl.load(output + panel_offsets)
left_hi = left.to(tl.bfloat16)
right_hi = right.to(tl.bfloat16)
left_lo = (left - left_hi.to(tl.float32)).to(tl.bfloat16)
right_lo = (right - right_hi.to(tl.float32)).to(tl.bfloat16)
update = tl.dot(
left_hi,
tl.trans(right_hi),
out_dtype=tl.float32,
)
update += tl.dot(
left_hi,
tl.trans(right_lo),
out_dtype=tl.float32,
)
update += tl.dot(
left_lo,
tl.trans(right_hi),
out_dtype=tl.float32,
)
final_offsets = (
matrix_base
+ (192 + rows[:, None]) * N
+ 192
+ cols[None, :]
)
lower = cols[None, :] <= rows[:, None]
factor = tl.load(
output + final_offsets,
mask=lower,
other=0.0,
).to(tl.float32)
factor -= update
for pivot in tl.range(0, BLOCK_SIZE):
diagonal_rows = tl.sum(
tl.where(
(rows[:, None] == pivot) & (cols[None, :] == pivot),
factor,
0.0,
),
axis=1,
)
diagonal = tl.sum(diagonal_rows, axis=0)
diagonal = tl.sqrt(
tl.maximum(diagonal, 1.1754943508222875e-38)
)
column = tl.sum(
tl.where(cols[None, :] == pivot, factor, 0.0),
axis=1,
)
column = column / diagonal
column = tl.where(rows < pivot, 0.0, column)
factor = tl.where(
(cols[None, :] == pivot) & (rows[:, None] >= pivot),
column[:, None],
factor,
)
trailing_mask = (
(rows[:, None] > pivot)
& (cols[None, :] > pivot)
& (cols[None, :] <= rows[:, None])
)
factor = tl.where(
trailing_mask,
factor - column[:, None] * column[None, :],
factor,
)
tl.store(output + final_offsets, factor, mask=lower)
@triton.jit
def _factor_solve_panel_128_256(output, n):
matrix = tl.program_id(0)
rr = tl.arange(0, 64)[:, None]
cc = tl.arange(0, 64)[None, :]
matrix_base = matrix * n * n
diagonal_offsets = matrix_base + (128 + rr) * n + 128 + cc
solve_offsets = matrix_base + (192 + rr) * n + 128 + cc
lower = cc <= rr
factor = tl.load(
output + diagonal_offsets,
mask=lower,
other=0.0,
).to(tl.float32)
solved = tl.load(output + solve_offsets).to(tl.float32)
for p in tl.range(0, 64):
diagonal = tl.sum(
tl.sum(
tl.where((rr == p) & (cc == p), factor, 0.0),
axis=1,
),
axis=0,
)
diagonal = tl.sqrt(
tl.maximum(diagonal, 1.1754943508222875e-38)
)
column = tl.sum(
tl.where(cc == p, factor, 0.0),
axis=1,
)[:, None]
column = tl.where(
rr > p,
column / diagonal,
tl.where(rr == p, diagonal, 0.0),
)
factor = tl.where(
cc == p,
column,
factor,
)
factor = tl.where(
(rr > p) & (cc > p) & (cc <= rr),
factor - column * tl.trans(column),
factor,
)
solved = tl.where(cc == p, solved / diagonal, solved)
pivot = tl.sum(
tl.where(cc == p, solved, 0.0),
axis=1,
)[:, None]
coefficients = tl.sum(
tl.where(cc == p, factor, 0.0),
axis=1,
)[None, :]
solved = tl.where(
cc > p,
solved - pivot * coefficients,
solved,
)
tl.store(output + diagonal_offsets, factor, mask=lower)
tl.store(output + solve_offsets, solved)
@triton.jit
def _trailing_update_tf32(
output,
n,
panel_start,
BLOCK: tl.constexpr,
COMPENSATED: tl.constexpr,
):
triangular_tile = tl.program_id(0)
matrix = tl.program_id(1)
# Invert row*(row+1)/2 to map one-dimensional program IDs onto only the
# lower-triangular output tiles. This avoids launching/computing the unused
# upper half of the trailing matrix.
triangular_float = triangular_tile.to(tl.float32)
tile_row = tl.floor(
(tl.sqrt(8.0 * triangular_float + 1.0) - 1.0) * 0.5
).to(tl.int32)
tile_col = triangular_tile - tile_row * (tile_row + 1) // 2
trailing_start = panel_start + BLOCK
rows = trailing_start + tile_row * BLOCK + tl.arange(0, BLOCK)
cols = trailing_start + tile_col * BLOCK + tl.arange(0, BLOCK)
reduction = panel_start + tl.arange(0, BLOCK)
matrix_base = matrix * n * n
left = tl.load(
output + matrix_base + rows[:, None] * n + reduction[None, :]
)
right = tl.load(
output + matrix_base + cols[:, None] * n + reduction[None, :]
)
if COMPENSATED:
left_hi = left.to(tl.bfloat16)
right_hi = right.to(tl.bfloat16)
left_lo = (left - left_hi.to(tl.float32)).to(tl.bfloat16)
right_lo = (right - right_hi.to(tl.float32)).to(tl.bfloat16)
update = tl.dot(
left_hi, tl.trans(right_hi), out_dtype=tl.float32
)
update += tl.dot(
left_hi, tl.trans(right_lo), out_dtype=tl.float32
)
update += tl.dot(
left_lo, tl.trans(right_hi), out_dtype=tl.float32
)
else:
update = tl.dot(left, tl.trans(right), input_precision="tf32")
output_offsets = matrix_base + rows[:, None] * n + cols[None, :]
old = tl.load(output + output_offsets)
lower_mask = cols[None, :] <= rows[:, None]
tl.store(output + output_offsets, old - update, mask=lower_mask)
@triton.jit
def _trailing_update_tf32_panel128(
output,
n,
panel_start,
BLOCK: tl.constexpr,
PANEL: tl.constexpr,
):
triangular_tile = tl.program_id(0)
matrix = tl.program_id(1)
triangular_float = triangular_tile.to(tl.float32)
tile_row = tl.floor(
(tl.sqrt(8.0 * triangular_float + 1.0) - 1.0) * 0.5
).to(tl.int32)
tile_col = triangular_tile - tile_row * (tile_row + 1) // 2
trailing_start = panel_start + PANEL
rows = trailing_start + tile_row * BLOCK + tl.arange(0, BLOCK)
cols = trailing_start + tile_col * BLOCK + tl.arange(0, BLOCK)
reduction = panel_start + tl.arange(0, BLOCK // 2)
matrix_base = matrix * n * n
left = tl.load(
output + matrix_base + rows[:, None] * n + reduction[None, :]
)
right = tl.load(
output + matrix_base + cols[:, None] * n + reduction[None, :]
)
update = tl.dot(left, tl.trans(right), input_precision="tf32")
reduction += BLOCK // 2
left = tl.load(
output + matrix_base + rows[:, None] * n + reduction[None, :]
)
right = tl.load(
output + matrix_base + cols[:, None] * n + reduction[None, :]
)
update += tl.dot(left, tl.trans(right), input_precision="tf32")
reduction += BLOCK // 2
left = tl.load(
output + matrix_base + rows[:, None] * n + reduction[None, :]
)
right = tl.load(
output + matrix_base + cols[:, None] * n + reduction[None, :]
)
update += tl.dot(left, tl.trans(right), input_precision="tf32")
reduction += BLOCK // 2
left = tl.load(
output + matrix_base + rows[:, None] * n + reduction[None, :]
)
right = tl.load(
output + matrix_base + cols[:, None] * n + reduction[None, :]
)
update += tl.dot(left, tl.trans(right), input_precision="tf32")
output_offsets = matrix_base + rows[:, None] * n + cols[None, :]
old = tl.load(output + output_offsets)
lower_mask = cols[None, :] <= rows[:, None]
tl.store(output + output_offsets, old - update, mask=lower_mask)
# B640_FUSED_V1_BEGIN
@triton.jit
def _b640_fused_panel_schur_inverse_apply_v1(
data,
output,
inverse,
PANEL_START,
BLOCK_M: tl.constexpr,
):
row_block = tl.program_id(0)
matrix = tl.program_id(1)
rows = PANEL_START + 64 + row_block * BLOCK_M + tl.arange(0, BLOCK_M)
panel_cols = PANEL_START + tl.arange(0, 64)
matrix_base = matrix * 512 * 512
panel = tl.load(
data + matrix_base + rows[:, None] * 512 + panel_cols[None, :]
).to(tl.float32)
for previous_start in tl.static_range(0, 512, 64):
if previous_start < PANEL_START:
previous_cols = previous_start + tl.arange(0, 64)
left = tl.load(
output + matrix_base + rows[:, None] * 512 + previous_cols[None, :]
)
right = tl.load(
output
+ matrix_base
+ panel_cols[:, None] * 512
+ previous_cols[None, :]
)
panel -= tl.dot(left, tl.trans(right), input_precision="tf32")
inverse_rows = tl.arange(0, 64)
inverse_cols = tl.arange(0, 64)
upper_inverse = tl.load(
inverse
+ matrix * 64 * 64
+ inverse_rows[:, None] * 64
+ inverse_cols[None, :]
)
solved = tl.dot(panel, upper_inverse, input_precision="tf32")
tl.store(
output + matrix_base + rows[:, None] * 512 + panel_cols[None, :],
solved,
)
# B640_FUSED_V1_END
def _old_custom_kernel(data: input_t) -> output_t:
n = data.shape[-1]
if n < 256 or n % 64 != 0:
return _extension.raw_cholesky(data)
output = torch.empty_like(data)
batch = data.shape[0]
_copy_lower_2d[
(triton.cdiv(n, 16), triton.cdiv(n, 256), batch)
](
data,
output,
n,
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
num_stages=1,
)
if n == 256 and batch == 64:
for panel_start in (0, 64):
_extension.factor_panel(output, panel_start)
trailing_rows = n - panel_start - 64
tile_count = triton.cdiv(trailing_rows, 64)
triangular_tiles = tile_count * (tile_count + 1) // 2
_trailing_update_tf32[(triangular_tiles, batch)](
output,
n,
panel_start,
BLOCK=64,
COMPENSATED=True,
num_warps=8,
num_stages=2,
)
_factor_solve_panel_128_256[(batch,)](
output,
n,
num_warps=8,
num_stages=1,
)
_update_factor_final64_n256[(batch,)](
output,
N=256,
BLOCK_SIZE=64,
num_warps=8,
num_stages=2,
)
return output
if n == 32768:
for panel_start in range(0, n, 128):
_extension.factor_panel_128(output, panel_start)
trailing_rows = n - panel_start - 128
if trailing_rows > 0:
tile_count = triton.cdiv(trailing_rows, 64)
triangular_tiles = tile_count * (tile_count + 1) // 2
_trailing_update_tf32_panel128[(triangular_tiles, batch)](
output,
n,
panel_start,
BLOCK=64,
PANEL=128,
num_warps=8,
num_stages=2,
)
return output
use_inverse = n == 512 and batch >= 128
for panel_start in range(0, n, 64):
if use_inverse and panel_start + 64 < n:
_extension.factor_panel_inverse(output, panel_start)
else:
_extension.factor_panel(output, panel_start)
trailing_rows = n - panel_start - 64
if trailing_rows > 0:
tile_count = triton.cdiv(trailing_rows, 64)
triangular_tiles = tile_count * (tile_count + 1) // 2
_trailing_update_tf32[(triangular_tiles, batch)](
output,
n,
panel_start,
BLOCK=64,
COMPENSATED=n < 2048,
num_warps=8,
num_stages=1 if n >= 2048 else 2,
)
return output
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True
CPP_SRC = r"""torch::Tensor chol_small(torch::Tensor input);"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cmath>
__global__ void register_cholesky32(const float* input, float* output) {
const int lane = threadIdx.x;
const size_t base = static_cast<size_t>(blockIdx.x) * 1024;
float row[32];
// A matrix is 4 KiB aligned, so each lane can move its complete row as
// eight naturally aligned 16-byte transactions rather than scalar loads.
const float4* input4 = reinterpret_cast<const float4*>(input + base + lane * 32);
float4* output4 = reinterpret_cast<float4*>(output + base + lane * 32);
#pragma unroll
for (int vector = 0; vector < 8; ++vector) {
const float4 values = input4[vector];
const int col = vector * 4;
row[col + 0] = col + 0 <= lane ? values.x : 0.0f;
row[col + 1] = col + 1 <= lane ? values.y : 0.0f;
row[col + 2] = col + 2 <= lane ? values.z : 0.0f;
row[col + 3] = col + 3 <= lane ? values.w : 0.0f;
}
#pragma unroll
for (int k = 0; k < 32; ++k) {
if (lane == k) row[k] = sqrtf(fmaxf(row[k], 0.0f));
const float diagonal = __shfl_sync(0xffffffffu, row[k], k);
if (lane > k) row[k] /= diagonal;
const float own = row[k];
#pragma unroll
for (int col = k + 1; col < 32; ++col) {
const float other = __shfl_sync(0xffffffffu, row[k], col);
if (lane >= col) row[col] = fmaf(-own, other, row[col]);
}
}
#pragma unroll
for (int vector = 0; vector < 8; ++vector) {
const int col = vector * 4;
output4[vector] = make_float4(
row[col + 0], row[col + 1], row[col + 2], row[col + 3]);
}
}
__global__ void shared_cholesky64(const float* input, float* output) {
extern __shared__ float matrix[];
const int tid = threadIdx.x;
const size_t base = static_cast<size_t>(blockIdx.x) * 4096;
for (int index = tid; index < 4096; index += blockDim.x) {
const int row = index / 64, col = index - row * 64;
matrix[index] = row >= col ? input[base + index] : 0.0f;
}
__syncthreads();
#pragma unroll
for (int k = 0; k < 64; ++k) {
if (tid == 0) {
float value = matrix[k * 64 + k];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
const float x = matrix[k * 64 + j];
value = fmaf(-x, x, value);
}
matrix[k * 64 + k] = sqrtf(fmaxf(value, 0.0f));
}
__syncthreads();
if (tid > k && tid < 64) {
float value = matrix[tid * 64 + k];
#pragma unroll 4
for (int j = 0; j < k; ++j)
value = fmaf(-matrix[tid * 64 + j], matrix[k * 64 + j], value);
matrix[tid * 64 + k] = value / matrix[k * 64 + k];
}
__syncthreads();
}
for (int index = tid; index < 4096; index += blockDim.x)
output[base + index] = matrix[index];
}
torch::Tensor chol_small(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == at::kFloat && input.is_contiguous(),
"expected contiguous CUDA float32");
auto output = torch::empty_like(input);
const int64_t n = input.size(1), batch = input.size(0);
if (n == 32) register_cholesky32<<<batch, 32>>>(input.data_ptr<float>(), output.data_ptr<float>());
else if (n == 64) shared_cholesky64<<<batch, 64, 16384>>>(input.data_ptr<float>(), output.data_ptr<float>());
else TORCH_CHECK(false, "unsupported n");
cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
return output;
}
"""
_native = load_inline(
name="hopper_chol_merged_v5",
cpp_sources=[CPP_SRC], cuda_sources=[CUDA_SRC], functions=["chol_small"],
extra_cuda_cflags=["-O3"], verbose=False,
)
def _blocked_inverse(data: torch.Tensor, block_size: int) -> torch.Tensor:
batch, n, _ = data.shape
lower = torch.zeros_like(data)
identity = torch.eye(block_size, dtype=data.dtype, device=data.device).expand(
batch, block_size, block_size
)
for start in range(0, n, block_size):
end = start + block_size
diagonal = data[:, start:end, start:end]
if start:
previous = lower[:, start:end, :start]
diagonal = diagonal - previous @ previous.transpose(-1, -2)
factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
lower[:, start:end, start:end] = factor
if end < n:
panel = data[:, end:, start:end]
if start:
panel = panel - lower[:, end:, :start] @ lower[:, start:end, :start].transpose(-1, -2)
inverse = torch.linalg.solve_triangular(factor, identity, upper=False)
lower[:, end:, start:end] = panel @ inverse.transpose(-1, -2)
return lower
# B640_FUSED_V1_BEGIN
def _b640_fused_cholesky_v1(data: torch.Tensor) -> torch.Tensor:
output = torch.zeros_like(data)
inverse = torch.empty((640, 64, 64), dtype=torch.float32, device=data.device)
for panel_start in range(0, 512, 64):
panel_end = panel_start + 64
diagonal = data[:, panel_start:panel_end, panel_start:panel_end]
if panel_start:
previous = output[:, panel_start:panel_end, :panel_start]
diagonal = diagonal - previous @ previous.transpose(-1, -2)
diagonal_output = output[
:, panel_start:panel_end, panel_start:panel_end
]
diagonal_output.copy_(diagonal)
diagonal_output.tril_()
_extension.factor_panel_cutlass_inverse_b640(
output, inverse, panel_start
)
if panel_end < 512:
_b640_fused_panel_schur_inverse_apply_v1[
((512 - panel_end) // 64, 640)
](
data,
output,
inverse,
PANEL_START=panel_start,
BLOCK_M=64,
num_warps=8,
num_stages=2,
)
return output
# B640_FUSED_V1_END
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
if n in (32, 64):
return _native.chol_small(data)
if batch == 640 and n == 512:
# B640_FUSED_V1_BEGIN
return _b640_fused_cholesky_v1(data)
# B640_FUSED_V1_END
if batch == 60 and n == 1024:
return _blocked_inverse(data, 128)
if batch == 8 and n == 2048:
return _old_custom_kernel(data)
if batch == 1 and n == 8192:
return _blocked_inverse(data, 2048)
if batch == 1 and n == 16384:
return _blocked_inverse(data, 2048)
if batch == 1 and n == 32768:
return _blocked_inverse(data, 1024)
if batch == 2 and n in (2048, 4096):
return torch.cat([
torch.linalg.cholesky_ex(
data[i:i+1].transpose(-1, -2), check_errors=False
).L
for i in range(batch)
], dim=0)
if (batch, n) in (
(256, 128),
(64, 256),
(16, 512),
(4, 1024),
(1, 4096),
):
return torch.linalg.cholesky_ex(
data.transpose(-1, -2), check_errors=False
).L
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 1809 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON