submission 915822
Chanho Lee · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 5275 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-915822?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:62b76b9b54527dcc5eca19d11a3fb9cbada8dc6f5988a3e80a718115642db002
license declaredunknown
license concludedunknown
authorsChanho Lee
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;"mma
nvcuda::wmma::fragment<shared-memory
__shared__ float factor[4][n][n + 1];vector-width = float4
const float4 output_values = make_float4(Kernel source
submission.py5275 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
torch.backends.cuda.preferred_linalg_library("cusolver")
CUDA_SOURCE = r"""
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <mma.h>
#include <cusolverDn.h>
#include <cublas_v2.h>
#include <cublasLt.h>
#include <array>
#include <cstdint>
#include <stdexcept>
#include <vector>
__device__ __forceinline__ float sqrt_approx(float value) {
float result;
asm volatile("sqrt.approx.f32 %0, %1;" : "=f"(result) : "f"(value));
return result;
}
__device__ __forceinline__ float rcp_approx(float value) {
float result;
asm volatile("rcp.approx.f32 %0, %1;" : "=f"(result) : "f"(value));
return result;
}
// C560A_HELPERS_BEGIN
__device__ __forceinline__ unsigned int c560a_shared_address(
const void* pointer
) {
return static_cast<unsigned int>(__cvta_generic_to_shared(pointer));
}
__device__ __forceinline__ void c560a_mbar_init(
unsigned long long* barrier,
int arrivals
) {
const unsigned int address = c560a_shared_address(barrier);
asm volatile(
"mbarrier.init.shared::cta.b64 [%0], %1;"
:: "r"(address), "r"(arrivals)
: "memory"
);
}
__device__ __forceinline__ void c560a_mbar_arrive(
unsigned long long* barrier
) {
const unsigned int address = c560a_shared_address(barrier);
asm volatile(
"mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];"
:: "r"(address)
: "memory"
);
}
__device__ __forceinline__ void c560a_mbar_wait(
unsigned long long* barrier,
int parity
) {
const unsigned int address = c560a_shared_address(barrier);
asm volatile(
"{\n\t"
".reg .pred ready;\n\t"
"C560A_WAIT_%=: "
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 "
"ready, [%0], %1, 10000000;\n\t"
"@!ready bra.uni C560A_WAIT_%=;\n\t"
"}"
:: "r"(address), "r"(parity)
: "memory"
);
}
// C560A_HELPERS_END
__device__ __forceinline__ void panel_barrier128() {
asm volatile("bar.sync 1, 128;" ::: "memory");
}
template <typename Kernel, typename... Args>
cudaError_t launch_pdl(
Kernel kernel,
dim3 grid,
dim3 block,
size_t shared_bytes,
Args... args
) {
cudaLaunchConfig_t config = {};
config.gridDim = grid;
config.blockDim = block;
config.dynamicSmemBytes = shared_bytes;
cudaLaunchAttribute attribute;
attribute.id = static_cast<cudaLaunchAttributeID>(6);
*reinterpret_cast<int*>(&attribute.val) = 1;
config.attrs = &attribute;
config.numAttrs = 1;
return cudaLaunchKernelEx(&config, kernel, args...);
}
__global__ __launch_bounds__(128) void cholesky32_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
constexpr int n = 32;
__shared__ float factor[4][n][n + 1];
const int warp = threadIdx.x >> 5;
const int row = threadIdx.x & 31;
const int matrix_index = blockIdx.x * 4 + warp;
if (matrix_index >= batch) {
return;
}
float (*tile)[n + 1] = factor[warp];
const long long batch_offset =
static_cast<long long>(matrix_index) * n * n;
const float* matrix = input + batch_offset;
float* result = output + batch_offset;
#pragma unroll
for (int index = row; index < n * n; index += 32) {
const int load_row = index >> 5;
const int load_column = index & 31;
tile[load_row][load_column] = matrix[index];
}
__syncwarp();
const int trailing_column_lane = row & 7;
const int trailing_row_lane = row >> 3;
#pragma unroll
for (int factor_panel = 0; factor_panel < n; factor_panel += 8) {
#pragma unroll
for (int panel_column = 0; panel_column < 8; ++panel_column) {
const int column = factor_panel + panel_column;
float diagonal_value = tile[column][column];
#pragma unroll
for (int inner = factor_panel; inner < column; ++inner) {
const float value = tile[column][inner];
diagonal_value = fmaf(-value, value, diagonal_value);
}
diagonal_value = sqrt_approx(diagonal_value);
for (int factor_row = column + 1 + row;
factor_row < n;
factor_row += 32) {
float value = tile[factor_row][column];
#pragma unroll
for (int inner = factor_panel; inner < column; ++inner) {
value = fmaf(
-tile[factor_row][inner],
tile[column][inner],
value
);
}
tile[factor_row][column] =
value * rcp_approx(diagonal_value);
}
if (row == 0) {
tile[column][column] = diagonal_value;
}
__syncwarp();
}
for (int trailing_row = factor_panel + 8 + trailing_row_lane;
trailing_row < n;
trailing_row += 4) {
for (int trailing_column =
factor_panel + 8 + trailing_column_lane;
trailing_column <= trailing_row;
trailing_column += 8) {
float value = tile[trailing_row][trailing_column];
#pragma unroll
for (int panel_column = 0; panel_column < 8;
++panel_column) {
value = fmaf(
-tile[trailing_row][factor_panel + panel_column],
tile[trailing_column][factor_panel + panel_column],
value
);
}
tile[trailing_row][trailing_column] = value;
}
}
__syncwarp();
}
#pragma unroll
for (int index = row; index < n * n; index += 32) {
const int store_row = index >> 5;
const int store_column = index & 31;
result[index] = store_column <= store_row
? tile[store_row][store_column]
: 0.0f;
}
}
__global__ __launch_bounds__(64) void cholesky64_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
constexpr int n = 64;
__shared__ float factor[n][n + 1];
const int row = threadIdx.x;
const int matrix_index = blockIdx.x;
const bool active = matrix_index < batch;
float (*tile)[n + 1] = factor;
const long long batch_offset =
static_cast<long long>(matrix_index) * n * n;
const float* matrix = input + batch_offset;
float* result = output + batch_offset;
if (active) {
for (int index = row; index < n * n; index += n) {
const int load_row = index >> 6;
const int load_column = index & 63;
if (load_column <= load_row) {
tile[load_row][load_column] = matrix[index];
}
}
}
__syncthreads();
const int trailing_column_lane = row & 7;
const int trailing_row_lane = row >> 3;
for (int factor_panel = 0; factor_panel < n; factor_panel += 8) {
for (int panel_column = 0; panel_column < 8; ++panel_column) {
const int column = factor_panel + panel_column;
float diagonal_value = 1.0f;
if (active) {
diagonal_value = tile[column][column];
for (int inner = factor_panel; inner < column; ++inner) {
const float value = tile[column][inner];
diagonal_value = fmaf(-value, value, diagonal_value);
}
diagonal_value = sqrt_approx(diagonal_value);
for (int factor_row = column + 1 + row;
factor_row < n;
factor_row += 64) {
float value = tile[factor_row][column];
for (int inner = factor_panel; inner < column; ++inner) {
value = fmaf(
-tile[factor_row][inner],
tile[column][inner],
value
);
}
tile[factor_row][column] =
value * rcp_approx(diagonal_value);
}
if (row == 0) {
tile[column][column] = diagonal_value;
}
}
__syncthreads();
}
if (active) {
for (int trailing_row = factor_panel + 8 + trailing_row_lane;
trailing_row < n;
trailing_row += 8) {
for (int trailing_column =
factor_panel + 8 + trailing_column_lane;
trailing_column <= trailing_row;
trailing_column += 8) {
float value = tile[trailing_row][trailing_column];
#pragma unroll
for (int panel_column = 0; panel_column < 8;
++panel_column) {
value = fmaf(
-tile[trailing_row][
factor_panel + panel_column
],
tile[trailing_column][
factor_panel + panel_column
],
value
);
}
tile[trailing_row][trailing_column] = value;
}
}
}
__syncthreads();
}
if (active) {
for (int vector_index = row;
vector_index < n * n / 4;
vector_index += n) {
const int store_row = vector_index >> 4;
const int store_column = (vector_index & 15) * 4;
const float4 output_values = make_float4(
store_column <= store_row
? tile[store_row][store_column] : 0.0f,
store_column + 1 <= store_row
? tile[store_row][store_column + 1] : 0.0f,
store_column + 2 <= store_row
? tile[store_row][store_column + 2] : 0.0f,
store_column + 3 <= store_row
? tile[store_row][store_column + 3] : 0.0f
);
*reinterpret_cast<float4*>(
result + store_row * n + store_column
) = output_values;
}
}
}
__global__ __launch_bounds__(128) void cholesky64_wmma_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr int n = 64;
constexpr int stride = 68;
extern __shared__ float factor[];
const int thread = threadIdx.x;
const int warp = thread >> 5;
const long long batch_offset = static_cast<long long>(blockIdx.x) * n * n;
const float* matrix = input + batch_offset;
float* result = output + batch_offset;
for (int index = thread; index < n * n; index += blockDim.x) {
const int row = index >> 6;
const int column = index & 63;
if (column <= row) {
factor[row * stride + column] = matrix[index];
}
}
__syncthreads();
for (int panel = 0; panel < n; panel += 16) {
for (int panel_column = 0; panel_column < 16; ++panel_column) {
const int column = panel + panel_column;
float diagonal = factor[column * stride + column];
for (int inner = panel; inner < column; ++inner) {
const float value = factor[column * stride + inner];
diagonal = fmaf(-value, value, diagonal);
}
diagonal = sqrt_approx(diagonal);
for (int row = column + 1 + thread;
row < n;
row += blockDim.x) {
float value = factor[row * stride + column];
for (int inner = panel; inner < column; ++inner) {
value = fmaf(
-factor[row * stride + inner],
factor[column * stride + inner],
value
);
}
factor[row * stride + column] =
value * rcp_approx(diagonal);
}
if (thread == 0) {
factor[column * stride + column] = diagonal;
}
__syncthreads();
}
int tile_index = 0;
for (int tile_row = panel + 16; tile_row < n; tile_row += 16) {
for (int tile_column = panel + 16;
tile_column <= tile_row;
tile_column += 16, ++tile_index) {
if ((tile_index & 3) == warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator,
16, 16, 8,
float
> accumulated;
nvcuda::wmma::load_matrix_sync(
accumulated,
factor + tile_row * stride + tile_column,
stride,
nvcuda::wmma::mem_row_major
);
#pragma unroll
for (int panel_step = 0; panel_step < 16;
panel_step += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left_residual;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> right;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> right_residual;
nvcuda::wmma::load_matrix_sync(
left,
factor + tile_row * stride + panel + panel_step,
stride
);
nvcuda::wmma::load_matrix_sync(
right,
factor + tile_column * stride + panel + panel_step,
stride
);
#pragma unroll
for (int element = 0;
element < left.num_elements;
++element) {
const float full = left.x[element];
const float high = __uint_as_float(
__float_as_uint(full) & 0xFFFFE000u
);
left.x[element] = -high;
left_residual.x[element] = -(full - high);
}
#pragma unroll
for (int element = 0;
element < right.num_elements;
++element) {
const float full = right.x[element];
const float high = __uint_as_float(
__float_as_uint(full) & 0xFFFFE000u
);
right.x[element] = high;
right_residual.x[element] = full - high;
}
nvcuda::wmma::mma_sync(
accumulated, left, right, accumulated
);
nvcuda::wmma::mma_sync(
accumulated, left, right_residual, accumulated
);
nvcuda::wmma::mma_sync(
accumulated, left_residual, right, accumulated
);
}
nvcuda::wmma::store_matrix_sync(
factor + tile_row * stride + tile_column,
accumulated,
stride,
nvcuda::wmma::mem_row_major
);
}
}
}
__syncthreads();
}
for (int index = thread; index < n * n; index += blockDim.x) {
const int row = index >> 6;
const int column = index & 63;
result[index] = column <= row
? factor[row * stride + column]
: 0.0f;
}
}
// C560A_KERNEL_BEGIN
__global__ __launch_bounds__(256, 4)
void cholesky64_coalesced_wavefront_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr int tile_size = 64;
constexpr int warps_per_matrix = 8;
constexpr int columns_per_warp = 8;
static_assert(warps_per_matrix * columns_per_warp == tile_size);
__shared__ float published_col[tile_size][tile_size + 1];
__shared__ unsigned long long column_mbar[tile_size];
const int thread = threadIdx.x;
const int warp = thread >> 5;
const int lane = thread & 31;
const int owned_column_begin = warp * columns_per_warp;
const long long matrix_offset =
static_cast<long long>(blockIdx.x) * tile_size * tile_size;
const float* matrix = input + matrix_offset;
float* result = output + matrix_offset;
if (thread == 0) {
for (int column = 0; column < tile_size; ++column) {
c560a_mbar_init(column_mbar + column, 32);
}
}
// Input phase: use the padded allocation as [row][column]. Consecutive
// threads issue consecutive float4 loads; the next barrier exposes both
// this tile and thread 0's mbarrier initialization.
for (int vector_index = thread;
vector_index < tile_size * tile_size / 4;
vector_index += blockDim.x) {
const int load_row = vector_index >> 4;
const int load_column = (vector_index & 15) * 4;
const float4 input_values = *reinterpret_cast<const float4*>(
matrix + load_row * tile_size + load_column
);
published_col[load_row][load_column] = input_values.x;
published_col[load_row][load_column + 1] = input_values.y;
published_col[load_row][load_column + 2] = input_values.z;
published_col[load_row][load_column + 3] = input_values.w;
}
__syncthreads();
{
float columns[2][columns_per_warp];
#pragma unroll
for (int item = 0; item < 2; ++item) {
const int load_row = lane + item * 32;
#pragma unroll
for (int local = 0; local < columns_per_warp; ++local) {
columns[item][local] = published_col[
load_row
][owned_column_begin + local];
}
}
// Every raw input value is now register-resident. No factor owner may
// reinterpret/overwrite published_col as [column][row] before this.
__syncthreads();
// C560A_FACTOR_HOT_BEGIN
for (int inner = 0; inner < owned_column_begin; ++inner) {
c560a_mbar_wait(column_mbar + inner, 0);
#pragma unroll
for (int local = 0; local < columns_per_warp; ++local) {
const int column = owned_column_begin + local;
const float pivot_row_value =
published_col[inner][column];
#pragma unroll
for (int item = 0; item < 2; ++item) {
const int factor_row = lane + item * 32;
if (factor_row >= column) {
columns[item][local] = fmaf(
-published_col[inner][factor_row],
pivot_row_value,
columns[item][local]
);
}
}
}
}
#pragma unroll
for (int local = 0; local < columns_per_warp; ++local) {
const int column = owned_column_begin + local;
const float diagonal_lane_value =
warp < warps_per_matrix / 2
? columns[0][local]
: columns[1][local];
float diagonal_value = __shfl_sync(
0xFFFFFFFFu,
diagonal_lane_value,
column & 31
);
diagonal_value = sqrt_approx(diagonal_value);
const float inverse_diagonal = rcp_approx(diagonal_value);
#pragma unroll
for (int item = 0; item < 2; ++item) {
const int factor_row = lane + item * 32;
if (factor_row >= column) {
const float factor_value = factor_row == column
? diagonal_value
: columns[item][local] * inverse_diagonal;
columns[item][local] = factor_value;
published_col[column][factor_row] = factor_value;
}
}
c560a_mbar_arrive(column_mbar + column);
#pragma unroll
for (int trailing_local = local + 1;
trailing_local < columns_per_warp;
++trailing_local) {
const int trailing_column =
owned_column_begin + trailing_local;
const float pivot_row_value =
published_col[column][trailing_column];
#pragma unroll
for (int item = 0; item < 2; ++item) {
const int factor_row = lane + item * 32;
if (factor_row >= trailing_column) {
columns[item][trailing_local] = fmaf(
-published_col[column][factor_row],
pivot_row_value,
columns[item][trailing_local]
);
}
}
}
}
// C560A_FACTOR_HOT_END
}
__syncthreads();
for (int vector_index = thread;
vector_index < tile_size * tile_size / 4;
vector_index += blockDim.x) {
const int store_row = vector_index >> 4;
const int store_column = (vector_index & 15) * 4;
const float4 output_values = make_float4(
store_column <= store_row
? published_col[store_column][store_row] : 0.0f,
store_column + 1 <= store_row
? published_col[store_column + 1][store_row] : 0.0f,
store_column + 2 <= store_row
? published_col[store_column + 2][store_row] : 0.0f,
store_column + 3 <= store_row
? published_col[store_column + 3][store_row] : 0.0f
);
*reinterpret_cast<float4*>(
result + store_row * tile_size + store_column
) = output_values;
}
}
// C560A_KERNEL_END
__global__ __launch_bounds__(64) void cholesky64_diagonal_kernel(
float* __restrict__ matrices,
const float* __restrict__ input_matrices,
float* __restrict__ inverse_matrices,
int batch,
int n,
int offset
) {
constexpr int tile_size = 64;
__shared__ float factor[tile_size][tile_size + 4];
__shared__ float diagonal[tile_size];
__shared__ float inverse_product[32][36];
__shared__ bool exact_inverse_fallback;
const int row = threadIdx.x & 63;
const int group_warp = row >> 5;
const int matrix_index = blockIdx.x;
const bool active = matrix_index < batch;
float (*tile)[tile_size + 4] = factor;
const long long matrix_offset =
static_cast<long long>(matrix_index) * n * n;
float* matrix = matrices + matrix_offset;
const float* factor_source =
input_matrices != nullptr && (
offset == 0 || (n == 512 && batch == 640)
)
? input_matrices + matrix_offset
: matrix;
cudaGridDependencySynchronize();
if (active) {
for (int vector_index = row;
vector_index < tile_size * tile_size / 4;
vector_index += tile_size) {
const int load_row = vector_index >> 4;
const int load_column = (vector_index & 15) * 4;
const float4 input_values = *reinterpret_cast<const float4*>(
factor_source +
(offset + load_row) * n + offset + load_column
);
tile[load_row][load_column] = input_values.x;
tile[load_row][load_column + 1] = input_values.y;
tile[load_row][load_column + 2] = input_values.z;
tile[load_row][load_column + 3] = input_values.w;
}
}
__syncthreads();
for (int factor_panel = 0; factor_panel < tile_size;
factor_panel += 16) {
for (int panel_column = 0; panel_column < 16; ++panel_column) {
const int column = factor_panel + panel_column;
float diagonal_value = tile[column][column];
for (int inner = factor_panel; inner < column; ++inner) {
const float value = tile[column][inner];
diagonal_value = fmaf(-value, value, diagonal_value);
}
diagonal_value = sqrt_approx(diagonal_value);
if (active) {
for (int factor_row = column + 1 + row;
factor_row < tile_size;
factor_row += 64) {
float value = tile[factor_row][column];
for (int inner = factor_panel; inner < column; ++inner) {
value = fmaf(
-tile[factor_row][inner],
tile[column][inner],
value
);
}
tile[factor_row][column] =
value * rcp_approx(diagonal_value);
}
if (row == 0) {
tile[column][column] = diagonal_value;
}
}
__syncthreads();
}
int tile_index = 0;
for (int tile_row = factor_panel + 16;
tile_row < tile_size;
tile_row += 16) {
for (int tile_column = factor_panel + 16;
tile_column <= tile_row;
tile_column += 16, ++tile_index) {
if (active && (tile_index & 1) == group_warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator,
16, 16, 8,
float
> accumulated;
nvcuda::wmma::load_matrix_sync(
accumulated,
&tile[tile_row][tile_column],
tile_size + 4,
nvcuda::wmma::mem_row_major
);
#pragma unroll
for (int panel_step = 0; panel_step < 16;
panel_step += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> right;
nvcuda::wmma::load_matrix_sync(
left,
&tile[tile_row][factor_panel + panel_step],
tile_size + 4
);
nvcuda::wmma::load_matrix_sync(
right,
&tile[tile_column][factor_panel + panel_step],
tile_size + 4
);
#pragma unroll
for (int element = 0;
element < left.num_elements;
++element) {
left.x[element] = -left.x[element];
}
nvcuda::wmma::mma_sync(
accumulated, left, right, accumulated
);
}
nvcuda::wmma::store_matrix_sync(
&tile[tile_row][tile_column],
accumulated,
tile_size + 4,
nvcuda::wmma::mem_row_major
);
}
}
}
__syncthreads();
}
if (active) {
diagonal[row] = tile[row][row];
tile[row][tile_size] = rcp_approx(diagonal[row]);
}
__syncthreads();
if (active && row == 0) {
float minimum_diagonal = diagonal[0];
float maximum_diagonal = diagonal[0];
for (int index = 1; index < tile_size; ++index) {
minimum_diagonal = fminf(
minimum_diagonal, diagonal[index]
);
maximum_diagonal = fmaxf(
maximum_diagonal, diagonal[index]
);
}
exact_inverse_fallback =
minimum_diagonal < 0.5f * maximum_diagonal;
}
__syncthreads();
if (active) {
for (int vector_index = row;
vector_index < tile_size * tile_size / 4;
vector_index += tile_size) {
const int store_row = vector_index >> 4;
const int store_column = (vector_index & 15) * 4;
const float4 factor_values = make_float4(
store_column <= store_row
? (store_column == store_row
? diagonal[store_row]
: tile[store_row][store_column]) : 0.0f,
store_column + 1 <= store_row
? (store_column + 1 == store_row
? diagonal[store_row]
: tile[store_row][store_column + 1]) : 0.0f,
store_column + 2 <= store_row
? (store_column + 2 == store_row
? diagonal[store_row]
: tile[store_row][store_column + 2]) : 0.0f,
store_column + 3 <= store_row
? (store_column + 3 == store_row
? diagonal[store_row]
: tile[store_row][store_column + 3]) : 0.0f
);
*reinterpret_cast<float4*>(
matrix + (offset + store_row) * n + offset + store_column
) = factor_values;
}
}
const int inverse_column = row;
const int inverse_end = inverse_column < 32 ? 32 : tile_size;
if (active) {
for (int inverse_row = inverse_column;
inverse_row < inverse_end;
++inverse_row) {
float value = inverse_column == inverse_row ? 1.0f : 0.0f;
for (int inner = inverse_column; inner < inverse_row; ++inner) {
value = fmaf(
-tile[inverse_row][inner],
tile[inverse_column][inner],
value
);
}
tile[inverse_column][inverse_row] =
value * tile[inverse_row][tile_size];
}
}
__syncthreads();
if (active) {
for (int index = row; index < 2 * 32 * 32; index += 64) {
const int half = index >> 10;
const int local = index & 1023;
const int inverse_row = local >> 5;
const int inverse_column_to_zero = local & 31;
if (inverse_column_to_zero < inverse_row) {
tile[half * 32 + inverse_row]
[half * 32 + inverse_column_to_zero] = 0.0f;
}
}
}
__syncthreads();
int inverse_tile_index = 0;
for (int tile_row = 0; tile_row < 32; tile_row += 16) {
for (int tile_column = 0; tile_column < 32;
tile_column += 16, ++inverse_tile_index) {
if (active && (inverse_tile_index & 1) == group_warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator, 16, 16, 8, float
> product;
nvcuda::wmma::fill_fragment(product, 0.0f);
#pragma unroll
for (int inner = 0; inner < 32; inner += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> inverse_diagonal;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> lower_left;
nvcuda::wmma::load_matrix_sync(
inverse_diagonal,
&tile[32 + inner][32 + tile_row],
tile_size + 4
);
nvcuda::wmma::load_matrix_sync(
lower_left,
&tile[32 + inner][tile_column],
tile_size + 4
);
nvcuda::wmma::mma_sync(
product,
inverse_diagonal,
lower_left,
product
);
}
nvcuda::wmma::store_matrix_sync(
&inverse_product[tile_row][tile_column],
product,
36,
nvcuda::wmma::mem_row_major
);
}
}
}
__syncthreads();
inverse_tile_index = 0;
for (int tile_row = 0; tile_row < 32; tile_row += 16) {
for (int tile_column = 0; tile_column < 32;
tile_column += 16, ++inverse_tile_index) {
if (active && (inverse_tile_index & 1) == group_warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator, 16, 16, 8, float
> inverse_off_diagonal;
nvcuda::wmma::fill_fragment(inverse_off_diagonal, 0.0f);
#pragma unroll
for (int inner = 0; inner < 32; inner += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> product;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> inverse_diagonal;
nvcuda::wmma::load_matrix_sync(
product,
&inverse_product[tile_row][inner],
36
);
nvcuda::wmma::load_matrix_sync(
inverse_diagonal,
&tile[tile_column][inner],
tile_size + 4
);
#pragma unroll
for (int element = 0;
element < product.num_elements;
++element) {
product.x[element] = -product.x[element];
}
nvcuda::wmma::mma_sync(
inverse_off_diagonal,
product,
inverse_diagonal,
inverse_off_diagonal
);
}
nvcuda::wmma::store_matrix_sync(
&tile[tile_column][32 + tile_row],
inverse_off_diagonal,
tile_size + 4,
nvcuda::wmma::mem_col_major
);
}
}
}
__syncthreads();
if (active && exact_inverse_fallback) {
for (int vector_index = row;
vector_index < tile_size * tile_size / 4;
vector_index += tile_size) {
const int load_row = vector_index >> 4;
const int load_column = (vector_index & 15) * 4;
const float4 factor_values = *reinterpret_cast<const float4*>(
matrix + (offset + load_row) * n + offset + load_column
);
tile[load_row][load_column] = factor_values.x;
tile[load_row][load_column + 1] = factor_values.y;
tile[load_row][load_column + 2] = factor_values.z;
tile[load_row][load_column + 3] = factor_values.w;
}
}
__syncthreads();
if (active && exact_inverse_fallback) {
for (int inverse_row = inverse_column;
inverse_row < tile_size;
++inverse_row) {
float value = inverse_column == inverse_row ? 1.0f : 0.0f;
for (int inner = inverse_column; inner < inverse_row; ++inner) {
value = fmaf(
-tile[inverse_row][inner],
tile[inverse_column][inner],
value
);
}
tile[inverse_column][inverse_row] =
value * tile[inverse_row][tile_size];
}
}
__syncthreads();
if (active) {
for (int vector_index = row;
vector_index < tile_size * tile_size / 4;
vector_index += tile_size) {
const int store_row = vector_index >> 4;
const int store_column = (vector_index & 15) * 4;
const float4 inverse_values = make_float4(
store_column <= store_row
? tile[store_column][store_row] : 0.0f,
store_column + 1 <= store_row
? tile[store_column + 1][store_row] : 0.0f,
store_column + 2 <= store_row
? tile[store_column + 2][store_row] : 0.0f,
store_column + 3 <= store_row
? tile[store_column + 3][store_row] : 0.0f
);
*reinterpret_cast<float4*>(
inverse_matrices +
static_cast<long long>(matrix_index) * tile_size * tile_size +
store_row * tile_size + store_column
) = inverse_values;
}
}
}
// C563_BLOCK_JACOBI_TAIL_BEGIN
__global__ __launch_bounds__(256) void prepare_tail_block_diagonal(
const float* __restrict__ input,
float* __restrict__ output,
int n,
int tail_start,
int tail_size
) {
const long long vector_index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long vectors =
static_cast<long long>(tail_size) * tail_size / 4;
if (vector_index >= vectors) {
return;
}
const int local_row = vector_index / (tail_size / 4);
const int local_column = (vector_index % (tail_size / 4)) * 4;
const bool same_diagonal_block =
(local_row >> 7) == (local_column >> 7);
const long long index =
static_cast<long long>(tail_start + local_row) * n +
tail_start + local_column;
const float4 values = same_diagonal_block
? *reinterpret_cast<const float4*>(input + index)
: make_float4(0.0f, 0.0f, 0.0f, 0.0f);
*reinterpret_cast<float4*>(output + index) = values;
}
// C563_BLOCK_JACOBI_TAIL_END
// C598_TAIL_KERNELS_BEGIN
__global__ __launch_bounds__(256) void prepare_c598_tail_full(
const float* __restrict__ input,
float* __restrict__ output,
int n,
int tail_start,
int tail_size
) {
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long elements =
static_cast<long long>(tail_size) * tail_size;
if (index >= elements) {
return;
}
const int row = index / tail_size;
const int column = index - static_cast<long long>(row) * tail_size;
const long long matrix_index =
static_cast<long long>(tail_start + row) * n + tail_start + column;
output[matrix_index] = input[matrix_index];
}
__global__ __launch_bounds__(256) void restore_c598_original_diagonal(
const float* __restrict__ input,
float* __restrict__ output,
int n,
int tail_start,
int tail_blocks
) {
constexpr int block = 128;
const int index = blockIdx.x * blockDim.x + threadIdx.x;
const int elements = tail_blocks * block * block;
if (index >= elements) {
return;
}
const int block_index = index / (block * block);
const int local = index - block_index * block * block;
const int row = local / block;
const int column = local - row * block;
const int offset = tail_start + block_index * block;
const long long matrix_index =
static_cast<long long>(offset + row) * n + offset + column;
output[matrix_index] = input[matrix_index];
}
__global__ __launch_bounds__(256) void pack_c598_tail_factor_half(
const float* __restrict__ factor,
__half* __restrict__ staged_factor,
float* __restrict__ diagonal_storage,
int n,
int tail_start,
int tail_size
) {
constexpr int block = 128;
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long elements =
static_cast<long long>(tail_size) * tail_size;
if (index >= elements) {
return;
}
const int row = index / tail_size;
const int column = index - static_cast<long long>(row) * tail_size;
const long long matrix_index =
static_cast<long long>(tail_start + row) * n + tail_start + column;
const float value = column <= row ? factor[matrix_index] : 0.0f;
staged_factor[matrix_index] = __float2half_rn(value);
const int block_row = row / block;
const int block_column = column / block;
if (block_row == block_column) {
const int local_row = row - block_row * block;
const int local_column = column - block_column * block;
diagonal_storage[
static_cast<long long>(block_row) * block * block +
local_row * block + local_column
] = local_column <= local_row ? value : 0.0f;
}
}
__global__ __launch_bounds__(256) void mirror_c598_tail_schur_lower(
float* matrices,
int n,
int tail_start,
int tail_size
) {
constexpr int block = 128;
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long elements =
static_cast<long long>(tail_size) * tail_size;
if (index >= elements) {
return;
}
const int row = index / tail_size;
const int column = index - static_cast<long long>(row) * tail_size;
if (row / block > column / block) {
matrices[
static_cast<long long>(tail_start + row) * n +
tail_start + column
] = matrices[
static_cast<long long>(tail_start + column) * n +
tail_start + row
];
}
}
__global__ __launch_bounds__(256) void restore_c598_factor_diagonal(
const float* __restrict__ diagonal_storage,
float* __restrict__ factor,
int n,
int tail_start,
int tail_blocks
) {
constexpr int block = 128;
const int index = blockIdx.x * blockDim.x + threadIdx.x;
const int elements = tail_blocks * block * block;
if (index >= elements) {
return;
}
const int block_index = index / (block * block);
const int local = index - block_index * block * block;
const int row = local / block;
const int column = local - row * block;
const int offset = tail_start + block_index * block;
factor[static_cast<long long>(offset + row) * n + offset + column] =
diagonal_storage[index];
}
__global__ __launch_bounds__(256) void add_c598_staged_first_factor(
const __half* __restrict__ staged_factor,
float* __restrict__ correction,
int n,
int tail_start,
int tail_size
) {
constexpr int block = 128;
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long elements =
static_cast<long long>(tail_size) * tail_size;
if (index >= elements) {
return;
}
const int row = index / tail_size;
const int column = index - static_cast<long long>(row) * tail_size;
if (row / block > column / block) {
const long long matrix_index =
static_cast<long long>(tail_start + row) * n +
tail_start + column;
correction[matrix_index] += __half2float(staged_factor[matrix_index]);
}
}
// C598_TAIL_KERNELS_END
// C565_BLOCK_BANDED_TAIL_BEGIN
__global__ __launch_bounds__(256) void prepare_tail_block_band(
const float* __restrict__ input,
float* __restrict__ output,
int n,
int tail_start,
int tail_size
) {
const long long vector_index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long vectors =
static_cast<long long>(tail_size) * tail_size / 4;
if (vector_index >= vectors) {
return;
}
const int local_row = vector_index / (tail_size / 4);
const int local_column = (vector_index % (tail_size / 4)) * 4;
const int block_row = local_row >> 7;
const int block_column = local_column >> 7;
const bool in_lower_block_band =
block_row == block_column || block_row == block_column + 1;
const long long index =
static_cast<long long>(tail_start + local_row) * n +
tail_start + local_column;
const float4 values = in_lower_block_band
? *reinterpret_cast<const float4*>(input + index)
: make_float4(0.0f, 0.0f, 0.0f, 0.0f);
*reinterpret_cast<float4*>(output + index) = values;
}
// C565_BLOCK_BANDED_TAIL_END
__global__ __launch_bounds__(256) void cholesky128_diagonal_kernel(
float* __restrict__ matrices,
const float* __restrict__ input_matrices,
float* __restrict__ inverse_matrices,
int n,
int offset
) {
constexpr int tile_size = 128;
constexpr int stride = 132;
extern __shared__ float factor[];
__shared__ bool exact_inverse128_fallback;
const int thread = threadIdx.x;
const int warp = thread >> 5;
// C550_FUSED_INVERSE_GUARD_BEGIN
const bool c550_fused_final_inverse =
n == 512 && gridDim.x == 16 && offset == 256;
const bool fused_inverse128 =
inverse_matrices != nullptr && (
(n == 1024 && gridDim.x == 60) ||
(n == 2048 && gridDim.x == 8) ||
c550_fused_final_inverse
);
// C550_FUSED_INVERSE_GUARD_END
// C563_BLOCK_JACOBI_TAIL_BEGIN
const bool c563_parallel_tail_grid =
n == 32768 && gridDim.x == 32 && offset == 28672 &&
input_matrices == nullptr && inverse_matrices == nullptr;
// C565_BLOCK_BANDED_TAIL_BEGIN
const bool c565_parallel_tail_grid =
input_matrices == nullptr && inverse_matrices == nullptr &&
n == 16384 && gridDim.x == 10 && offset == 15104;
// C598_PARALLEL_TAIL_GUARD_BEGIN
const bool c598_parallel_tail_grid =
input_matrices == nullptr && inverse_matrices == nullptr &&
n == 16384 && gridDim.x == 22 && offset == 13568;
// C614_PARALLEL_TAIL_GUARD_BEGIN
const bool c614_parallel_tail_grid =
input_matrices == nullptr && inverse_matrices == nullptr &&
n == 32768 && gridDim.x == 80 && offset == 22528;
// C614_PARALLEL_TAIL_GUARD_END
const bool parallel_tail_grid =
c563_parallel_tail_grid || c565_parallel_tail_grid ||
c598_parallel_tail_grid || c614_parallel_tail_grid;
// C598_PARALLEL_TAIL_GUARD_END
// C565_BLOCK_BANDED_TAIL_END
const int factor_offset = parallel_tail_grid
? offset + static_cast<int>(blockIdx.x) * tile_size
: offset;
float* matrix = matrices + (
parallel_tail_grid
? 0
: static_cast<long long>(blockIdx.x) * n * n
);
// C563_BLOCK_JACOBI_TAIL_END
const float* factor_source =
input_matrices != nullptr && factor_offset == 0
? input_matrices + static_cast<long long>(blockIdx.x) * n * n
: matrix;
float* inverse_matrix = inverse_matrices == nullptr
? nullptr
: inverse_matrices +
static_cast<long long>(blockIdx.x) * tile_size * tile_size;
cudaGridDependencySynchronize();
for (int vector_index = thread;
vector_index < tile_size * tile_size / 4;
vector_index += blockDim.x) {
const int row = vector_index >> 5;
const int column = (vector_index & 31) * 4;
const float4 values = *reinterpret_cast<const float4*>(
factor_source +
(factor_offset + row) * n + factor_offset + column
);
factor[row * stride + column] = values.x;
factor[row * stride + column + 1] = values.y;
factor[row * stride + column + 2] = values.z;
factor[row * stride + column + 3] = values.w;
}
__syncthreads();
for (int factor_panel = 0; factor_panel < tile_size;
factor_panel += 16) {
for (int panel_column = 0; panel_column < 16; ++panel_column) {
const int column = factor_panel + panel_column;
float diagonal = factor[column * stride + column];
for (int inner = factor_panel; inner < column; ++inner) {
const float value = factor[column * stride + inner];
diagonal = fmaf(-value, value, diagonal);
}
diagonal = sqrt_approx(diagonal);
for (int row = column + 1 + thread;
row < tile_size;
row += blockDim.x) {
float value = factor[row * stride + column];
for (int inner = factor_panel; inner < column; ++inner) {
value = fmaf(
-factor[row * stride + inner],
factor[column * stride + inner],
value
);
}
factor[row * stride + column] =
value * rcp_approx(diagonal);
}
if (thread == 0) {
factor[column * stride + column] = diagonal;
}
__syncthreads();
}
int tile_index = 0;
for (int tile_row = factor_panel + 16;
tile_row < tile_size;
tile_row += 16) {
for (int tile_column = factor_panel + 16;
tile_column <= tile_row;
tile_column += 16, ++tile_index) {
if ((tile_index & 7) == warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator,
16, 16, 8,
float
> accumulated;
nvcuda::wmma::load_matrix_sync(
accumulated,
factor + tile_row * stride + tile_column,
stride,
nvcuda::wmma::mem_row_major
);
#pragma unroll
for (int panel_step = 0; panel_step < 16;
panel_step += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> right;
nvcuda::wmma::load_matrix_sync(
left,
factor + tile_row * stride +
factor_panel + panel_step,
stride
);
nvcuda::wmma::load_matrix_sync(
right,
factor + tile_column * stride +
factor_panel + panel_step,
stride
);
#pragma unroll
for (int element = 0;
element < left.num_elements;
++element) {
left.x[element] = -left.x[element];
}
nvcuda::wmma::mma_sync(
accumulated, left, right, accumulated
);
}
nvcuda::wmma::store_matrix_sync(
factor + tile_row * stride + tile_column,
accumulated,
stride,
nvcuda::wmma::mem_row_major
);
}
}
}
__syncthreads();
}
for (int vector_index = thread;
vector_index < tile_size * tile_size / 4;
vector_index += blockDim.x) {
const int row = vector_index >> 5;
const int column = (vector_index & 31) * 4;
if (column + 3 <= row) {
const float4 output_values = make_float4(
factor[row * stride + column],
factor[row * stride + column + 1],
factor[row * stride + column + 2],
factor[row * stride + column + 3]
);
*reinterpret_cast<float4*>(
matrix +
(factor_offset + row) * n + factor_offset + column
) = output_values;
} else {
#pragma unroll
for (int lane = 0; lane < 4; ++lane) {
if (column + lane <= row) {
matrix[(factor_offset + row) * n +
factor_offset + column + lane] =
factor[row * stride + column + lane];
}
}
}
if (inverse_matrices != nullptr && !fused_inverse128) {
const float4 identity_values = make_float4(
row == column ? 1.0f : 0.0f,
row == column + 1 ? 1.0f : 0.0f,
row == column + 2 ? 1.0f : 0.0f,
row == column + 3 ? 1.0f : 0.0f
);
*reinterpret_cast<float4*>(
inverse_matrix + row * tile_size + column
) = identity_values;
}
}
__syncthreads();
if (!fused_inverse128) {
return;
}
if (thread < tile_size) {
factor[thread * stride + tile_size] =
rcp_approx(factor[thread * stride + thread]);
}
if (thread == 0) {
float minimum_diagonal = factor[0];
float maximum_diagonal = factor[0];
for (int index = 1; index < tile_size; ++index) {
const float value = factor[index * stride + index];
minimum_diagonal = fminf(minimum_diagonal, value);
maximum_diagonal = fmaxf(maximum_diagonal, value);
}
// C550D_GUARD055_BEGIN
const float guard_ratio = c550_fused_final_inverse
? 0.55f
: (n == 2048 ? 0.75f : 0.5f);
// C550D_GUARD055_END
exact_inverse128_fallback =
minimum_diagonal < guard_ratio * maximum_diagonal;
}
__syncthreads();
const int inverse_column = thread;
if (inverse_column < tile_size) {
const int inverse_end =
inverse_column < 64 ? 64 : tile_size;
for (int inverse_row = inverse_column;
inverse_row < inverse_end;
++inverse_row) {
float value = inverse_column == inverse_row ? 1.0f : 0.0f;
for (int inner = inverse_column; inner < inverse_row; ++inner) {
value = fmaf(
-factor[inverse_row * stride + inner],
factor[inverse_column * stride + inner],
value
);
}
factor[inverse_column * stride + inverse_row] =
value * factor[inverse_row * stride + tile_size];
}
}
__syncthreads();
for (int index = thread; index < 2 * 64 * 64;
index += blockDim.x) {
const int half = index >> 12;
const int local = index & 4095;
const int inverse_row = local >> 6;
const int inverse_column_to_zero = local & 63;
if (inverse_column_to_zero < inverse_row) {
factor[(half * 64 + inverse_row) * stride +
half * 64 + inverse_column_to_zero] = 0.0f;
}
}
__syncthreads();
int inverse_tile_index = 0;
for (int tile_row = 0; tile_row < 64; tile_row += 16) {
for (int tile_column = 0; tile_column < 64;
tile_column += 16, ++inverse_tile_index) {
if ((inverse_tile_index & 7) == warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator, 16, 16, 8, float
> inverse128_product;
nvcuda::wmma::fill_fragment(inverse128_product, 0.0f);
#pragma unroll
for (int inner = 0; inner < 64; inner += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> inverse_diagonal;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> lower_left;
nvcuda::wmma::load_matrix_sync(
inverse_diagonal,
factor + (64 + inner) * stride + 64 + tile_row,
stride
);
nvcuda::wmma::load_matrix_sync(
lower_left,
factor + (64 + inner) * stride + tile_column,
stride
);
nvcuda::wmma::mma_sync(
inverse128_product,
inverse_diagonal,
lower_left,
inverse128_product
);
}
nvcuda::wmma::store_matrix_sync(
inverse_matrix + tile_row * tile_size + tile_column,
inverse128_product,
tile_size,
nvcuda::wmma::mem_row_major
);
}
}
}
__syncthreads();
inverse_tile_index = 0;
for (int tile_row = 0; tile_row < 64; tile_row += 16) {
for (int tile_column = 0; tile_column < 64;
tile_column += 16, ++inverse_tile_index) {
if ((inverse_tile_index & 7) == warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator, 16, 16, 8, float
> inverse_off_diagonal;
nvcuda::wmma::fill_fragment(inverse_off_diagonal, 0.0f);
#pragma unroll
for (int inner = 0; inner < 64; inner += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> product;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> inverse_diagonal;
nvcuda::wmma::load_matrix_sync(
product,
inverse_matrix + tile_row * tile_size + inner,
tile_size
);
nvcuda::wmma::load_matrix_sync(
inverse_diagonal,
factor + tile_column * stride + inner,
stride
);
#pragma unroll
for (int element = 0;
element < product.num_elements;
++element) {
product.x[element] = -product.x[element];
}
nvcuda::wmma::mma_sync(
inverse_off_diagonal,
product,
inverse_diagonal,
inverse_off_diagonal
);
}
nvcuda::wmma::store_matrix_sync(
factor + tile_column * stride + 64 + tile_row,
inverse_off_diagonal,
stride,
nvcuda::wmma::mem_col_major
);
}
}
}
__syncthreads();
if (exact_inverse128_fallback) {
for (int vector_index = thread;
vector_index < tile_size * tile_size / 4;
vector_index += blockDim.x) {
const int row = vector_index >> 5;
const int column = (vector_index & 31) * 4;
const float4 values = *reinterpret_cast<const float4*>(
matrix +
(factor_offset + row) * n + factor_offset + column
);
factor[row * stride + column] = values.x;
factor[row * stride + column + 1] = values.y;
factor[row * stride + column + 2] = values.z;
factor[row * stride + column + 3] = values.w;
}
}
__syncthreads();
if (exact_inverse128_fallback && inverse_column < tile_size) {
for (int inverse_row = inverse_column;
inverse_row < tile_size;
++inverse_row) {
float value = inverse_column == inverse_row ? 1.0f : 0.0f;
for (int inner = inverse_column; inner < inverse_row; ++inner) {
value = fmaf(
-factor[inverse_row * stride + inner],
factor[inverse_column * stride + inner],
value
);
}
factor[inverse_column * stride + inverse_row] =
value * factor[inverse_row * stride + tile_size];
}
}
__syncthreads();
for (int vector_index = thread;
vector_index < tile_size * tile_size / 4;
vector_index += blockDim.x) {
const int row = vector_index >> 5;
const int column = (vector_index & 31) * 4;
const float4 inverse_values = make_float4(
column <= row ? factor[column * stride + row] : 0.0f,
column + 1 <= row
? factor[(column + 1) * stride + row] : 0.0f,
column + 2 <= row
? factor[(column + 2) * stride + row] : 0.0f,
column + 3 <= row
? factor[(column + 3) * stride + row] : 0.0f
);
*reinterpret_cast<float4*>(
inverse_matrix + row * tile_size + column
) = inverse_values;
}
}
__global__ __launch_bounds__(64) void invert128_diagonal_halves_kernel(
const float* __restrict__ matrices,
float* __restrict__ inverse_matrices,
int n,
int offset
) {
constexpr int half_size = 64;
constexpr int stride = 68;
__shared__ float factor[half_size][stride];
const int column = threadIdx.x;
const int matrix_index = blockIdx.x;
const int half_offset = blockIdx.y * half_size;
const float* matrix = matrices +
static_cast<long long>(matrix_index) * n * n;
float* inverse = inverse_matrices +
static_cast<long long>(matrix_index) * 128 * 128;
for (int row = column; row < half_size; ++row) {
factor[row][column] = matrix[
(offset + half_offset + row) * n +
offset + half_offset + column
];
}
factor[column][half_size] = rcp_approx(matrix[
(offset + half_offset + column) * n +
offset + half_offset + column
]);
__syncthreads();
for (int inverse_row = column;
inverse_row < half_size;
++inverse_row) {
float value = column == inverse_row ? 1.0f : 0.0f;
for (int inner = column; inner < inverse_row; ++inner) {
value = fmaf(
-factor[inverse_row][inner],
factor[column][inner],
value
);
}
factor[column][inverse_row] =
value * factor[inverse_row][half_size];
}
for (int row = 0; row < half_size; ++row) {
inverse[(half_offset + row) * 128 + half_offset + column] =
column <= row ? factor[column][row] : 0.0f;
}
}
__global__ __launch_bounds__(32) void inverse128_offdiagonal_stage1_kernel(
const float* __restrict__ matrices,
float* __restrict__ inverse_matrices,
int n,
int offset
) {
const int tile_row = (blockIdx.x >> 2) * 16;
const int tile_column = (blockIdx.x & 3) * 16;
const int matrix_index = blockIdx.y;
const float* matrix = matrices +
static_cast<long long>(matrix_index) * n * n;
float* inverse = inverse_matrices +
static_cast<long long>(matrix_index) * 128 * 128;
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator, 16, 16, 8, float
> product;
nvcuda::wmma::fill_fragment(product, 0.0f);
#pragma unroll
for (int inner = 0; inner < 64; inner += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> inverse_lower;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> lower_left;
nvcuda::wmma::load_matrix_sync(
inverse_lower,
inverse + (64 + tile_row) * 128 + 64 + inner,
128
);
nvcuda::wmma::load_matrix_sync(
lower_left,
matrix + (offset + 64 + inner) * n +
offset + tile_column,
n
);
nvcuda::wmma::mma_sync(
product, inverse_lower, lower_left, product
);
}
nvcuda::wmma::store_matrix_sync(
inverse + tile_row * 128 + 64 + tile_column,
product,
128,
nvcuda::wmma::mem_row_major
);
}
__global__ __launch_bounds__(32) void inverse128_offdiagonal_stage2_kernel(
float* __restrict__ inverse_matrices
) {
const int tile_row = (blockIdx.x >> 2) * 16;
const int tile_column = (blockIdx.x & 3) * 16;
float* inverse = inverse_matrices +
static_cast<long long>(blockIdx.y) * 128 * 128;
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator, 16, 16, 8, float
> product;
nvcuda::wmma::fill_fragment(product, 0.0f);
#pragma unroll
for (int inner = 0; inner < 64; inner += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> right;
nvcuda::wmma::load_matrix_sync(
left,
inverse + tile_row * 128 + 64 + inner,
128
);
nvcuda::wmma::load_matrix_sync(
right,
inverse + inner * 128 + tile_column,
128
);
#pragma unroll
for (int element = 0; element < left.num_elements; ++element) {
left.x[element] = -left.x[element];
}
nvcuda::wmma::mma_sync(product, left, right, product);
}
nvcuda::wmma::store_matrix_sync(
inverse + (64 + tile_row) * 128 + tile_column,
product,
128,
nvcuda::wmma::mem_row_major
);
}
__global__ __launch_bounds__(128) void inverse128_exact_fallback_kernel(
const float* __restrict__ matrices,
float* __restrict__ inverse_matrices,
int n,
int offset
) {
constexpr int tile_size = 128;
constexpr int stride = 129;
extern __shared__ float factor[];
__shared__ bool use_exact;
const int column = threadIdx.x;
const int matrix_index = blockIdx.x;
const float* matrix = matrices +
static_cast<long long>(matrix_index) * n * n;
float* inverse = inverse_matrices +
static_cast<long long>(matrix_index) * tile_size * tile_size;
if (column == 0) {
float minimum_diagonal = matrix[offset * n + offset];
float maximum_diagonal = minimum_diagonal;
for (int index = 1; index < tile_size; ++index) {
const float value = matrix[
(offset + index) * n + offset + index
];
minimum_diagonal = fminf(minimum_diagonal, value);
maximum_diagonal = fmaxf(maximum_diagonal, value);
}
const float guard_ratio = n == 2048 ? 0.75f : 0.5f;
use_exact = minimum_diagonal < guard_ratio * maximum_diagonal;
}
__syncthreads();
if (!use_exact) {
for (int index = column; index < 64 * 64;
index += blockDim.x) {
const int row = index >> 6;
const int local_column = index & 63;
inverse[row * tile_size + 64 + local_column] = 0.0f;
}
return;
}
for (int row = column; row < tile_size; ++row) {
factor[row * stride + column] =
matrix[(offset + row) * n + offset + column];
}
factor[column * stride + tile_size] = rcp_approx(
matrix[(offset + column) * n + offset + column]
);
__syncthreads();
for (int inverse_row = column;
inverse_row < tile_size;
++inverse_row) {
float value = column == inverse_row ? 1.0f : 0.0f;
for (int inner = column; inner < inverse_row; ++inner) {
value = fmaf(
-factor[inverse_row * stride + inner],
factor[column * stride + inner],
value
);
}
factor[column * stride + inverse_row] =
value * factor[inverse_row * stride + tile_size];
}
for (int index = column; index < tile_size * tile_size;
index += blockDim.x) {
const int row = index >> 7;
const int local_column = index & 127;
inverse[index] = local_column <= row
? factor[local_column * stride + row]
: 0.0f;
}
}
__global__ __launch_bounds__(256) void cholesky128_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr int n = 128;
constexpr int stride = n + 1;
extern __shared__ float factor[];
const int thread = threadIdx.x;
const int tile_column = thread & 15;
const int tile_row = thread >> 4;
const long long batch_offset = static_cast<long long>(blockIdx.x) * n * n;
const float* matrix = input + batch_offset;
float* result = output + batch_offset;
for (int index = threadIdx.x; index < n * n; index += blockDim.x) {
const int load_row = index >> 7;
const int load_column = index & 127;
if (load_column <= load_row) {
factor[load_row * stride + load_column] = matrix[index];
}
}
__syncthreads();
for (int panel = 0; panel < n; panel += 8) {
for (int panel_column = 0; panel_column < 8; ++panel_column) {
const int column = panel + panel_column;
float diagonal = factor[column * stride + column];
for (int inner = panel; inner < column; ++inner) {
const float value = factor[column * stride + inner];
diagonal = fmaf(-value, value, diagonal);
}
diagonal = sqrt_approx(diagonal);
for (int row = column + 1 + thread;
row < n;
row += blockDim.x) {
float value = factor[row * stride + column];
for (int inner = panel; inner < column; ++inner) {
value = fmaf(
-factor[row * stride + inner],
factor[column * stride + inner],
value
);
}
factor[row * stride + column] =
value * rcp_approx(diagonal);
}
if (thread == 0) {
factor[column * stride + column] = diagonal;
}
__syncthreads();
}
for (int trailing_row = panel + 8 + tile_row;
trailing_row < n;
trailing_row += 16) {
for (int trailing_column = panel + 8 + tile_column;
trailing_column <= trailing_row;
trailing_column += 16) {
float value = factor[
trailing_row * stride + trailing_column
];
#pragma unroll
for (int panel_column = 0; panel_column < 8;
++panel_column) {
value = fmaf(
-factor[
trailing_row * stride + panel + panel_column
],
factor[
trailing_column * stride + panel + panel_column
],
value
);
}
factor[trailing_row * stride + trailing_column] = value;
}
}
__syncthreads();
}
for (int index = threadIdx.x; index < n * n; index += blockDim.x) {
const int store_row = index >> 7;
const int store_column = index & 127;
result[index] = store_column <= store_row
? factor[store_row * stride + store_column]
: 0.0f;
}
}
__global__ __launch_bounds__(256) void cholesky128_wmma_kernel(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr int n = 128;
constexpr int stride = 132;
extern __shared__ float factor[];
const int thread = threadIdx.x;
const int warp = thread >> 5;
const long long batch_offset = static_cast<long long>(blockIdx.x) * n * n;
const float* matrix = input + batch_offset;
float* result = output + batch_offset;
for (int index = thread; index < n * n; index += blockDim.x) {
const int row = index >> 7;
const int column = index & 127;
if (column <= row) {
factor[row * stride + column] = matrix[index];
}
}
__syncthreads();
for (int panel = 0; panel < n; panel += 16) {
if (thread < n) {
for (int panel_column = 0; panel_column < 16; ++panel_column) {
const int column = panel + panel_column;
float diagonal = factor[column * stride + column];
for (int inner = panel; inner < column; ++inner) {
const float value = factor[column * stride + inner];
diagonal = fmaf(-value, value, diagonal);
}
diagonal = sqrt_approx(diagonal);
for (int row = column + 1 + thread;
row < n;
row += n) {
float value = factor[row * stride + column];
for (int inner = panel; inner < column; ++inner) {
value = fmaf(
-factor[row * stride + inner],
factor[column * stride + inner],
value
);
}
factor[row * stride + column] =
value * rcp_approx(diagonal);
}
if (thread == 0) {
factor[column * stride + column] = diagonal;
}
panel_barrier128();
}
}
__syncthreads();
int tile_index = 0;
for (int tile_row = panel + 16; tile_row < n; tile_row += 16) {
for (int tile_column = panel + 16;
tile_column <= tile_row;
tile_column += 16, ++tile_index) {
if ((tile_index & 7) == warp) {
nvcuda::wmma::fragment<
nvcuda::wmma::accumulator,
16, 16, 8,
float
> accumulated;
nvcuda::wmma::load_matrix_sync(
accumulated,
factor + tile_row * stride + tile_column,
stride,
nvcuda::wmma::mem_row_major
);
#pragma unroll
for (int panel_step = 0; panel_step < 16;
panel_step += 8) {
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_a,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::row_major
> left_residual;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> right;
nvcuda::wmma::fragment<
nvcuda::wmma::matrix_b,
16, 16, 8,
nvcuda::wmma::precision::tf32,
nvcuda::wmma::col_major
> right_residual;
nvcuda::wmma::load_matrix_sync(
left,
factor + tile_row * stride + panel + panel_step,
stride
);
nvcuda::wmma::load_matrix_sync(
right,
factor + tile_column * stride + panel + panel_step,
stride
);
#pragma unroll
for (int element = 0;
element < left.num_elements;
++element) {
const float full = left.x[element];
const float high = __uint_as_float(
__float_as_uint(full) & 0xFFFFE000u
);
left.x[element] = -high;
left_residual.x[element] = -(full - high);
}
#pragma unroll
for (int element = 0;
element < right.num_elements;
++element) {
const float full = right.x[element];
const float high = __uint_as_float(
__float_as_uint(full) & 0xFFFFE000u
);
right.x[element] = high;
right_residual.x[element] = full - high;
}
nvcuda::wmma::mma_sync(
accumulated, left, right, accumulated
);
nvcuda::wmma::mma_sync(
accumulated, left, right_residual, accumulated
);
nvcuda::wmma::mma_sync(
accumulated, left_residual, right, accumulated
);
}
nvcuda::wmma::store_matrix_sync(
factor + tile_row * stride + tile_column,
accumulated,
stride,
nvcuda::wmma::mem_row_major
);
}
}
}
__syncthreads();
}
for (int index = thread; index < n * n; index += blockDim.x) {
const int row = index >> 7;
const int column = index & 127;
result[index] = column <= row
? factor[row * stride + column]
: 0.0f;
}
}
void cholesky32(std::uint64_t input, std::uint64_t output, int batch) {
cholesky32_kernel<<<(batch + 3) / 4, 128>>>(
reinterpret_cast<const float*>(input),
reinterpret_cast<float*>(output),
batch
);
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
void cholesky64(std::uint64_t input, std::uint64_t output, int batch) {
// C560A_ROUTE_BEGIN
if (batch == 1024) {
cholesky64_coalesced_wavefront_kernel<<<batch, 256>>>(
reinterpret_cast<const float*>(input),
reinterpret_cast<float*>(output)
);
} else {
constexpr int shared_bytes = 64 * 68 * sizeof(float);
cholesky64_wmma_kernel<<<batch, 128, shared_bytes>>>(
reinterpret_cast<const float*>(input),
reinterpret_cast<float*>(output)
);
}
// C560A_ROUTE_END
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
void cholesky128(std::uint64_t input, std::uint64_t output, int batch) {
constexpr int shared_bytes = 128 * 132 * sizeof(float);
static const cudaError_t configured = cudaFuncSetAttribute(
cholesky128_wmma_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shared_bytes
);
if (configured != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(configured));
}
cholesky128_wmma_kernel<<<batch, 256, shared_bytes>>>(
reinterpret_cast<const float*>(input),
reinterpret_cast<float*>(output)
);
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
__global__ void fill_matrix_pointers(
float** pointers, float* matrices, int stride, int batch
) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index < batch) {
pointers[index] = matrices + static_cast<long long>(index) * stride;
}
}
__global__ void fill_all_block_pointers(
float** pointer_workspace,
float* matrices,
float* inverse_workspace,
float* solved_workspace,
int n,
int batch,
int block,
int block_count,
int solution_stride
) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index < block_count * batch) {
const int block_index = index / batch;
const int matrix_index = index - block_index * batch;
const int offset = block_index * block;
float** diagonal =
pointer_workspace + block_index * 4 * batch;
float** panel = diagonal + batch;
float** inverse = panel + batch;
float** solved = inverse + batch;
float* matrix = matrices +
static_cast<long long>(matrix_index) * n * n;
diagonal[matrix_index] = matrix + offset + offset * n;
panel[matrix_index] = matrix + offset + (offset + block) * n;
inverse[matrix_index] = inverse_workspace +
static_cast<long long>(matrix_index) * block * block;
solved[matrix_index] = solved_workspace +
static_cast<long long>(matrix_index) * solution_stride;
}
}
// C614_CLASSIFIER_KERNEL_BEGIN
__global__ __launch_bounds__(256) void classify_c614_first_block_copy(
const float* __restrict__ copied,
double* __restrict__ statistics,
int n
) {
__shared__ double diagonal_sum[256];
__shared__ double diagonal_square_sum[256];
__shared__ double off_diagonal_square_sum[256];
const int thread = threadIdx.x;
double local_off_diagonal = 0.0;
for (int row = thread; row < n; row += blockDim.x) {
if (row != 0) {
const float* row_pointer =
copied + static_cast<long long>(row) * n;
const double value = static_cast<double>(row_pointer[0]);
local_off_diagonal += value * value;
}
}
double local_diagonal = 0.0;
if (thread < 128) {
local_diagonal = static_cast<double>(
copied[static_cast<long long>(thread) * n + thread]
);
}
diagonal_sum[thread] = local_diagonal;
diagonal_square_sum[thread] = local_diagonal * local_diagonal;
off_diagonal_square_sum[thread] = local_off_diagonal;
__syncthreads();
for (int offset = 128; offset != 0; offset >>= 1) {
if (thread < offset) {
diagonal_sum[thread] += diagonal_sum[thread + offset];
diagonal_square_sum[thread] +=
diagonal_square_sum[thread + offset];
off_diagonal_square_sum[thread] +=
off_diagonal_square_sum[thread + offset];
}
__syncthreads();
}
if (thread == 0) {
statistics[0] = diagonal_sum[0];
statistics[1] = diagonal_square_sum[0];
statistics[2] = off_diagonal_square_sum[0];
}
}
__global__ void fill_c614_hierarchy_pointers(
const __half** left_prefix_pointers,
const __half** right_prefix_pointers,
float** output_pointers,
const __half* half_factor,
float* factor,
int n,
int tail_start,
int tile_blocks,
int tile_count,
bool cross_rectangle
) {
constexpr int block = 128;
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index >= tile_count) {
return;
}
int row_block;
int column_block;
if (cross_rectangle) {
row_block = 64;
column_block = 0;
} else {
const int first64_count = 64 / (2 * tile_blocks);
const int begin = index < first64_count
? index * 2 * tile_blocks
: 64 + (index - first64_count) * 2 * tile_blocks;
row_block = begin + tile_blocks;
column_block = begin;
}
const int row_offset = tail_start + row_block * block;
const int column_offset = tail_start + column_block * block;
// Column-major cuBLAS aliases transpose each row-major rectangle.
left_prefix_pointers[index] =
half_factor + static_cast<long long>(column_offset) * n;
right_prefix_pointers[index] =
half_factor + static_cast<long long>(row_offset) * n;
output_pointers[index] =
factor + static_cast<long long>(row_offset) * n + column_offset;
}
// C614_CLASSIFIER_KERNEL_END
template <int Block>
__global__ void copy_first_block_column_float4(
const float* input,
float* output,
int n
) {
constexpr int vectors_per_row = Block / 4;
const int local = blockIdx.x * blockDim.x + threadIdx.x;
const int vectors = n * vectors_per_row;
if (local < vectors) {
const int row = local / vectors_per_row;
const int column = (local % vectors_per_row) * 4;
const long long batch_offset =
static_cast<long long>(blockIdx.y) * n * n;
*reinterpret_cast<float4*>(
output + batch_offset + static_cast<long long>(row) * n + column
) = *reinterpret_cast<const float4*>(
input + batch_offset + static_cast<long long>(row) * n + column
);
}
}
// C565_BLOCK_BANDED_TAIL_BEGIN
__global__ void fill_tail_band_pointers(
float** diagonal_pointers,
float** panel_pointers,
float* matrices,
int n,
int tail_start,
int block,
int panel_count
) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index < panel_count) {
const int offset = tail_start + index * block;
diagonal_pointers[index] =
matrices + static_cast<long long>(offset) * n + offset;
panel_pointers[index] =
matrices + static_cast<long long>(offset + block) * n + offset;
}
}
// C565_BLOCK_BANDED_TAIL_END
// C598_TAIL_POINTERS_BEGIN
__global__ void fill_c598_tail_pair_pointers(
float** diagonal_pointers,
float** panel_pointers,
float* matrices,
int n,
int tail_start,
int block,
int pair_count
) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index >= pair_count) {
return;
}
int block_row = 1;
int block_column = index;
while (block_column >= block_row) {
block_column -= block_row;
++block_row;
}
const int diagonal_offset = tail_start + block_column * block;
const int panel_row = tail_start + block_row * block;
diagonal_pointers[index] =
matrices + static_cast<long long>(diagonal_offset) * n +
diagonal_offset;
panel_pointers[index] =
matrices + static_cast<long long>(panel_row) * n +
diagonal_offset;
}
// C598_TAIL_POINTERS_END
template <int Block>
__global__ void scatter_solved_panel_float4(
float* matrices,
const float* solved,
int n,
int offset,
int remaining,
int solution_stride
) {
constexpr int vectors_per_column = Block / 4;
constexpr int vector_shift = Block == 64 ? 4 : 5;
const int local =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const int matrix_index = blockIdx.y;
const int panel_vectors = vectors_per_column * remaining;
if (local < panel_vectors) {
const int row_vector = local & (vectors_per_column - 1);
const int column = local >> vector_shift;
float* matrix = matrices +
static_cast<long long>(matrix_index) * n * n;
*reinterpret_cast<float4*>(
matrix + offset + row_vector * 4 +
static_cast<long long>(offset + Block + column) * n
) = reinterpret_cast<const float4*>(
solved + static_cast<long long>(matrix_index) * solution_stride
)[local];
}
}
template <int Block>
void launch_scatter_solved_panel_float4(
float* matrices,
const float* solved,
int n,
int batch,
int offset,
int remaining,
int solution_stride
) {
const int panel_vectors = (Block / 4) * remaining;
dim3 grid((panel_vectors + 255) / 256, batch);
scatter_solved_panel_float4<Block><<<grid, 256>>>(
matrices,
solved,
n,
offset,
remaining,
solution_stride
);
}
__global__ void pack_unsolved_panel_float4(
const float* panel,
float* packed,
int n,
int block,
int remaining
) {
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const int vectors_per_column = block / 4;
const long long panel_vectors =
static_cast<long long>(vectors_per_column) * remaining;
if (index < panel_vectors) {
const int row_vector =
static_cast<int>(index % vectors_per_column);
const int column =
static_cast<int>(index / vectors_per_column);
reinterpret_cast<float4*>(packed)[index] =
*reinterpret_cast<const float4*>(
panel + static_cast<long long>(column) * n + row_vector * 4
);
}
}
// C550_PACK_KERNEL_BEGIN
__global__ void pack_batched_unsolved_panel_float4(
const float* matrices,
float* packed,
int n,
int offset,
int block,
int remaining,
int solution_stride
) {
const int local = blockIdx.x * blockDim.x + threadIdx.x;
const int matrix_index = blockIdx.y;
const int vectors_per_column = block / 4;
const int panel_vectors = vectors_per_column * remaining;
if (local < panel_vectors) {
const int row_vector = local % vectors_per_column;
const int panel_column = local / vectors_per_column;
const long long matrix_stride = static_cast<long long>(n) * n;
const float* matrix = matrices +
static_cast<long long>(matrix_index) * matrix_stride;
float* packed_matrix = packed +
static_cast<long long>(matrix_index) * solution_stride;
reinterpret_cast<float4*>(packed_matrix)[local] =
*reinterpret_cast<const float4*>(
matrix +
static_cast<long long>(offset + block + panel_column) * n +
offset + row_vector * 4
);
}
}
// C550_PACK_KERNEL_END
__global__ void pack_factor_panel_half(
const float* matrices,
__half* half_matrices,
int n,
int offset,
int block
) {
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const int vectors_per_row = block / 4;
const long long vectors =
static_cast<long long>(n - offset) * vectors_per_row;
if (index < vectors) {
const int matrix = blockIdx.y;
const int row =
offset + static_cast<int>(index / vectors_per_row);
const int column =
offset + static_cast<int>(index % vectors_per_row) * 4;
const long long matrix_index =
static_cast<long long>(matrix) * n * n +
static_cast<long long>(row) * n + column;
const float4 values = *reinterpret_cast<const float4*>(
matrices + matrix_index
);
*reinterpret_cast<__half2*>(half_matrices + matrix_index) =
__floats2half2_rn(values.x, values.y);
*reinterpret_cast<__half2*>(half_matrices + matrix_index + 2) =
__floats2half2_rn(values.z, values.w);
}
}
__global__ void fill_batched_identity(
float* matrices, int batch, int n
) {
const long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long elements = static_cast<long long>(batch) * n * n;
if (index < elements) {
const int local = static_cast<int>(index % (n * n));
matrices[index] = local / n == local % n ? 1.0f : 0.0f;
}
}
__device__ __forceinline__ void zero_upper_row_float4(
float* matrix, int n, int row
) {
const int aligned_column = (row + 4) & ~3;
for (int column = row + 1 + threadIdx.x;
column < aligned_column && column < n;
column += blockDim.x) {
matrix[static_cast<long long>(row) * n + column] = 0.0f;
}
const float4 zeros = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
for (int column = aligned_column + 4 * threadIdx.x;
column + 3 < n;
column += 4 * blockDim.x) {
*reinterpret_cast<float4*>(
matrix + static_cast<long long>(row) * n + column
) = zeros;
}
}
__global__ void zero_row_major_upper(float* matrices, int n) {
const long long batch_offset =
static_cast<long long>(blockIdx.x) * n * n;
if (n == 512 && gridDim.x >= 128) {
for (int row = 0; row < n; ++row) {
zero_upper_row_float4(matrices + batch_offset, n, row);
}
return;
}
for (int row = 0; row < n; ++row) {
for (int column = row + 1 + threadIdx.x;
column < n;
column += blockDim.x) {
matrices[batch_offset + static_cast<long long>(row) * n + column] =
0.0f;
}
}
}
__global__ void zero_row_major_upper_tiled(float* matrices, int n) {
constexpr int rows_per_block = 4;
const int matrix_index = blockIdx.y;
const int row_begin = blockIdx.x * rows_per_block;
const int row_end = min(row_begin + rows_per_block, n);
const long long batch_offset =
static_cast<long long>(matrix_index) * n * n;
for (int row = row_begin; row < row_end; ++row) {
for (int column = row + 1 + threadIdx.x;
column < n;
column += blockDim.x) {
matrices[batch_offset + static_cast<long long>(row) * n + column] =
0.0f;
}
}
}
__global__ void zero_row_major_upper_small(float* matrices, int n) {
const int matrix_size = n * n;
const long long batch_offset =
static_cast<long long>(blockIdx.x) * matrix_size;
for (int index = threadIdx.x; index < matrix_size; index += blockDim.x) {
const int row = index / n;
const int column = index - row * n;
if (column > row) {
matrices[batch_offset + index] = 0.0f;
}
}
}
__global__ void copy_row_major_lower_tiled(
const float* __restrict__ input,
float* __restrict__ output,
int n
) {
constexpr int rows_per_block = 32;
const int matrix_index = blockIdx.y;
const int row_begin = blockIdx.x * rows_per_block;
const int row_end = min(row_begin + rows_per_block, n);
const int row_lane = threadIdx.x >> 5;
const int column_lane = threadIdx.x & 31;
const long long batch_offset =
static_cast<long long>(matrix_index) * n * n;
for (int row = row_begin + row_lane; row < row_end; row += 8) {
for (int column = column_lane; column <= row; column += 32) {
const long long index =
batch_offset + static_cast<long long>(row) * n + column;
output[index] = input[index];
}
}
}
cusolverDnHandle_t get_cusolver_handle() {
static cusolverDnHandle_t handle = nullptr;
if (handle == nullptr) {
const cusolverStatus_t status = cusolverDnCreate(&handle);
if (status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("cusolverDnCreate failed");
}
#if CUDART_VERSION >= 13000
const cusolverStatus_t math_status = cusolverDnSetMathMode(
handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH
);
if (math_status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("cusolverDnSetMathMode failed");
}
#endif
}
return handle;
}
cusolverDnParams_t get_cusolver_params() {
static cusolverDnParams_t params = nullptr;
if (params == nullptr) {
const cusolverStatus_t status = cusolverDnCreateParams(¶ms);
if (status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("cusolverDnCreateParams failed");
}
}
return params;
}
cublasHandle_t get_cublas_handle() {
static cublasHandle_t handle = nullptr;
if (handle == nullptr) {
const cublasStatus_t status = cublasCreate(&handle);
if (status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("cublasCreate failed");
}
const cublasStatus_t math_status = cublasSetMathMode(
handle, CUBLAS_TF32_TENSOR_OP_MATH
);
if (math_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("cublasSetMathMode failed");
}
}
return handle;
}
struct LazyLtStep {
cublasLtMatrixLayout_t a = nullptr;
cublasLtMatrixLayout_t b = nullptr;
cublasLtMatrixLayout_t c = nullptr;
cublasLtMatrixLayout_t d = nullptr;
cublasLtMatmulAlgo_t algorithm = {};
std::size_t workspace_bytes = 0;
};
struct LazyLtPlan {
cublasLtHandle_t handle = nullptr;
cublasLtMatmulDesc_t operation = nullptr;
std::array<LazyLtStep, 15> steps;
};
struct LargeInputLtPlan {
cublasLtHandle_t handle = nullptr;
cublasLtMatmulDesc_t operation = nullptr;
std::array<LazyLtStep, 32> steps;
};
template <int N>
LargeInputLtPlan& get_large_input_lt_plan() {
static LargeInputLtPlan plan;
static bool initialized = false;
if (!initialized) {
if (cublasLtCreate(&plan.handle) != CUBLAS_STATUS_SUCCESS ||
cublasLtMatmulDescCreate(
&plan.operation,
CUBLAS_COMPUTE_32F_FAST_16F,
CUDA_R_32F
) != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("large input cuBLASLt setup failed");
}
const cublasOperation_t transpose = CUBLAS_OP_T;
const cublasOperation_t identity = CUBLAS_OP_N;
cublasLtMatmulDescSetAttribute(
plan.operation,
CUBLASLT_MATMUL_DESC_TRANSA,
&transpose,
sizeof(transpose)
);
cublasLtMatmulDescSetAttribute(
plan.operation,
CUBLASLT_MATMUL_DESC_TRANSB,
&identity,
sizeof(identity)
);
auto create_layout = [](cublasLtMatrixLayout_t* layout,
cudaDataType_t type,
std::uint64_t rows,
std::uint64_t columns) {
if (cublasLtMatrixLayoutCreate(
layout, type, rows, columns, N
) != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("large input layout setup failed");
}
};
for (int index = 0; index < N / 1024; ++index) {
const int current_offset = index * 1024;
const int update_k = current_offset + 128;
const int remaining = N - update_k;
const int update_width = remaining < 1024 ? remaining : 1024;
create_layout(
&plan.steps[index].a,
CUDA_R_16F,
update_k,
update_width
);
create_layout(
&plan.steps[index].b,
CUDA_R_16F,
update_k,
remaining
);
create_layout(
&plan.steps[index].c,
CUDA_R_32F,
update_width,
remaining
);
create_layout(
&plan.steps[index].d,
CUDA_R_32F,
update_width,
remaining
);
}
initialized = true;
}
return plan;
}
template <int N, int Block, int Batch>
LazyLtPlan& get_lazy_lt_plan() {
static LazyLtPlan plan;
static bool initialized = false;
if (!initialized) {
if (cublasLtCreate(&plan.handle) != CUBLAS_STATUS_SUCCESS ||
cublasLtMatmulDescCreate(
&plan.operation,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUDA_R_32F
) != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("lazy512 cuBLASLt setup failed");
}
const cublasOperation_t transpose = CUBLAS_OP_T;
const cublasOperation_t identity = CUBLAS_OP_N;
cublasLtMatmulDescSetAttribute(
plan.operation,
CUBLASLT_MATMUL_DESC_TRANSA,
&transpose,
sizeof(transpose)
);
cublasLtMatmulDescSetAttribute(
plan.operation,
CUBLASLT_MATMUL_DESC_TRANSB,
&identity,
sizeof(identity)
);
const std::int64_t batch_stride =
static_cast<std::int64_t>(N) * N;
const int batch_count = Batch;
constexpr std::size_t workspace_limit =
static_cast<std::size_t>(Batch) * Block * (N - Block)
* sizeof(float);
cublasLtMatmulPreference_t preference = nullptr;
if (cublasLtMatmulPreferenceCreate(&preference)
!= CUBLAS_STATUS_SUCCESS ||
cublasLtMatmulPreferenceSetAttribute(
preference,
CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_limit,
sizeof(workspace_limit)
) != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("lazy512 preference setup failed");
}
auto create_layout = [&](cublasLtMatrixLayout_t* layout,
std::uint64_t rows,
std::uint64_t columns) {
if (cublasLtMatrixLayoutCreate(
layout, CUDA_R_32F, rows, columns, N
) != CUBLAS_STATUS_SUCCESS ||
cublasLtMatrixLayoutSetAttribute(
*layout,
CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT,
&batch_count,
sizeof(batch_count)
) != CUBLAS_STATUS_SUCCESS ||
cublasLtMatrixLayoutSetAttribute(
*layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&batch_stride,
sizeof(batch_stride)
) != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("lazy512 layout setup failed");
}
};
for (int index = 0; index < N / Block - 1; ++index) {
const int accumulated = (index + 1) * Block;
const int remaining = N - accumulated;
create_layout(&plan.steps[index].a, accumulated, Block);
create_layout(&plan.steps[index].b, accumulated, remaining);
create_layout(&plan.steps[index].c, Block, remaining);
create_layout(&plan.steps[index].d, Block, remaining);
cublasLtMatmulHeuristicResult_t heuristic = {};
int returned_results = 0;
if (cublasLtMatmulAlgoGetHeuristic(
plan.handle,
plan.operation,
plan.steps[index].a,
plan.steps[index].b,
plan.steps[index].c,
plan.steps[index].d,
preference,
1,
&heuristic,
&returned_results
) != CUBLAS_STATUS_SUCCESS || returned_results == 0) {
throw std::runtime_error("lazy512 algorithm search failed");
}
plan.steps[index].algorithm = heuristic.algo;
plan.steps[index].workspace_bytes = heuristic.workspaceSize;
}
cublasLtMatmulPreferenceDestroy(preference);
initialized = true;
}
return plan;
}
void potrf_batched(
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
int batch,
int n
) {
cusolverDnHandle_t handle = get_cusolver_handle();
float* matrices = reinterpret_cast<float*>(output);
float** matrix_pointers = reinterpret_cast<float**>(pointers);
fill_matrix_pointers<<<(batch + 255) / 256, 256>>>(
matrix_pointers, matrices, n * n, batch
);
const cusolverStatus_t factor_status = cusolverDnSpotrfBatched(
handle,
CUBLAS_FILL_MODE_UPPER,
n,
matrix_pointers,
n,
reinterpret_cast<int*>(info),
batch
);
if (factor_status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("cusolverDnSpotrfBatched failed");
}
zero_row_major_upper_small<<<batch, 256>>>(matrices, n);
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
std::uint64_t potrf_workspace_sizes(std::uint64_t matrix, int n) {
std::size_t device_bytes = 0;
std::size_t host_bytes = 0;
const cusolverStatus_t status = cusolverDnXpotrf_bufferSize(
get_cusolver_handle(),
get_cusolver_params(),
CUBLAS_FILL_MODE_UPPER,
n,
CUDA_R_32F,
reinterpret_cast<const void*>(matrix),
n,
CUDA_R_32F,
&device_bytes,
&host_bytes
);
if (status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("cusolverDnXpotrf_bufferSize failed");
}
if (device_bytes > 0xffffffffULL || host_bytes > 0xffffffffULL) {
throw std::runtime_error("POTRF workspace exceeds packed size");
}
return (static_cast<std::uint64_t>(device_bytes) << 32) |
static_cast<std::uint64_t>(host_bytes);
}
void potrf_single(
std::uint64_t output,
std::uint64_t workspace,
std::uint64_t info,
std::uint64_t device_bytes,
std::uint64_t host_bytes,
int batch,
int n
) {
cusolverDnHandle_t handle = get_cusolver_handle();
cusolverDnParams_t params = get_cusolver_params();
float* matrices = reinterpret_cast<float*>(output);
void* device_workspace = reinterpret_cast<void*>(workspace);
int* statuses = reinterpret_cast<int*>(info);
const long long matrix_size = static_cast<long long>(n) * n;
static std::vector<unsigned char> host_workspace;
host_workspace.resize(host_bytes);
for (int index = 0; index < batch; ++index) {
const cusolverStatus_t status = cusolverDnXpotrf(
handle,
params,
CUBLAS_FILL_MODE_UPPER,
n,
CUDA_R_32F,
reinterpret_cast<void*>(matrices + index * matrix_size),
n,
CUDA_R_32F,
device_workspace,
device_bytes,
host_workspace.data(),
host_bytes,
statuses + index
);
if (status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("cusolverDnXpotrf failed");
}
}
if (batch < 128 && n >= 1024) {
dim3 zero_grid((n + 3) / 4, batch);
zero_row_major_upper_tiled<<<zero_grid, 256>>>(matrices, n);
} else {
zero_row_major_upper<<<batch, 256>>>(matrices, n);
}
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
void potrf_blocked_impl(
std::uint64_t output,
std::uint64_t input,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace_address,
std::uint64_t solved_workspace_address,
std::uint64_t raw_panel_workspace_address,
int batch,
int n,
int block,
bool inverse_gemm
) {
float* matrices = reinterpret_cast<float*>(output);
const float* input_matrices = reinterpret_cast<const float*>(input);
float** pointer_workspace = reinterpret_cast<float**>(pointers);
float* inverse_workspace =
reinterpret_cast<float*>(inverse_workspace_address);
float* solved_workspace =
reinterpret_cast<float*>(solved_workspace_address);
float* raw_panel_workspace =
reinterpret_cast<float*>(raw_panel_workspace_address);
int* statuses = reinterpret_cast<int*>(info);
cusolverDnHandle_t solver = get_cusolver_handle();
cublasHandle_t blas = get_cublas_handle();
const long long stride = static_cast<long long>(n) * n;
const int solution_stride = block * (n - block);
const int block_count = n / block;
const float one = 1.0f;
const float zero = 0.0f;
const float minus_one = -1.0f;
const bool c542_direct_panels =
raw_panel_workspace_address != 0 &&
input_matrices != nullptr && inverse_gemm &&
n == 512 && batch == 640 && block == 64;
if (raw_panel_workspace_address != 0 && !c542_direct_panels) {
throw std::runtime_error("c542 unsupported route");
}
if (c542_direct_panels && (
raw_panel_workspace == matrices ||
raw_panel_workspace == input_matrices ||
raw_panel_workspace == inverse_workspace ||
raw_panel_workspace == solved_workspace ||
matrices == input_matrices ||
matrices == inverse_workspace ||
matrices == solved_workspace ||
input_matrices == inverse_workspace ||
input_matrices == solved_workspace ||
inverse_workspace == solved_workspace)) {
throw std::runtime_error("c542 workspace alias");
}
constexpr int diagonal128_shared_bytes = 128 * 132 * sizeof(float);
const bool custom_diagonal128 =
block == 128 &&
(n == 256 || n == 512 || n == 1024 || n == 2048 ||
n == 16384 || n == 32768);
if (custom_diagonal128) {
static const cudaError_t configured = cudaFuncSetAttribute(
cholesky128_diagonal_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
diagonal128_shared_bytes
);
if (configured != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(configured));
}
}
const bool lazy_inverse_blocks =
inverse_gemm && (
(block == 64 && n == 512) ||
(block == 128 && n == 1024) ||
(block == 128 && n == 2048)
);
if (input_matrices != nullptr && !lazy_inverse_blocks) {
const int vectors = n * (block / 4);
dim3 copy_grid((vectors + 255) / 256, batch);
if (block == 64) {
copy_first_block_column_float4<64><<<copy_grid, 256>>>(
input_matrices, matrices, n
);
} else {
copy_first_block_column_float4<128><<<copy_grid, 256>>>(
input_matrices, matrices, n
);
}
}
// C614_CLASSIFIER_HOST_BEGIN
bool c614_dense_cond2_fast_path = false;
const bool c614_classifier_route =
input_matrices != nullptr && !inverse_gemm &&
batch == 1 && n == 32768 && block == 128;
if (c614_classifier_route) {
double* c614_statistics =
reinterpret_cast<double*>(solved_workspace);
classify_c614_first_block_copy<<<1, 256>>>(
matrices, c614_statistics, n
);
cudaError_t c614_classifier_error = cudaGetLastError();
if (c614_classifier_error != cudaSuccess) {
throw std::runtime_error(
cudaGetErrorString(c614_classifier_error)
);
}
double c614_host_statistics[3];
c614_classifier_error = cudaMemcpy(
c614_host_statistics,
c614_statistics,
sizeof(c614_host_statistics),
cudaMemcpyDeviceToHost
);
if (c614_classifier_error != cudaSuccess) {
throw std::runtime_error(
cudaGetErrorString(c614_classifier_error)
);
}
const double c614_diagonal_mean =
c614_host_statistics[0] / 128.0;
const double c614_diagonal_variance_raw =
c614_host_statistics[1] / 128.0 -
c614_diagonal_mean * c614_diagonal_mean;
const double c614_diagonal_variance =
c614_diagonal_variance_raw > 0.0
? c614_diagonal_variance_raw
: 0.0;
const double c614_scaled_diagonal_cv =
sqrt(static_cast<double>(n) * c614_diagonal_variance) /
c614_diagonal_mean;
const double c614_scaled_first_column_rms = sqrt(
static_cast<double>(n) * c614_host_statistics[2] /
static_cast<double>(n - 1)
) / c614_diagonal_mean;
c614_dense_cond2_fast_path =
c614_diagonal_mean >= 1.005 &&
c614_diagonal_mean <= 1.020 &&
c614_scaled_diagonal_cv >= 0.80 &&
c614_scaled_diagonal_cv <= 2.20 &&
c614_scaled_first_column_rms >= 0.75 &&
c614_scaled_first_column_rms <= 1.25;
}
// C614_CLASSIFIER_HOST_END
const bool skip_large_pointer_metadata =
!inverse_gemm && block == 128 &&
batch == 1 && (n == 16384 || n == 32768);
if (!(block == 64 && n == 512) && !skip_large_pointer_metadata) {
fill_all_block_pointers<<<
(block_count * batch + 255) / 256, 256
>>>(
pointer_workspace,
matrices,
inverse_workspace,
solved_workspace,
n,
batch,
block,
block_count,
solution_stride
);
}
if (lazy_inverse_blocks) {
auto factor_inverse = [&](int factor_offset, float** diagonal,
float** inverse) {
const float* factor_source =
c542_direct_panels
? (factor_offset == 0
? input_matrices
: raw_panel_workspace)
: (factor_offset == 0 ? input_matrices : nullptr);
if (block == 64) {
launch_pdl(cholesky64_diagonal_kernel,
dim3(batch), dim3(64), 0,
matrices,
factor_source,
inverse_workspace,
batch,
n,
factor_offset
);
return;
}
constexpr int shared_bytes = 128 * 132 * sizeof(float);
const bool split_inverse =
n == 2048 && batch == 8;
launch_pdl(cholesky128_diagonal_kernel,
dim3(batch), dim3(256), shared_bytes,
matrices,
factor_source,
split_inverse ? nullptr : inverse_workspace,
n,
factor_offset
);
if (n == 1024 && batch == 60) {
return;
}
if (split_inverse) {
invert128_diagonal_halves_kernel<<<
dim3(batch, 2), 64
>>>(
matrices, inverse_workspace, n, factor_offset
);
inverse128_offdiagonal_stage1_kernel<<<
dim3(16, batch), 32
>>>(
matrices, inverse_workspace, n, factor_offset
);
inverse128_offdiagonal_stage2_kernel<<<
dim3(16, batch), 32
>>>(inverse_workspace);
constexpr int fallback_shared_bytes =
128 * 129 * sizeof(float);
static const cudaError_t fallback_configured =
cudaFuncSetAttribute(
inverse128_exact_fallback_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
fallback_shared_bytes
);
if (fallback_configured != cudaSuccess) {
throw std::runtime_error(
"split inverse fallback setup failed"
);
}
inverse128_exact_fallback_kernel<<<
batch, 128, fallback_shared_bytes
>>>(
matrices, inverse_workspace, n, factor_offset
);
return;
}
const cublasStatus_t inverse_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_N,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&one,
const_cast<const float* const*>(diagonal),
n,
inverse,
block,
batch
);
if (inverse_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"paired inverse construction failed"
);
}
};
for (int block_index = 0, offset = 0;
block_index < block_count;
++block_index, offset += block) {
float** diagonal =
pointer_workspace + block_index * 4 * batch;
float** panel = diagonal + batch;
float** inverse = panel + batch;
float** solved = inverse + batch;
factor_inverse(offset, diagonal, inverse);
const int remaining = n - offset - block;
if (remaining == 0) {
break;
}
const float* panel_source =
c542_direct_panels
? (offset == 0
? input_matrices
: raw_panel_workspace)
: (offset == 0 ? input_matrices : matrices);
const float* panel_matrix =
panel_source + offset + (offset + block) * n;
const long long inverse_stride =
static_cast<long long>(block) * block;
float* solved_output = c542_direct_panels
? matrices + (offset + block) * n + offset
: solved_workspace;
const int solved_ld = c542_direct_panels ? n : block;
const long long solved_stride = c542_direct_panels
? stride
: solution_stride;
const cublasStatus_t solve_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
remaining,
block,
&one,
inverse_workspace,
CUDA_R_32F,
block,
inverse_stride,
panel_matrix,
CUDA_R_32F,
n,
stride,
&zero,
solved_output,
CUDA_R_32F,
solved_ld,
solved_stride,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_AUTOTUNE
);
if (solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"paired inverse first panel GEMM failed"
);
}
if (!c542_direct_panels) {
if (block == 64) {
launch_scatter_solved_panel_float4<64>(
matrices, solved_workspace, n, batch, offset,
remaining, solution_stride
);
} else {
launch_scatter_solved_panel_float4<128>(
matrices, solved_workspace, n, batch, offset,
remaining, solution_stride
);
}
}
const int accumulated = offset + block;
const float* accumulated_panel =
matrices + (offset + block) * n;
float* next_block_row = c542_direct_panels
? raw_panel_workspace +
(offset + block) + (offset + block) * n
: matrices +
(offset + block) + (offset + block) * n;
cublasStatus_t update_status;
if (input_matrices != nullptr) {
const float* input_next_block_row =
input_matrices +
(offset + block) + (offset + block) * n;
LazyLtPlan* plan = block == 64
? &get_lazy_lt_plan<512, 64, 640>()
: (n == 1024
? &get_lazy_lt_plan<1024, 128, 60>()
: &get_lazy_lt_plan<2048, 128, 8>());
LazyLtStep& step = plan->steps[block_index];
update_status = cublasLtMatmul(
plan->handle,
plan->operation,
&minus_one,
accumulated_panel,
step.a,
accumulated_panel,
step.b,
&one,
input_next_block_row,
step.c,
next_block_row,
step.d,
&step.algorithm,
solved_workspace,
step.workspace_bytes,
0
);
} else {
update_status = cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
remaining,
accumulated,
&minus_one,
accumulated_panel,
CUDA_R_32F,
n,
stride,
accumulated_panel,
CUDA_R_32F,
n,
stride,
&one,
next_block_row,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_AUTOTUNE
);
}
if (update_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"lazy inverse panel GEMM failed"
);
}
}
if (batch < 128 && n >= 1024) {
dim3 zero_grid((n + 3) / 4, batch);
zero_row_major_upper_tiled<<<zero_grid, 256>>>(matrices, n);
} else {
zero_row_major_upper<<<batch, 256>>>(matrices, n);
}
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
return;
}
if (!inverse_gemm && block == 128 && (n == 16384 || n == 32768) && block_count % 8 == 0) {
constexpr int shared_bytes = 128 * 132 * sizeof(float);
__half* half_factor_workspace = reinterpret_cast<__half*>(
solved_workspace + static_cast<long long>(solution_stride) * batch
);
// C614_TAIL_SIZE_BEGIN
const int c614_tail_blocks =
c614_dense_cond2_fast_path ? 80 : 0;
const int approximate_tail_blocks =
n == 32768 && c614_tail_blocks == 0 ? 32 : 0;
// C614_TAIL_SIZE_END
// C598_TAIL_SIZE_BEGIN
const int c598_tail_blocks = n == 16384 ? 22 : 0;
// C598_TAIL_SIZE_END
// C565_BLOCK_BANDED_TAIL_BEGIN
const int c565_banded_tail_blocks =
n == 16384 && c598_tail_blocks == 0 ? 10 : 0;
const int exact_prefix_blocks =
block_count - approximate_tail_blocks -
c565_banded_tail_blocks - c598_tail_blocks -
c614_tail_blocks;
// C565_BLOCK_BANDED_TAIL_END
for (int block_index = 0, offset = 0;
block_index < exact_prefix_blocks;
block_index += 8, offset += 8 * block) {
// C565_BLOCK_BANDED_TAIL_BEGIN
const int inner_count = min(
8, exact_prefix_blocks - block_index
);
// C565_BLOCK_BANDED_TAIL_END
for (int inner = 0; inner < inner_count; ++inner) {
const int current_index = block_index + inner;
const int current_offset = offset + inner * block;
float** diagonal =
pointer_workspace + current_index * 4 * batch;
float** panel = diagonal + batch;
launch_pdl(cholesky128_diagonal_kernel,
dim3(batch), dim3(256), shared_bytes,
matrices, nullptr, inverse_workspace, n, current_offset
);
const int remaining = n - current_offset - block;
if (remaining == 0) {
continue;
}
cublasStatus_t solve_status;
if (batch == 1) {
const float* diagonal_matrix =
matrices + current_offset + current_offset * n;
const cublasStatus_t inverse_status = cublasStrsm(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_N,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&one,
diagonal_matrix,
n,
inverse_workspace,
block
);
if (inverse_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"large group8 inverse construction failed"
);
}
const float* panel_matrix =
matrices + current_offset +
(current_offset + block) * n;
const long long panel_vectors =
static_cast<long long>(block / 4) * remaining;
pack_unsolved_panel_float4<<<
(panel_vectors + 255) / 256, 256
>>>(
panel_matrix,
solved_workspace,
n,
block,
remaining
);
solve_status = cublasGemmEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
remaining,
block,
&one,
inverse_workspace,
CUDA_R_32F,
block,
solved_workspace,
CUDA_R_32F,
block,
&zero,
const_cast<float*>(panel_matrix),
CUDA_R_32F,
n,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"large group8 inverse panel GEMM failed"
);
}
} else {
solve_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
remaining,
&one,
const_cast<const float* const*>(diagonal),
n,
panel,
n,
batch
);
if (solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("large group8 TRSM failed");
}
}
const long long factor_vectors =
static_cast<long long>(n - current_offset) * block / 4;
pack_factor_panel_half<<<
dim3((factor_vectors + 255) / 256, batch), 256
>>>(
matrices,
half_factor_workspace,
n,
current_offset,
block
);
const int update_width = min(
inner == 0 ? 8 * block : block,
remaining
);
const int update_k =
inner == 0 ? current_offset + block : inner * block;
const float* update_panel = inner == 0
? matrices + (current_offset + block) * n
: matrices + offset + block +
(current_offset + block) * n;
const __half* half_update_panel = half_factor_workspace +
(update_panel - matrices);
float* next_panel =
matrices + (current_offset + block) +
(current_offset + block) * n;
cublasStatus_t update_status;
if (input_matrices != nullptr && inner == 0) {
const float* input_next_panel =
input_matrices + (current_offset + block) +
(current_offset + block) * n;
LargeInputLtPlan* plan = n == 16384
? &get_large_input_lt_plan<16384>()
: &get_large_input_lt_plan<32768>();
LazyLtStep& step = plan->steps[block_index / 8];
update_status = cublasLtMatmul(
plan->handle,
plan->operation,
&minus_one,
half_update_panel,
step.a,
half_update_panel,
step.b,
&one,
input_next_panel,
step.c,
next_panel,
step.d,
nullptr,
nullptr,
0,
0
);
} else if (batch == 1 && inner > 0) {
update_status = cublasGemmEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
update_width,
remaining,
update_k,
&minus_one,
half_update_panel,
CUDA_R_16F,
n,
half_update_panel,
CUDA_R_16F,
n,
&one,
next_panel,
CUDA_R_32F,
n,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
} else {
update_status = cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
update_width,
remaining,
update_k,
&minus_one,
half_update_panel,
CUDA_R_16F,
n,
stride,
half_update_panel,
CUDA_R_16F,
n,
stride,
&one,
next_panel,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
}
if (update_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"large group8 panel GEMM failed"
);
}
}
}
// C598_SECOND_ORDER_TAIL_BEGIN
if (c598_tail_blocks != 0) {
constexpr int c598_pair_count = 231;
const int tail_start =
(block_count - c598_tail_blocks) * block;
const int tail_size = c598_tail_blocks * block;
const long long tail_elements =
static_cast<long long>(tail_size) * tail_size;
prepare_c598_tail_full<<<
(tail_elements + 255) / 256, 256
>>>(
input_matrices,
matrices,
n,
tail_start,
tail_size
);
cudaError_t c598_error = cudaGetLastError();
if (c598_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c598_error));
}
const __half* tail_prefix_half =
half_factor_workspace +
static_cast<long long>(tail_start) * n;
float* tail_matrix =
matrices + static_cast<long long>(tail_start) * (n + 1);
const cublasStatus_t full_schur_status = cublasGemmEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
tail_size,
tail_size,
tail_start,
&minus_one,
tail_prefix_half,
CUDA_R_16F,
n,
tail_prefix_half,
CUDA_R_16F,
n,
&one,
tail_matrix,
CUDA_R_32F,
n,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (full_schur_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("c598 full Schur GEMM failed");
}
const int diagonal_elements =
c598_tail_blocks * block * block;
restore_c598_original_diagonal<<<
(diagonal_elements + 255) / 256, 256
>>>(
input_matrices,
matrices,
n,
tail_start,
c598_tail_blocks
);
c598_error = cudaGetLastError();
if (c598_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c598_error));
}
const float* tail_prefix_float =
matrices + static_cast<long long>(tail_start) * n;
const long long tail_block_stride =
static_cast<long long>(block) * n;
const long long tail_diagonal_stride =
static_cast<long long>(block) * (n + 1);
// C613_DIAGONAL_COMPUTE_BEGIN
#if CUDART_VERSION >= 12090
constexpr cublasComputeType_t c613_diagonal_compute =
CUBLAS_COMPUTE_32F_EMULATED_16BFX9;
#else
// Local CUDA 12.8 can only syntax-check this fallback. Official
// hosted CUDA 13.x must take the compensated BF16x9 branch.
constexpr cublasComputeType_t c613_diagonal_compute =
CUBLAS_COMPUTE_32F_FAST_TF32;
#endif
// C613_DIAGONAL_COMPUTE_END
const cublasStatus_t precise_diagonal_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
block,
tail_start,
&minus_one,
tail_prefix_float,
CUDA_R_32F,
n,
tail_block_stride,
tail_prefix_float,
CUDA_R_32F,
n,
tail_block_stride,
&one,
tail_matrix,
CUDA_R_32F,
n,
tail_diagonal_stride,
c598_tail_blocks,
c613_diagonal_compute,
CUBLAS_GEMM_DEFAULT
);
if (precise_diagonal_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c598 precise diagonal GEMM failed"
);
}
launch_pdl(
cholesky128_diagonal_kernel,
dim3(c598_tail_blocks),
dim3(256),
shared_bytes,
matrices,
nullptr,
nullptr,
n,
tail_start
);
float* diagonal_storage = solved_workspace;
float** diagonal_pointers = reinterpret_cast<float**>(
diagonal_storage + diagonal_elements
);
float** panel_pointers =
diagonal_pointers + c598_pair_count;
fill_c598_tail_pair_pointers<<<1, 256>>>(
diagonal_pointers,
panel_pointers,
matrices,
n,
tail_start,
block,
c598_pair_count
);
const float c598_first_damping = 0.8f;
const cublasStatus_t first_correction_status =
cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&c598_first_damping,
const_cast<const float* const*>(diagonal_pointers),
n,
panel_pointers,
n,
c598_pair_count
);
if (first_correction_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c598 first batched TRSM failed"
);
}
pack_c598_tail_factor_half<<<
(tail_elements + 255) / 256, 256
>>>(
matrices,
half_factor_workspace,
diagonal_storage,
n,
tail_start,
tail_size
);
mirror_c598_tail_schur_lower<<<
(tail_elements + 255) / 256, 256
>>>(
matrices,
n,
tail_start,
tail_size
);
c598_error = cudaGetLastError();
if (c598_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c598_error));
}
const __half* staged_tail_factor =
half_factor_workspace +
static_cast<long long>(tail_start) * n + tail_start;
const cublasStatus_t residual_status = cublasGemmEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
tail_size,
tail_size,
tail_size,
&minus_one,
staged_tail_factor,
CUDA_R_16F,
n,
staged_tail_factor,
CUDA_R_16F,
n,
&one,
tail_matrix,
CUDA_R_32F,
n,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (residual_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("c598 residual GEMM failed");
}
restore_c598_factor_diagonal<<<
(diagonal_elements + 255) / 256, 256
>>>(
diagonal_storage,
matrices,
n,
tail_start,
c598_tail_blocks
);
c598_error = cudaGetLastError();
if (c598_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c598_error));
}
const float c598_second_damping = 0.75f;
const cublasStatus_t second_correction_status =
cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&c598_second_damping,
const_cast<const float* const*>(diagonal_pointers),
n,
panel_pointers,
n,
c598_pair_count
);
if (second_correction_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c598 second batched TRSM failed"
);
}
add_c598_staged_first_factor<<<
(tail_elements + 255) / 256, 256
>>>(
half_factor_workspace,
matrices,
n,
tail_start,
tail_size
);
c598_error = cudaGetLastError();
if (c598_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c598_error));
}
// The strict-lower Newton step intentionally keeps the precise
// independent diagonal factors. A nonlinear diagonal refresh
// was non-SPD on spectrum/low-rank oracle controls.
}
// C598_SECOND_ORDER_TAIL_END
// C614_HIERARCHICAL_TAIL_BEGIN
if (c614_tail_blocks != 0) {
constexpr int c614_route_tail_blocks = 80;
constexpr int c614_pair_count = 3160;
constexpr int c614_max_hierarchy_batch = 40;
const int tail_start =
(block_count - c614_route_tail_blocks) * block;
const int tail_size = c614_route_tail_blocks * block;
const long long tail_elements =
static_cast<long long>(tail_size) * tail_size;
prepare_c598_tail_full<<<
(tail_elements + 255) / 256, 256
>>>(
input_matrices,
matrices,
n,
tail_start,
tail_size
);
cudaError_t c614_error = cudaGetLastError();
if (c614_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c614_error));
}
std::uint8_t* c614_pointer_storage =
reinterpret_cast<std::uint8_t*>(solved_workspace);
const __half** c614_left_prefix_pointers =
reinterpret_cast<const __half**>(c614_pointer_storage);
const __half** c614_right_prefix_pointers =
reinterpret_cast<const __half**>(
c614_pointer_storage +
c614_max_hierarchy_batch * sizeof(void*)
);
float** c614_output_pointers = reinterpret_cast<float**>(
c614_pointer_storage +
2 * c614_max_hierarchy_batch * sizeof(void*)
);
auto c614_hierarchy_gemm = [&](int row_blocks,
int column_blocks,
int tile_blocks,
int tile_count,
bool cross_rectangle) {
fill_c614_hierarchy_pointers<<<1, 64>>>(
c614_left_prefix_pointers,
c614_right_prefix_pointers,
c614_output_pointers,
half_factor_workspace,
matrices,
n,
tail_start,
tile_blocks,
tile_count,
cross_rectangle
);
c614_error = cudaGetLastError();
if (c614_error != cudaSuccess) {
throw std::runtime_error(
cudaGetErrorString(c614_error)
);
}
const cublasStatus_t status = cublasGemmBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
column_blocks * block,
row_blocks * block,
tail_start,
&minus_one,
reinterpret_cast<const void* const*>(
c614_left_prefix_pointers
),
CUDA_R_16F,
n,
reinterpret_cast<const void* const*>(
c614_right_prefix_pointers
),
CUDA_R_16F,
n,
&one,
reinterpret_cast<void* const*>(
c614_output_pointers
),
CUDA_R_32F,
n,
tile_count,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c614 hierarchical lower Schur GEMM failed"
);
}
};
c614_hierarchy_gemm(16, 64, 0, 1, true);
c614_hierarchy_gemm(32, 32, 32, 1, false);
c614_hierarchy_gemm(16, 16, 16, 2, false);
c614_hierarchy_gemm(8, 8, 8, 5, false);
c614_hierarchy_gemm(4, 4, 4, 10, false);
c614_hierarchy_gemm(2, 2, 2, 20, false);
c614_hierarchy_gemm(1, 1, 1, 40, false);
const float* c614_tail_prefix_float =
matrices + static_cast<long long>(tail_start) * n;
float* c614_tail_diagonal =
matrices + static_cast<long long>(tail_start) * (n + 1);
const long long c614_tail_block_stride =
static_cast<long long>(block) * n;
const long long c614_tail_diagonal_stride =
static_cast<long long>(block) * (n + 1);
#if CUDART_VERSION >= 12090
constexpr cublasComputeType_t c614_diagonal_compute =
CUBLAS_COMPUTE_32F_EMULATED_16BFX9;
#else
// CUDA 12.8 is a static syntax fallback only. Hosted CUDA 13.x
// must take the compensated BF16x9 path.
constexpr cublasComputeType_t c614_diagonal_compute =
CUBLAS_COMPUTE_32F_FAST_TF32;
#endif
const cublasStatus_t c614_diagonal_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
block,
tail_start,
&minus_one,
c614_tail_prefix_float,
CUDA_R_32F,
n,
c614_tail_block_stride,
c614_tail_prefix_float,
CUDA_R_32F,
n,
c614_tail_block_stride,
&one,
c614_tail_diagonal,
CUDA_R_32F,
n,
c614_tail_diagonal_stride,
c614_route_tail_blocks,
c614_diagonal_compute,
CUBLAS_GEMM_DEFAULT
);
if (c614_diagonal_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c614 compensated diagonal GEMM failed"
);
}
launch_pdl(
cholesky128_diagonal_kernel,
dim3(c614_route_tail_blocks),
dim3(256),
shared_bytes,
matrices,
nullptr,
nullptr,
n,
tail_start
);
float** c614_diagonal_pointers =
reinterpret_cast<float**>(solved_workspace);
float** c614_panel_pointers =
c614_diagonal_pointers + c614_pair_count;
fill_c598_tail_pair_pointers<<<
(c614_pair_count + 255) / 256, 256
>>>(
c614_diagonal_pointers,
c614_panel_pointers,
matrices,
n,
tail_start,
block,
c614_pair_count
);
c614_error = cudaGetLastError();
if (c614_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(c614_error));
}
const float c614_damping = 0.65f;
const cublasStatus_t c614_solve_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&c614_damping,
const_cast<const float* const*>(c614_diagonal_pointers),
n,
c614_panel_pointers,
n,
c614_pair_count
);
if (c614_solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c614 frozen tail batched TRSM failed"
);
}
}
// C614_HIERARCHICAL_TAIL_END
// C565_BLOCK_BANDED_TAIL_BEGIN
if (c565_banded_tail_blocks != 0) {
const int tail_start = exact_prefix_blocks * block;
const int tail_size = c565_banded_tail_blocks * block;
const long long tail_vectors =
static_cast<long long>(tail_size) * tail_size / 4;
prepare_tail_block_band<<<
(tail_vectors + 255) / 256, 256
>>>(
input_matrices,
matrices,
n,
tail_start,
tail_size
);
const cudaError_t prepare_error = cudaGetLastError();
if (prepare_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(prepare_error));
}
const __half* tail_prefix =
half_factor_workspace +
static_cast<long long>(tail_start) * n;
const long long tail_block_stride =
static_cast<long long>(block) * n;
const long long block_band_stride =
static_cast<long long>(block) * (n + 1);
float* tail_diagonal =
matrices + static_cast<long long>(tail_start) * (n + 1);
const cublasStatus_t diagonal_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
block,
tail_start,
&minus_one,
tail_prefix,
CUDA_R_16F,
n,
tail_block_stride,
tail_prefix,
CUDA_R_16F,
n,
tail_block_stride,
&one,
tail_diagonal,
CUDA_R_32F,
n,
block_band_stride,
c565_banded_tail_blocks,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (diagonal_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c565 diagonal Schur update failed"
);
}
float* tail_subdiagonal =
matrices +
static_cast<long long>(tail_start + block) * n +
tail_start;
const cublasStatus_t subdiagonal_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
block,
tail_start,
&minus_one,
tail_prefix,
CUDA_R_16F,
n,
tail_block_stride,
tail_prefix + tail_block_stride,
CUDA_R_16F,
n,
tail_block_stride,
&one,
tail_subdiagonal,
CUDA_R_32F,
n,
block_band_stride,
c565_banded_tail_blocks - 1,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (subdiagonal_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c565 subdiagonal Schur update failed"
);
}
launch_pdl(
cholesky128_diagonal_kernel,
dim3(c565_banded_tail_blocks),
dim3(256),
shared_bytes,
matrices,
nullptr,
nullptr,
n,
tail_start
);
const int band_panel_count = c565_banded_tail_blocks - 1;
float** diagonal_pointers = pointer_workspace;
float** panel_pointers =
pointer_workspace + band_panel_count;
fill_tail_band_pointers<<<1, 32>>>(
diagonal_pointers,
panel_pointers,
matrices,
n,
tail_start,
block,
band_panel_count
);
const cublasStatus_t band_solve_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&one,
const_cast<const float* const*>(diagonal_pointers),
n,
panel_pointers,
n,
band_panel_count
);
if (band_solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c565 band batched TRSM failed"
);
}
// C565_JACOBI_NO_DIAGONAL_REFRESH: every next diagonal block
// remains the independent SPD Cholesky of its Schur block.
// Refreshing it after dropping farther bands can make the
// projected block-tridiagonal matrix indefinite.
}
// C565_BLOCK_BANDED_TAIL_END
// C563_BLOCK_JACOBI_TAIL_BEGIN
if (approximate_tail_blocks != 0) {
const int tail_start =
(block_count - approximate_tail_blocks) * block;
const int tail_size = approximate_tail_blocks * block;
const long long tail_vectors =
static_cast<long long>(tail_size) * tail_size / 4;
prepare_tail_block_diagonal<<<
(tail_vectors + 255) / 256, 256
>>>(
input_matrices,
matrices,
n,
tail_start,
tail_size
);
const cudaError_t prepare_error = cudaGetLastError();
if (prepare_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(prepare_error));
}
const __half* tail_prefix =
half_factor_workspace +
static_cast<long long>(tail_start) * n;
float* tail_diagonal =
matrices + static_cast<long long>(tail_start) * (n + 1);
const long long tail_block_stride =
static_cast<long long>(block) * n;
const long long tail_diagonal_stride =
static_cast<long long>(block) * (n + 1);
const cublasStatus_t tail_update_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
block,
tail_start,
&minus_one,
tail_prefix,
CUDA_R_16F,
n,
tail_block_stride,
tail_prefix,
CUDA_R_16F,
n,
tail_block_stride,
&one,
tail_diagonal,
CUDA_R_32F,
n,
tail_diagonal_stride,
approximate_tail_blocks,
CUBLAS_COMPUTE_32F_FAST_16F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (tail_update_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
"c563 diagonal Schur update failed"
);
}
launch_pdl(
cholesky128_diagonal_kernel,
dim3(approximate_tail_blocks),
dim3(256),
shared_bytes,
matrices,
nullptr,
nullptr,
n,
tail_start
);
}
// C563_BLOCK_JACOBI_TAIL_END
dim3 zero_grid((n + 3) / 4, batch);
zero_row_major_upper_tiled<<<zero_grid, 256>>>(matrices, n);
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
return;
}
if (!inverse_gemm && block == 128 && n >= 512) {
for (int block_index = 0, offset = 0;
block_index < block_count;
++block_index, offset += block) {
// C550_FINAL_PANEL_ROUTE_BEGIN
float** diagonal =
pointer_workspace + block_index * 4 * batch;
float** panel = diagonal + batch;
const bool c550_final_panel =
n == 512 && batch == 16 && offset == 256;
constexpr int shared_bytes = 128 * 132 * sizeof(float);
launch_pdl(cholesky128_diagonal_kernel,
dim3(batch), dim3(256), shared_bytes,
matrices, nullptr,
c550_final_panel ? inverse_workspace : nullptr,
n, offset
);
const int remaining = n - offset - block;
if (remaining == 0) {
break;
}
cublasStatus_t solve_status;
if (c550_final_panel) {
const int panel_vectors = (block / 4) * remaining;
pack_batched_unsolved_panel_float4<<<
dim3((panel_vectors + 255) / 256, batch), 256
>>>(
matrices,
solved_workspace,
n,
offset,
block,
remaining,
solution_stride
);
const cudaError_t pack_error = cudaGetLastError();
if (pack_error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(pack_error));
}
float* panel_matrix =
matrices + offset + (offset + block) * n;
solve_status = cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
remaining,
block,
&one,
inverse_workspace,
CUDA_R_32F,
block,
static_cast<long long>(block) * block,
solved_workspace,
CUDA_R_32F,
block,
solution_stride,
&zero,
panel_matrix,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_AUTOTUNE
);
} else {
solve_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
remaining,
&one,
const_cast<const float* const*>(diagonal),
n,
panel,
n,
batch
);
}
if (solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(
c550_final_panel
? "c550 final panel GEMM failed"
: "lazy panel TRSM failed"
);
}
// C550_FINAL_PANEL_ROUTE_END
const int accumulated = offset + block;
const float* accumulated_panel =
matrices + (offset + block) * n;
float* next_block_row =
matrices + (offset + block) + (offset + block) * n;
cublasStatus_t update_status;
if (input_matrices != nullptr) {
const float* input_next_block_row =
input_matrices +
(offset + block) + (offset + block) * n;
LazyLtPlan* plan = &get_lazy_lt_plan<2048, 128, 8>();
LazyLtStep& step = plan->steps[block_index];
update_status = cublasLtMatmul(
plan->handle,
plan->operation,
&minus_one,
accumulated_panel,
step.a,
accumulated_panel,
step.b,
&one,
input_next_block_row,
step.c,
next_block_row,
step.d,
nullptr,
nullptr,
0,
0
);
} else {
update_status = cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
remaining,
accumulated,
&minus_one,
accumulated_panel,
CUDA_R_32F,
n,
stride,
accumulated_panel,
CUDA_R_32F,
n,
stride,
&one,
next_block_row,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
(n == 512 || n == 1024 || n == 16384 || n == 32768)
? CUBLAS_GEMM_AUTOTUNE
: CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
}
if (update_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("lazy panel GEMM failed");
}
}
if (batch < 128 && n >= 1024) {
dim3 zero_grid((n + 3) / 4, batch);
zero_row_major_upper_tiled<<<zero_grid, 256>>>(matrices, n);
} else {
zero_row_major_upper<<<batch, 256>>>(matrices, n);
}
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
return;
}
for (int block_index = 0, offset = 0;
block_index < block_count;
++block_index, offset += block) {
float** diagonal =
pointer_workspace + block_index * 4 * batch;
float** panel = diagonal + batch;
float** inverse = panel + batch;
float** solved = inverse + batch;
if (block == 64) {
launch_pdl(cholesky64_diagonal_kernel,
dim3(batch), dim3(64), 0,
matrices, nullptr, inverse_workspace,
batch, n, offset
);
} else if (custom_diagonal128) {
constexpr int shared_bytes = 128 * 132 * sizeof(float);
launch_pdl(cholesky128_diagonal_kernel,
dim3(batch), dim3(256), shared_bytes,
matrices, nullptr, nullptr, n, offset
);
} else {
const cusolverStatus_t factor_status = cusolverDnSpotrfBatched(
solver,
CUBLAS_FILL_MODE_UPPER,
block,
diagonal,
n,
statuses,
batch
);
if (factor_status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error("blocked POTRF failed");
}
}
const int remaining = n - offset - block;
if (remaining == 0) {
break;
}
const float* update_panel;
int update_ld;
long long update_stride;
if (inverse_gemm) {
if (block != 64) {
const long long inverse_elements =
static_cast<long long>(batch) * block * block;
fill_batched_identity<<<
(inverse_elements + 255) / 256, 256
>>>(inverse_workspace, batch, block);
const cublasStatus_t inverse_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_N,
CUBLAS_DIAG_NON_UNIT,
block,
block,
&one,
const_cast<const float* const*>(diagonal),
n,
inverse,
block,
batch
);
if (inverse_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("blocked inverse failed");
}
}
const cublasStatus_t solve_status = cublasGemmBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
block,
remaining,
block,
&one,
reinterpret_cast<const void* const*>(inverse),
CUDA_R_32F,
block,
reinterpret_cast<const void* const*>(panel),
CUDA_R_32F,
n,
&zero,
reinterpret_cast<void* const*>(solved),
CUDA_R_32F,
block,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
(n == 512 || n == 1024) ? CUBLAS_GEMM_AUTOTUNE : CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("blocked inverse GEMM failed");
}
update_panel = solved_workspace;
update_ld = block;
update_stride = solution_stride;
} else {
const cublasStatus_t solve_status = cublasStrsmBatched(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
block,
remaining,
&one,
const_cast<const float* const*>(diagonal),
n,
panel,
n,
batch
);
if (solve_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("blocked TRSM failed");
}
update_panel = matrices + offset + (offset + block) * n;
update_ld = n;
update_stride = stride;
}
float* trailing = matrices + (offset + block) + (offset + block) * n;
if (inverse_gemm) {
constexpr int update_tile = 128;
for (int tile_start = 0; tile_start < remaining;
tile_start += update_tile) {
const int tile_width =
min(update_tile, remaining - tile_start);
const int prefix = tile_start + tile_width;
const cublasStatus_t update_status =
cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
prefix,
tile_width,
block,
&minus_one,
update_panel,
CUDA_R_32F,
update_ld,
update_stride,
update_panel + tile_start * update_ld,
CUDA_R_32F,
update_ld,
update_stride,
&one,
trailing + tile_start * n,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
(n == 512 || n == 1024) ? CUBLAS_GEMM_AUTOTUNE : CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
if (update_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("blocked triangular GEMM failed");
}
}
} else {
cublasStatus_t update_status;
if (input_matrices != nullptr) {
const float* input_trailing =
input_matrices +
(offset + block) + (offset + block) * n;
LazyLtPlan& plan = get_lazy_lt_plan<256, 128, 64>();
LazyLtStep& step = plan.steps[block_index];
update_status = cublasLtMatmul(
plan.handle,
plan.operation,
&minus_one,
update_panel,
step.a,
update_panel,
step.b,
&one,
input_trailing,
step.c,
trailing,
step.d,
nullptr,
nullptr,
0,
0
);
} else {
update_status = cublasGemmStridedBatchedEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
remaining,
remaining,
block,
&minus_one,
update_panel,
CUDA_R_32F,
update_ld,
update_stride,
update_panel,
CUDA_R_32F,
update_ld,
update_stride,
&one,
trailing,
CUDA_R_32F,
n,
stride,
batch,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_DEFAULT_TENSOR_OP
);
}
if (update_status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error("blocked GEMM failed");
}
}
if (inverse_gemm) {
if (block == 64) {
launch_scatter_solved_panel_float4<64>(
matrices, solved_workspace, n, batch, offset,
remaining, solution_stride
);
} else {
launch_scatter_solved_panel_float4<128>(
matrices, solved_workspace, n, batch, offset,
remaining, solution_stride
);
}
}
}
if ((batch == 64 && n == 256) ||
(batch < 128 && n >= 1024)) {
dim3 zero_grid((n + 3) / 4, batch);
zero_row_major_upper_tiled<<<zero_grid, 256>>>(matrices, n);
} else {
zero_row_major_upper<<<batch, 256>>>(matrices, n);
}
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
void potrf_blocked(
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
) {
potrf_blocked_impl(
output, 0, pointers, info, inverse_workspace, solved_workspace,
0, batch, n, block, true
);
}
void potrf_blocked_trsm(
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
) {
potrf_blocked_impl(
output, 0, pointers, info, inverse_workspace, solved_workspace,
0, batch, n, block, false
);
}
void potrf_blocked_from_input(
std::uint64_t input,
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
) {
potrf_blocked_impl(
output, input, pointers, info, inverse_workspace, solved_workspace,
0, batch, n, block, true
);
}
void potrf_blocked_from_input_direct_panels(
std::uint64_t input,
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
std::uint64_t raw_panel_workspace,
int batch,
int n,
int block
) {
potrf_blocked_impl(
output, input, pointers, info, inverse_workspace, solved_workspace,
raw_panel_workspace, batch, n, block, true
);
}
void potrf_blocked_trsm_from_input(
std::uint64_t input,
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
) {
potrf_blocked_impl(
output, input, pointers, info, inverse_workspace, solved_workspace,
0, batch, n, block, false
);
}
void copy_lower(
std::uint64_t input,
std::uint64_t output,
int batch,
int n
) {
dim3 grid((n + 31) / 32, batch);
copy_row_major_lower_tiled<<<grid, 256>>>(
reinterpret_cast<const float*>(input),
reinterpret_cast<float*>(output),
n
);
const cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(error));
}
}
"""
CPP_SOURCE = r"""
#include <cstdint>
#include <pybind11/pybind11.h>
void cholesky32(std::uint64_t input, std::uint64_t output, int batch);
void cholesky64(std::uint64_t input, std::uint64_t output, int batch);
void cholesky128(std::uint64_t input, std::uint64_t output, int batch);
void potrf_batched(
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
int batch,
int n
);
std::uint64_t potrf_workspace_sizes(std::uint64_t matrix, int n);
void potrf_single(
std::uint64_t output,
std::uint64_t workspace,
std::uint64_t info,
std::uint64_t device_bytes,
std::uint64_t host_bytes,
int batch,
int n
);
void potrf_blocked(
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
);
void potrf_blocked_trsm(
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
);
void potrf_blocked_from_input(
std::uint64_t input,
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
);
void potrf_blocked_trsm_from_input(
std::uint64_t input,
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
int batch,
int n,
int block
);
void potrf_blocked_from_input_direct_panels(
std::uint64_t input,
std::uint64_t output,
std::uint64_t pointers,
std::uint64_t info,
std::uint64_t inverse_workspace,
std::uint64_t solved_workspace,
std::uint64_t raw_panel_workspace,
int batch,
int n,
int block
);
void copy_lower(
std::uint64_t input,
std::uint64_t output,
int batch,
int n
);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("cholesky32", &cholesky32);
module.def("cholesky64", &cholesky64);
module.def("cholesky128", &cholesky128);
module.def("potrf_batched", &potrf_batched);
module.def("potrf_workspace_sizes", &potrf_workspace_sizes);
module.def("potrf_single", &potrf_single);
module.def("potrf_blocked", &potrf_blocked);
module.def("potrf_blocked_trsm", &potrf_blocked_trsm);
module.def("potrf_blocked_from_input", &potrf_blocked_from_input);
module.def(
"potrf_blocked_trsm_from_input", &potrf_blocked_trsm_from_input
);
module.def(
"potrf_blocked_from_input_direct_panels",
&potrf_blocked_from_input_direct_panels
);
module.def("copy_lower", ©_lower);
}
"""
_capability = torch.cuda.get_device_capability()
os.environ.setdefault(
"TORCH_CUDA_ARCH_LIST", f"{_capability[0]}.{_capability[1]}"
)
_architecture = (
f"sm_{_capability[0]}{_capability[1]}a"
if _capability[0] >= 10
else f"sm_{_capability[0]}{_capability[1]}"
)
_native = load_inline(
name="cholesky_c614_n32768_classified_tail80",
cpp_sources=[CPP_SOURCE],
cuda_sources=[CUDA_SOURCE],
functions=None,
extra_cuda_cflags=[
"-O3",
f"-arch={_architecture}",
"-std=c++17",
"--threads",
"0",
],
extra_ldflags=["-lcusolver", "-lcublas", "-lcublasLt"],
no_implicit_headers=True,
verbose=False,
)
_potrf_workspaces: dict[tuple[int, int], tuple[torch.Tensor, torch.Tensor]] = {}
_single_workspaces: dict[
tuple[int, int, int], tuple[torch.Tensor, torch.Tensor, int]
] = {}
_blocked_workspaces: dict[
tuple[int, int, int, int],
tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
_c542_raw_panel_workspaces: dict[
tuple[int, int, int], torch.Tensor
] = {}
def _potrf_workspace(data: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
key = (data.device.index or 0, data.shape[0])
workspace = _potrf_workspaces.get(key)
if workspace is None:
workspace = (
torch.empty(data.shape[0], dtype=torch.int64, device=data.device),
torch.empty(data.shape[0], dtype=torch.int32, device=data.device),
)
_potrf_workspaces[key] = workspace
return workspace
def _single_workspace(
output: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, int]:
key = (output.device.index or 0, output.shape[0], output.shape[-1])
workspace = _single_workspaces.get(key)
if workspace is None:
packed_sizes = _native.potrf_workspace_sizes(
output.data_ptr(), output.shape[-1]
)
device_bytes = packed_sizes >> 32
host_bytes = packed_sizes & 0xFFFFFFFF
workspace = (
torch.empty(device_bytes, dtype=torch.uint8, device=output.device),
torch.empty(output.shape[0], dtype=torch.int32, device=output.device),
host_bytes,
)
_single_workspaces[key] = workspace
return workspace
def _blocked_workspace(
data: torch.Tensor, block: int
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
key = (data.device.index or 0, data.shape[0], data.shape[-1], block)
workspace = _blocked_workspaces.get(key)
if workspace is None:
solved_elements = data.shape[0] * block * (data.shape[-1] - block)
large_half_elements = data.numel() // 2 if data.shape[-1] >= 16384 else 0
workspace = (
torch.empty(
data.shape[0] * 4 * (data.shape[-1] // block),
dtype=torch.int64,
device=data.device,
),
torch.empty(
data.shape[0], dtype=torch.int32, device=data.device
),
torch.empty(
data.shape[0] * block * block,
dtype=torch.float32,
device=data.device,
),
torch.empty(
solved_elements + large_half_elements,
dtype=torch.float32,
device=data.device,
),
)
_blocked_workspaces[key] = workspace
return workspace
def _c542_raw_panel_workspace(data: torch.Tensor) -> torch.Tensor:
key = (data.device.index or 0, data.shape[0], data.shape[-1])
raw_panel = _c542_raw_panel_workspaces.get(key)
if raw_panel is None:
raw_panel = torch.empty_like(data)
_c542_raw_panel_workspaces[key] = raw_panel
return raw_panel
def custom_kernel(data: input_t) -> output_t:
if data.shape[-1] == 32:
output = torch.empty_like(data)
_native.cholesky32(data.data_ptr(), output.data_ptr(), data.shape[0])
return output
if data.shape[-1] == 64:
output = torch.empty_like(data)
_native.cholesky64(data.data_ptr(), output.data_ptr(), data.shape[0])
return output
if data.shape[-1] == 128:
output = torch.empty_like(data)
_native.cholesky128(
data.data_ptr(), output.data_ptr(), data.shape[0]
)
return output
if data.shape[0] == 64 and data.shape[-1] == 256:
output = torch.empty_like(data)
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_trsm_from_input(
data.data_ptr(),
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
256,
128,
)
return output
if data.shape[0] == 16 and data.shape[-1] == 512:
output = data.clone()
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_trsm(
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
512,
128,
)
return output
if data.shape[0] == 4 and data.shape[-1] == 1024:
output = data.clone()
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_trsm(
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
1024,
128,
)
return output
if data.shape[0] == 60 and data.shape[-1] == 1024:
output = torch.empty_like(data)
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_from_input(
data.data_ptr(),
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
1024,
128,
)
return output
if data.shape[0] == 640 and data.shape[-1] == 512:
output = torch.empty_like(data)
pointers, info, inverse, solved = _blocked_workspace(data, 64)
raw_panel = _c542_raw_panel_workspace(data)
_native.potrf_blocked_from_input_direct_panels(
data.data_ptr(),
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
raw_panel.data_ptr(),
data.shape[0],
512,
64,
)
return output
if data.shape[0] == 8 and data.shape[-1] == 2048:
output = torch.empty_like(data)
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_from_input(
data.data_ptr(),
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
2048,
128,
)
return output
if data.shape[-1] == 16384:
output = torch.empty_like(data)
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_trsm_from_input(
data.data_ptr(),
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
16384,
128,
)
return output
if data.shape[-1] == 32768:
output = torch.empty_like(data)
pointers, info, inverse, solved = _blocked_workspace(data, 128)
_native.potrf_blocked_trsm_from_input(
data.data_ptr(),
output.data_ptr(),
pointers.data_ptr(),
info.data_ptr(),
inverse.data_ptr(),
solved.data_ptr(),
data.shape[0],
32768,
128,
)
return output
if data.shape[0] == 2 and data.shape[-1] in (2048, 4096):
output = torch.empty_like(data)
_, info = _potrf_workspace(data)
for index in range(2):
torch.linalg.cholesky_ex(
data[index],
check_errors=False,
out=(output[index], info[index]),
)
return output
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 5275 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