submission 899876
leanyoshi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3999 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-899876?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:1ff203f9f5b46d764251648bd8c993a0fa413644f195b8726c6140cc9291bb40
license declaredunknown
license concludedunknown
authorsleanyoshi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
const __nv_fp8_e4m3* primary,mma
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, 16, 16, 16, float>shared-memory
__shared__ float tile[kMatrixSize][kSharedStride];vector-width = float4
const float4 packed = *reinterpret_cast<const float4*>(Kernel source
submission.py3999 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
_CHOLESKY64_CPP = r"""
#include <torch/extension.h>
torch::Tensor cholesky64_block32_hacc4safe_cuda(torch::Tensor input);
torch::Tensor cholesky32_warp_cuda(torch::Tensor input);
torch::Tensor cholesky128_wmma_schur_cuda(torch::Tensor input);
torch::Tensor cholesky128_wmma_schur_exact_cuda(torch::Tensor input);
torch::Tensor cholesky128_block32_warp16_lowerio_anybatch_cuda(
torch::Tensor input);
torch::Tensor cholesky256_batch64_register_panel64_cuda(torch::Tensor input);
torch::Tensor cholesky256_combined_batched_cuda(torch::Tensor input);
torch::Tensor cholesky512_batch16_register_panel64_cuda(torch::Tensor input);
torch::Tensor cholesky512_combined_batched_cuda(torch::Tensor input);
torch::Tensor cholesky512_highbatch_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky1024_batch4_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky1024_batch60_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky2048_batch2_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky2048_batch8_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky4096_batch2_truebatched_fast16bf_cuda(
torch::Tensor input);
torch::Tensor cholesky_large_explicit_bf16_cuda(torch::Tensor input);
torch::Tensor cholesky_large_combined_bf16x9_cuda(torch::Tensor input);
"""
_CHOLESKY64_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cublasLt.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <algorithm>
#include <cstdint>
#include <limits>
#include <memory>
#include <mma.h>
#include <mutex>
#include <vector>
namespace {
constexpr int kMatrixSize = 64;
constexpr int kSharedStride = 65;
constexpr int kMatrixElements = kMatrixSize * kMatrixSize;
constexpr int kBlockSize = 32;
constexpr int kWarpsPerBlock = 4;
constexpr int kThreads = 32 * kWarpsPerBlock;
constexpr int kSchurHalfTerms = 14;
constexpr int kSchurHalfStride = 16;
// Accumulate a product of two binary16 inputs directly into binary32. On
// sm_100 this maps to the scalar mixed-precision FMA without explicit unpack
// instructions or binary16 rounding of the product.
__device__ __forceinline__ float n64_fhfma(
unsigned short a, unsigned short b, float accumulator) {
float result;
asm("{.reg .b16 ha, hb;\n\t"
"mov.b16 ha, %1;\n\t"
"mov.b16 hb, %2;\n\t"
"fma.rn.f32.f16 %0, ha, hb, %3;}\n"
: "=f"(result)
: "h"(a), "h"(b), "f"(accumulator));
return result;
}
__device__ __forceinline__ float n64_fhfma2_acc(
__half2 a, __half2 b, float accumulator) {
const unsigned int a_bits = *reinterpret_cast<const unsigned int*>(&a);
const unsigned int b_bits = *reinterpret_cast<const unsigned int*>(&b);
accumulator = n64_fhfma(
static_cast<unsigned short>(a_bits),
static_cast<unsigned short>(b_bits), accumulator);
return n64_fhfma(
static_cast<unsigned short>(a_bits >> 16),
static_cast<unsigned short>(b_bits >> 16), accumulator);
}
constexpr int kMatrixSize32 = 32;
constexpr int kMatrixElements32 = kMatrixSize32 * kMatrixSize32;
constexpr int kBlockSize32 = 16;
constexpr int kIoStride32 = 33;
constexpr int kIoTileElements32 = kMatrixSize32 * kIoStride32;
constexpr int kWarpsPerBlock32 = 4;
constexpr int kThreads32 = 32 * kWarpsPerBlock32;
constexpr int kSharedBytes32 =
kWarpsPerBlock32 * kIoTileElements32 * sizeof(float);
static_assert(kSharedBytes32 == 16896,
"n=32 register-stage I/O shared-memory size changed");
constexpr int kMatrixSize128 = 128;
constexpr int kSharedStride128 = 129;
// WMMA accumulator loads of float require a leading dimension divisible by
// four. Padding by four also keeps every 16-column tile origin aligned.
constexpr int kWmmaSharedStride128 = 132;
constexpr int kMatrixElements128 = kMatrixSize128 * kMatrixSize128;
constexpr int kWarpsPerBlock128 = 16;
constexpr int kThreads128 = 32 * kWarpsPerBlock128;
constexpr int kFloatSharedBytes128 =
kMatrixSize128 * kSharedStride128 * static_cast<int>(sizeof(float));
constexpr int kWmmaFloatSharedBytes128 =
kMatrixSize128 * kWmmaSharedStride128 * static_cast<int>(sizeof(float));
constexpr int kHalfPanelRows128 = kMatrixSize128 - kBlockSize;
constexpr int kHalfPanelStride128 = 40;
constexpr int kHalfPanelBytes128 =
kHalfPanelRows128 * kHalfPanelStride128 * static_cast<int>(sizeof(__half));
constexpr int kHalfPanelsBytes128 = 2 * kHalfPanelBytes128;
constexpr int kWmmaSharedBytes128 =
kWmmaFloatSharedBytes128 + kHalfPanelsBytes128;
template <int... Indices>
struct N64IndexSequence {};
template <int Count, int... Indices>
struct N64MakeIndexSequence
: N64MakeIndexSequence<Count - 1, Count - 1, Indices...> {};
template <int... Indices>
struct N64MakeIndexSequence<0, Indices...> {
using type = N64IndexSequence<Indices...>;
};
constexpr int kNestedBlockSize64 = 16;
using N64AllIndices =
typename N64MakeIndexSequence<kNestedBlockSize64>::type;
struct N64RegisterRow {
float v0;
float v1;
float v2;
float v3;
float v4;
float v5;
float v6;
float v7;
float v8;
float v9;
float v10;
float v11;
float v12;
float v13;
float v14;
float v15;
};
template <int Index>
struct N64RegisterAccess;
#define N64_REGISTER_ACCESS(Index) \
template <> \
struct N64RegisterAccess<Index> { \
__device__ __forceinline__ static float get( \
const N64RegisterRow& row) { \
return row.v##Index; \
} \
__device__ __forceinline__ static void set( \
N64RegisterRow& row, float value) { \
row.v##Index = value; \
} \
};
N64_REGISTER_ACCESS(0)
N64_REGISTER_ACCESS(1)
N64_REGISTER_ACCESS(2)
N64_REGISTER_ACCESS(3)
N64_REGISTER_ACCESS(4)
N64_REGISTER_ACCESS(5)
N64_REGISTER_ACCESS(6)
N64_REGISTER_ACCESS(7)
N64_REGISTER_ACCESS(8)
N64_REGISTER_ACCESS(9)
N64_REGISTER_ACCESS(10)
N64_REGISTER_ACCESS(11)
N64_REGISTER_ACCESS(12)
N64_REGISTER_ACCESS(13)
N64_REGISTER_ACCESS(14)
N64_REGISTER_ACCESS(15)
#undef N64_REGISTER_ACCESS
template <int Stride, int... Indices>
__device__ __forceinline__ void n64_load_register_row(
N64RegisterRow& values,
float (*tile)[Stride],
int row,
int col_base,
N64IndexSequence<Indices...>) {
int unused[] = {
0, (N64RegisterAccess<Indices>::set(
values, tile[row][col_base + Indices]),
0)...};
(void)unused;
}
template <int Stride, int... Indices>
__device__ __forceinline__ void n64_store_register_row(
const N64RegisterRow& values,
float (*tile)[Stride],
int row,
int col_base,
N64IndexSequence<Indices...>) {
int unused[] = {
0, ((tile[row][col_base + Indices] =
N64RegisterAccess<Indices>::get(values)),
0)...};
(void)unused;
}
template <int Pivot, int... Previous>
__device__ __forceinline__ void n64_factor_step(
N64RegisterRow& values,
int lane,
N64IndexSequence<Previous...>) {
float diagonal = 0.0f;
if (lane == Pivot) {
diagonal = N64RegisterAccess<Pivot>::get(values);
int diagonal_terms[] = {
0, ((diagonal = fmaf(
-N64RegisterAccess<Previous>::get(values),
N64RegisterAccess<Previous>::get(values), diagonal)),
0)...};
(void)diagonal_terms;
diagonal = sqrtf(fmaxf(diagonal, 0.0f));
}
diagonal = __shfl_sync(0xffffffffu, diagonal, Pivot);
float value = N64RegisterAccess<Pivot>::get(values);
int update_terms[] = {
0, ((value = fmaf(
-N64RegisterAccess<Previous>::get(values),
__shfl_sync(0xffffffffu,
N64RegisterAccess<Previous>::get(values), Pivot),
value)),
0)...};
(void)update_terms;
if (lane == Pivot) {
N64RegisterAccess<Pivot>::set(values, diagonal);
} else if (lane > Pivot && lane < kNestedBlockSize64) {
N64RegisterAccess<Pivot>::set(values, value / diagonal);
}
}
template <int... Pivots>
__device__ __forceinline__ void n64_factor_register_block(
N64RegisterRow& values,
int lane,
N64IndexSequence<Pivots...>) {
int unused[] = {
0, (n64_factor_step<Pivots>(
values, lane,
typename N64MakeIndexSequence<Pivots>::type{}),
0)...};
(void)unused;
}
template <int Pivot, int... Previous>
__device__ __forceinline__ void n64_nested_panel_step(
N64RegisterRow& values,
int lane,
N64IndexSequence<Previous...>) {
float value = N64RegisterAccess<Pivot>::get(values);
int update_terms[] = {
0, ((value = fmaf(
-N64RegisterAccess<Previous>::get(values),
__shfl_sync(0xffffffffu,
N64RegisterAccess<Previous>::get(values), Pivot),
value)),
0)...};
(void)update_terms;
const float diagonal = __shfl_sync(
0xffffffffu, N64RegisterAccess<Pivot>::get(values), Pivot);
if (lane >= kNestedBlockSize64) {
N64RegisterAccess<Pivot>::set(values, value / diagonal);
}
}
template <int... Pivots>
__device__ __forceinline__ void n64_solve_nested_panel(
N64RegisterRow& values,
int lane,
N64IndexSequence<Pivots...>) {
int unused[] = {
0, (n64_nested_panel_step<Pivots>(
values, lane,
typename N64MakeIndexSequence<Pivots>::type{}),
0)...};
(void)unused;
}
template <int... Terms>
__device__ __forceinline__ float n64_nested_panel_dot(
const N64RegisterRow& values,
int row,
int col,
float value,
N64IndexSequence<Terms...>) {
int unused[] = {
0, ((value = fmaf(
-__shfl_sync(0xffffffffu,
N64RegisterAccess<Terms>::get(values),
kNestedBlockSize64 + row),
__shfl_sync(0xffffffffu,
N64RegisterAccess<Terms>::get(values),
kNestedBlockSize64 + col),
value)),
0)...};
(void)unused;
return value;
}
template <int Stride>
__device__ __forceinline__ void n64_factor_nested32(
float (*tile)[Stride],
int lane,
int base) {
N64RegisterRow values;
if (lane < kNestedBlockSize64) {
n64_load_register_row(values, tile, base + lane, base, N64AllIndices{});
} else {
values = {};
}
n64_factor_register_block(values, lane, N64AllIndices{});
if (lane < kNestedBlockSize64) {
n64_store_register_row(values, tile, base + lane, base, N64AllIndices{});
} else {
n64_load_register_row(values, tile, base + lane, base, N64AllIndices{});
}
n64_solve_nested_panel(values, lane, N64AllIndices{});
if (lane >= kNestedBlockSize64) {
n64_store_register_row(values, tile, base + lane, base, N64AllIndices{});
}
#pragma unroll
for (int linear = lane;
linear < kNestedBlockSize64 * kNestedBlockSize64;
linear += 32) {
const int row = linear >> 4;
const int col = linear & 15;
float value = tile[base + kNestedBlockSize64 + row]
[base + kNestedBlockSize64 + col];
value = n64_nested_panel_dot(
values, row, col, value, N64AllIndices{});
if (row >= col) {
tile[base + kNestedBlockSize64 + row]
[base + kNestedBlockSize64 + col] = value;
}
}
__syncwarp();
if (lane < kNestedBlockSize64) {
n64_load_register_row(
values, tile, base + kNestedBlockSize64 + lane,
base + kNestedBlockSize64, N64AllIndices{});
}
n64_factor_register_block(values, lane, N64AllIndices{});
if (lane < kNestedBlockSize64) {
n64_store_register_row(
values, tile, base + kNestedBlockSize64 + lane,
base + kNestedBlockSize64, N64AllIndices{});
}
}
template <int Pivot, int Stride, int... Previous>
__device__ __forceinline__ void n64_outer_panel_first_step(
N64RegisterRow& values,
float (*tile)[Stride],
int factor_base,
N64IndexSequence<Previous...>) {
float value = N64RegisterAccess<Pivot>::get(values);
int update_terms[] = {
0, ((value = fmaf(
-N64RegisterAccess<Previous>::get(values),
tile[factor_base + Pivot][factor_base + Previous], value)),
0)...};
(void)update_terms;
N64RegisterAccess<Pivot>::set(
values, value / tile[factor_base + Pivot][factor_base + Pivot]);
}
template <int Stride, int... Pivots>
__device__ __forceinline__ void n64_solve_outer_panel_first(
N64RegisterRow& values,
float (*tile)[Stride],
int factor_base,
N64IndexSequence<Pivots...>) {
int unused[] = {
0, (n64_outer_panel_first_step<Pivots>(
values, tile, factor_base,
typename N64MakeIndexSequence<Pivots>::type{}),
0)...};
(void)unused;
}
template <int Pivot, int Stride, int... History, int... Previous>
__device__ __forceinline__ void n64_outer_panel_second_step(
N64RegisterRow& values,
float (*tile)[Stride],
int row,
int factor_base,
N64IndexSequence<History...>,
N64IndexSequence<Previous...>) {
float value = N64RegisterAccess<Pivot>::get(values);
int history_terms[] = {
0, ((value = fmaf(
-tile[row][factor_base + History],
tile[factor_base + kNestedBlockSize64 + Pivot]
[factor_base + History],
value)),
0)...};
(void)history_terms;
int update_terms[] = {
0, ((value = fmaf(
-N64RegisterAccess<Previous>::get(values),
tile[factor_base + kNestedBlockSize64 + Pivot]
[factor_base + kNestedBlockSize64 + Previous],
value)),
0)...};
(void)update_terms;
N64RegisterAccess<Pivot>::set(
values,
value /
tile[factor_base + kNestedBlockSize64 + Pivot]
[factor_base + kNestedBlockSize64 + Pivot]);
}
template <int Stride, int... Pivots>
__device__ __forceinline__ void n64_solve_outer_panel_second(
N64RegisterRow& values,
float (*tile)[Stride],
int row,
int factor_base,
N64IndexSequence<Pivots...>) {
int unused[] = {
0, (n64_outer_panel_second_step<Pivots>(
values, tile, row, factor_base, N64AllIndices{},
typename N64MakeIndexSequence<Pivots>::type{}),
0)...};
(void)unused;
}
// Factor a 64x64 matrix as two 32x32 diagonal blocks. FP32 shared memory is
// authoritative throughout. A compact, aligned FP16 mirror holds selected
// finalized L10 entries for grouped half2 work in the rank-32 Schur update.
__global__ __launch_bounds__(kThreads) void cholesky64_block32_hacc4safe_kernel(
const float* __restrict__ input,
float* __restrict__ output) {
const int matrix_offset = static_cast<int>(blockIdx.x) * kMatrixElements;
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
__shared__ float tile[kMatrixSize][kSharedStride];
__shared__ float schur_diagonal[kBlockSize];
__shared__ __align__(16) __half
schur_half[kBlockSize][kSchurHalfStride];
#pragma unroll
for (int vector = tid; vector < kMatrixElements / 4;
vector += kThreads) {
const int load_row = vector >> 4;
const int col = (vector & 15) << 2;
const float4 packed = *reinterpret_cast<const float4*>(
input + matrix_offset + load_row * kMatrixSize + col);
tile[load_row][col + 0] = load_row >= col + 0 ? packed.x : 0.0f;
tile[load_row][col + 1] = load_row >= col + 1 ? packed.y : 0.0f;
tile[load_row][col + 2] = load_row >= col + 2 ? packed.z : 0.0f;
tile[load_row][col + 3] = load_row >= col + 3 ? packed.w : 0.0f;
}
__syncthreads();
// Phase 1: warp zero factors A00 as two bounded 16-column register stages.
// Every lane carries at most one explicit 16-float row, while shuffles move
// finalized recurrence values between row owners.
if (warp == 0) {
n64_factor_nested32(tile, lane, 0);
}
__syncthreads();
// Phase 2: two warps each solve sixteen L10 rows. A lane owns one bounded
// 16-float row at a time; the second stage consumes the first stage from
// shared memory and preserves the original left-to-right arithmetic order.
if (warp < 2 && lane < kNestedBlockSize64) {
const int row = kBlockSize + warp * kNestedBlockSize64 + lane;
N64RegisterRow values;
n64_load_register_row(values, tile, row, 0, N64AllIndices{});
n64_solve_outer_panel_first(values, tile, 0, N64AllIndices{});
n64_store_register_row(values, tile, row, 0, N64AllIndices{});
n64_load_register_row(
values, tile, row, kNestedBlockSize64, N64AllIndices{});
n64_solve_outer_panel_second(
values, tile, row, 0, N64AllIndices{});
n64_store_register_row(
values, tile, row, kNestedBlockSize64, N64AllIndices{});
}
__syncthreads();
if (warp == 0) {
const int diagonal = kBlockSize + lane;
schur_diagonal[lane] = tile[diagonal][diagonal];
}
// Compact odd columns 5..31 from the finalized panel. Seven adjacent half2
// loads cover fourteen products; the 32-byte row stride keeps every load
// naturally aligned and avoids per-product half-to-float conversion.
#pragma unroll
for (int linear = tid; linear < kBlockSize * kSchurHalfTerms;
linear += kThreads) {
const int local_row = linear / kSchurHalfTerms;
const int slot = linear - local_row * kSchurHalfTerms;
schur_half[local_row][slot] =
__float2half_rn(tile[kBlockSize + local_row][5 + 2 * slot]);
}
__syncthreads();
// Phase 3: all threads update the lower triangle of A11 in shared memory.
#pragma unroll
for (int linear = tid; linear < kBlockSize * kBlockSize;
linear += kThreads) {
const int local_row = linear >> 5;
const int local_col = linear & 31;
if (local_row >= local_col) {
const int row = kBlockSize + local_row;
const int col = kBlockSize + local_col;
float value = tile[row][col];
// Keep eighteen products in authoritative FP32: all columns 0..4 and
// the thirteen even columns 6..30. The remaining fourteen odd-column
// products use binary16 inputs with four independent FP32 FMA chains.
#pragma unroll 4
for (int k = 0; k < 5; ++k) {
value = fmaf(-tile[row][k], tile[col][k], value);
}
#pragma unroll
for (int k = 6; k < kBlockSize; k += 2) {
value = fmaf(-tile[row][k], tile[col][k], value);
}
const __half2 row0 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][0]);
const __half2 row1 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][2]);
const __half2 row2 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][4]);
const __half2 row3 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][6]);
const __half2 col0 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][0]);
const __half2 col1 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][2]);
const __half2 col2 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][4]);
const __half2 col3 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][6]);
float half_dot0 = n64_fhfma2_acc(row0, col0, 0.0f);
half_dot0 = n64_fhfma2_acc(row1, col1, half_dot0);
float half_dot1 = n64_fhfma2_acc(row2, col2, 0.0f);
half_dot1 = n64_fhfma2_acc(row3, col3, half_dot1);
const __half2 row4 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][8]);
const __half2 row5 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][10]);
const __half2 row6 = *reinterpret_cast<const __half2*>(
&schur_half[local_row][12]);
const __half2 col4 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][8]);
const __half2 col5 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][10]);
const __half2 col6 = *reinterpret_cast<const __half2*>(
&schur_half[local_col][12]);
float half_dot2 = n64_fhfma2_acc(row4, col4, 0.0f);
half_dot2 = n64_fhfma2_acc(row5, col5, half_dot2);
const float half_dot3 = n64_fhfma2_acc(row6, col6, 0.0f);
value -= (half_dot0 + half_dot1) + (half_dot2 + half_dot3);
tile[row][col] = value;
}
}
__syncthreads();
// Restore all Schur diagonal entries from their original values with exact
// FP32 rank-32 dots. Two warps each own sixteen diagonals, and adjacent
// lanes split one diagonal into two independent sixteen-term chains.
if (warp < 2) {
const int pair = lane >> 1;
const int half = lane & 1;
const int local_diagonal = warp * 16 + pair;
const int row = kBlockSize + local_diagonal;
float value = half == 0 ? schur_diagonal[local_diagonal] : 0.0f;
#pragma unroll 4
for (int k = half * 16; k < half * 16 + 16; ++k) {
const float panel_value = tile[row][k];
value = fmaf(-panel_value, panel_value, value);
}
value += __shfl_xor_sync(0xffffffffu, value, 1);
if (half == 0) {
tile[row][row] = value;
}
}
__syncthreads();
// Phase 4: warp zero factors the repaired A11 with the same two bounded
// 16-column register stages used for A00.
if (warp == 0) {
n64_factor_nested32(tile, lane, kBlockSize);
}
__syncthreads();
#pragma unroll
for (int vector = tid; vector < kMatrixElements / 4;
vector += kThreads) {
const int store_row = vector >> 4;
const int col = (vector & 15) << 2;
float4 packed;
packed.x = store_row >= col + 0 ? tile[store_row][col + 0] : 0.0f;
packed.y = store_row >= col + 1 ? tile[store_row][col + 1] : 0.0f;
packed.z = store_row >= col + 2 ? tile[store_row][col + 2] : 0.0f;
packed.w = store_row >= col + 3 ? tile[store_row][col + 3] : 0.0f;
*reinterpret_cast<float4*>(
output + matrix_offset + store_row * kMatrixSize + col) = packed;
}
}
template <int... Indices>
struct N32IndexSequence {};
template <int Count, int... Indices>
struct N32MakeIndexSequence
: N32MakeIndexSequence<Count - 1, Count - 1, Indices...> {};
template <int... Indices>
struct N32MakeIndexSequence<0, Indices...> {
using type = N32IndexSequence<Indices...>;
};
using N32AllIndices = typename N32MakeIndexSequence<kBlockSize32>::type;
struct N32RegisterRow {
float v0;
float v1;
float v2;
float v3;
float v4;
float v5;
float v6;
float v7;
float v8;
float v9;
float v10;
float v11;
float v12;
float v13;
float v14;
float v15;
};
template <int Index>
struct N32RegisterAccess;
#define N32_REGISTER_ACCESS(Index) \
template <> \
struct N32RegisterAccess<Index> { \
__device__ __forceinline__ static float get( \
const N32RegisterRow& row) { \
return row.v##Index; \
} \
__device__ __forceinline__ static void set( \
N32RegisterRow& row, float value) { \
row.v##Index = value; \
} \
};
N32_REGISTER_ACCESS(0)
N32_REGISTER_ACCESS(1)
N32_REGISTER_ACCESS(2)
N32_REGISTER_ACCESS(3)
N32_REGISTER_ACCESS(4)
N32_REGISTER_ACCESS(5)
N32_REGISTER_ACCESS(6)
N32_REGISTER_ACCESS(7)
N32_REGISTER_ACCESS(8)
N32_REGISTER_ACCESS(9)
N32_REGISTER_ACCESS(10)
N32_REGISTER_ACCESS(11)
N32_REGISTER_ACCESS(12)
N32_REGISTER_ACCESS(13)
N32_REGISTER_ACCESS(14)
N32_REGISTER_ACCESS(15)
#undef N32_REGISTER_ACCESS
template <int... Indices>
__device__ __forceinline__ void n32_load_register_row(
N32RegisterRow& values,
float (*io)[kIoStride32],
int row,
int col_base,
N32IndexSequence<Indices...>) {
int unused[] = {
0, (N32RegisterAccess<Indices>::set(
values, io[row][col_base + Indices]),
0)...};
(void)unused;
}
template <int... Indices>
__device__ __forceinline__ void n32_store_register_row(
const N32RegisterRow& values,
float (*io)[kIoStride32],
int row,
int col_base,
N32IndexSequence<Indices...>) {
int unused[] = {
0, ((io[row][col_base + Indices] =
N32RegisterAccess<Indices>::get(values)),
0)...};
(void)unused;
}
template <int Pivot, int... Previous>
__device__ __forceinline__ void n32_factor_step(
N32RegisterRow& values,
int lane,
N32IndexSequence<Previous...>) {
float diagonal = 0.0f;
if (lane == Pivot) {
diagonal = N32RegisterAccess<Pivot>::get(values);
int diagonal_terms[] = {
0, ((diagonal = fmaf(
-N32RegisterAccess<Previous>::get(values),
N32RegisterAccess<Previous>::get(values), diagonal)),
0)...};
(void)diagonal_terms;
diagonal = sqrtf(fmaxf(diagonal, 0.0f));
}
diagonal = __shfl_sync(0xffffffffu, diagonal, Pivot);
float value = N32RegisterAccess<Pivot>::get(values);
int update_terms[] = {
0, ((value = fmaf(
-N32RegisterAccess<Previous>::get(values),
__shfl_sync(0xffffffffu,
N32RegisterAccess<Previous>::get(values), Pivot),
value)),
0)...};
(void)update_terms;
if (lane == Pivot) {
N32RegisterAccess<Pivot>::set(values, diagonal);
} else if (lane > Pivot && lane < kBlockSize32) {
N32RegisterAccess<Pivot>::set(values, value / diagonal);
}
}
template <int... Pivots>
__device__ __forceinline__ void n32_factor_register_block(
N32RegisterRow& values,
int lane,
N32IndexSequence<Pivots...>) {
int unused[] = {
0, (n32_factor_step<Pivots>(
values, lane,
typename N32MakeIndexSequence<Pivots>::type{}),
0)...};
(void)unused;
}
template <int Pivot, int... Previous>
__device__ __forceinline__ void n32_panel_step(
N32RegisterRow& values,
int lane,
N32IndexSequence<Previous...>) {
float value = N32RegisterAccess<Pivot>::get(values);
int update_terms[] = {
0, ((value = fmaf(
-N32RegisterAccess<Previous>::get(values),
__shfl_sync(0xffffffffu,
N32RegisterAccess<Previous>::get(values), Pivot),
value)),
0)...};
(void)update_terms;
const float diagonal = __shfl_sync(
0xffffffffu, N32RegisterAccess<Pivot>::get(values), Pivot);
if (lane >= kBlockSize32) {
N32RegisterAccess<Pivot>::set(values, value / diagonal);
}
}
template <int... Pivots>
__device__ __forceinline__ void n32_solve_register_panel(
N32RegisterRow& values,
int lane,
N32IndexSequence<Pivots...>) {
int unused[] = {
0, (n32_panel_step<Pivots>(
values, lane,
typename N32MakeIndexSequence<Pivots>::type{}),
0)...};
(void)unused;
}
template <int... Terms>
__device__ __forceinline__ float n32_panel_dot(
const N32RegisterRow& values,
int row,
int col,
float value,
N32IndexSequence<Terms...>) {
int unused[] = {
0, ((value = fmaf(
-__shfl_sync(0xffffffffu,
N32RegisterAccess<Terms>::get(values),
kBlockSize32 + row),
__shfl_sync(0xffffffffu,
N32RegisterAccess<Terms>::get(values),
kBlockSize32 + col),
value)),
0)...};
(void)unused;
return value;
}
// Each warp factors one matrix as two exact 16-column register stages. Lanes
// 0..15 first own D00 rows while lanes 16..31 take over for L10, then those
// roles reverse for D11. Thus every lane carries at most one 16-float row.
// Shared memory is only a padded whole-matrix I/O tile; all recurrence values
// move between lanes through full-warp shuffles.
__global__ __launch_bounds__(kThreads32) void
cholesky32_lower_blocktiles_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch) {
const int lane = static_cast<int>(threadIdx.x) & 31;
const int warp = static_cast<int>(threadIdx.x) >> 5;
const int matrix = static_cast<int>(blockIdx.x) * kWarpsPerBlock32 + warp;
if (matrix >= batch) {
return;
}
const int matrix_offset = matrix * kMatrixElements32;
__shared__ float io_tiles[kWarpsPerBlock32][kMatrixSize32][kIoStride32];
float (*io)[kIoStride32] = io_tiles[warp];
// Coalesced whole-matrix load. Upper entries become the final explicit
// zeros immediately and are never touched by the factorization.
#pragma unroll
for (int linear = lane; linear < kMatrixElements32;
linear += 32) {
const int row = linear >> 5;
const int col = linear & 31;
io[row][col] = row >= col ? input[matrix_offset + linear] : 0.0f;
}
__syncwarp();
N32RegisterRow values;
if (lane < kBlockSize32) {
n32_load_register_row(values, io, lane, 0, N32AllIndices{});
} else {
values = {};
}
// Stage 1: exact FP32 Cholesky of D00. Every lane executes every shuffle;
// only lanes 0..15 retain the row recurrence in their private registers.
n32_factor_register_block(values, lane, N32AllIndices{});
if (lane < kBlockSize32) {
n32_store_register_row(values, io, lane, 0, N32AllIndices{});
} else {
n32_load_register_row(values, io, lane, 0, N32AllIndices{});
}
// Lanes 16..31 now own L10 rows. Lanes 0..15 keep D00 resident so the
// triangular solve also consumes recurrence data only through shuffles.
n32_solve_register_panel(values, lane, N32AllIndices{});
if (lane >= kBlockSize32) {
n32_store_register_row(values, io, lane, 0, N32AllIndices{});
}
// Form D11 with its 256 entries striped across the full warp. The solved
// panel remains in lanes 16..31; every lane executes every shuffle, and the
// completed lower entries return to the sole shared I/O tile.
#pragma unroll
for (int linear = lane; linear < kBlockSize32 * kBlockSize32;
linear += 32) {
const int row = linear >> 4;
const int col = linear & 15;
float value = io[kBlockSize32 + row][kBlockSize32 + col];
value = n32_panel_dot(values, row, col, value, N32AllIndices{});
if (row >= col) {
io[kBlockSize32 + row][kBlockSize32 + col] = value;
}
}
__syncwarp();
if (lane < kBlockSize32) {
n32_load_register_row(
values, io, kBlockSize32 + lane, kBlockSize32, N32AllIndices{});
}
// Stage 2: exact FP32 Cholesky of the register-resident D11 tile.
n32_factor_register_block(values, lane, N32AllIndices{});
if (lane < kBlockSize32) {
n32_store_register_row(
values, io, kBlockSize32 + lane, kBlockSize32, N32AllIndices{});
}
__syncwarp();
// Coalesced whole-matrix store includes all explicitly zeroed upper entries.
#pragma unroll
for (int linear = lane; linear < kMatrixElements32;
linear += 32) {
const int row = linear >> 5;
const int col = linear & 31;
output[matrix_offset + linear] = io[row][col];
}
}
// Four 32x32 blocked stages factor one 128x128 matrix per CTA. FP32 remains
// authoritative while a padded FP16 panel mirror feeds rank-32 WMMA updates.
// Full-FP32 fallback for non-public n=128 batch sizes. Four 32x32 blocked
// stages factor one matrix per CTA while transferring only the lower triangle.
__global__ __launch_bounds__(kThreads128) void
cholesky128_block32_warp16_lowerio_anybatch_kernel(
const float* __restrict__ input,
float* __restrict__ output) {
const int matrix_offset = static_cast<int>(blockIdx.x) * kMatrixElements128;
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
extern __shared__ float fp32_shared_values[];
float (*tile)[kSharedStride128] =
reinterpret_cast<float (*)[kSharedStride128]>(fp32_shared_values);
#pragma unroll
for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
for (int col = lane; col <= row; col += 32) {
tile[row][col] = input[matrix_offset + row * kMatrixSize128 + col];
}
}
__syncthreads();
#pragma unroll
for (int block = 0; block < 4; ++block) {
const int base = block * kBlockSize;
if (warp == 0) {
const int row = base + lane;
#pragma unroll 1
for (int k = 0; k < kBlockSize; ++k) {
const int col = base + k;
float diagonal = 0.0f;
if (lane == k) {
diagonal = tile[row][col];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
const float value = tile[row][base + j];
diagonal = fmaf(-value, value, diagonal);
}
diagonal = sqrtf(fmaxf(diagonal, 0.0f));
}
diagonal = __shfl_sync(0xffffffffu, diagonal, k);
if (lane == k) {
tile[row][col] = diagonal;
} else if (lane > k) {
float value = tile[row][col];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
value = fmaf(-tile[row][base + j], tile[col][base + j], value);
}
tile[row][col] = value / diagonal;
}
__syncwarp();
}
}
__syncthreads();
const int remaining_blocks = 3 - block;
if (warp < remaining_blocks) {
const int row = base + kBlockSize * (warp + 1) + lane;
#pragma unroll 1
for (int k = 0; k < kBlockSize; ++k) {
const int col = base + k;
float value = tile[row][col];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
value = fmaf(-tile[row][base + j], tile[col][base + j], value);
}
tile[row][col] = value / tile[col][col];
}
}
__syncthreads();
const int trailing_base = base + kBlockSize;
const int trailing_size = kMatrixSize128 - trailing_base;
const int trailing_elements = trailing_size * trailing_size;
for (int linear = tid; linear < trailing_elements;
linear += kThreads128) {
const int local_row = linear / trailing_size;
const int local_col = linear - local_row * trailing_size;
if (local_row >= local_col) {
const int row = trailing_base + local_row;
const int col = trailing_base + local_col;
float value = tile[row][col];
#pragma unroll
for (int k = 0; k < kBlockSize; ++k) {
value = fmaf(-tile[row][base + k], tile[col][base + k], value);
}
tile[row][col] = value;
}
}
__syncthreads();
}
#pragma unroll
for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
for (int col = lane; col <= row; col += 32) {
output[matrix_offset + row * kMatrixSize128 + col] = tile[row][col];
}
#pragma unroll
for (int col = row + 1 + lane; col < kMatrixSize128; col += 32) {
output[matrix_offset + row * kMatrixSize128 + col] = 0.0f;
}
}
}
template <bool Compensated>
__global__ __launch_bounds__(kThreads128) void cholesky128_wmma_schur_kernel(
const float* __restrict__ input,
float* __restrict__ output) {
const int matrix_offset = static_cast<int>(blockIdx.x) * kMatrixElements128;
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
extern __shared__ __align__(32) unsigned char shared_values[];
float (*tile)[kWmmaSharedStride128] =
reinterpret_cast<float (*)[kWmmaSharedStride128]>(shared_values);
__half (*factor_half_hi)[kHalfPanelStride128] =
reinterpret_cast<__half (*)[kHalfPanelStride128]>(
shared_values + kWmmaFloatSharedBytes128);
__half (*factor_half_lo)[kHalfPanelStride128] =
reinterpret_cast<__half (*)[kHalfPanelStride128]>(
shared_values + kWmmaFloatSharedBytes128 + kHalfPanelBytes128);
// Direct accumulator loads consume complete 16x16 diagonal tiles, so every
// shared element must be initialized. The input contract is symmetric.
#pragma unroll
for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
for (int col = lane;
col < (Compensated ? kMatrixSize128 : row + 1); col += 32) {
tile[row][col] =
input[matrix_offset + row * kMatrixSize128 + col];
}
}
__syncthreads();
#pragma unroll
for (int block = 0; block < 4; ++block) {
const int base = block * kBlockSize;
// The exact public path factors each 32x32 diagonal tile as two bounded
// 16-column register stages. The compensated fallback retains its
// established shared-memory recurrence.
if (warp == 0) {
if constexpr (!Compensated) {
n64_factor_nested32(tile, lane, base);
} else {
const int row = base + lane;
#pragma unroll 1
for (int k = 0; k < kBlockSize; ++k) {
const int col = base + k;
float diagonal = 0.0f;
if (lane == k) {
diagonal = tile[row][col];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
const float value = tile[row][base + j];
diagonal = fmaf(-value, value, diagonal);
}
diagonal = sqrtf(fmaxf(diagonal, 0.0f));
}
diagonal = __shfl_sync(0xffffffffu, diagonal, k);
if (lane == k) {
tile[row][col] = diagonal;
} else if (lane > k) {
float value = tile[row][col];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
value = fmaf(-tile[row][base + j],
tile[col][base + j], value);
}
tile[row][col] = value / diagonal;
}
__syncwarp();
}
}
}
__syncthreads();
// One warp owns each remaining 32-row panel block. The exact path keeps
// each half-row in registers and performs the two triangular solves with
// fully expanded dependencies; no inter-stage polling is required.
const int remaining_blocks = 3 - block;
if (warp < remaining_blocks) {
const int row = base + kBlockSize * (warp + 1) + lane;
if constexpr (!Compensated) {
N64RegisterRow values;
n64_load_register_row(values, tile, row, base, N64AllIndices{});
n64_solve_outer_panel_first(
values, tile, base, N64AllIndices{});
n64_store_register_row(values, tile, row, base, N64AllIndices{});
n64_load_register_row(
values, tile, row, base + kNestedBlockSize64, N64AllIndices{});
n64_solve_outer_panel_second(
values, tile, row, base, N64AllIndices{});
n64_store_register_row(
values, tile, row, base + kNestedBlockSize64, N64AllIndices{});
} else {
#pragma unroll 1
for (int k = 0; k < kBlockSize; ++k) {
const int col = base + k;
float value = tile[row][col];
#pragma unroll 4
for (int j = 0; j < k; ++j) {
value = fmaf(-tile[row][base + j],
tile[col][base + j], value);
}
tile[row][col] = value / tile[col][col];
}
}
}
__syncthreads();
const int trailing_base = base + kBlockSize;
const int trailing_size = kMatrixSize128 - trailing_base;
if (trailing_size == 0) {
continue;
}
// Split each finalized panel value into a rounded half and a rounded-half
// residual. The 40-half leading dimension satisfies WMMA alignment.
const int panel_elements = trailing_size * kBlockSize;
for (int linear = tid; linear < panel_elements; linear += kThreads128) {
const int local_row = linear >> 5;
const int k = linear & 31;
const int row = trailing_base + local_row;
const float value = tile[row][base + k];
const __half high = __float2half_rn(value);
factor_half_hi[local_row][k] = high;
if constexpr (Compensated) {
factor_half_lo[local_row][k] =
__float2half_rn(value - __half2float(high));
}
}
__syncthreads();
// Enumerate lower 16x16 trailing tiles exactly once. Each warp loads its
// current FP32 tile directly into an accumulator, negates the FP16 A
// operand, and computes C + (-A) * B before storing directly back to the
// authoritative tile. Full diagonal tiles are updated; their upper
// entries are never consumed by the lower-triangular factorization.
const int trailing_tiles = trailing_size >> 4;
const int task_count = trailing_tiles * (trailing_tiles + 1) / 2;
const int task_rounds = (task_count + kWarpsPerBlock128 - 1) /
kWarpsPerBlock128;
for (int round = 0; round < task_rounds; ++round) {
const int triangular_task = round * kWarpsPerBlock128 + warp;
if (triangular_task < task_count) {
int tile_row = 0;
int tile_col = triangular_task;
#pragma unroll
for (int candidate_row = 1; candidate_row < 8; ++candidate_row) {
if (tile_col >= candidate_row) {
tile_col -= candidate_row;
tile_row = candidate_row;
}
}
const int row = trailing_base + (tile_row << 4);
const int col = trailing_base + (tile_col << 4);
nvcuda::wmma::fragment<nvcuda::wmma::accumulator, 16, 16, 16, float>
updated;
nvcuda::wmma::load_matrix_sync(
updated, &tile[row][col], kWmmaSharedStride128,
nvcuda::wmma::mem_row_major);
#pragma unroll
for (int k = 0; k < kBlockSize; k += 16) {
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a, 16, 16, 16, __half,
nvcuda::wmma::row_major>
panel_a;
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b, 16, 16, 16, __half,
nvcuda::wmma::col_major>
panel_b;
nvcuda::wmma::load_matrix_sync(
panel_a, &factor_half_hi[tile_row << 4][k],
kHalfPanelStride128);
#pragma unroll
for (int element = 0; element < panel_a.num_elements; ++element) {
panel_a.x[element] = __hneg(panel_a.x[element]);
}
nvcuda::wmma::load_matrix_sync(
panel_b, &factor_half_hi[tile_col << 4][k],
kHalfPanelStride128);
nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
if constexpr (Compensated) {
nvcuda::wmma::load_matrix_sync(
panel_b, &factor_half_lo[tile_col << 4][k],
kHalfPanelStride128);
nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
nvcuda::wmma::load_matrix_sync(
panel_a, &factor_half_lo[tile_row << 4][k],
kHalfPanelStride128);
#pragma unroll
for (int element = 0; element < panel_a.num_elements; ++element) {
panel_a.x[element] = __hneg(panel_a.x[element]);
}
nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
nvcuda::wmma::load_matrix_sync(
panel_b, &factor_half_hi[tile_col << 4][k],
kHalfPanelStride128);
nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
}
}
nvcuda::wmma::store_matrix_sync(
&tile[row][col], updated, kWmmaSharedStride128,
nvcuda::wmma::mem_row_major);
}
}
__syncthreads();
}
// Read back only the factor. Upper entries bypass shared memory and are
// written directly as zero, preserving the exact dense-output contract.
#pragma unroll
for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
for (int col = lane; col <= row; col += 32) {
output[matrix_offset + row * kMatrixSize128 + col] = tile[row][col];
}
#pragma unroll
for (int col = row + 1 + lane; col < kMatrixSize128; col += 32) {
output[matrix_offset + row * kMatrixSize128 + col] = 0.0f;
}
}
}
} // namespace
torch::Tensor cholesky64_block32_hacc4safe_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(),
"cholesky64_block32_hacc4safe_cuda expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky64_block32_hacc4safe_cuda expects float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
input.size(1) == kMatrixSize &&
input.size(2) == kMatrixSize,
"cholesky64_block32_hacc4safe_cuda expects a nonempty "
"[batch, 64, 64] input");
TORCH_CHECK(input.is_contiguous(),
"cholesky64_block32_hacc4safe_cuda expects contiguous input");
const auto batch = input.size(0);
auto output = torch::empty_like(input);
cholesky64_block32_hacc4safe_kernel<<<batch, kThreads>>>(
input.data_ptr<float>(), output.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor cholesky32_warp_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "cholesky32_warp_cuda expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky32_warp_cuda expects float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
input.size(1) == kMatrixSize32 &&
input.size(2) == kMatrixSize32,
"cholesky32_warp_cuda expects a nonempty [batch, 32, 32] input");
TORCH_CHECK(input.is_contiguous(),
"cholesky32_warp_cuda expects contiguous input");
const auto batch = input.size(0);
TORCH_CHECK(batch <= std::numeric_limits<int>::max(),
"cholesky32_warp_cuda batch is too large");
const auto blocks =
(batch + kWarpsPerBlock32 - 1) / kWarpsPerBlock32;
auto output = torch::empty_like(input);
cholesky32_lower_blocktiles_kernel<<<blocks, kThreads32>>>(
input.data_ptr<float>(), output.data_ptr<float>(),
static_cast<int>(batch));
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor cholesky128_wmma_schur_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(),
"cholesky128_wmma_schur_cuda expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky128_wmma_schur_cuda expects float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
input.size(1) == kMatrixSize128 &&
input.size(2) == kMatrixSize128,
"cholesky128_wmma_schur_cuda expects a nonempty "
"[batch, 128, 128] input");
TORCH_CHECK(input.is_contiguous(),
"cholesky128_wmma_schur_cuda expects contiguous input");
static const cudaError_t attribute_status = cudaFuncSetAttribute(
cholesky128_wmma_schur_kernel<true>,
cudaFuncAttributeMaxDynamicSharedMemorySize, kWmmaSharedBytes128);
TORCH_CHECK(attribute_status == cudaSuccess,
"cholesky128 shared-memory opt-in failed: ",
cudaGetErrorString(attribute_status));
const auto batch = input.size(0);
auto output = torch::empty_like(input);
cholesky128_wmma_schur_kernel<true>
<<<batch, kThreads128, kWmmaSharedBytes128>>>(
input.data_ptr<float>(), output.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor cholesky128_wmma_schur_exact_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(),
"cholesky128_wmma_schur_exact_cuda expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky128_wmma_schur_exact_cuda expects float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) == 256 &&
input.size(1) == kMatrixSize128 &&
input.size(2) == kMatrixSize128,
"cholesky128_wmma_schur_exact_cuda expects "
"[256, 128, 128] input");
TORCH_CHECK(input.is_contiguous(),
"cholesky128_wmma_schur_exact_cuda expects contiguous input");
constexpr int shared_bytes =
kWmmaFloatSharedBytes128 + kHalfPanelBytes128;
static const cudaError_t attribute_status = cudaFuncSetAttribute(
cholesky128_wmma_schur_kernel<false>,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared_bytes);
TORCH_CHECK(attribute_status == cudaSuccess,
"cholesky128 exact shared-memory opt-in failed: ",
cudaGetErrorString(attribute_status));
auto output = torch::empty_like(input);
cholesky128_wmma_schur_kernel<false>
<<<256, kThreads128, shared_bytes>>>(
input.data_ptr<float>(), output.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor cholesky128_block32_warp16_lowerio_anybatch_cuda(
torch::Tensor input) {
TORCH_CHECK(input.is_cuda(),
"cholesky128_block32_warp16_lowerio_anybatch_cuda expects a "
"CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky128_block32_warp16_lowerio_anybatch_cuda expects "
"float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
input.size(1) == kMatrixSize128 &&
input.size(2) == kMatrixSize128,
"cholesky128_block32_warp16_lowerio_anybatch_cuda expects a "
"nonempty [batch, 128, 128] input");
TORCH_CHECK(input.is_contiguous(),
"cholesky128_block32_warp16_lowerio_anybatch_cuda expects "
"contiguous input");
static const cudaError_t attribute_status = cudaFuncSetAttribute(
cholesky128_block32_warp16_lowerio_anybatch_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, kFloatSharedBytes128);
TORCH_CHECK(attribute_status == cudaSuccess,
"cholesky128 FP32 shared-memory opt-in failed: ",
cudaGetErrorString(attribute_status));
const auto batch = input.size(0);
auto output = torch::empty_like(input);
cholesky128_block32_warp16_lowerio_anybatch_kernel
<<<batch, kThreads128, kFloatSharedBytes128>>>(
input.data_ptr<float>(), output.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
namespace {
constexpr int kCombinedSize256 = 256;
constexpr int kCombinedSize512 = 512;
inline void check_combined_solver(cusolverStatus_t status,
const char* operation) {
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, operation,
" failed with cuSOLVER status ", static_cast<int>(status));
}
inline void check_combined_blas(cublasStatus_t status,
const char* operation) {
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, operation,
" failed with cuBLAS status ", static_cast<int>(status));
}
void lt_packed_bf16_block_column(
const __nv_bfloat16* packed,
const __nv_bfloat16* packed_j,
const float* source_j,
float* destination_j,
int rows,
int cols,
int inner,
int output_ld,
const float* alpha,
const float* beta) {
cublasLtMatmulDescOpaque_t operation_storage{};
cublasLtMatmulDesc_t operation = &operation_storage;
check_combined_blas(
cublasLtMatmulDescInit(
operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
"cublasLtMatmulDescInit BF16 block-column update");
const cublasOperation_t transpose = CUBLAS_OP_T;
const cublasOperation_t identity = CUBLAS_OP_N;
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_TRANSA,
&transpose, sizeof(transpose)),
"cublasLt set transpose-A BF16 block-column update");
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_TRANSB,
&identity, sizeof(identity)),
"cublasLt set identity-B BF16 block-column update");
cublasLtMatrixLayoutOpaque_t a_storage{};
cublasLtMatrixLayoutOpaque_t b_storage{};
cublasLtMatrixLayoutOpaque_t c_storage{};
cublasLtMatrixLayoutOpaque_t d_storage{};
cublasLtMatrixLayout_t a_layout = &a_storage;
cublasLtMatrixLayout_t b_layout = &b_storage;
cublasLtMatrixLayout_t c_layout = &c_storage;
cublasLtMatrixLayout_t d_layout = &d_storage;
check_combined_blas(
cublasLtMatrixLayoutInit(a_layout, CUDA_R_16BF, inner, rows, inner),
"cublasLt A-layout init BF16 block-column update");
check_combined_blas(
cublasLtMatrixLayoutInit(b_layout, CUDA_R_16BF, inner, cols, inner),
"cublasLt B-layout init BF16 block-column update");
check_combined_blas(
cublasLtMatrixLayoutInit(
c_layout, CUDA_R_32F, rows, cols, output_ld),
"cublasLt C-layout init BF16 block-column update");
check_combined_blas(
cublasLtMatrixLayoutInit(
d_layout, CUDA_R_32F, rows, cols, output_ld),
"cublasLt D-layout init BF16 block-column update");
cublasLtMatmulPreferenceOpaque_t preference_storage{};
cublasLtMatmulPreference_t preference = &preference_storage;
check_combined_blas(
cublasLtMatmulPreferenceInit(preference),
"cublasLt preference init BF16 block-column update");
const size_t scratch_bytes = 0;
check_combined_blas(
cublasLtMatmulPreferenceSetAttribute(
preference, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&scratch_bytes, sizeof(scratch_bytes)),
"cublasLt scratch limit BF16 block-column update");
cublasLtMatmulHeuristicResult_t heuristic{};
int returned = 0;
cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();
check_combined_blas(
cublasLtMatmulAlgoGetHeuristic(
handle, operation, a_layout, b_layout,
c_layout, d_layout, preference, 1,
&heuristic, &returned),
"cublasLt heuristic BF16 block-column update");
TORCH_CHECK(
returned == 1 && heuristic.state == CUBLAS_STATUS_SUCCESS &&
heuristic.workspaceSize == 0,
"cublasLt found no scratch-free BF16 block-column algorithm");
check_combined_blas(
cublasLtMatmul(
handle, operation, alpha, packed, a_layout,
packed_j, b_layout, beta, source_j, c_layout,
destination_j, d_layout, &heuristic.algo, nullptr, 0, nullptr),
"cublasLtMatmul BF16 block-column update");
}
struct Mxfp8LtCacheKey {
int device;
int rows;
int cols;
int inner;
int output_ld;
cudaDataType_t a_type;
cudaDataType_t b_type;
cudaDataType_t c_type;
cudaDataType_t d_type;
cublasComputeType_t compute_type;
cudaDataType_t scale_type;
cublasOperation_t transa;
cublasOperation_t transb;
cublasLtMatmulMatrixScale_t a_scale_mode;
cublasLtMatmulMatrixScale_t b_scale_mode;
int8_t fast_accumulation;
int a_alignment;
int b_alignment;
int c_alignment;
int d_alignment;
int a_scale_alignment;
int b_scale_alignment;
int workspace_alignment;
int alpha_alignment;
int beta_alignment;
size_t scratch_bytes;
bool operator==(const Mxfp8LtCacheKey& other) const {
return device == other.device && rows == other.rows &&
cols == other.cols && inner == other.inner &&
output_ld == other.output_ld && a_type == other.a_type &&
b_type == other.b_type && c_type == other.c_type &&
d_type == other.d_type && compute_type == other.compute_type &&
scale_type == other.scale_type && transa == other.transa &&
transb == other.transb && a_scale_mode == other.a_scale_mode &&
b_scale_mode == other.b_scale_mode &&
fast_accumulation == other.fast_accumulation &&
a_alignment == other.a_alignment &&
b_alignment == other.b_alignment &&
c_alignment == other.c_alignment &&
d_alignment == other.d_alignment &&
a_scale_alignment == other.a_scale_alignment &&
b_scale_alignment == other.b_scale_alignment &&
workspace_alignment == other.workspace_alignment &&
alpha_alignment == other.alpha_alignment &&
beta_alignment == other.beta_alignment &&
scratch_bytes == other.scratch_bytes;
}
};
struct Mxfp8LtCacheEntry {
Mxfp8LtCacheKey key;
cublasLtMatrixLayoutOpaque_t a_storage{};
cublasLtMatrixLayoutOpaque_t b_storage{};
cublasLtMatrixLayoutOpaque_t c_storage{};
cublasLtMatrixLayoutOpaque_t d_storage{};
cublasLtMatmulPreferenceOpaque_t preference_storage{};
cublasLtMatmulHeuristicResult_t heuristic{};
cublasLtMatrixLayout_t a_layout() { return &a_storage; }
cublasLtMatrixLayout_t b_layout() { return &b_storage; }
cublasLtMatrixLayout_t c_layout() { return &c_storage; }
cublasLtMatrixLayout_t d_layout() { return &d_storage; }
cublasLtMatmulPreference_t preference() { return &preference_storage; }
};
constexpr size_t kMxfp8LtCacheCapacity = 64;
int capped_pointer_alignment(const void* pointer) {
const uintptr_t address = reinterpret_cast<uintptr_t>(pointer);
int alignment = 1;
while (alignment < 256 && (address & (alignment * 2 - 1)) == 0) {
alignment *= 2;
}
return alignment;
}
Mxfp8LtCacheEntry* find_mxfp8_lt_cache_entry(
std::vector<std::unique_ptr<Mxfp8LtCacheEntry>>& cache,
const Mxfp8LtCacheKey& key) {
for (auto& entry : cache) {
if (entry->key == key) {
return entry.get();
}
}
return nullptr;
}
Mxfp8LtCacheEntry* create_mxfp8_lt_cache_entry(
std::vector<std::unique_ptr<Mxfp8LtCacheEntry>>& cache,
const Mxfp8LtCacheKey& key,
cublasLtHandle_t handle,
cublasLtMatmulDesc_t operation) {
auto entry = std::make_unique<Mxfp8LtCacheEntry>();
entry->key = key;
check_combined_blas(
cublasLtMatrixLayoutInit(
entry->a_layout(), key.a_type, key.inner, key.rows, key.inner),
"cublasLt cached A-layout init MXFP8 block-column update");
check_combined_blas(
cublasLtMatrixLayoutInit(
entry->b_layout(), key.b_type, key.inner, key.cols, key.inner),
"cublasLt cached B-layout init MXFP8 block-column update");
check_combined_blas(
cublasLtMatrixLayoutInit(
entry->c_layout(), key.c_type, key.rows, key.cols, key.output_ld),
"cublasLt cached C-layout init MXFP8 block-column update");
check_combined_blas(
cublasLtMatrixLayoutInit(
entry->d_layout(), key.d_type, key.rows, key.cols, key.output_ld),
"cublasLt cached D-layout init MXFP8 block-column update");
check_combined_blas(
cublasLtMatmulPreferenceInit(entry->preference()),
"cublasLt cached preference init MXFP8 block-column update");
check_combined_blas(
cublasLtMatmulPreferenceSetAttribute(
entry->preference(), CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&key.scratch_bytes, sizeof(key.scratch_bytes)),
"cublasLt cached scratch limit MXFP8 block-column update");
int returned = 0;
check_combined_blas(
cublasLtMatmulAlgoGetHeuristic(
handle, operation, entry->a_layout(), entry->b_layout(),
entry->c_layout(), entry->d_layout(), entry->preference(), 1,
&entry->heuristic, &returned),
"cublasLt cached heuristic MXFP8 block-column update");
TORCH_CHECK(
returned == 1 && entry->heuristic.state == CUBLAS_STATUS_SUCCESS &&
entry->heuristic.workspaceSize <= key.scratch_bytes,
"cublasLt found no cached bounded-workspace MXFP8 algorithm");
if (cache.size() == kMxfp8LtCacheCapacity) {
cache.erase(cache.begin());
}
Mxfp8LtCacheEntry* const result = entry.get();
cache.push_back(std::move(entry));
return result;
}
void lt_packed_mxfp8_block_column(
const __nv_fp8_e4m3* primary,
const __nv_fp8_e4m3* primary_j,
const __nv_fp8_e8m0* primary_scale,
const __nv_fp8_e8m0* primary_scale_j,
const float* source_j,
float* destination_j,
int rows,
int cols,
int inner,
int output_ld,
int device,
std::vector<std::unique_ptr<Mxfp8LtCacheEntry>>& cache,
void* scratch,
size_t scratch_bytes,
const float* alpha,
const float* beta) {
cublasLtMatmulDescOpaque_t operation_storage{};
cublasLtMatmulDesc_t operation = &operation_storage;
check_combined_blas(
cublasLtMatmulDescInit(
operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
"cublasLtMatmulDescInit MXFP8 block-column update");
const cublasOperation_t transpose = CUBLAS_OP_T;
const cublasOperation_t identity = CUBLAS_OP_N;
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_TRANSA,
&transpose, sizeof(transpose)),
"cublasLt set transpose-A MXFP8 block-column update");
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_TRANSB,
&identity, sizeof(identity)),
"cublasLt set identity-B MXFP8 block-column update");
const cublasLtMatmulMatrixScale_t block_scale =
CUBLASLT_MATMUL_MATRIX_SCALE_VEC32_UE8M0;
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_A_SCALE_MODE,
&block_scale, sizeof(block_scale)),
"cublasLt set A block scale mode MXFP8 block-column update");
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_B_SCALE_MODE,
&block_scale, sizeof(block_scale)),
"cublasLt set B block scale mode MXFP8 block-column update");
const int8_t standard_accumulation = 0;
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_FAST_ACCUM,
&standard_accumulation, sizeof(standard_accumulation)),
"cublasLt set FP32 accumulation MXFP8 block-column update");
const __nv_fp8_e8m0* const a_scale_pointer = primary_scale;
const __nv_fp8_e8m0* const b_scale_pointer = primary_scale_j;
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_A_SCALE_POINTER,
&a_scale_pointer, sizeof(a_scale_pointer)),
"cublasLt set A scale pointer MXFP8 block-column update");
check_combined_blas(
cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_B_SCALE_POINTER,
&b_scale_pointer, sizeof(b_scale_pointer)),
"cublasLt set B scale pointer MXFP8 block-column update");
cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();
const Mxfp8LtCacheKey cache_key{
device,
rows,
cols,
inner,
output_ld,
CUDA_R_8F_E4M3,
CUDA_R_8F_E4M3,
CUDA_R_32F,
CUDA_R_32F,
CUBLAS_COMPUTE_32F,
CUDA_R_32F,
transpose,
identity,
block_scale,
block_scale,
standard_accumulation,
capped_pointer_alignment(primary),
capped_pointer_alignment(primary_j),
capped_pointer_alignment(source_j),
capped_pointer_alignment(destination_j),
capped_pointer_alignment(primary_scale),
capped_pointer_alignment(primary_scale_j),
capped_pointer_alignment(scratch),
capped_pointer_alignment(alpha),
capped_pointer_alignment(beta),
scratch_bytes};
Mxfp8LtCacheEntry* cache_entry =
find_mxfp8_lt_cache_entry(cache, cache_key);
if (cache_entry == nullptr) {
cache_entry = create_mxfp8_lt_cache_entry(
cache, cache_key, handle, operation);
}
cublasLtMatrixLayout_t a_layout = cache_entry->a_layout();
cublasLtMatrixLayout_t b_layout = cache_entry->b_layout();
cublasLtMatrixLayout_t c_layout = cache_entry->c_layout();
cublasLtMatrixLayout_t d_layout = cache_entry->d_layout();
const cublasLtMatmulAlgo_t* const algorithm =
&cache_entry->heuristic.algo;
check_combined_blas(
cublasLtMatmul(
handle, operation, alpha, primary, a_layout,
primary_j, b_layout, beta, source_j, c_layout,
destination_j, d_layout, algorithm, scratch,
scratch_bytes, nullptr),
"cublasLtMatmul MXFP8 primary block-column update");
}
struct CombinedSolverState {
cusolverDnHandle_t batched_handle = nullptr;
cusolverDnHandle_t large_handle = nullptr;
torch::Tensor pointers;
torch::Tensor batched_info;
torch::Tensor workspace;
torch::Tensor large_info;
int pointer_capacity = 0;
int workspace_elements = 0;
int device = -1;
std::mutex mutex;
~CombinedSolverState() {
if (batched_handle != nullptr) {
cusolverDnDestroy(batched_handle);
}
if (large_handle != nullptr) {
cusolverDnDestroy(large_handle);
}
}
void prepare_device(const torch::Tensor& input) {
const int input_device = input.get_device();
if (device == input_device && batched_handle != nullptr &&
large_handle != nullptr) {
return;
}
if (batched_handle != nullptr) {
cusolverDnDestroy(batched_handle);
batched_handle = nullptr;
}
if (large_handle != nullptr) {
cusolverDnDestroy(large_handle);
large_handle = nullptr;
}
check_combined_solver(cusolverDnCreate(&batched_handle),
"cusolverDnCreate batched");
check_combined_solver(cusolverDnCreate(&large_handle),
"cusolverDnCreate large");
#if CUSOLVER_VERSION >= 12000
check_combined_solver(
cusolverDnSetMathMode(
large_handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
"cusolverDnSetMathMode");
check_combined_solver(
cusolverDnSetEmulationStrategy(
large_handle, CUDA_EMULATION_STRATEGY_EAGER),
"cusolverDnSetEmulationStrategy");
#endif
pointers = torch::Tensor();
batched_info = torch::Tensor();
workspace = torch::Tensor();
large_info = torch::Tensor();
pointer_capacity = 0;
workspace_elements = 0;
device = input_device;
}
void prepare_batched(const torch::Tensor& input, int batch) {
prepare_device(input);
if (pointer_capacity < batch) {
pointers = torch::empty(
{batch}, input.options().dtype(torch::kInt64));
batched_info = torch::empty(
{batch}, input.options().dtype(torch::kInt32));
pointer_capacity = batch;
}
}
};
CombinedSolverState& combined_solver_state() {
static CombinedSolverState state;
return state;
}
struct ExplicitLowPrecisionState {
cusolverDnHandle_t solver_handle = nullptr;
cublasHandle_t blas_handle = nullptr;
torch::Tensor workspace;
torch::Tensor info;
torch::Tensor panel_bf16;
torch::Tensor mxfp8_primary_scale;
torch::Tensor mxfp8_diagonal_correction;
torch::Tensor mxfp8_scratch;
std::vector<std::unique_ptr<Mxfp8LtCacheEntry>> mxfp8_lt_cache;
std::vector<torch::Tensor> retired_buffers;
int workspace_elements = 0;
long long panel_capacity = 0;
int device = -1;
std::mutex mutex;
~ExplicitLowPrecisionState() {
if (solver_handle != nullptr) {
cusolverDnDestroy(solver_handle);
}
if (blas_handle != nullptr) {
cublasDestroy(blas_handle);
}
}
void prepare(const torch::Tensor& input) {
const int input_device = input.get_device();
if (device == input_device && solver_handle != nullptr &&
blas_handle != nullptr) {
return;
}
if (solver_handle != nullptr) {
cusolverDnDestroy(solver_handle);
}
if (blas_handle != nullptr) {
cublasDestroy(blas_handle);
}
check_combined_solver(cusolverDnCreate(&solver_handle),
"cusolverDnCreate explicit BF16");
check_combined_blas(cublasCreate(&blas_handle),
"cublasCreate explicit BF16");
workspace = torch::Tensor();
info = torch::Tensor();
panel_bf16 = torch::Tensor();
mxfp8_primary_scale = torch::Tensor();
mxfp8_diagonal_correction = torch::Tensor();
mxfp8_scratch = torch::Tensor();
mxfp8_lt_cache.clear();
retired_buffers.clear();
workspace_elements = 0;
panel_capacity = 0;
device = input_device;
}
};
ExplicitLowPrecisionState& explicit_low_precision_state() {
static ExplicitLowPrecisionState state;
return state;
}
struct TrueBatchedFast16BFState {
cusolverDnHandle_t solver = nullptr;
cublasHandle_t blas = nullptr;
torch::Tensor pointers;
torch::Tensor info;
std::vector<torch::Tensor> retired_buffers;
int pointer_capacity = 0;
int device = -1;
std::mutex mutex;
~TrueBatchedFast16BFState() {
if (solver != nullptr) {
cusolverDnDestroy(solver);
}
if (blas != nullptr) {
cublasDestroy(blas);
}
}
void prepare(const torch::Tensor& input, int batch) {
const int input_device = input.get_device();
if (device == -1) {
check_combined_solver(cusolverDnCreate(&solver),
"cusolverDnCreate true-batched fast16bf");
check_combined_blas(cublasCreate(&blas),
"cublasCreate true-batched fast16bf");
check_combined_blas(
cublasSetPointerMode(blas, CUBLAS_POINTER_MODE_HOST),
"cublasSetPointerMode true-batched fast16bf");
device = input_device;
} else {
TORCH_CHECK(device == input_device && solver != nullptr &&
blas != nullptr,
"true-batched fast16bf state cannot migrate between CUDA "
"devices");
}
if (pointer_capacity < batch) {
if (pointers.defined()) {
retired_buffers.push_back(pointers);
}
if (info.defined()) {
retired_buffers.push_back(info);
}
pointers = torch::empty(
{2, batch}, input.options().dtype(torch::kInt64));
info = torch::empty(
{batch}, input.options().dtype(torch::kInt32));
pointer_capacity = batch;
}
}
};
TrueBatchedFast16BFState& true_batched_fast16bf_state() {
static TrueBatchedFast16BFState state;
return state;
}
template <int MatrixSize>
__global__ void initialize_combined_batched_triangle(
const float* __restrict__ input, float* __restrict__ output,
float** pointers, int batch, long long total) {
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (long long matrix = linear; matrix < batch; matrix += stride) {
pointers[matrix] = output + matrix * matrix_elements;
}
for (; linear < total; linear += stride) {
const int within = static_cast<int>(linear % matrix_elements);
const int row = within / MatrixSize;
const int col = within - row * MatrixSize;
if (col <= row) {
output[linear] = input[linear];
}
}
}
template <int MatrixSize>
__global__ void fill_highbatch_panel_pointers(
float* base, float** diagonal_pointers, float** panel_pointers,
int diagonal_offset, int panel_offset, int batch) {
const int matrix =
static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x;
if (matrix < batch) {
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
float* const matrix_base =
base + static_cast<long long>(matrix) * matrix_elements;
diagonal_pointers[matrix] = matrix_base + diagonal_offset;
panel_pointers[matrix] = matrix_base + panel_offset;
}
}
template <int MatrixSize>
__global__ void copy_highbatch_lower_zero_upper(
const float* __restrict__ input, float* __restrict__ output,
long long total) {
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int within = static_cast<int>(linear % matrix_elements);
const int row = within / MatrixSize;
const int col = within - row * MatrixSize;
output[linear] = col <= row ? input[linear] : 0.0f;
}
}
template <int MatrixSize>
__global__ void copy_highbatch_lower_only(
const float* __restrict__ input, float* __restrict__ output,
long long total) {
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int within = static_cast<int>(linear % matrix_elements);
const int row = within / MatrixSize;
const int col = within - row * MatrixSize;
if (col <= row) {
output[linear] = input[linear];
}
}
}
template <int MatrixSize>
__global__ void copy_highbatch_lower_only_float4(
const float* __restrict__ input, float* __restrict__ output,
long long vector_total) {
static_assert((MatrixSize & 3) == 0,
"float4 triangular copy requires four-column rows");
constexpr int vectors_per_row = MatrixSize / 4;
constexpr long long vectors_per_matrix =
static_cast<long long>(MatrixSize) * vectors_per_row;
const auto* input4 = reinterpret_cast<const float4*>(input);
auto* output4 = reinterpret_cast<float4*>(output);
long long vector_linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; vector_linear < vector_total; vector_linear += stride) {
const int within =
static_cast<int>(vector_linear % vectors_per_matrix);
const int row = within / vectors_per_row;
const int col = (within - row * vectors_per_row) * 4;
if (col + 3 <= row) {
output4[vector_linear] = input4[vector_linear];
} else if (col <= row) {
const long long scalar = vector_linear * 4;
#pragma unroll
for (int lane = 0; lane < 4; ++lane) {
if (col + lane <= row) {
output[scalar + lane] = input[scalar + lane];
}
}
}
}
}
template <int MatrixSize>
__global__ void clear_combined_batched_row_major_upper(float* output,
long long total) {
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int within = static_cast<int>(linear % matrix_elements);
const int row = within / MatrixSize;
const int col = within - row * MatrixSize;
if (col > row) {
output[linear] = 0.0f;
}
}
}
template <int MatrixSize>
__global__ void clear_highbatch_upper_float4(float* output,
long long vector_total) {
static_assert((MatrixSize & 3) == 0,
"float4 triangular clear requires four-column rows");
constexpr int vectors_per_row = MatrixSize / 4;
constexpr long long vectors_per_matrix =
static_cast<long long>(MatrixSize) * vectors_per_row;
auto* output4 = reinterpret_cast<float4*>(output);
long long vector_linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; vector_linear < vector_total; vector_linear += stride) {
const int within =
static_cast<int>(vector_linear % vectors_per_matrix);
const int row = within / vectors_per_row;
const int col = (within - row * vectors_per_row) * 4;
if (col > row) {
output4[vector_linear] = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
} else if (col + 3 > row) {
const long long scalar = vector_linear * 4;
#pragma unroll
for (int lane = 0; lane < 4; ++lane) {
if (col + lane > row) {
output[scalar + lane] = 0.0f;
}
}
}
}
}
__global__ void clear_combined_large_upper(float* output, int n) {
const long long total = static_cast<long long>(n) * n;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int row = static_cast<int>(linear / n);
const int col =
static_cast<int>(linear - static_cast<long long>(row) * n);
if (col > row) {
output[linear] = 0.0f;
}
}
}
__global__ void clear_explicit_batch_upper(float* output, int n,
long long total) {
const long long matrix_elements = static_cast<long long>(n) * n;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const long long within = linear % matrix_elements;
const int row = static_cast<int>(within / n);
const int col = static_cast<int>(
within - static_cast<long long>(row) * n);
if (col > row) {
output[linear] = 0.0f;
}
}
}
__global__ void clear_explicit_large_upper_vectorized(float* output, int n) {
const int row = static_cast<int>(blockIdx.x);
const int upper_begin = row + 1;
const int vector_begin = (upper_begin + 3) & ~3;
const int prefix_end = vector_begin < n ? vector_begin : n;
const long long row_offset = static_cast<long long>(row) * n;
float* const output_row = output + row_offset;
for (int col = upper_begin + static_cast<int>(threadIdx.x);
col < prefix_end; col += static_cast<int>(blockDim.x)) {
output_row[col] = 0.0f;
}
const int vector_count = (n - vector_begin) >> 2;
auto* const output_vectors =
reinterpret_cast<float4*>(output_row + vector_begin);
for (int vector = static_cast<int>(threadIdx.x);
vector < vector_count;
vector += static_cast<int>(blockDim.x)) {
output_vectors[vector] = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
const int tail_begin = vector_begin + (vector_count << 2);
for (int col = tail_begin + static_cast<int>(threadIdx.x); col < n;
col += static_cast<int>(blockDim.x)) {
output_row[col] = 0.0f;
}
}
__global__ void initialize_explicit_large_first_panel_lower(
const float* input, float* output, int n, int initialized_cols) {
const int row = static_cast<int>(blockIdx.x);
const long long row_offset = static_cast<long long>(row) * n;
for (int col = static_cast<int>(threadIdx.x) * 4;
col < initialized_cols;
col += static_cast<int>(blockDim.x) * 4) {
const long long first = row_offset + col;
if (col + 3 < initialized_cols && col + 3 <= row) {
reinterpret_cast<float4*>(output + first)[0] =
reinterpret_cast<const float4*>(input + first)[0];
} else {
#pragma unroll
for (int lane = 0; lane < 4; ++lane) {
const int scalar_col = col + lane;
if (scalar_col < initialized_cols && scalar_col <= row) {
output[row_offset + scalar_col] = input[row_offset + scalar_col];
}
}
}
}
}
__global__ void pack_panel_bf16(const float* panel, __nv_bfloat16* packed,
int rows, int cols, int source_ld) {
const long long total = static_cast<long long>(rows) * cols;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int row = static_cast<int>(linear % rows);
const int col = static_cast<int>(linear / rows);
packed[linear] = __float2bfloat16_rn(
panel[row + static_cast<long long>(col) * source_ld]);
}
}
__global__ void pack_outer_panel_bf16(
const float* panel, __nv_bfloat16* packed, int rows, int cols,
int source_ld, int packed_ld, int packed_row, int packed_col) {
const long long total = static_cast<long long>(rows) * cols;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int row = static_cast<int>(linear % rows);
const int col = static_cast<int>(linear / rows);
packed[packed_row + row +
static_cast<long long>(packed_col + col) * packed_ld] =
__float2bfloat16_rn(
panel[row + static_cast<long long>(col) * source_ld]);
}
}
union alignas(16) PackedOuterBf16x8 {
struct {
__nv_bfloat162 pair0;
__nv_bfloat162 pair1;
__nv_bfloat162 pair2;
__nv_bfloat162 pair3;
} pairs;
uint4 vector;
};
static_assert(sizeof(PackedOuterBf16x8) == 16,
"eight packed BF16 values must occupy 16 bytes");
__global__ void pack_outer_panel_bf16_vector8(
const float* panel, __nv_bfloat16* packed, int rows, int cols,
int source_ld, int packed_ld, int packed_row, int packed_col) {
const int row_groups = rows / 8;
const long long total = static_cast<long long>(row_groups) * cols;
long long linear =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride =
static_cast<long long>(gridDim.x) * blockDim.x;
for (; linear < total; linear += stride) {
const int row_group = static_cast<int>(linear % row_groups);
const int col = static_cast<int>(linear / row_groups);
const int row = row_group * 8;
const float* const source =
panel + row + static_cast<long long>(col) * source_ld;
const float4 first = reinterpret_cast<const float4*>(source)[0];
const float4 second = reinterpret_cast<const float4*>(source)[1];
PackedOuterBf16x8 converted;
converted.pairs.pair0 = __floats2bfloat162_rn(first.x, first.y);
converted.pairs.pair1 = __floats2bfloat162_rn(first.z, first.w);
converted.pairs.pair2 = __floats2bfloat162_rn(second.x, second.y);
converted.pairs.pair3 = __floats2bfloat162_rn(second.z, second.w);
auto* const destination = reinterpret_cast<uint4*>(
packed + packed_row + row +
static_cast<long long>(packed_col + col) * packed_ld);
destination[0] = converted.vector;
}
}
__device__ __forceinline__ int mxfp8_e8m0_exponent(float maximum) {
const unsigned bits = __float_as_uint(maximum);
const int exponent = static_cast<int>((bits >> 23) & 255U);
const unsigned mantissa = bits & 0x7FFFFFU;
const int rounded = exponent + (mantissa > 0x600000U) - 8;
return max(1, min(253, rounded));
}
__device__ __forceinline__ long long mxfp8_scale_offset(
int group, int column, int inner) {
constexpr int scale_block_elements = 512;
const int scaled_rows = (inner + 31) / 32;
const int scaled_ld = (scaled_rows + 3) & ~3;
const long long block =
static_cast<long long>(group / 4) +
static_cast<long long>(column / 128) * (scaled_ld / 4);
return block * scale_block_elements +
static_cast<long long>(column & 31) * 16 +
static_cast<long long>((column / 32) & 3) * 4 +
(group & 3);
}
inline long long mxfp8_scale_column_base(int inner, int column) {
const int scaled_rows = (inner + 31) / 32;
const int scaled_ld = (scaled_rows + 3) & ~3;
return static_cast<long long>(column / 128) *
(scaled_ld / 4) * 512;
}
__global__ void clear_mxfp8_diagonal_correction(
float* correction, int count) {
for (int index = static_cast<int>(blockIdx.x) * blockDim.x +
static_cast<int>(threadIdx.x);
index < count;
index += static_cast<int>(gridDim.x) * blockDim.x) {
correction[index] = 0.0f;
}
}
__global__ void pack_outer_far_mxfp8_block32(
const float* panel,
__nv_fp8_e4m3* primary,
__nv_fp8_e8m0* primary_scale,
float* diagonal_correction,
int inner,
int columns,
int source_ld) {
const int column = static_cast<int>(blockIdx.x);
if (column >= columns) {
return;
}
const int lane = static_cast<int>(threadIdx.x) & 31;
const int warp = static_cast<int>(threadIdx.x) >> 5;
const int warps = static_cast<int>(blockDim.x) >> 5;
const int groups = inner / 32;
for (int group = warp; group < groups; group += warps) {
const int row = group * 32 + lane;
const long long packed_offset =
row + static_cast<long long>(column) * inner;
const float value =
panel[row + static_cast<long long>(column) * source_ld];
float maximum = fabsf(value);
#pragma unroll
for (int delta = 16; delta > 0; delta >>= 1) {
maximum = fmaxf(
maximum, __shfl_down_sync(0xFFFFFFFFU, maximum, delta));
}
maximum = __shfl_sync(0xFFFFFFFFU, maximum, 0);
const int primary_exponent = mxfp8_e8m0_exponent(maximum);
const float primary_value_scale =
__uint_as_float(static_cast<unsigned>(primary_exponent) << 23);
const float primary_inverse = __uint_as_float(
static_cast<unsigned>(254 - primary_exponent) << 23);
const __nv_fp8_e4m3 primary_value(value * primary_inverse);
primary[packed_offset] = primary_value;
const float primary_reconstruction =
static_cast<float>(primary_value) * primary_value_scale;
if (lane == 0) {
const long long scale_offset =
mxfp8_scale_offset(group, column, inner);
__nv_fp8_e8m0 primary_scale_value;
primary_scale_value.__x =
static_cast<__nv_fp8_storage_t>(primary_exponent);
primary_scale[scale_offset] = primary_scale_value;
}
float correction =
value * value - primary_reconstruction * primary_reconstruction;
#pragma unroll
for (int delta = 16; delta > 0; delta >>= 1) {
correction +=
__shfl_down_sync(0xFFFFFFFFU, correction, delta);
}
if (lane == 0) {
atomicAdd(diagonal_correction + column, correction);
}
}
}
__global__ void add_mxfp8_diagonal_correction(
float* diagonal, const float* correction, int count, int ld) {
for (int index = static_cast<int>(blockIdx.x) * blockDim.x +
static_cast<int>(threadIdx.x);
index < count;
index += static_cast<int>(gridDim.x) * blockDim.x) {
diagonal[static_cast<long long>(index) * (ld + 1)] +=
correction[index];
}
}
__global__ void copy_combined_large_lower(const float* input, float* output,
int n) {
const int row = static_cast<int>(blockIdx.x);
const long long row_offset = static_cast<long long>(row) * n;
for (int col = static_cast<int>(threadIdx.x) * 4; col <= row;
col += static_cast<int>(blockDim.x) * 4) {
const long long first = row_offset + col;
if ((n & 3) == 0 && col + 3 <= row) {
reinterpret_cast<float4*>(output + first)[0] =
reinterpret_cast<const float4*>(input + first)[0];
} else {
#pragma unroll
for (int lane = 0; lane < 4; ++lane) {
const int scalar_col = col + lane;
if (scalar_col <= row) {
output[row_offset + scalar_col] = input[row_offset + scalar_col];
}
}
}
}
}
template <int MatrixSize>
torch::Tensor cholesky_combined_batched_impl(torch::Tensor input) {
c10::cuda::CUDAGuard device_guard(input.device());
const int batch = static_cast<int>(input.size(0));
auto output = torch::empty_like(input);
auto& state = combined_solver_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare_batched(input, batch);
constexpr int threads = 256;
auto pointer_data = reinterpret_cast<float**>(
state.pointers.data_ptr<int64_t>());
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
const long long total = static_cast<long long>(batch) * matrix_elements;
const int blocks = static_cast<int>(
std::min<long long>(4096, (total + threads - 1) / threads));
initialize_combined_batched_triangle<MatrixSize><<<blocks, threads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), pointer_data, batch,
total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
check_combined_solver(
cusolverDnSpotrfBatched(
state.batched_handle, CUBLAS_FILL_MODE_UPPER, MatrixSize,
pointer_data, MatrixSize, state.batched_info.data_ptr<int>(), batch),
"cusolverDnSpotrfBatched");
clear_combined_batched_row_major_upper<MatrixSize><<<blocks, threads>>>(
output.data_ptr<float>(), total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
void check_combined_batched_input(const torch::Tensor& input, int n,
const char* operation) {
TORCH_CHECK(input.is_cuda(), operation, " expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
operation, " expects float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
input.size(1) == n && input.size(2) == n,
operation, " received an invalid shape");
TORCH_CHECK(input.is_contiguous(), operation, " expects contiguous input");
TORCH_CHECK(input.size(0) <= std::numeric_limits<int>::max(),
operation, " batch is too large");
static_assert(sizeof(void*) == sizeof(int64_t),
"pointer storage requires 64-bit pointers");
}
struct TrueBatchedBlockedState {
cusolverDnHandle_t solver = nullptr;
cublasHandle_t blas = nullptr;
torch::Tensor pointers;
torch::Tensor info;
int pointer_capacity = 0;
int device = -1;
std::mutex mutex;
~TrueBatchedBlockedState() {
if (solver != nullptr) {
cusolverDnDestroy(solver);
}
if (blas != nullptr) {
cublasDestroy(blas);
}
}
void prepare(const torch::Tensor& input, int batch) {
const int input_device = input.get_device();
if (device == -1) {
check_combined_solver(cusolverDnCreate(&solver),
"cusolverDnCreate true-batched blocked");
check_combined_blas(cublasCreate(&blas),
"cublasCreate true-batched blocked");
check_combined_blas(
cublasSetPointerMode(blas, CUBLAS_POINTER_MODE_HOST),
"cublasSetPointerMode true-batched blocked");
device = input_device;
} else {
TORCH_CHECK(device == input_device && solver != nullptr &&
blas != nullptr,
"true-batched blocked state cannot migrate between CUDA "
"devices");
}
if (pointer_capacity < batch) {
pointers = torch::empty(
{2, batch}, input.options().dtype(torch::kInt64));
info = torch::empty({batch}, input.options().dtype(torch::kInt32));
pointer_capacity = batch;
}
}
};
TrueBatchedBlockedState& true_batched_blocked_state() {
static TrueBatchedBlockedState state;
return state;
}
constexpr int kRegisterPanel64 = 64;
constexpr int kRegisterPanelStride64 = 65;
constexpr int kRegisterPanelThreads64 = 128;
constexpr int kRegisterPanelSharedBytes64 =
kRegisterPanel64 * kRegisterPanelStride64 * sizeof(float);
static_assert(kRegisterPanelSharedBytes64 < 17 * 1024,
"register panel shared storage exceeds its bound");
constexpr int kRegisterDiagonal128 = 128;
constexpr int kRegisterDiagonalRows128 = 128;
constexpr int kRegisterDiagonalStride128 = 65;
constexpr int kRegisterDiagonalThreads128 = 128;
constexpr int kRegisterDiagonalSharedBytes128 =
kRegisterDiagonalRows128 * kRegisterDiagonalStride128 * sizeof(float);
static_assert(kRegisterDiagonalSharedBytes128 < 33 * 1024,
"register diagonal shared storage exceeds its bound");
template <int Base>
__device__ __forceinline__ void factor_register_panel64_stage(
float (*tile)[kRegisterPanelStride64], int tid, int lane, int warp) {
if (warp == 0) {
N64RegisterRow values;
if (lane < kNestedBlockSize64) {
n64_load_register_row(values, tile, Base + lane, Base, N64AllIndices{});
} else {
values = {};
}
n64_factor_register_block(values, lane, N64AllIndices{});
if (lane < kNestedBlockSize64) {
n64_store_register_row(values, tile, Base + lane, Base,
N64AllIndices{});
}
}
__syncthreads();
constexpr int next = Base + kNestedBlockSize64;
for (int row = next + tid; row < kRegisterPanel64;
row += kRegisterPanelThreads64) {
N64RegisterRow values;
n64_load_register_row(values, tile, row, Base, N64AllIndices{});
n64_solve_outer_panel_first(values, tile, Base, N64AllIndices{});
n64_store_register_row(values, tile, row, Base, N64AllIndices{});
}
__syncthreads();
constexpr int remaining = kRegisterPanel64 - next;
if constexpr (remaining > 0) {
for (int linear = tid; linear < remaining * remaining;
linear += kRegisterPanelThreads64) {
const int row = next + linear / remaining;
const int col = next + linear % remaining;
if (row >= col) {
float value = tile[row][col];
#pragma unroll
for (int k = Base; k < next; ++k) {
value = fmaf(-tile[row][k], tile[col][k], value);
}
tile[row][col] = value;
}
}
}
__syncthreads();
}
template <int Base>
__device__ __forceinline__ void factor_register_diagonal128_first_stage(
float (*tile)[kRegisterDiagonalStride128], int tid, int lane, int warp) {
if (warp == 0) {
N64RegisterRow values;
if (lane < kNestedBlockSize64) {
n64_load_register_row(values, tile, Base + lane, Base, N64AllIndices{});
} else {
values = {};
}
n64_factor_register_block(values, lane, N64AllIndices{});
if (lane < kNestedBlockSize64) {
n64_store_register_row(values, tile, Base + lane, Base,
N64AllIndices{});
}
}
__syncthreads();
constexpr int next = Base + kNestedBlockSize64;
for (int row = next + tid; row < kRegisterDiagonalRows128;
row += kRegisterDiagonalThreads128) {
N64RegisterRow values;
n64_load_register_row(values, tile, row, Base, N64AllIndices{});
n64_solve_outer_panel_first(values, tile, Base, N64AllIndices{});
n64_store_register_row(values, tile, row, Base, N64AllIndices{});
}
__syncthreads();
constexpr int remaining_columns = kRegisterPanel64 - next;
constexpr int remaining_rows = kRegisterDiagonalRows128 - next;
if constexpr (remaining_columns > 0) {
for (int linear = tid;
linear < remaining_rows * remaining_columns;
linear += kRegisterDiagonalThreads128) {
const int row = next + linear / remaining_columns;
const int col = next + linear % remaining_columns;
if (row >= col) {
float value = tile[row][col];
#pragma unroll
for (int k = Base; k < next; ++k) {
value = fmaf(-tile[row][k], tile[col][k], value);
}
tile[row][col] = value;
}
}
}
__syncthreads();
}
template <int MatrixSize>
__global__ __launch_bounds__(kRegisterDiagonalThreads128)
void factor_register_diagonal128_kernel(
float* matrices, float** diagonal_pointers, float** panel_pointers,
int offset) {
static_assert(MatrixSize == 512 || MatrixSize == 1024 ||
MatrixSize == 2048,
"register diagonal kernel supports the high-batch paths");
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
float* const matrix_base =
matrices + static_cast<long long>(matrix) * matrix_elements;
__shared__ float tile[kRegisterDiagonalRows128]
[kRegisterDiagonalStride128];
for (int linear = tid;
linear < kRegisterDiagonalRows128 * kRegisterPanel64;
linear += kRegisterDiagonalThreads128) {
const int row = linear >> 6;
const int col = linear & 63;
tile[row][col] =
row >= col
? matrix_base[(offset + row) * MatrixSize + offset + col]
: 0.0f;
}
__syncthreads();
factor_register_diagonal128_first_stage<0>(tile, tid, lane, warp);
factor_register_diagonal128_first_stage<16>(tile, tid, lane, warp);
factor_register_diagonal128_first_stage<32>(tile, tid, lane, warp);
factor_register_diagonal128_first_stage<48>(tile, tid, lane, warp);
for (int linear = tid;
linear < kRegisterDiagonalRows128 * kRegisterPanel64;
linear += kRegisterDiagonalThreads128) {
const int row = linear >> 6;
const int col = linear & 63;
if (row >= col) {
matrix_base[(offset + row) * MatrixSize + offset + col] =
tile[row][col];
}
}
for (int linear = tid;
linear < kRegisterPanel64 * kRegisterPanel64;
linear += kRegisterDiagonalThreads128) {
const int row = linear >> 6;
const int col = linear & 63;
float value = 0.0f;
if (row >= col) {
value = matrix_base[(offset + kRegisterPanel64 + row) * MatrixSize +
offset + kRegisterPanel64 + col];
#pragma unroll
for (int k = 0; k < kRegisterPanel64; ++k) {
value = fmaf(-tile[kRegisterPanel64 + row][k],
tile[kRegisterPanel64 + col][k], value);
}
}
tile[row][col] = value;
}
__syncthreads();
factor_register_panel64_stage<0>(tile, tid, lane, warp);
factor_register_panel64_stage<16>(tile, tid, lane, warp);
factor_register_panel64_stage<32>(tile, tid, lane, warp);
factor_register_panel64_stage<48>(tile, tid, lane, warp);
for (int linear = tid;
linear < kRegisterPanel64 * kRegisterPanel64;
linear += kRegisterDiagonalThreads128) {
const int row = linear >> 6;
const int col = linear & 63;
if (row >= col) {
matrix_base[(offset + kRegisterPanel64 + row) * MatrixSize +
offset + kRegisterPanel64 + col] = tile[row][col];
}
}
if (tid == 0) {
float* const diagonal = matrix_base + offset * (MatrixSize + 1);
diagonal_pointers[matrix] = diagonal;
panel_pointers[matrix] =
offset + kRegisterDiagonal128 < MatrixSize
? matrix_base + (offset + kRegisterDiagonal128) * MatrixSize +
offset
: diagonal;
}
}
template <int MatrixSize>
__global__ __launch_bounds__(kRegisterPanelThreads64)
void factor_register_panel64_kernel(
float* matrices, float** diagonal_pointers, float** panel_pointers,
int offset) {
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
const int matrix = static_cast<int>(blockIdx.x);
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
float* const matrix_base =
matrices + static_cast<long long>(matrix) * matrix_elements;
__shared__ float tile[kRegisterPanel64][kRegisterPanelStride64];
for (int linear = tid;
linear < kRegisterPanel64 * kRegisterPanel64;
linear += kRegisterPanelThreads64) {
const int row = linear >> 6;
const int col = linear & 63;
tile[row][col] = row >= col
? matrix_base[(offset + row) * MatrixSize +
offset + col]
: 0.0f;
}
__syncthreads();
factor_register_panel64_stage<0>(tile, tid, lane, warp);
factor_register_panel64_stage<16>(tile, tid, lane, warp);
factor_register_panel64_stage<32>(tile, tid, lane, warp);
factor_register_panel64_stage<48>(tile, tid, lane, warp);
for (int linear = tid;
linear < kRegisterPanel64 * kRegisterPanel64;
linear += kRegisterPanelThreads64) {
const int row = linear >> 6;
const int col = linear & 63;
if (row >= col) {
matrix_base[(offset + row) * MatrixSize + offset + col] =
tile[row][col];
}
}
if (tid == 0) {
float* const diagonal =
matrix_base + offset * (MatrixSize + 1);
diagonal_pointers[matrix] = diagonal;
panel_pointers[matrix] =
offset + kRegisterPanel64 < MatrixSize
? matrix_base + (offset + kRegisterPanel64) * MatrixSize + offset
: diagonal;
}
}
torch::Tensor cholesky256_batch64_register_panel64_impl(torch::Tensor input) {
constexpr int matrix_size = 256;
constexpr int batch = 64;
constexpr long long matrix_elements =
static_cast<long long>(matrix_size) * matrix_size;
c10::cuda::CUDAGuard device_guard(input.device());
auto output = input.clone();
auto& state = true_batched_blocked_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare(input, batch);
auto diagonal_pointers = reinterpret_cast<float**>(
state.pointers.data_ptr<int64_t>());
auto panel_pointers = diagonal_pointers + batch;
const float one = 1.0f;
const float minus_one = -1.0f;
float* const base = output.data_ptr<float>();
for (int offset = 0; offset < matrix_size; offset += kRegisterPanel64) {
const int trailing = matrix_size - offset - kRegisterPanel64;
factor_register_panel64_kernel<matrix_size>
<<<batch, kRegisterPanelThreads64>>>(
base, diagonal_pointers, panel_pointers, offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if (trailing == 0) {
continue;
}
check_combined_blas(
cublasStrsmBatched(
state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, kRegisterPanel64, trailing,
&one,
reinterpret_cast<const float* const*>(diagonal_pointers),
matrix_size, panel_pointers, matrix_size, batch),
"cublasStrsmBatched register panel");
float* const panel =
base + offset + (offset + kRegisterPanel64) * matrix_size;
float* const trailing_matrix =
base + (offset + kRegisterPanel64) * (matrix_size + 1);
check_combined_blas(
cublasSgemmStridedBatched(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, trailing, trailing,
kRegisterPanel64, &minus_one, panel, matrix_size, matrix_elements,
panel, matrix_size, matrix_elements, &one, trailing_matrix,
matrix_size, matrix_elements, batch),
"cublasSgemmStridedBatched register trailing update");
}
constexpr int threads = 256;
constexpr long long total = batch * matrix_elements;
constexpr int blocks = 4096;
clear_combined_batched_row_major_upper<matrix_size>
<<<blocks, threads>>>(base, total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor cholesky512_batch16_register_panel64_impl(torch::Tensor input) {
constexpr int matrix_size = 512;
constexpr int batch = 16;
constexpr long long matrix_elements =
static_cast<long long>(matrix_size) * matrix_size;
c10::cuda::CUDAGuard device_guard(input.device());
auto output = input.clone();
auto& state = true_batched_blocked_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare(input, batch);
auto diagonal_pointers = reinterpret_cast<float**>(
state.pointers.data_ptr<int64_t>());
auto panel_pointers = diagonal_pointers + state.pointer_capacity;
const float one = 1.0f;
const float minus_one = -1.0f;
float* const base = output.data_ptr<float>();
for (int offset = 0; offset < matrix_size; offset += kRegisterPanel64) {
const int trailing = matrix_size - offset - kRegisterPanel64;
factor_register_panel64_kernel<matrix_size>
<<<batch, kRegisterPanelThreads64>>>(
base, diagonal_pointers, panel_pointers, offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if (trailing == 0) {
continue;
}
check_combined_blas(
cublasStrsmBatched(
state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, kRegisterPanel64, trailing,
&one,
reinterpret_cast<const float* const*>(diagonal_pointers),
matrix_size, panel_pointers, matrix_size, batch),
"cublasStrsmBatched n512 batch16 register panel");
float* const panel =
base + offset + (offset + kRegisterPanel64) * matrix_size;
float* const trailing_matrix =
base + (offset + kRegisterPanel64) * (matrix_size + 1);
check_combined_blas(
cublasSgemmStridedBatched(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, trailing, trailing,
kRegisterPanel64, &minus_one, panel, matrix_size, matrix_elements,
panel, matrix_size, matrix_elements, &one, trailing_matrix,
matrix_size, matrix_elements, batch),
"cublasSgemmStridedBatched n512 batch16 register trailing update");
}
constexpr int threads = 256;
constexpr long long total = batch * matrix_elements;
constexpr int blocks = 4096;
clear_combined_batched_row_major_upper<matrix_size>
<<<blocks, threads>>>(base, total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
template <int MatrixSize, int Batch, int PanelSize, int ColumnTile = 0>
torch::Tensor cholesky_true_batched_blocked_impl(
torch::Tensor input, const char* operation) {
static_assert(MatrixSize % PanelSize == 0,
"panel size must divide matrix size");
static_assert(ColumnTile == 0 ||
(ColumnTile > 0 && ColumnTile <= MatrixSize),
"column tile must be disabled or fit the matrix");
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
check_combined_batched_input(input, MatrixSize, operation);
TORCH_CHECK(input.size(0) == Batch, operation, " received an invalid batch");
c10::cuda::CUDAGuard device_guard(input.device());
auto output = torch::empty_like(input);
auto& state = true_batched_blocked_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare(input, Batch);
constexpr int threads = 256;
constexpr long long total = static_cast<long long>(Batch) * matrix_elements;
const int matrix_blocks = static_cast<int>(
std::min<long long>(4096, (total + threads - 1) / threads));
auto diagonal_pointers = reinterpret_cast<float**>(
state.pointers.data_ptr<int64_t>());
auto panel_pointers = diagonal_pointers + Batch;
if constexpr ((MatrixSize == 512 && Batch == 640) ||
(MatrixSize == 1024 && Batch == 60)) {
copy_highbatch_lower_only<MatrixSize>
<<<matrix_blocks, threads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), total);
} else {
copy_highbatch_lower_zero_upper<MatrixSize>
<<<matrix_blocks, threads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), total);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
constexpr int pointer_blocks = (Batch + threads - 1) / threads;
const float one = 1.0f;
const float minus_one = -1.0f;
float* const base = output.data_ptr<float>();
for (int offset = 0; offset < MatrixSize; offset += PanelSize) {
const int trailing = MatrixSize - offset - PanelSize;
const int diagonal_offset = offset + offset * MatrixSize;
const int panel_offset =
trailing > 0 ? offset + (offset + PanelSize) * MatrixSize
: diagonal_offset;
fill_highbatch_panel_pointers<MatrixSize>
<<<pointer_blocks, threads>>>(
base, diagonal_pointers, panel_pointers, diagonal_offset,
panel_offset, Batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
check_combined_solver(
cusolverDnSpotrfBatched(
state.solver, CUBLAS_FILL_MODE_UPPER, PanelSize,
diagonal_pointers, MatrixSize, state.info.data_ptr<int>(), Batch),
"cusolverDnSpotrfBatched true-batched blocked diagonal");
if (trailing == 0) {
continue;
}
check_combined_blas(
cublasStrsmBatched(
state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, PanelSize, trailing, &one,
reinterpret_cast<const float* const*>(diagonal_pointers),
MatrixSize, panel_pointers, MatrixSize, Batch),
"cublasStrsmBatched true-batched blocked panel");
float* const panel = base + panel_offset;
float* const trailing_matrix =
base + (offset + PanelSize) * (MatrixSize + 1);
if constexpr (ColumnTile > 0) {
for (int column = 0; column < trailing; column += ColumnTile) {
const int width = std::min(ColumnTile, trailing - column);
const int rows = trailing - column;
float* const panel_column =
panel + static_cast<long long>(column) * MatrixSize;
float* const trailing_column =
trailing_matrix + static_cast<long long>(column) *
(MatrixSize + 1);
check_combined_blas(
cublasGemmStridedBatchedEx(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, width, rows,
PanelSize, &minus_one, panel_column, CUDA_R_32F, MatrixSize,
matrix_elements, panel_column, CUDA_R_32F, MatrixSize,
matrix_elements, &one, trailing_column, CUDA_R_32F,
MatrixSize, matrix_elements, Batch,
CUBLAS_COMPUTE_32F_FAST_TF32, CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx true-batched blocked block-column");
}
} else {
check_combined_blas(
cublasGemmStridedBatchedEx(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, trailing, trailing,
PanelSize, &minus_one, panel, CUDA_R_32F, MatrixSize,
matrix_elements, panel, CUDA_R_32F, MatrixSize, matrix_elements,
&one, trailing_matrix, CUDA_R_32F, MatrixSize, matrix_elements,
Batch, CUBLAS_COMPUTE_32F_FAST_TF32, CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx true-batched blocked trailing");
}
}
clear_combined_batched_row_major_upper<MatrixSize>
<<<matrix_blocks, threads>>>(base, total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
template <int MatrixSize, int Batch, int InnerPanel, int OuterPanel,
bool FastBf16, bool CloneInput = false>
torch::Tensor cholesky_true_batched_two_level_impl(
torch::Tensor input, const char* operation) {
static_assert(InnerPanel > 0 && OuterPanel >= InnerPanel,
"invalid two-level panel sizes");
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
check_combined_batched_input(input, MatrixSize, operation);
TORCH_CHECK(input.size(0) == Batch, operation, " received an invalid batch");
c10::cuda::CUDAGuard device_guard(input.device());
torch::Tensor output;
if constexpr (CloneInput) {
output = input.clone();
} else {
output = torch::empty_like(input);
}
auto& state = true_batched_blocked_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare(input, Batch);
constexpr int threads = 256;
constexpr long long total = static_cast<long long>(Batch) * matrix_elements;
const int matrix_blocks = static_cast<int>(
std::min<long long>(4096, (total + threads - 1) / threads));
if constexpr (!CloneInput) {
copy_highbatch_lower_zero_upper<MatrixSize>
<<<matrix_blocks, threads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
auto diagonal_pointers = reinterpret_cast<float**>(
state.pointers.data_ptr<int64_t>());
auto panel_pointers = diagonal_pointers + Batch;
constexpr int pointer_blocks = (Batch + threads - 1) / threads;
const float one = 1.0f;
const float minus_one = -1.0f;
float* const base = output.data_ptr<float>();
constexpr cublasComputeType_t update_compute =
FastBf16 ? CUBLAS_COMPUTE_32F_FAST_16BF
: CUBLAS_COMPUTE_32F_FAST_TF32;
for (int outer = 0; outer < MatrixSize; outer += OuterPanel) {
const int outer_end = std::min(MatrixSize, outer + OuterPanel);
const int outer_width = outer_end - outer;
for (int offset = outer; offset < outer_end; offset += InnerPanel) {
const int jb = std::min(InnerPanel, outer_end - offset);
const int trailing = MatrixSize - offset - jb;
const int inner_remaining = outer_end - offset - jb;
const int diagonal_offset = offset + offset * MatrixSize;
const int panel_offset =
trailing > 0 ? offset + (offset + jb) * MatrixSize
: diagonal_offset;
if constexpr ((((MatrixSize == 2048 || MatrixSize == 4096) &&
Batch == 2) ||
(MatrixSize == 2048 && Batch == 8)) &&
InnerPanel == kRegisterPanel64) {
factor_register_panel64_kernel<MatrixSize>
<<<Batch, kRegisterPanelThreads64>>>(
base, diagonal_pointers, panel_pointers, offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
} else if constexpr (MatrixSize == 2048 && Batch == 8 &&
InnerPanel == kRegisterDiagonal128) {
factor_register_diagonal128_kernel<MatrixSize>
<<<Batch, kRegisterDiagonalThreads128>>>(
base, diagonal_pointers, panel_pointers, offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
} else {
fill_highbatch_panel_pointers<MatrixSize>
<<<pointer_blocks, threads>>>(
base, diagonal_pointers, panel_pointers, diagonal_offset,
panel_offset, Batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
check_combined_solver(
cusolverDnSpotrfBatched(
state.solver, CUBLAS_FILL_MODE_UPPER, jb,
diagonal_pointers, MatrixSize, state.info.data_ptr<int>(),
Batch),
"cusolverDnSpotrfBatched two-level diagonal");
}
if (trailing == 0) {
continue;
}
check_combined_blas(
cublasStrsmBatched(
state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, jb, trailing, &one,
reinterpret_cast<const float* const*>(diagonal_pointers),
MatrixSize, panel_pointers, MatrixSize, Batch),
"cublasStrsmBatched two-level panel");
if (inner_remaining > 0) {
float* const panel = base + panel_offset;
float* const next_strip =
base + (offset + jb) * (MatrixSize + 1);
check_combined_blas(
cublasGemmStridedBatchedEx(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, inner_remaining,
trailing, jb, &minus_one, panel, CUDA_R_32F, MatrixSize,
matrix_elements, panel, CUDA_R_32F, MatrixSize,
matrix_elements, &one, next_strip, CUDA_R_32F, MatrixSize,
matrix_elements, Batch, update_compute,
CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx two-level panel strip");
}
}
const int far = MatrixSize - outer_end;
if (far > 0) {
float* const outer_to_far =
base + outer + static_cast<long long>(outer_end) * MatrixSize;
float* const far_diagonal =
base + static_cast<long long>(outer_end) * (MatrixSize + 1);
check_combined_blas(
cublasGemmStridedBatchedEx(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, far, far, outer_width,
&minus_one, outer_to_far, CUDA_R_32F, MatrixSize,
matrix_elements, outer_to_far, CUDA_R_32F, MatrixSize,
matrix_elements, &one, far_diagonal, CUDA_R_32F, MatrixSize,
matrix_elements, Batch, update_compute, CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx two-level wide trailing");
}
}
clear_combined_batched_row_major_upper<MatrixSize>
<<<matrix_blocks, threads>>>(base, total);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
template <int MatrixSize, int Batch, int InnerPanel, int OuterPanel>
torch::Tensor cholesky_highbatch_two_level_impl(
torch::Tensor input, const char* operation) {
static_assert(InnerPanel > 0 && OuterPanel >= InnerPanel,
"invalid two-level panel sizes");
static_assert(OuterPanel % InnerPanel == 0,
"outer panel must contain whole inner panels");
static_assert(MatrixSize % OuterPanel == 0,
"outer panel must divide matrix size");
constexpr long long matrix_elements =
static_cast<long long>(MatrixSize) * MatrixSize;
check_combined_batched_input(input, MatrixSize, operation);
TORCH_CHECK(input.size(0) == Batch, operation, " received an invalid batch");
c10::cuda::CUDAGuard device_guard(input.device());
auto output = torch::empty_like(input);
auto& state = true_batched_blocked_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare(input, Batch);
constexpr int threads = 256;
constexpr long long total = static_cast<long long>(Batch) * matrix_elements;
const int matrix_blocks = static_cast<int>(
std::min<long long>(4096, (total + threads - 1) / threads));
constexpr long long vector_total = total / 4;
const int vector_blocks = static_cast<int>(
std::min<long long>(4096, (vector_total + threads - 1) / threads));
if constexpr (MatrixSize == 1024 && (Batch == 4 || Batch == 60)) {
copy_highbatch_lower_only_float4<MatrixSize>
<<<vector_blocks, threads>>>(input.data_ptr<float>(),
output.data_ptr<float>(), vector_total);
} else {
copy_highbatch_lower_only<MatrixSize>
<<<matrix_blocks, threads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), total);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
auto diagonal_pointers = reinterpret_cast<float**>(
state.pointers.data_ptr<int64_t>());
auto panel_pointers = diagonal_pointers + Batch;
constexpr int pointer_blocks = (Batch + threads - 1) / threads;
const float one = 1.0f;
const float minus_one = -1.0f;
float* const base = output.data_ptr<float>();
for (int outer = 0; outer < MatrixSize; outer += OuterPanel) {
constexpr int outer_width = OuterPanel;
const int outer_end = outer + OuterPanel;
for (int offset = outer; offset < outer_end; offset += InnerPanel) {
constexpr int jb = InnerPanel;
const int trailing = MatrixSize - offset - jb;
const int inner_remaining = outer_end - offset - jb;
const int diagonal_offset = offset + offset * MatrixSize;
const int panel_offset =
trailing > 0 ? offset + (offset + jb) * MatrixSize
: diagonal_offset;
if constexpr (MatrixSize == 1024 &&
(Batch == 4 || Batch == 60) &&
InnerPanel == kRegisterPanel64) {
factor_register_panel64_kernel<MatrixSize>
<<<Batch, kRegisterPanelThreads64>>>(
base, diagonal_pointers, panel_pointers, offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
} else if constexpr (MatrixSize == 512 && Batch == 640 &&
InnerPanel == kRegisterDiagonal128) {
factor_register_diagonal128_kernel<MatrixSize>
<<<Batch, kRegisterDiagonalThreads128>>>(
base, diagonal_pointers, panel_pointers, offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
} else {
fill_highbatch_panel_pointers<MatrixSize>
<<<pointer_blocks, threads>>>(
base, diagonal_pointers, panel_pointers, diagonal_offset,
panel_offset, Batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
check_combined_solver(
cusolverDnSpotrfBatched(
state.solver, CUBLAS_FILL_MODE_UPPER, jb,
diagonal_pointers, MatrixSize, state.info.data_ptr<int>(),
Batch),
"cusolverDnSpotrfBatched high-batch two-level diagonal");
}
if (trailing == 0) {
continue;
}
check_combined_blas(
cublasStrsmBatched(
state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, jb, trailing, &one,
reinterpret_cast<const float* const*>(diagonal_pointers),
MatrixSize, panel_pointers, MatrixSize, Batch),
"cublasStrsmBatched high-batch two-level panel");
if (inner_remaining > 0) {
float* const panel = base + panel_offset;
float* const next_strip =
base + (offset + jb) * (MatrixSize + 1);
check_combined_blas(
cublasGemmStridedBatchedEx(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, inner_remaining,
trailing, jb, &minus_one, panel, CUDA_R_32F, MatrixSize,
matrix_elements, panel, CUDA_R_32F, MatrixSize,
matrix_elements, &one, next_strip, CUDA_R_32F, MatrixSize,
matrix_elements, Batch, CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx high-batch two-level panel strip");
}
}
const int far = MatrixSize - outer_end;
if (far > 0) {
float* const outer_to_far =
base + outer + static_cast<long long>(outer_end) * MatrixSize;
float* const far_diagonal =
base + static_cast<long long>(outer_end) * (MatrixSize + 1);
check_combined_blas(
cublasGemmStridedBatchedEx(
state.blas, CUBLAS_OP_T, CUBLAS_OP_N, far, far, outer_width,
&minus_one, outer_to_far, CUDA_R_32F, MatrixSize,
matrix_elements, outer_to_far, CUDA_R_32F, MatrixSize,
matrix_elements, &one, far_diagonal, CUDA_R_32F, MatrixSize,
matrix_elements, Batch, CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx high-batch two-level wide trailing");
}
}
if constexpr (MatrixSize == 1024 && (Batch == 4 || Batch == 60)) {
clear_highbatch_upper_float4<MatrixSize>
<<<vector_blocks, threads>>>(base, vector_total);
} else {
clear_combined_batched_row_major_upper<MatrixSize>
<<<matrix_blocks, threads>>>(base, total);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
} // namespace
torch::Tensor cholesky256_batch64_register_panel64_cuda(torch::Tensor input) {
check_combined_batched_input(
input, kCombinedSize256,
"cholesky256_batch64_register_panel64_cuda");
TORCH_CHECK(input.size(0) == 64,
"cholesky256_batch64_register_panel64_cuda expects batch 64");
return cholesky256_batch64_register_panel64_impl(input);
}
torch::Tensor cholesky256_combined_batched_cuda(torch::Tensor input) {
check_combined_batched_input(
input, kCombinedSize256, "cholesky256_combined_batched_cuda");
return cholesky_combined_batched_impl<kCombinedSize256>(input);
}
torch::Tensor cholesky512_batch16_register_panel64_cuda(torch::Tensor input) {
check_combined_batched_input(
input, kCombinedSize512,
"cholesky512_batch16_register_panel64_cuda");
TORCH_CHECK(input.size(0) == 16,
"cholesky512_batch16_register_panel64_cuda expects batch 16");
return cholesky512_batch16_register_panel64_impl(input);
}
torch::Tensor cholesky512_combined_batched_cuda(torch::Tensor input) {
check_combined_batched_input(
input, kCombinedSize512, "cholesky512_combined_batched_cuda");
return cholesky_combined_batched_impl<kCombinedSize512>(input);
}
torch::Tensor cholesky512_highbatch_blocked_cuda(torch::Tensor input) {
return cholesky_highbatch_two_level_impl<512, 640, 128, 256>(
input, "cholesky512_highbatch_blocked_cuda");
}
torch::Tensor cholesky1024_batch4_blocked_cuda(torch::Tensor input) {
return cholesky_highbatch_two_level_impl<1024, 4, 64, 128>(
input, "cholesky1024_batch4_blocked_cuda");
}
torch::Tensor cholesky1024_batch60_blocked_cuda(torch::Tensor input) {
return cholesky_highbatch_two_level_impl<1024, 60, 64, 256>(
input, "cholesky1024_batch60_blocked_cuda");
}
torch::Tensor cholesky2048_batch2_blocked_cuda(torch::Tensor input) {
return cholesky_true_batched_two_level_impl<2048, 2, 64, 256, true>(
input, "cholesky2048_batch2_blocked_cuda");
}
torch::Tensor cholesky2048_batch8_blocked_cuda(torch::Tensor input) {
return cholesky_true_batched_two_level_impl<2048, 8, 64, 512, false>(
input, "cholesky2048_batch8_blocked_cuda");
}
torch::Tensor cholesky4096_batch2_truebatched_fast16bf_cuda(
torch::Tensor input) {
return cholesky_true_batched_two_level_impl<4096, 2, 64, 512, true,
true>(
input, "cholesky4096_batch2_truebatched_fast16bf_cuda");
}
torch::Tensor cholesky_large_explicit_bf16_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(),
"cholesky_large_explicit_bf16_cuda expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky_large_explicit_bf16_cuda expects float32 input");
TORCH_CHECK(
input.dim() == 3 && input.size(1) == input.size(2) &&
((input.size(0) == 2 && input.size(1) == 4096) ||
(input.size(0) == 1 &&
(input.size(1) == 16384 || input.size(1) == 32768))),
"cholesky_large_explicit_bf16_cuda expects batch/n in "
"{(2,4096),(1,16384),(1,32768)}");
TORCH_CHECK(input.is_contiguous(),
"cholesky_large_explicit_bf16_cuda expects contiguous input");
c10::cuda::CUDAGuard device_guard(input.device());
const int device = input.get_device();
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
const long long matrix_elements = static_cast<long long>(n) * n;
constexpr int standard_nb = 512;
constexpr int champion_inner_nb = 128;
constexpr int wide_inner_nb = 384;
constexpr int panel_correction_nb = 96;
constexpr int wide_prefix_end = 23040;
constexpr int outer_nb = 1536;
constexpr int trailing_tile = 3072;
const bool lower_only_initialization = n >= 16384;
const bool regime_schedule = n == 32768;
const int factor_nb = regime_schedule
? wide_inner_nb
: (lower_only_initialization ? champion_inner_nb
: standard_nb);
auto output = lower_only_initialization ? torch::empty_like(input)
: input.clone();
float* const base = output.data_ptr<float>();
if (lower_only_initialization) {
constexpr int initialization_threads = 256;
initialize_explicit_large_first_panel_lower
<<<n, initialization_threads>>>(
input.data_ptr<float>(), base, n, outer_nb);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
auto& state = explicit_low_precision_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare(input);
int required_elements = 0;
check_combined_solver(
cusolverDnSpotrf_bufferSize(
state.solver_handle, CUBLAS_FILL_MODE_UPPER, factor_nb, base, n,
&required_elements),
"cusolverDnSpotrf_bufferSize explicit BF16");
if (state.workspace_elements < required_elements) {
if (state.workspace.defined()) {
state.retired_buffers.push_back(state.workspace);
}
state.workspace = torch::empty(
{required_elements}, input.options().dtype(torch::kFloat32));
state.workspace_elements = required_elements;
}
if (!state.info.defined()) {
state.info = torch::empty({1}, input.options().dtype(torch::kInt32));
}
const int packed_rows = lower_only_initialization ? outer_nb : standard_nb;
const long long required_panel_capacity =
static_cast<long long>(packed_rows) * n;
if (state.panel_capacity < required_panel_capacity) {
if (state.panel_bf16.defined()) {
state.retired_buffers.push_back(state.panel_bf16);
}
state.panel_bf16 = torch::empty(
{required_panel_capacity}, input.options().dtype(torch::kBFloat16));
state.panel_capacity = required_panel_capacity;
}
auto* const packed = reinterpret_cast<__nv_bfloat16*>(
state.panel_bf16.data_ptr<at::BFloat16>());
constexpr long long mxfp8_scratch_bytes = 32LL * 1024 * 1024;
if (regime_schedule) {
const long long scale_elements =
static_cast<long long>((outer_nb + 127) / 128) *
((n + 127) / 128) * 512;
if (!state.mxfp8_primary_scale.defined() ||
state.mxfp8_primary_scale.numel() < scale_elements) {
if (state.mxfp8_primary_scale.defined()) {
state.retired_buffers.push_back(state.mxfp8_primary_scale);
}
state.mxfp8_primary_scale = torch::empty(
{scale_elements}, input.options().dtype(torch::kUInt8));
}
if (!state.mxfp8_diagonal_correction.defined() ||
state.mxfp8_diagonal_correction.numel() < n) {
if (state.mxfp8_diagonal_correction.defined()) {
state.retired_buffers.push_back(state.mxfp8_diagonal_correction);
}
state.mxfp8_diagonal_correction = torch::empty(
{n}, input.options().dtype(torch::kFloat32));
}
if (!state.mxfp8_scratch.defined() ||
state.mxfp8_scratch.numel() < mxfp8_scratch_bytes) {
if (state.mxfp8_scratch.defined()) {
state.retired_buffers.push_back(state.mxfp8_scratch);
}
state.mxfp8_scratch = torch::empty(
{mxfp8_scratch_bytes}, input.options().dtype(torch::kUInt8));
}
}
const float one = 1.0f;
const float minus_one = -1.0f;
for (int matrix = 0; matrix < batch; ++matrix) {
const float* const input_matrix_base =
input.data_ptr<float>() +
static_cast<long long>(matrix) * matrix_elements;
float* const matrix_base =
base + static_cast<long long>(matrix) * matrix_elements;
if (lower_only_initialization) {
constexpr int threads = 256;
for (int outer = 0; outer < n; outer += outer_nb) {
const int outer_end = std::min(n, outer + outer_nb);
const int outer_width = outer_end - outer;
const bool wide_prefix =
regime_schedule && outer < wide_prefix_end;
const int inner_nb =
wide_prefix ? wide_inner_nb : champion_inner_nb;
for (int k = outer; k < outer_end; k += inner_nb) {
const int jb = std::min(inner_nb, outer_end - k);
const int r = n - k - jb;
const int inner_remaining = outer_end - k - jb;
float* const diagonal =
matrix_base + static_cast<long long>(k) * (n + 1);
check_combined_solver(
cusolverDnSpotrf(
state.solver_handle, CUBLAS_FILL_MODE_UPPER, jb, diagonal,
n, state.workspace.data_ptr<float>(),
state.workspace_elements, state.info.data_ptr<int>()),
"cusolverDnSpotrf two-level explicit BF16");
if (r == 0) {
continue;
}
const int packed_row = k - outer;
const int packed_col = k + jb - outer;
if (wide_prefix) {
for (int q = 0; q < jb; q += panel_correction_nb) {
const int qb = std::min(panel_correction_nb, jb - q);
float* const correction_diagonal =
matrix_base + (k + q) +
static_cast<long long>(k + q) * n;
float* const correction_panel =
matrix_base + (k + q) +
static_cast<long long>(k + jb) * n;
check_combined_blas(
cublasStrsm(
state.blas_handle, CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER, CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT, qb, r, &one,
correction_diagonal, n, correction_panel, n),
"cublasStrsm FP32 diagonal correction for BF16 panel");
const int source_col = k + q + qb;
const int packed_source_col = source_col - outer;
const int packed_source_row = k + q - outer;
const int source_cols = n - source_col;
const float* const source =
matrix_base + (k + q) +
static_cast<long long>(source_col) * n;
const long long packed_offset =
packed_source_row +
static_cast<long long>(packed_source_col) * outer_width;
const bool vector_pack =
(qb & 7) == 0 && (n & 3) == 0 &&
(outer_width & 7) == 0 &&
(reinterpret_cast<std::uintptr_t>(source) & 15) == 0 &&
(reinterpret_cast<std::uintptr_t>(packed + packed_offset) &
15) == 0;
if (vector_pack) {
const long long source_groups =
static_cast<long long>(qb / 8) * source_cols;
const int blocks = static_cast<int>(std::min<long long>(
4096, (source_groups + threads - 1) / threads));
pack_outer_panel_bf16_vector8<<<blocks, threads>>>(
source, packed, qb, source_cols, n, outer_width,
packed_source_row, packed_source_col);
} else {
const long long source_elements =
static_cast<long long>(qb) * source_cols;
const int blocks = static_cast<int>(std::min<long long>(
4096, (source_elements + threads - 1) / threads));
pack_outer_panel_bf16<<<blocks, threads>>>(
source, packed, qb, source_cols, n, outer_width,
packed_source_row, packed_source_col);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
const int unsolved_rows = jb - q - qb;
if (unsolved_rows > 0) {
const __nv_bfloat16* const packed_diagonal_row =
packed + packed_source_row +
static_cast<long long>(packed_source_col) * outer_width;
const __nv_bfloat16* const packed_panel_row =
packed + packed_source_row +
static_cast<long long>(packed_col) * outer_width;
float* const unsolved_panel =
matrix_base + (k + q + qb) +
static_cast<long long>(k + jb) * n;
check_combined_blas(
cublasGemmEx(
state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N,
unsolved_rows, r, qb, &minus_one,
packed_diagonal_row, CUDA_R_16BF, outer_width,
packed_panel_row, CUDA_R_16BF, outer_width, &one,
unsolved_panel, CUDA_R_32F, n, CUBLAS_COMPUTE_32F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx BF16 tensor panel solve update");
}
}
} else {
float* const panel =
matrix_base + k + static_cast<long long>(k + jb) * n;
check_combined_blas(
cublasStrsm(
state.blas_handle, CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER, CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT, jb, r, &one, diagonal, n, panel,
n),
"cublasStrsm two-level explicit BF16 panel");
const long long packed_offset =
packed_row +
static_cast<long long>(packed_col) * outer_width;
const bool vector_pack =
(jb & 7) == 0 && (n & 3) == 0 &&
(outer_width & 7) == 0 &&
(reinterpret_cast<std::uintptr_t>(panel) & 15) == 0 &&
(reinterpret_cast<std::uintptr_t>(packed + packed_offset) &
15) == 0;
if (vector_pack) {
const long long panel_groups =
static_cast<long long>(jb / 8) * r;
const int blocks = static_cast<int>(std::min<long long>(
4096, (panel_groups + threads - 1) / threads));
pack_outer_panel_bf16_vector8<<<blocks, threads>>>(
panel, packed, jb, r, n, outer_width, packed_row,
packed_col);
} else {
const long long panel_elements =
static_cast<long long>(jb) * r;
const int blocks = static_cast<int>(std::min<long long>(
4096, (panel_elements + threads - 1) / threads));
pack_outer_panel_bf16<<<blocks, threads>>>(
panel, packed, jb, r, n, outer_width, packed_row,
packed_col);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
if (inner_remaining > 0) {
const __nv_bfloat16* const packed_panel =
packed + packed_row +
static_cast<long long>(packed_col) * outer_width;
float* const next_strip =
matrix_base + static_cast<long long>(k + jb) * (n + 1);
check_combined_blas(
cublasGemmEx(
state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N,
inner_remaining, r, jb, &minus_one, packed_panel,
CUDA_R_16BF, outer_width, packed_panel, CUDA_R_16BF,
outer_width, &one, next_strip, CUDA_R_32F, n,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx two-level BF16 unfinished outer strip");
}
}
const int far = n - outer_end;
if (far > 0) {
const __nv_bfloat16* const packed_far =
packed + static_cast<long long>(outer_width) * outer_width;
float* const far_diagonal =
matrix_base + static_cast<long long>(outer_end) * (n + 1);
const float* const input_far_diagonal =
input_matrix_base +
static_cast<long long>(outer_end) * (n + 1);
if (regime_schedule) {
auto* const primary = reinterpret_cast<__nv_fp8_e4m3*>(packed);
auto* const primary_scale =
reinterpret_cast<__nv_fp8_e8m0*>(
state.mxfp8_primary_scale.data_ptr<uint8_t>());
float* const diagonal_correction =
state.mxfp8_diagonal_correction.data_ptr<float>();
const int clear_blocks =
std::min(4096, (far + threads - 1) / threads);
clear_mxfp8_diagonal_correction<<<clear_blocks, threads>>>(
diagonal_correction, far);
const float* const finalized_far_panel =
matrix_base + outer +
static_cast<long long>(outer_end) * n;
pack_outer_far_mxfp8_block32<<<far, threads>>>(
finalized_far_panel, primary, primary_scale,
diagonal_correction, outer_width, far, n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
for (int j = 0; j < far; j += trailing_tile) {
const int nj = std::min(trailing_tile, far - j);
const int m = j + nj;
const long long packed_offset =
static_cast<long long>(j) * outer_width;
const long long scale_offset =
mxfp8_scale_column_base(outer_width, j);
float* const trailing_j =
far_diagonal + static_cast<long long>(j) * n;
const float* const source_j =
outer == 0
? input_far_diagonal + static_cast<long long>(j) * n
: trailing_j;
lt_packed_mxfp8_block_column(
primary, primary + packed_offset,
primary_scale, primary_scale + scale_offset,
source_j, trailing_j, m, nj, outer_width, n,
device, state.mxfp8_lt_cache,
state.mxfp8_scratch.data_ptr<uint8_t>(),
mxfp8_scratch_bytes, &minus_one, &one);
}
add_mxfp8_diagonal_correction<<<clear_blocks, threads>>>(
far_diagonal, diagonal_correction, far, n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
continue;
}
for (int j = 0; j < far; j += trailing_tile) {
const int nj = std::min(trailing_tile, far - j);
const int m = j + nj;
const __nv_bfloat16* const packed_j =
packed_far + static_cast<long long>(j) * outer_width;
float* const trailing_j =
far_diagonal + static_cast<long long>(j) * n;
if (outer == 0) {
const float* const source_j =
input_far_diagonal + static_cast<long long>(j) * n;
lt_packed_bf16_block_column(
packed_far, packed_j, source_j, trailing_j, m, nj,
outer_width, n, &minus_one, &one);
} else {
check_combined_blas(
cublasGemmEx(
state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N, m, nj,
outer_width, &minus_one, packed_far, CUDA_R_16BF,
outer_width, packed_j, CUDA_R_16BF, outer_width, &one,
trailing_j, CUDA_R_32F, n, CUBLAS_COMPUTE_32F,
CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx two-level wide packed BF16 update");
}
}
}
}
} else {
for (int k = 0; k < n; k += standard_nb) {
const int jb = std::min(standard_nb, n - k);
const int r = n - k - jb;
float* const diagonal =
matrix_base + static_cast<long long>(k) * (n + 1);
check_combined_solver(
cusolverDnSpotrf(
state.solver_handle, CUBLAS_FILL_MODE_UPPER, jb, diagonal, n,
state.workspace.data_ptr<float>(), state.workspace_elements,
state.info.data_ptr<int>()),
"cusolverDnSpotrf explicit BF16");
if (r == 0) {
continue;
}
float* const panel =
matrix_base + k + static_cast<long long>(k + jb) * n;
check_combined_blas(
cublasStrsm(
state.blas_handle, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, jb, r, &one, diagonal, n,
panel, n),
"cublasStrsm explicit BF16 panel");
constexpr int threads = 256;
const long long panel_elements = static_cast<long long>(jb) * r;
const int blocks = static_cast<int>(
std::min<long long>(4096,
(panel_elements + threads - 1) / threads));
pack_panel_bf16<<<blocks, threads>>>(panel, packed, jb, r, n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
float* const trailing =
matrix_base + static_cast<long long>(k + jb) * (n + 1);
for (int j = 0; j < r; j += trailing_tile) {
const int nj = std::min(trailing_tile, r - j);
const int m = j + nj;
float* const trailing_j = trailing + static_cast<long long>(j) * n;
const __nv_bfloat16* const packed_j =
packed + static_cast<long long>(j) * jb;
check_combined_blas(
cublasGemmEx(
state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N, m, nj, jb,
&minus_one, packed, CUDA_R_16BF, jb, packed_j,
CUDA_R_16BF, jb, &one, trailing_j, CUDA_R_32F, n,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmEx packed BF16 block-column update");
}
}
}
}
constexpr int threads = 256;
const long long total = static_cast<long long>(batch) * matrix_elements;
if (lower_only_initialization) {
clear_explicit_large_upper_vectorized<<<n, threads>>>(base, n);
} else {
const int blocks = static_cast<int>(
std::min<long long>(4096, (total + threads - 1) / threads));
clear_explicit_batch_upper<<<blocks, threads>>>(base, n, total);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
torch::Tensor cholesky_large_combined_bf16x9_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(),
"cholesky_large_combined_bf16x9_cuda expects a CUDA tensor");
TORCH_CHECK(input.scalar_type() == torch::kFloat32,
"cholesky_large_combined_bf16x9_cuda expects float32 input");
TORCH_CHECK(input.dim() == 3 && input.size(0) == 1 &&
input.size(1) == input.size(2) && input.size(1) >= 24576,
"cholesky_large_combined_bf16x9_cuda expects [1, n, n] with "
"n >= 24576");
TORCH_CHECK(input.is_contiguous(),
"cholesky_large_combined_bf16x9_cuda expects contiguous input");
TORCH_CHECK(input.size(1) <= std::numeric_limits<int>::max(),
"cholesky_large_combined_bf16x9_cuda matrix is too large");
c10::cuda::CUDAGuard device_guard(input.device());
const int n = static_cast<int>(input.size(1));
auto output = torch::empty_like(input);
constexpr int threads = 256;
const long long total = static_cast<long long>(n) * n;
copy_combined_large_lower<<<n, threads>>>(
input.data_ptr<float>(), output.data_ptr<float>(), n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
auto& state = combined_solver_state();
std::lock_guard<std::mutex> lock(state.mutex);
state.prepare_device(input);
int required_elements = 0;
check_combined_solver(
cusolverDnSpotrf_bufferSize(
state.large_handle, CUBLAS_FILL_MODE_UPPER, n,
output.data_ptr<float>(), n, &required_elements),
"cusolverDnSpotrf_bufferSize");
if (state.workspace_elements < required_elements) {
state.workspace = torch::empty(
{required_elements}, input.options().dtype(torch::kFloat32));
state.workspace_elements = required_elements;
}
if (!state.large_info.defined()) {
state.large_info = torch::empty(
{1}, input.options().dtype(torch::kInt32));
}
check_combined_solver(
cusolverDnSpotrf(
state.large_handle, CUBLAS_FILL_MODE_UPPER, n,
output.data_ptr<float>(), n, state.workspace.data_ptr<float>(),
required_elements, state.large_info.data_ptr<int>()),
"cusolverDnSpotrf");
const int blocks = static_cast<int>(
std::min<long long>(4096, (total + threads - 1) / threads));
clear_combined_large_upper<<<blocks, threads>>>(output.data_ptr<float>(), n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return output;
}
"""
_cholesky64_left_extension = load_inline(
name="popcorn_cholesky_integration_medium_schedules_v3",
cpp_sources=_CHOLESKY64_CPP,
cuda_sources=_CHOLESKY64_CUDA,
functions=[
"cholesky64_block32_hacc4safe_cuda",
"cholesky32_warp_cuda",
"cholesky128_wmma_schur_cuda",
"cholesky128_wmma_schur_exact_cuda",
"cholesky128_block32_warp16_lowerio_anybatch_cuda",
"cholesky256_batch64_register_panel64_cuda",
"cholesky256_combined_batched_cuda",
"cholesky512_batch16_register_panel64_cuda",
"cholesky512_combined_batched_cuda",
"cholesky512_highbatch_blocked_cuda",
"cholesky1024_batch4_blocked_cuda",
"cholesky1024_batch60_blocked_cuda",
"cholesky2048_batch2_blocked_cuda",
"cholesky2048_batch8_blocked_cuda",
"cholesky4096_batch2_truebatched_fast16bf_cuda",
"cholesky_large_explicit_bf16_cuda",
"cholesky_large_combined_bf16x9_cuda",
],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3", "--use_fast_math"],
extra_ldflags=["-lcusolver", "-lcublas", "-lcublasLt"],
with_cuda=True,
verbose=False,
)
@triton.jit
def _cholesky32_left_kernel(input_ptr, output_ptr, matrix_stride: tl.constexpr):
matrix = tl.program_id(0)
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
offsets = matrix * matrix_stride + rows * 32 + cols
values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)
for k in range(32):
row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(tl.where(col_ids == k, row, 0.0), axis=0)
diagonal -= tl.sum(tl.where(col_ids < k, row * row, 0.0), axis=0)
diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
products = tl.where(cols < k, values * row[None, :], 0.0)
column = (column - tl.sum(products, axis=1)) / diagonal
values = tl.where((rows == k) & (cols == k), diagonal, values)
values = tl.where((rows > k) & (cols == k), column[:, None], values)
tl.store(output_ptr + offsets, values)
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
if batch == 2 and n == 4096:
return (
_cholesky64_left_extension
.cholesky4096_batch2_truebatched_fast16bf_cuda(data)
)
if batch == 1 and n in (16384, 32768):
return _cholesky64_left_extension.cholesky_large_explicit_bf16_cuda(
data
)
if batch == 1 and n >= 24576:
return _cholesky64_left_extension.cholesky_large_combined_bf16x9_cuda(
data
)
if batch == 64 and n == 256:
return (
_cholesky64_left_extension
.cholesky256_batch64_register_panel64_cuda(data)
)
if batch >= 128 and n == 256:
return _cholesky64_left_extension.cholesky256_combined_batched_cuda(
data
)
if batch == 640 and n == 512:
return (
_cholesky64_left_extension
.cholesky512_highbatch_blocked_cuda(data)
)
if batch == 16 and n == 512:
return (
_cholesky64_left_extension
.cholesky512_batch16_register_panel64_cuda(data)
)
if batch == 4 and n == 1024:
return (
_cholesky64_left_extension
.cholesky1024_batch4_blocked_cuda(data)
)
if batch == 60 and n == 1024:
return (
_cholesky64_left_extension
.cholesky1024_batch60_blocked_cuda(data)
)
if batch == 2 and n == 2048:
return (
_cholesky64_left_extension
.cholesky2048_batch2_blocked_cuda(data)
)
if batch == 8 and n == 2048:
return (
_cholesky64_left_extension
.cholesky2048_batch8_blocked_cuda(data)
)
if batch >= 256 and n == 512:
return _cholesky64_left_extension.cholesky512_combined_batched_cuda(
data
)
if batch > 0 and n == 32:
return _cholesky64_left_extension.cholesky32_warp_cuda(data)
if batch > 0 and n == 64:
return _cholesky64_left_extension.cholesky64_block32_hacc4safe_cuda(data)
if batch == 256 and n == 128:
return _cholesky64_left_extension.cholesky128_wmma_schur_exact_cuda(
data
)
if batch >= 16 and n == 128:
return (
_cholesky64_left_extension
.cholesky128_block32_warp16_lowerio_anybatch_cuda(data)
)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 3999 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