submission 892316
badelsteinlelbach · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 10592 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-892316?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:cfbe9a7d488dc5ed3ec53ec816b9be6a6959d5fd76a71db1795145931fc9974a
license declaredunknown
license concludedunknown
authorsbadelsteinlelbach
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
void dx_cp_async_pair(mbarrier
"mbarrier.init.shared::cta.b64 [%0], 1;\n"mma
namespace wmma = nvcuda::wmma;persistent-kernel
from cute.blackwell.kernel.dense_gemm.dense_gemm_alpha_beta_persistent import (shared-memory
__shared__ float tile[32][33];vector-width = float4
const float4 values = *reinterpret_cast<const float4*>(Kernel source
submission.py10592 lines
import glob
import os
import sys
import numpy as np
import torch
from numba import cuda, types
from numba.cuda.cudadrv import driver as numba_driver
from nvmath.device.common_numba import get_array_ptr
from nvmath.device import CholeskySolver, Matmul, TriangularSolver
from torch.utils.cpp_extension import CUDA_HOME, load_inline
from task import input_t, output_t
_CUTE_ROOT = next(
path
for path in (
"/opt/cutlass/examples/python/CuTeDSL",
"/tmp/cutlass-v4.5.2/examples/python/CuTeDSL",
)
if os.path.isdir(path)
)
sys.path.insert(0, _CUTE_ROOT)
import cuda.bindings.driver as cuda_driver
import cutlass
import cutlass.cute as cute
import cutlass.cute.runtime as cute_runtime
import cutlass.utils as cutlass_utils
from cutlass.cutlass_dsl import dsl_user_op
from cute.blackwell.kernel.dense_gemm.dense_gemm_alpha_beta_persistent import (
SM100PersistentDenseGemmAlphaBetaKernel as _AlphaBetaGemm,
)
from cute.blackwell.kernel.dense_gemm.dense_gemm_persistent import (
PersistentDenseGemmKernel as _DenseGemm,
)
_DX_ASYNC_SOURCE = cuda.CUSource(
r"""
#include <cuda_fp16.h>
extern "C" __device__
void dx_cp_async_pair(
float* first,
float* second,
const float* source,
long long first_index,
long long second_index) {
#pragma unroll
for (int load = 0; load < 8; ++load) {
constexpr int fp32_ld = 36;
constexpr int rows_per_load = 4;
const int row = threadIdx.x / 32 + load * rows_per_load;
const int col = threadIdx.x % 32;
const int shared_index = row * fp32_ld + col;
unsigned first_shared = static_cast<unsigned>(
__cvta_generic_to_shared(first + shared_index));
unsigned second_shared = static_cast<unsigned>(
__cvta_generic_to_shared(second + shared_index));
const float* first_global = source + first_index + load * 2048;
const float* second_global = source + second_index + load * 2048;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 4;\n"
:: "r"(first_shared), "l"(first_global));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 4;\n"
:: "r"(second_shared), "l"(second_global));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void dx_b1024_cp_async_fp32_tile(
float* destination,
const float* source,
long long source_index,
int global_ld) {
// Let the final two warps stage the next 64x64 operand while the other six
// warps enter the MathDx operation that consumes the current shared bank.
if (threadIdx.x >= 192) {
constexpr int shared_ld = 68;
const int lane = threadIdx.x & 63;
#pragma unroll
for (int load = 0; load < 16; ++load) {
const int chunk = lane + load * 64;
const int row = chunk >> 4;
const int col = (chunk & 15) << 2;
const int shared_index = row * shared_ld + col;
unsigned destination_shared = static_cast<unsigned>(
__cvta_generic_to_shared(destination + shared_index));
const float* source_global =
source + source_index + row * global_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(destination_shared), "l"(source_global));
if ((load & 7) == 7) {
asm volatile("cp.async.commit_group;\n" ::);
}
}
}
}
extern "C" __device__
void dx_b1024_cp_async_half_tile(
__half* destination,
const __half* source,
long long source_index,
int global_ld) {
constexpr int shared_ld = 72;
#pragma unroll
for (int load = 0; load < 2; ++load) {
const int chunk = threadIdx.x + load * 256;
const int row = chunk >> 3;
const int col = (chunk & 7) << 3;
unsigned destination_shared = static_cast<unsigned>(
__cvta_generic_to_shared(destination + row * shared_ld + col));
const __half* source_global =
source + source_index + row * global_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(destination_shared), "l"(source_global));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void dx_b1024_load_convert_half_tile(
__half* destination,
const float* source,
long long source_index,
int global_ld) {
constexpr int shared_ld = 72;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) << 2;
#pragma unroll
for (int load = 0; load < 4; ++load) {
const int current_row = row + load * 16;
const float4 values = *reinterpret_cast<const float4*>(
source + source_index + current_row * global_ld + col);
__half2* packed = reinterpret_cast<__half2*>(
destination + current_row * shared_ld + col);
packed[0] = __floats2half2_rn(values.x, values.y);
packed[1] = __floats2half2_rn(values.z, values.w);
}
}
extern "C" __device__
void dx_copy_float8(
float* destination,
const float* source,
long long index) {
float4* destination_vectors = reinterpret_cast<float4*>(
destination + index);
const float4* source_vectors = reinterpret_cast<const float4*>(
source + index);
destination_vectors[0] = source_vectors[0];
destination_vectors[1] = source_vectors[1];
}
extern "C" __device__
void g2048_store_upper_zero_lower(
const float* source,
float* destination,
long long destination_index,
long long leading) {
const int row = threadIdx.x >> 3;
const int col = (threadIdx.x & 7) << 2;
float4 values = *reinterpret_cast<const float4*>(
source + row * 32 + col);
if (row > col) values.x = 0.0f;
if (row > col + 1) values.y = 0.0f;
if (row > col + 2) values.z = 0.0f;
if (row > col + 3) values.w = 0.0f;
*reinterpret_cast<float4*>(
destination + destination_index + row * leading + col) = values;
}
extern "C" __device__
void g2048_store_panel_float4(
const float* source,
float* destination,
long long destination_index,
long long symmetric_index,
long long leading,
int clear_symmetric) {
if (threadIdx.x < 128) {
const int row = threadIdx.x >> 2;
const int col = (threadIdx.x & 3) << 2;
const float4 values = *reinterpret_cast<const float4*>(
source + row * 16 + col);
*reinterpret_cast<float4*>(
destination + destination_index + row * leading + col) = values;
if (clear_symmetric) {
const int symmetric_row = threadIdx.x >> 3;
const int symmetric_col = (threadIdx.x & 7) << 2;
*reinterpret_cast<float4*>(
destination + symmetric_index
+ symmetric_row * leading + symmetric_col) =
make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
}
}
extern "C" __device__
void g2048_load_tile_float4(
float* destination,
const float* source,
long long source_index,
long long leading) {
const int row = threadIdx.x >> 3;
const int col = (threadIdx.x & 7) << 2;
const float4 values = *reinterpret_cast<const float4*>(
source + source_index + row * leading + col);
*reinterpret_cast<float4*>(destination + row * 32 + col) = values;
}
extern "C" __device__
void dx_cp_async_half_tile(
__half* destination,
const __half* source,
long long source_index) {
constexpr int destination_ld = 40;
constexpr int source_ld = 32;
const int row = threadIdx.x >> 2;
const int col = (threadIdx.x & 3) * 8;
const int destination_index = row * destination_ld + col;
const int packed_index = row * source_ld + col;
unsigned destination_shared = static_cast<unsigned>(
__cvta_generic_to_shared(destination + destination_index));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(destination_shared),
"l"(source + source_index + packed_index));
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void dx_cp_async_half_pair(
__half* first,
__half* second,
const __half* source,
long long first_index,
long long second_index) {
constexpr int destination_ld = 40;
constexpr int source_ld = 32;
const int row = threadIdx.x >> 2;
const int col = (threadIdx.x & 3) * 8;
const int destination_index = row * destination_ld + col;
const int packed_index = row * source_ld + col;
unsigned first_shared = static_cast<unsigned>(
__cvta_generic_to_shared(first + destination_index));
unsigned second_shared = static_cast<unsigned>(
__cvta_generic_to_shared(second + destination_index));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(first_shared),
"l"(source + first_index + packed_index));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(second_shared),
"l"(source + second_index + packed_index));
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void dx_cp_async_bulk_half_triple(
__half* first,
__half* second,
__half* third,
const __half* source,
long long first_index,
long long second_index,
long long third_index,
unsigned long long* barriers) {
constexpr int destination_ld = 40;
constexpr int source_ld = 32;
const int row = threadIdx.x >> 2;
const int col = (threadIdx.x & 3) * 8;
const int destination_index = row * destination_ld + col;
const int packed_index = row * source_ld + col;
unsigned first_shared = static_cast<unsigned>(
__cvta_generic_to_shared(first + destination_index));
unsigned second_shared = static_cast<unsigned>(
__cvta_generic_to_shared(second + destination_index));
unsigned third_shared = static_cast<unsigned>(
__cvta_generic_to_shared(third + destination_index));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(first_shared),
"l"(source + first_index + packed_index));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(second_shared),
"l"(source + second_index + packed_index));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(third_shared),
"l"(source + third_index + packed_index));
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d256_cp_async_fp32_tile(
float* destination,
const float* source,
long long source_index) {
constexpr int shared_ld = 68;
constexpr int global_ld = 256;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
unsigned destination_shared = static_cast<unsigned>(
__cvta_generic_to_shared(
destination + current_row * shared_ld + col));
const float* source_global =
source + source_index + current_row * global_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(destination_shared), "l"(source_global));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d256_cp_async_fp32_pair(
float* first,
float* second,
const float* first_source,
const float* second_source,
long long first_index,
long long second_index) {
constexpr int shared_ld = 68;
constexpr int global_ld = 256;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const int shared_index = current_row * shared_ld + col;
unsigned first_shared = static_cast<unsigned>(
__cvta_generic_to_shared(first + shared_index));
unsigned second_shared = static_cast<unsigned>(
__cvta_generic_to_shared(second + shared_index));
const int global_offset = current_row * global_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(first_shared),
"l"(first_source + first_index + global_offset));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(second_shared),
"l"(second_source + second_index + global_offset));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d256_cp_async_half_tile(
__half* destination,
const __half* source,
long long source_index) {
constexpr int shared_ld = 72;
constexpr int source_ld = 64;
const int row = threadIdx.x >> 3;
const int col = (threadIdx.x & 7) * 8;
#pragma unroll
for (int load = 0; load < 4; ++load) {
const int current_row = row + load * 16;
unsigned destination_shared = static_cast<unsigned>(
__cvta_generic_to_shared(
destination + current_row * shared_ld + col));
const __half* source_global =
source + source_index + current_row * source_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(destination_shared), "l"(source_global));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d256_cp_async_half_pair(
__half* first,
__half* second,
const __half* source,
long long first_index,
long long second_index) {
constexpr int shared_ld = 72;
constexpr int source_ld = 64;
const int row = threadIdx.x >> 3;
const int col = (threadIdx.x & 7) * 8;
#pragma unroll
for (int load = 0; load < 4; ++load) {
const int current_row = row + load * 16;
const int shared_index = current_row * shared_ld + col;
unsigned first_shared = static_cast<unsigned>(
__cvta_generic_to_shared(first + shared_index));
unsigned second_shared = static_cast<unsigned>(
__cvta_generic_to_shared(second + shared_index));
const int source_offset = current_row * source_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(first_shared),
"l"(source + first_index + source_offset));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(second_shared),
"l"(source + second_index + source_offset));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d256_store_fp32_and_half_tile(
const float* source,
float* output,
__half* sidecar,
long long output_index,
long long sidecar_index) {
constexpr int source_ld = 68;
constexpr int output_ld = 256;
constexpr int sidecar_ld = 64;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
*reinterpret_cast<float4*>(
output + output_index + current_row * output_ld + col) = values;
__half2* packed = reinterpret_cast<__half2*>(
sidecar + sidecar_index + current_row * sidecar_ld + col);
packed[0] = __floats2half2_rn(values.x, values.y);
packed[1] = __floats2half2_rn(values.z, values.w);
}
}
extern "C" __device__
void d256_store_fp32_tile(
const float* source,
float* output,
long long output_index) {
constexpr int source_ld = 68;
constexpr int output_ld = 256;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
*reinterpret_cast<float4*>(
output + output_index + current_row * output_ld + col) = values;
}
}
extern "C" __device__
void d256_store_fp32_upper_tile(
const float* source,
float* output,
long long output_index) {
constexpr int source_ld = 68;
constexpr int output_ld = 256;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
float* destination =
output + output_index + current_row * output_ld + col;
if (current_row <= col) {
*reinterpret_cast<float4*>(destination) = values;
} else if (current_row <= col + 3) {
if (current_row <= col + 1) destination[1] = values.y;
if (current_row <= col + 2) destination[2] = values.z;
destination[3] = values.w;
}
}
}
extern "C" __device__
void d256_store_fp32_and_shared_half_prefetch(
float* source,
float* output,
__half* destination,
const float* accumulator_source,
long long output_index,
long long accumulator_index) {
constexpr int source_ld = 68;
constexpr int output_ld = 256;
constexpr int destination_ld = 72;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
*reinterpret_cast<float4*>(
output + output_index + current_row * output_ld + col) = values;
__half2* packed = reinterpret_cast<__half2*>(
destination + current_row * destination_ld + col);
packed[0] = __floats2half2_rn(values.x, values.y);
packed[1] = __floats2half2_rn(values.z, values.w);
unsigned accumulator_shared = static_cast<unsigned>(
__cvta_generic_to_shared(
source + current_row * source_ld + col));
const float* accumulator_global =
accumulator_source
+ accumulator_index
+ current_row * output_ld
+ col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(accumulator_shared), "l"(accumulator_global)
: "memory");
}
asm volatile("cp.async.commit_group;\n" ::);
}
// The 512-wide owner-computes DAG uses the same 64x64 MathDx operands as the
// direct 256 path, but keeps the matrix leading dimension explicit. These
// vectorized movers are deliberately separate from the fixed-ld helpers so
// the panel/update kernels can own individual tiles without scalar Numba I/O.
extern "C" __device__
void d64_cp_async_fp32_tile_ld(
float* destination,
const float* source,
long long source_index,
int global_ld) {
constexpr int shared_ld = 68;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
unsigned destination_shared = static_cast<unsigned>(
__cvta_generic_to_shared(
destination + current_row * shared_ld + col));
const float* source_global =
source + source_index + current_row * global_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(destination_shared), "l"(source_global));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d64_cp_async_fp32_pair_ld(
float* first,
float* second,
const float* first_source,
const float* second_source,
long long first_index,
long long second_index,
int global_ld) {
constexpr int shared_ld = 68;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const int shared_index = current_row * shared_ld + col;
unsigned first_shared = static_cast<unsigned>(
__cvta_generic_to_shared(first + shared_index));
unsigned second_shared = static_cast<unsigned>(
__cvta_generic_to_shared(second + shared_index));
const int global_offset = current_row * global_ld + col;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(first_shared),
"l"(first_source + first_index + global_offset));
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(second_shared),
"l"(second_source + second_index + global_offset));
}
asm volatile("cp.async.commit_group;\n" ::);
}
extern "C" __device__
void d64_store_fp32_and_half_tile_ld(
const float* source,
float* output,
__half* sidecar,
long long output_index,
long long sidecar_index,
int output_ld) {
constexpr int source_ld = 68;
constexpr int sidecar_ld = 64;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
*reinterpret_cast<float4*>(
output + output_index + current_row * output_ld + col) = values;
__half2* packed = reinterpret_cast<__half2*>(
sidecar + sidecar_index + current_row * sidecar_ld + col);
packed[0] = __floats2half2_rn(values.x, values.y);
packed[1] = __floats2half2_rn(values.z, values.w);
}
}
extern "C" __device__
void d64_store_fp32_tile_ld(
const float* source,
float* output,
long long output_index,
int output_ld) {
constexpr int source_ld = 68;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
*reinterpret_cast<float4*>(
output + output_index + current_row * output_ld + col) = values;
}
}
extern "C" __device__
void d64_store_fp32_upper_tile_ld(
const float* source,
float* output,
long long output_index,
int output_ld) {
constexpr int source_ld = 68;
const int row = threadIdx.x >> 4;
const int col = (threadIdx.x & 15) * 4;
#pragma unroll
for (int load = 0; load < 8; ++load) {
const int current_row = row + load * 8;
const float4 values = *reinterpret_cast<const float4*>(
source + current_row * source_ld + col);
float* destination =
output + output_index + current_row * output_ld + col;
if (current_row <= col) {
*reinterpret_cast<float4*>(destination) = values;
} else if (current_row <= col + 3) {
if (current_row <= col + 1) destination[1] = values.y;
if (current_row <= col + 2) destination[2] = values.z;
destination[3] = values.w;
}
}
}
extern "C" __device__
void dx_async_bulk_init(unsigned long long* barriers) {
if (threadIdx.x == 0) {
unsigned barrier_shared = static_cast<unsigned>(
__cvta_generic_to_shared(barriers));
asm volatile(
"mbarrier.init.shared::cta.b64 [%0], 1;\n"
:: "r"(barrier_shared) : "memory");
}
__syncthreads();
}
extern "C" __device__
void dx_async_bulk_wait(
unsigned long long* barriers,
int phase) {
if (threadIdx.x == 0) {
unsigned barrier_shared = static_cast<unsigned>(
__cvta_generic_to_shared(barriers));
asm volatile(
"{\n"
".reg .pred ready;\n"
"DX_BULK_WAIT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 ready, [%0], %1;\n"
"@ready bra DX_BULK_DONE;\n"
"bra DX_BULK_WAIT;\n"
"DX_BULK_DONE:\n"
"}\n"
:: "r"(barrier_shared), "r"(phase) : "memory");
}
__syncthreads();
}
extern "C" __device__
void dx_cp_async_wait() {
asm volatile("cp.async.wait_group 0;\n" ::);
}
""",
name="dx_async.cu",
)
_dx_cp_async_pair = cuda.declare_device(
"dx_cp_async_pair",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_b1024_cp_async_fp32_tile = cuda.declare_device(
"dx_b1024_cp_async_fp32_tile",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_b1024_cp_async_half_tile = cuda.declare_device(
"dx_b1024_cp_async_half_tile",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float16),
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_b1024_load_convert_half_tile = cuda.declare_device(
"dx_b1024_load_convert_half_tile",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float32),
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_copy_float8 = cuda.declare_device(
"dx_copy_float8",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_g2048_store_upper_zero_lower = cuda.declare_device(
"g2048_store_upper_zero_lower",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_g2048_store_panel_float4 = cuda.declare_device(
"g2048_store_panel_float4",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int64,
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_g2048_load_tile_float4 = cuda.declare_device(
"g2048_load_tile_float4",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_cp_async_half_tile = cuda.declare_device(
"dx_cp_async_half_tile",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float16),
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_cp_async_half_pair = cuda.declare_device(
"dx_cp_async_half_pair",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float16),
types.CPointer(types.float16),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_cp_async_bulk_half_triple = cuda.declare_device(
"dx_cp_async_bulk_half_triple",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float16),
types.CPointer(types.float16),
types.CPointer(types.float16),
types.int64,
types.int64,
types.int64,
types.CPointer(types.uint64),
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_cp_async_fp32_tile = cuda.declare_device(
"d256_cp_async_fp32_tile",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_cp_async_fp32_pair = cuda.declare_device(
"d256_cp_async_fp32_pair",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_cp_async_half_tile = cuda.declare_device(
"d256_cp_async_half_tile",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float16),
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_cp_async_half_pair = cuda.declare_device(
"d256_cp_async_half_pair",
types.void(
types.CPointer(types.float16),
types.CPointer(types.float16),
types.CPointer(types.float16),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_store_fp32_and_half_tile = cuda.declare_device(
"d256_store_fp32_and_half_tile",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float16),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_store_fp32_tile = cuda.declare_device(
"d256_store_fp32_tile",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_store_fp32_upper_tile = cuda.declare_device(
"d256_store_fp32_upper_tile",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d256_store_fp32_and_shared_half_prefetch = cuda.declare_device(
"d256_store_fp32_and_shared_half_prefetch",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float16),
types.CPointer(types.float32),
types.int64,
types.int64,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d64_cp_async_fp32_tile_ld = cuda.declare_device(
"d64_cp_async_fp32_tile_ld",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d64_cp_async_fp32_pair_ld = cuda.declare_device(
"d64_cp_async_fp32_pair_ld",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d64_store_fp32_and_half_tile_ld = cuda.declare_device(
"d64_store_fp32_and_half_tile_ld",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.CPointer(types.float16),
types.int64,
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d64_store_fp32_tile_ld = cuda.declare_device(
"d64_store_fp32_tile_ld",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_d64_store_fp32_upper_tile_ld = cuda.declare_device(
"d64_store_fp32_upper_tile_ld",
types.void(
types.CPointer(types.float32),
types.CPointer(types.float32),
types.int64,
types.int32,
),
link=_DX_ASYNC_SOURCE,
abi="c",
)
_dx_async_bulk_init = cuda.declare_device(
"dx_async_bulk_init",
types.void(types.CPointer(types.uint64)),
abi="c",
)
_dx_async_bulk_wait = cuda.declare_device(
"dx_async_bulk_wait",
types.void(types.CPointer(types.uint64), types.int32),
abi="c",
)
_dx_cp_async_wait = cuda.declare_device(
"dx_cp_async_wait", types.void(), abi="c"
)
_CPP = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContextLight.h>
#include <ATen/ops/triu.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <algorithm>
torch::Tensor cholesky32_cuda(torch::Tensor input);
torch::Tensor cholesky64_cuda(torch::Tensor input);
torch::Tensor cholesky128_cuda(torch::Tensor input);
torch::Tensor grouped_pointer_table_cuda(
torch::Tensor output, int64_t block_size);
namespace {
void check_cusolver(cusolverStatus_t status, const char* what) {
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, what, " failed with status ",
static_cast<int>(status));
}
void check_cublas(cublasStatus_t status, const char* what) {
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, what, " failed with status ",
static_cast<int>(status));
}
} // namespace
torch::Tensor direct_cholesky(
torch::Tensor input,
torch::Tensor output,
torch::Tensor pointers,
torch::Tensor info,
torch::Tensor workspace) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
TORCH_CHECK(input.dim() == 3 && input.size(1) == input.size(2),
"input must be a batch of square matrices");
TORCH_CHECK(output.is_cuda() && output.scalar_type() == at::kFloat,
"output must be CUDA float32");
TORCH_CHECK(output.is_contiguous() && output.sizes() == input.sizes(),
"output must be a same-sized contiguous tensor");
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
// cuSOLVER's column-major lower triangle is the row-major upper triangle.
// Copy just that source triangle and zero the unused half in one kernel.
at::triu_out(output, input, 0);
TORCH_CHECK(info.is_cuda() && info.scalar_type() == at::kInt &&
info.numel() == batch, "invalid info tensor");
auto handle = at::cuda::getCurrentCUDASolverDnHandle();
check_cusolver(
cusolverDnSetMathMode(handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
"cusolverDnSetMathMode");
const bool use_looped = batch <= 4 && n >= 1024;
if (!use_looped) {
TORCH_CHECK(pointers.is_cuda() &&
pointers.scalar_type() == at::kLong &&
pointers.numel() == batch, "invalid pointer array");
check_cusolver(
cusolverDnSpotrfBatched(
handle, CUBLAS_FILL_MODE_LOWER, n,
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()), n,
info.data_ptr<int>(), batch),
"cusolverDnSpotrfBatched");
} else {
TORCH_CHECK(workspace.is_cuda() &&
workspace.scalar_type() == at::kFloat,
"invalid workspace tensor");
const int64_t matrix_stride = static_cast<int64_t>(n) * n;
for (int i = 0; i < batch; ++i) {
check_cusolver(
cusolverDnSpotrf(
handle, CUBLAS_FILL_MODE_LOWER, n,
output.data_ptr<float>() + i * matrix_stride, n,
workspace.data_ptr<float>(), workspace.numel(),
info.data_ptr<int>() + i),
"cusolverDnSpotrf");
}
}
// The row-major upper factor becomes an F-contiguous lower-triangular view.
return output.transpose(-2, -1);
}
torch::Tensor grouped_blocked_cholesky(
torch::Tensor input,
torch::Tensor output,
torch::Tensor pointers,
torch::Tensor info,
int64_t block_size) {
TORCH_CHECK(input.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(input.scalar_type() == at::kFloat &&
output.scalar_type() == at::kFloat, "FP32 tensors required");
TORCH_CHECK(input.is_contiguous() && output.is_contiguous() &&
input.sizes() == output.sizes(), "invalid matrix buffers");
TORCH_CHECK(input.dim() == 3 && input.size(1) == input.size(2),
"square matrix batches required");
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
const int block = static_cast<int>(block_size);
const int panels = (n + block - 1) / block;
TORCH_CHECK(block > 0 && block <= n, "invalid block size");
TORCH_CHECK(pointers.is_cuda() &&
pointers.scalar_type() == at::kLong &&
pointers.numel() == static_cast<int64_t>(panels) * 2 * batch,
"invalid grouped pointer tables");
TORCH_CHECK(info.is_cuda() && info.scalar_type() == at::kInt &&
info.numel() == batch, "invalid info tensor");
// Treat row-major upper storage as column-major lower storage.
at::triu_out(output, input, 0);
auto solver = at::cuda::getCurrentCUDASolverDnHandle();
auto blas = at::cuda::getCurrentCUDABlasHandle();
check_cusolver(
cusolverDnSetMathMode(solver, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
"cusolverDnSetMathMode");
const float one = 1.0f;
const float minus_one = -1.0f;
const int64_t matrix_stride = static_cast<int64_t>(n) * n;
auto pointer_data = pointers.data_ptr<int64_t>();
float* output_data = output.data_ptr<float>();
for (int panel_index = 0, k = 0; k < n;
++panel_index, k += block) {
const int width = std::min(block, n - k);
const int stop = k + width;
float** diagonal = reinterpret_cast<float**>(
pointer_data + static_cast<int64_t>(panel_index) * 2 * batch
);
check_cusolver(
cusolverDnSpotrfBatched(
solver, CUBLAS_FILL_MODE_LOWER, width, diagonal, n,
info.data_ptr<int>(), batch),
"grouped diagonal potrf");
if (stop == n) {
continue;
}
const int trailing = n - stop;
float* const* panel = reinterpret_cast<float* const*>(
pointer_data +
static_cast<int64_t>(panel_index) * 2 * batch + batch
);
check_cublas(
cublasStrsmBatched(
blas, CUBLAS_SIDE_RIGHT, CUBLAS_FILL_MODE_LOWER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
trailing, width, &one,
reinterpret_cast<const float* const*>(diagonal), n,
panel, n, batch),
"grouped panel trsm");
float* first_panel = output_data +
static_cast<int64_t>(k) * n + stop;
float* first_trailing = output_data +
static_cast<int64_t>(stop) * n + stop;
check_cublas(
cublasGemmStridedBatchedEx(
blas, CUBLAS_OP_N, CUBLAS_OP_T,
trailing, trailing, width,
&minus_one,
first_panel, CUDA_R_32F, n, matrix_stride,
first_panel, CUDA_R_32F, n, matrix_stride,
&one,
first_trailing, CUDA_R_32F, n, matrix_stride,
batch, CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"grouped trailing gemm");
}
// Full GEMMs touched both halves; retain only the physical upper factor.
at::triu_out(output, output, 0);
return output.transpose(-2, -1);
}
void grouped_upper_copy_cuda(
torch::Tensor input, torch::Tensor output, torch::Tensor flags);
void grouped_convert_panel_cuda(
torch::Tensor panel, torch::Tensor converted);
void grouped_prepare_(
torch::Tensor input,
torch::Tensor output,
torch::Tensor flags) {
TORCH_CHECK(input.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(input.scalar_type() == at::kFloat &&
output.scalar_type() == at::kFloat, "FP32 tensors required");
TORCH_CHECK(input.is_contiguous() && output.is_contiguous() &&
input.sizes() == output.sizes(), "invalid matrix buffers");
TORCH_CHECK(input.dim() == 3 && input.size(1) == input.size(2),
"square matrix batches required");
TORCH_CHECK(flags.is_cuda() && flags.scalar_type() == at::kInt,
"CUDA int32 flags required");
grouped_upper_copy_cuda(input, output, flags);
}
void grouped_panel_update_(
torch::Tensor output,
torch::Tensor pointers,
int64_t block_size,
int64_t panel_index_value) {
TORCH_CHECK(output.is_cuda() && output.scalar_type() == at::kFloat &&
output.is_contiguous(), "contiguous CUDA FP32 output required");
TORCH_CHECK(output.dim() == 3 && output.size(1) == output.size(2),
"square matrix batches required");
const int batch = static_cast<int>(output.size(0));
const int n = static_cast<int>(output.size(1));
const int block = static_cast<int>(block_size);
const int panels = (n + block - 1) / block;
const int panel_index = static_cast<int>(panel_index_value);
TORCH_CHECK(block > 0 && block <= n &&
panel_index >= 0 && panel_index < panels,
"invalid grouped panel");
TORCH_CHECK(pointers.is_cuda() &&
pointers.scalar_type() == at::kLong &&
pointers.numel() == static_cast<int64_t>(panels) * 2 * batch,
"invalid grouped pointer tables");
const int k = panel_index * block;
const int width = std::min(block, n - k);
const int stop = k + width;
if (stop == n) {
return;
}
const int trailing = n - stop;
auto pointer_data = pointers.data_ptr<int64_t>();
float** diagonal = reinterpret_cast<float**>(
pointer_data + static_cast<int64_t>(panel_index) * 2 * batch
);
float* const* panel = reinterpret_cast<float* const*>(
pointer_data + static_cast<int64_t>(panel_index) * 2 * batch + batch
);
auto blas = at::cuda::getCurrentCUDABlasHandle();
const float one = 1.0f;
check_cublas(
cublasStrsmBatched(
blas, CUBLAS_SIDE_RIGHT, CUBLAS_FILL_MODE_LOWER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
trailing, width, &one,
reinterpret_cast<const float* const*>(diagonal), n,
panel, n, batch),
"grouped panel trsm");
}
void grouped_panel_update_one_(
torch::Tensor output,
int64_t block_size,
int64_t panel_index_value) {
TORCH_CHECK(output.is_cuda() && output.scalar_type() == at::kFloat &&
output.is_contiguous(), "contiguous CUDA FP32 output required");
TORCH_CHECK(output.dim() == 3 && output.size(0) == 1 &&
output.size(1) == output.size(2),
"one square matrix required");
const int n = static_cast<int>(output.size(1));
const int block = static_cast<int>(block_size);
const int panels = (n + block - 1) / block;
const int panel_index = static_cast<int>(panel_index_value);
TORCH_CHECK(block > 0 && block <= n &&
panel_index >= 0 && panel_index < panels,
"invalid grouped panel");
const int k = panel_index * block;
const int width = std::min(block, n - k);
const int stop = k + width;
if (stop == n) {
return;
}
const int trailing = n - stop;
float* output_data = output.data_ptr<float>();
const float* diagonal = output_data + static_cast<int64_t>(k) * n + k;
float* panel = output_data + static_cast<int64_t>(k) * n + stop;
auto blas = at::cuda::getCurrentCUDABlasHandle();
const float one = 1.0f;
cublasMath_t prior_math = CUBLAS_DEFAULT_MATH;
check_cublas(cublasGetMathMode(blas, &prior_math),
"get panel math mode");
check_cublas(cublasSetMathMode(blas, CUBLAS_TF32_TENSOR_OP_MATH),
"set panel TF32 math mode");
check_cublas(
cublasStrsm(
blas, CUBLAS_SIDE_RIGHT, CUBLAS_FILL_MODE_LOWER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
trailing, width, &one, diagonal, n, panel, n),
"single grouped panel trsm");
check_cublas(cublasSetMathMode(blas, prior_math),
"restore panel math mode");
}
void grouped_zero_lower_cuda(torch::Tensor output);
torch::Tensor grouped_finish(torch::Tensor output) {
grouped_zero_lower_cuda(output);
return output.transpose(-2, -1);
}
void fp16_update_(torch::Tensor trailing, torch::Tensor panel) {
TORCH_CHECK(trailing.is_cuda() && panel.is_cuda(), "CUDA tensors required");
TORCH_CHECK(trailing.scalar_type() == at::kFloat,
"FP32 trailing matrix required");
TORCH_CHECK(panel.scalar_type() == at::kHalf,
"FP16 panel required");
TORCH_CHECK(trailing.dim() == 3 && panel.dim() == 3,
"batched matrices required");
const int batch = static_cast<int>(panel.size(0));
const int n = static_cast<int>(panel.size(1));
const int k = static_cast<int>(panel.size(2));
const int lda = static_cast<int>(panel.stride(1));
const int ldc = static_cast<int>(trailing.stride(1));
const int64_t panel_batch_stride = panel.stride(0);
const int64_t trailing_batch_stride = trailing.stride(0);
const float alpha = -1.0f;
const float beta = 1.0f;
auto handle = at::cuda::getCurrentCUDABlasHandle();
check_cublas(
cublasGemmStridedBatchedEx(
handle, CUBLAS_OP_T, CUBLAS_OP_N, n, n, k,
&alpha,
panel.data_ptr<at::Half>(), CUDA_R_16F, lda,
panel_batch_stride,
panel.data_ptr<at::Half>(), CUDA_R_16F, lda,
panel_batch_stride,
&beta,
trailing.data_ptr<float>(), CUDA_R_32F, ldc,
trailing_batch_stride,
batch, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT),
"cublasGemmStridedBatchedEx");
}
void trsm_(torch::Tensor factor, torch::Tensor panel) {
TORCH_CHECK(factor.is_cuda() && panel.is_cuda(), "CUDA tensors required");
TORCH_CHECK(factor.scalar_type() == at::kFloat &&
panel.scalar_type() == at::kFloat,
"FP32 factor and panel required");
TORCH_CHECK(factor.dim() == 3 && panel.dim() == 3 &&
factor.size(0) == 1 && panel.size(0) == 1,
"single batched matrices required");
TORCH_CHECK(factor.size(1) == factor.size(2) &&
panel.size(2) == factor.size(1),
"incompatible triangular solve shapes");
const int k = static_cast<int>(factor.size(1));
const int m = static_cast<int>(panel.size(1));
TORCH_CHECK(factor.stride(1) == 1,
"column-major triangular factor required");
const int lda = static_cast<int>(factor.stride(2));
const int ldb = static_cast<int>(panel.stride(1));
const float alpha = 1.0f;
auto handle = at::cuda::getCurrentCUDABlasHandle();
check_cublas(
cublasStrsm(
handle, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_LOWER,
CUBLAS_OP_N, CUBLAS_DIAG_NON_UNIT, k, m, &alpha,
factor.data_ptr<float>(), lda, panel.data_ptr<float>(), ldb),
"cublasStrsm");
}
void lower_copy_(torch::Tensor input, torch::Tensor output);
void lower_panel_copy_(
torch::Tensor input, torch::Tensor output, int64_t width);
void copy_panel_to_factor_cuda(
torch::Tensor input, torch::Tensor factor);
void copy_factor_stack_lower_(
torch::Tensor factors, torch::Tensor output);
void prepare_batch_factor_and_pointers_cuda(
torch::Tensor factor,
torch::Tensor diagonal,
torch::Tensor panel,
torch::Tensor pointers);
void batch_factor_copy_cuda(
torch::Tensor factor,
torch::Tensor diagonal,
torch::Tensor output);
void convert_batch_panel_cuda(
torch::Tensor panel, torch::Tensor converted);
void batch_lower_copy_cuda(
torch::Tensor input, torch::Tensor output);
void gather_batch_factor_and_pointers_cuda(
torch::Tensor diagonal,
torch::Tensor factor,
torch::Tensor pointers);
void panel_cholesky_(
torch::Tensor input,
torch::Tensor factor,
torch::Tensor info,
torch::Tensor workspace) {
TORCH_CHECK(input.is_cuda() && factor.is_cuda(), "CUDA tensors required");
TORCH_CHECK(input.scalar_type() == at::kFloat &&
factor.scalar_type() == at::kFloat,
"FP32 tensors required");
TORCH_CHECK(input.dim() == 3 && factor.dim() == 3 &&
input.sizes() == factor.sizes() &&
input.size(0) == 1 && input.size(1) == 4096 &&
input.stride(2) == 1 && factor.stride(1) == 1 &&
factor.stride(2) == 4096,
"compatible 4096-square panel buffers required");
TORCH_CHECK(info.is_cuda() && info.scalar_type() == at::kInt &&
info.numel() == 1, "invalid panel info tensor");
TORCH_CHECK(workspace.is_cuda() &&
workspace.scalar_type() == at::kFloat,
"invalid panel workspace tensor");
copy_panel_to_factor_cuda(input, factor);
auto handle = at::cuda::getCurrentCUDASolverDnHandle();
check_cusolver(
cusolverDnSetMathMode(handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
"cusolverDnSetMathMode");
check_cusolver(
cusolverDnSpotrf(
handle, CUBLAS_FILL_MODE_LOWER, 4096,
factor.data_ptr<float>(), 4096,
workspace.data_ptr<float>(), workspace.numel(),
info.data_ptr<int>()),
"panel cusolverDnSpotrf");
}
void compact_factor_cholesky_(
torch::Tensor factor,
torch::Tensor info,
torch::Tensor workspace) {
TORCH_CHECK(factor.is_cuda() && factor.scalar_type() == at::kFloat,
"CUDA FP32 factor required");
TORCH_CHECK(factor.dim() == 3 && factor.size(0) == 1 &&
factor.size(1) == factor.size(2) &&
factor.stride(1) == 1 &&
factor.stride(2) >= factor.size(1),
"column-major square factor required");
TORCH_CHECK(info.is_cuda() && info.scalar_type() == at::kInt &&
info.numel() == 1, "invalid compact factor info");
TORCH_CHECK(workspace.is_cuda() &&
workspace.scalar_type() == at::kFloat,
"invalid compact factor workspace");
const int n = static_cast<int>(factor.size(1));
const int lda = static_cast<int>(factor.stride(2));
auto handle = at::cuda::getCurrentCUDASolverDnHandle();
check_cusolver(
cusolverDnSetMathMode(handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
"compact cusolverDnSetMathMode");
check_cusolver(
cusolverDnSpotrf(
handle, CUBLAS_FILL_MODE_LOWER, n,
factor.data_ptr<float>(), lda,
workspace.data_ptr<float>(), workspace.numel(),
info.data_ptr<int>()),
"compact cusolverDnSpotrf");
}
void batch_factor_gather_(
torch::Tensor diagonal,
torch::Tensor factor,
torch::Tensor pointers
) {
TORCH_CHECK(
factor.is_cuda() && factor.scalar_type() == at::kFloat
&& factor.dim() == 3 && factor.size(1) == 256
&& factor.size(2) == 256 && factor.stride(1) == 1
&& factor.stride(2) == 256,
"column-major batch256 factor workspace required"
);
gather_batch_factor_and_pointers_cuda(diagonal, factor, pointers);
}
void batch_panel_solve_(
torch::Tensor factor,
torch::Tensor panel,
torch::Tensor pointers
) {
TORCH_CHECK(factor.is_cuda() && panel.is_cuda(), "CUDA tensors required");
TORCH_CHECK(
factor.scalar_type() == at::kFloat
&& panel.scalar_type() == at::kFloat,
"FP32 tensors required"
);
TORCH_CHECK(
factor.dim() == 3 && panel.dim() == 3
&& factor.size(0) == panel.size(0)
&& factor.size(1) == factor.size(2)
&& factor.size(1) == panel.size(2),
"invalid factor/panel shapes"
);
TORCH_CHECK(panel.stride(2) == 1, "panel columns must be contiguous");
const int batch = static_cast<int>(factor.size(0));
const int k = static_cast<int>(factor.size(1));
const int columns = static_cast<int>(panel.size(1));
cublasFillMode_t uplo;
int lda;
if (factor.stride(1) == 1) {
uplo = CUBLAS_FILL_MODE_LOWER;
lda = static_cast<int>(factor.stride(2));
} else {
TORCH_CHECK(
factor.stride(2) == 1,
"factor must be row-major or column-major contiguous"
);
uplo = CUBLAS_FILL_MODE_UPPER;
lda = static_cast<int>(factor.stride(1));
}
const int ldb = static_cast<int>(panel.stride(1));
const float alpha = 1.0f;
TORCH_CHECK(
pointers.is_cuda() && pointers.scalar_type() == at::kLong
&& pointers.numel() == 2 * batch,
"invalid TRSM pointer storage"
);
auto factor_pointers = reinterpret_cast<const float* const*>(
pointers.data_ptr<int64_t>()
);
auto panel_pointers = reinterpret_cast<float* const*>(
pointers.data_ptr<int64_t>() + batch
);
auto handle = at::cuda::getCurrentCUDABlasHandle();
check_cublas(
cublasStrsmBatched(
handle,
CUBLAS_SIDE_LEFT,
uplo,
CUBLAS_OP_N,
CUBLAS_DIAG_NON_UNIT,
k,
columns,
&alpha,
factor_pointers,
lda,
panel_pointers,
ldb,
batch
),
"cublasStrsmBatched"
);
}
void batch_panel_prepare_(
torch::Tensor factor,
torch::Tensor diagonal,
torch::Tensor panel,
torch::Tensor converted,
torch::Tensor pointers
) {
prepare_batch_factor_and_pointers_cuda(
factor, diagonal, panel, pointers
);
batch_panel_solve_(factor, panel, pointers);
convert_batch_panel_cuda(panel, converted);
}
int64_t workspace_size(torch::Tensor input) {
const int n = static_cast<int>(input.size(-1));
int lwork = 0;
auto handle = at::cuda::getCurrentCUDASolverDnHandle();
check_cusolver(
cusolverDnSetMathMode(handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
"cusolverDnSetMathMode");
check_cusolver(
cusolverDnSpotrf_bufferSize(
handle, CUBLAS_FILL_MODE_LOWER, n, input.data_ptr<float>(), n,
&lwork),
"cusolverDnSpotrf_bufferSize");
return lwork;
}
void convert_panel_(torch::Tensor panel, torch::Tensor compact);
void pack_panel_half_transpose_(
torch::Tensor panel, torch::Tensor destination);
void emit_solved_panel_(
torch::Tensor solved, torch::Tensor panel, torch::Tensor compact);
"""
_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <mma.h>
#include <algorithm>
#include <array>
#include <cstdint>
#define AC_CAT_I(a, b) a##b
#define AC_CAT(a, b) AC_CAT_I(a, b)
namespace wmma = nvcuda::wmma;
__global__ void batch_lower_copy_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int n,
int batch
) {
const int total_rows = batch * n;
const int warp = static_cast<int>(threadIdx.x) >> 5;
const int lane = static_cast<int>(threadIdx.x) & 31;
constexpr int warps_per_block = 8;
for (int linear_row =
static_cast<int>(blockIdx.x) * warps_per_block + warp;
linear_row < total_rows;
linear_row += static_cast<int>(gridDim.x) * warps_per_block) {
const int row = linear_row % n;
const long long base = static_cast<long long>(linear_row) * n;
const int full_vectors = (row + 1) / 4;
const float4* input_vectors = reinterpret_cast<const float4*>(
input + base
);
float4* output_vectors = reinterpret_cast<float4*>(output + base);
for (int vector = lane; vector < full_vectors;
vector += 32) {
output_vectors[vector] = input_vectors[vector];
}
const int tail = full_vectors * 4;
for (int col = tail + lane; col <= row; col += 32) {
output[base + col] = input[base + col];
}
}
}
void batch_lower_copy_cuda(torch::Tensor input, torch::Tensor output) {
TORCH_CHECK(input.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(
input.scalar_type() == torch::kFloat32
&& output.scalar_type() == torch::kFloat32,
"FP32 tensors required"
);
TORCH_CHECK(
input.dim() == 3 && input.size(1) == input.size(2)
&& input.is_contiguous() && output.is_contiguous()
&& input.sizes() == output.sizes(),
"matching contiguous batched square matrices required"
);
const int n = static_cast<int>(input.size(1));
const int batch = static_cast<int>(input.size(0));
TORCH_CHECK((n & 3) == 0, "matrix width must be float4 aligned");
const int rows = batch * n;
const int row_groups = (rows + 7) / 8;
// NCU measured six resident blocks per SM and a partial-wave tail on B200.
const int blocks = row_groups < 888 ? row_groups : 888;
batch_lower_copy_kernel<<<
blocks, 256, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
input.data_ptr<float>(), output.data_ptr<float>(), n, batch
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void grouped_upper_copy_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int* __restrict__ flags,
int flag_count,
int n,
int batch
) {
for (int index = static_cast<int>(blockIdx.x) * blockDim.x
+ static_cast<int>(threadIdx.x);
index < flag_count;
index += static_cast<int>(gridDim.x) * blockDim.x) {
flags[index] = 0;
}
const int total_rows = batch * n;
const int warp = static_cast<int>(threadIdx.x) >> 5;
const int lane = static_cast<int>(threadIdx.x) & 31;
for (int linear_row = static_cast<int>(blockIdx.x) * 16 + warp;
linear_row < total_rows;
linear_row += static_cast<int>(gridDim.x) * 16) {
const int row = linear_row % n;
const long long base = static_cast<long long>(linear_row) * n;
const int aligned_col = (row + 3) & ~3;
const int scalar_stop = aligned_col < n ? aligned_col : n;
for (int col = row + lane; col < scalar_stop; col += 32) {
output[base + col] = input[base + col];
}
const int first_vector = aligned_col / 4;
const int vector_count = n / 4;
const float4* input_vectors = reinterpret_cast<const float4*>(
input + base
);
float4* output_vectors = reinterpret_cast<float4*>(output + base);
for (int vector = first_vector + lane;
vector < vector_count; vector += 64) {
output_vectors[vector] = input_vectors[vector];
const int second = vector + 32;
if (second < vector_count) {
output_vectors[second] = input_vectors[second];
}
}
}
}
void grouped_upper_copy_cuda(
torch::Tensor input,
torch::Tensor output,
torch::Tensor flags) {
TORCH_CHECK(input.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(
input.scalar_type() == torch::kFloat32
&& output.scalar_type() == torch::kFloat32,
"FP32 tensors required"
);
TORCH_CHECK(
input.dim() == 3 && input.size(1) == input.size(2)
&& input.is_contiguous() && output.is_contiguous()
&& input.sizes() == output.sizes(),
"matching contiguous batched square matrices required"
);
TORCH_CHECK(
flags.is_cuda() && flags.scalar_type() == at::kInt,
"CUDA int32 flags required"
);
const int n = static_cast<int>(input.size(1));
const int batch = static_cast<int>(input.size(0));
TORCH_CHECK((n & 3) == 0, "matrix width must be float4 aligned");
const int rows = batch * n;
const int row_groups = (rows + 15) / 16;
const int blocks = row_groups < 1024 ? row_groups : 1024;
grouped_upper_copy_kernel<<<
blocks, 512, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
input.data_ptr<float>(), output.data_ptr<float>(),
flags.data_ptr<int>(), static_cast<int>(flags.numel()), n, batch
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void grouped_convert_panel_kernel(
const float* __restrict__ panel,
__half* __restrict__ converted,
int panel_rows,
int trailing_rows,
long long panel_stride_0,
long long panel_stride_1,
long long converted_stride_0,
long long converted_stride_1
) {
__shared__ float tile[32][33];
const int batch = static_cast<int>(blockIdx.z);
const int source_col = static_cast<int>(blockIdx.x) * 32 + threadIdx.x;
const int source_row = static_cast<int>(blockIdx.y) * 32 + threadIdx.y;
for (int offset = 0; offset < 32; offset += 8) {
if (source_row + offset < panel_rows
&& source_col < trailing_rows) {
tile[threadIdx.y + offset][threadIdx.x] = panel[
static_cast<long long>(batch) * panel_stride_0
+ static_cast<long long>(source_row + offset) * panel_stride_1
+ source_col
];
}
}
__syncthreads();
const int output_col = static_cast<int>(blockIdx.y) * 32 + threadIdx.x;
const int output_row = static_cast<int>(blockIdx.x) * 32 + threadIdx.y;
for (int offset = 0; offset < 32; offset += 8) {
if (output_row + offset < trailing_rows
&& output_col < panel_rows) {
converted[
static_cast<long long>(batch) * converted_stride_0
+ static_cast<long long>(output_row + offset)
* converted_stride_1
+ output_col
] = __float2half_rn(tile[threadIdx.x][threadIdx.y + offset]);
}
}
}
void grouped_convert_panel_cuda(
torch::Tensor panel, torch::Tensor converted) {
TORCH_CHECK(panel.is_cuda() && converted.is_cuda(),
"CUDA tensors required");
TORCH_CHECK(panel.scalar_type() == torch::kFloat32
&& converted.scalar_type() == torch::kFloat16,
"FP32 panel and FP16 output required");
TORCH_CHECK(panel.dim() == 3 && converted.dim() == 3
&& panel.size(0) == converted.size(0)
&& panel.size(1) == converted.size(2)
&& panel.size(2) == converted.size(1)
&& panel.stride(2) == 1 && converted.stride(2) == 1,
"transposed panel shapes required");
const int panel_rows = static_cast<int>(panel.size(1));
const int trailing_rows = static_cast<int>(panel.size(2));
const dim3 threads(32, 8, 1);
const dim3 blocks(
(trailing_rows + 31) / 32,
(panel_rows + 31) / 32,
static_cast<unsigned int>(panel.size(0))
);
grouped_convert_panel_kernel<<<
blocks, threads, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
panel.data_ptr<float>(),
reinterpret_cast<__half*>(converted.data_ptr<at::Half>()),
panel_rows,
trailing_rows,
panel.stride(0),
panel.stride(1),
converted.stride(0),
converted.stride(1)
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void grouped_zero_lower_kernel(
float* __restrict__ output,
int n,
int batch
) {
const int total_rows = batch * n;
const float4 zero = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
const int warps_per_block = static_cast<int>(blockDim.x) >> 5;
const int warp = static_cast<int>(threadIdx.x) >> 5;
const int lane = static_cast<int>(threadIdx.x) & 31;
for (int linear_row =
static_cast<int>(blockIdx.x) * warps_per_block + warp;
linear_row < total_rows;
linear_row += static_cast<int>(gridDim.x) * warps_per_block) {
const int row = linear_row % n;
const long long base = static_cast<long long>(linear_row) * n;
const int diagonal_start = row & ~255;
const int diagonal_row = row - diagonal_start;
const int full_vectors = diagonal_row / 4;
float4* output_vectors = reinterpret_cast<float4*>(
output + base + diagonal_start
);
for (int vector = lane; vector < full_vectors; vector += 64) {
output_vectors[vector] = zero;
const int second = vector + 32;
if (second < full_vectors) {
output_vectors[second] = zero;
}
}
const int tail = full_vectors * 4;
for (int col = tail + lane; col < diagonal_row; col += 32) {
output[base + diagonal_start + col] = 0.0f;
}
}
}
void grouped_zero_lower_cuda(torch::Tensor output) {
TORCH_CHECK(output.is_cuda(), "CUDA tensor required");
TORCH_CHECK(output.scalar_type() == torch::kFloat32,
"FP32 tensor required");
TORCH_CHECK(output.dim() == 3 && output.size(1) == output.size(2)
&& output.is_contiguous(),
"contiguous batched square matrix required");
const int n = static_cast<int>(output.size(1));
const int batch = static_cast<int>(output.size(0));
TORCH_CHECK((n & 3) == 0, "matrix width must be float4 aligned");
constexpr int threads = 128;
constexpr int rows_per_block = threads / 32;
const int rows = batch * n;
const int row_groups = (rows + rows_per_block - 1) / rows_per_block;
const int blocks = row_groups < 1024 ? row_groups : 1024;
grouped_zero_lower_kernel<<<
blocks, threads, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
output.data_ptr<float>(), n, batch
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void gather_batch_factor_and_pointers_kernel(
const float* __restrict__ diagonal,
float* __restrict__ factor,
int64_t* __restrict__ pointers,
int size,
int batch,
long long diagonal_stride,
long long diagonal_row_stride,
long long factor_stride,
long long factor_row_stride,
long long factor_col_stride
) {
int tile_col = static_cast<int>(blockIdx.x);
int tile_row = 0;
while (tile_col > tile_row) {
tile_col -= tile_row + 1;
++tile_row;
}
__shared__ float tile[32][33];
const int batch_index = static_cast<int>(blockIdx.z);
const int source_row_base = tile_row * 32;
const int source_col = tile_col * 32 + threadIdx.x;
for (int offset = 0; offset < 32; offset += 8) {
const int source_row = source_row_base + threadIdx.y + offset;
if (source_row < size && source_col < size
&& source_row >= source_col) {
tile[threadIdx.y + offset][threadIdx.x] = diagonal[
static_cast<long long>(batch_index) * diagonal_stride
+ static_cast<long long>(source_row) * diagonal_row_stride
+ source_col
];
}
}
__syncthreads();
const int factor_row = source_row_base + threadIdx.x;
const int factor_col_base = tile_col * 32 + threadIdx.y;
for (int offset = 0; offset < 32; offset += 8) {
const int factor_col = factor_col_base + offset;
if (factor_row < size && factor_col < size
&& factor_row >= factor_col) {
factor[
static_cast<long long>(batch_index) * factor_stride
+ static_cast<long long>(factor_row) * factor_row_stride
+ static_cast<long long>(factor_col) * factor_col_stride
] = tile[threadIdx.x][threadIdx.y + offset];
}
}
if (blockIdx.x == 0 && threadIdx.x == 0 && threadIdx.y == 0) {
pointers[batch_index] = reinterpret_cast<int64_t>(
factor + static_cast<long long>(batch_index) * factor_stride
);
}
}
void gather_batch_factor_and_pointers_cuda(
torch::Tensor diagonal,
torch::Tensor factor,
torch::Tensor pointers
) {
const int batch = static_cast<int>(factor.size(0));
TORCH_CHECK(
diagonal.is_cuda() && factor.is_cuda() && pointers.is_cuda()
&& diagonal.scalar_type() == torch::kFloat32
&& factor.scalar_type() == torch::kFloat32
&& pointers.scalar_type() == torch::kInt64,
"invalid batch factor gather types"
);
TORCH_CHECK(
diagonal.dim() == 3 && factor.dim() == 3
&& diagonal.sizes() == factor.sizes()
&& diagonal.stride(2) == 1
&& pointers.numel() == 2 * batch,
"invalid batch factor gather shapes"
);
const int size = static_cast<int>(factor.size(1));
const int tiles = (size + 31) / 32;
const dim3 threads(32, 8);
const dim3 blocks(tiles * (tiles + 1) / 2, 1, batch);
gather_batch_factor_and_pointers_kernel<<<
blocks, threads, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
diagonal.data_ptr<float>(),
factor.data_ptr<float>(),
pointers.data_ptr<int64_t>(),
size,
batch,
diagonal.stride(0),
diagonal.stride(1),
factor.stride(0),
factor.stride(1),
factor.stride(2)
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void prepare_batch_factor_and_pointers_kernel(
const float* factor,
float* diagonal,
float* panel,
int64_t* pointers,
int size,
int batch,
long long factor_stride,
long long factor_row_stride,
long long factor_col_stride,
long long diagonal_stride,
long long diagonal_row_stride,
long long panel_stride,
float* full_output,
int full_n
) {
__shared__ float tile[32][33];
const int batch_index = static_cast<int>(blockIdx.z);
if (full_output != nullptr) {
const int lane = static_cast<int>(threadIdx.x);
const int warp_index =
(static_cast<int>(blockIdx.y) * static_cast<int>(gridDim.x)
+ static_cast<int>(blockIdx.x))
* static_cast<int>(blockDim.y)
+ static_cast<int>(threadIdx.y);
const int warp_stride =
static_cast<int>(gridDim.x * gridDim.y * blockDim.y);
const long long output_base =
static_cast<long long>(batch_index) * full_n * full_n;
for (int row = warp_index; row < full_n; row += warp_stride) {
const long long row_base = output_base
+ static_cast<long long>(row) * full_n;
const int first = row + 1;
const int aligned = (first + 3) & ~3;
for (int col = first + lane;
col < aligned && col < full_n;
col += 32) {
full_output[row_base + col] = 0.0f;
}
float4* vectors = reinterpret_cast<float4*>(
full_output + row_base
);
for (int vector = aligned / 4 + lane;
vector < full_n / 4;
vector += 32) {
vectors[vector] = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
}
}
}
if (blockIdx.y > blockIdx.x) {
return;
}
const int factor_row = static_cast<int>(blockIdx.x) * 32 + threadIdx.x;
const int factor_col_base = static_cast<int>(blockIdx.y) * 32 + threadIdx.y;
for (int offset = 0; offset < 32; offset += 8) {
const int factor_col = factor_col_base + offset;
if (factor_row < size && factor_col < size
&& (blockIdx.x != blockIdx.y || factor_row >= factor_col)) {
tile[threadIdx.y + offset][threadIdx.x] = factor[
static_cast<long long>(batch_index) * factor_stride
+ static_cast<long long>(factor_row) * factor_row_stride
+ static_cast<long long>(factor_col) * factor_col_stride
];
}
}
__syncthreads();
const int diagonal_row = static_cast<int>(blockIdx.x) * 32 + threadIdx.y;
const int diagonal_col_base =
static_cast<int>(blockIdx.y) * 32 + threadIdx.x;
for (int offset = 0; offset < 32; offset += 8) {
if (diagonal_row + offset < size && diagonal_col_base < size
&& (blockIdx.x != blockIdx.y
|| diagonal_row + offset >= diagonal_col_base)) {
diagonal[
static_cast<long long>(batch_index) * diagonal_stride
+ static_cast<long long>(diagonal_row + offset)
* diagonal_row_stride
+ diagonal_col_base
] = tile[threadIdx.x][threadIdx.y + offset];
}
}
if (pointers != nullptr && blockIdx.x == 0 && blockIdx.y == 0
&& threadIdx.x == 0 && threadIdx.y == 0) {
pointers[batch_index] = reinterpret_cast<int64_t>(
factor + static_cast<long long>(batch_index) * factor_stride
);
pointers[batch + batch_index] = reinterpret_cast<int64_t>(
panel + static_cast<long long>(batch_index) * panel_stride
);
}
}
void prepare_batch_factor_and_pointers_cuda(
torch::Tensor factor,
torch::Tensor diagonal,
torch::Tensor panel,
torch::Tensor pointers
) {
const int batch = static_cast<int>(factor.size(0));
TORCH_CHECK(
factor.is_cuda() && diagonal.is_cuda() && panel.is_cuda()
&& pointers.is_cuda(),
"CUDA tensors required"
);
TORCH_CHECK(
factor.scalar_type() == torch::kFloat32
&& diagonal.scalar_type() == torch::kFloat32
&& panel.scalar_type() == torch::kFloat32
&& pointers.scalar_type() == torch::kInt64,
"invalid preparation tensor types"
);
TORCH_CHECK(
factor.dim() == 3 && diagonal.dim() == 3 && panel.dim() == 3
&& factor.sizes() == diagonal.sizes()
&& factor.size(0) == panel.size(0)
&& pointers.numel() == 2 * batch,
"invalid preparation tensor shapes"
);
TORCH_CHECK(diagonal.stride(2) == 1, "diagonal columns must be contiguous");
const int size = static_cast<int>(factor.size(1));
const dim3 threads(32, 8);
const dim3 blocks((size + 31) / 32, (size + 31) / 32, batch);
prepare_batch_factor_and_pointers_kernel<<<
blocks, threads, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
factor.data_ptr<float>(),
diagonal.data_ptr<float>(),
panel.data_ptr<float>(),
pointers.data_ptr<int64_t>(),
size,
batch,
factor.stride(0),
factor.stride(1),
factor.stride(2),
diagonal.stride(0),
diagonal.stride(1),
panel.stride(0),
nullptr,
0
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void batch_factor_copy_cuda(
torch::Tensor factor,
torch::Tensor diagonal,
torch::Tensor output
) {
const int batch = static_cast<int>(factor.size(0));
TORCH_CHECK(
factor.is_cuda() && diagonal.is_cuda() && output.is_cuda()
&& factor.scalar_type() == torch::kFloat32
&& diagonal.scalar_type() == torch::kFloat32
&& output.scalar_type() == torch::kFloat32,
"CUDA FP32 factors required"
);
TORCH_CHECK(
factor.dim() == 3 && diagonal.dim() == 3
&& factor.sizes() == diagonal.sizes()
&& diagonal.stride(2) == 1
&& output.dim() == 3 && output.size(0) == batch
&& output.size(1) == output.size(2) && output.is_contiguous(),
"invalid factor-copy tensors"
);
const int size = static_cast<int>(factor.size(1));
const dim3 threads(32, 8);
const dim3 blocks((size + 31) / 32, (size + 31) / 32, batch);
prepare_batch_factor_and_pointers_kernel<<<
blocks, threads, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
factor.data_ptr<float>(),
diagonal.data_ptr<float>(),
nullptr,
nullptr,
size,
batch,
factor.stride(0),
factor.stride(1),
factor.stride(2),
diagonal.stride(0),
diagonal.stride(1),
0,
output.data_ptr<float>(),
static_cast<int>(output.size(1))
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void convert_batch_panel_kernel(
__half* __restrict__ converted,
const float* __restrict__ source,
long long group_count,
int rows,
int col_groups,
long long converted_stride_0,
long long converted_stride_1,
long long source_stride_0,
long long source_stride_1
) {
const long long step =
static_cast<long long>(gridDim.x) * blockDim.x;
for (long long index =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
index < group_count;
index += step) {
const int col_group = static_cast<int>(index % col_groups);
const long long row_index = index / col_groups;
const int row = static_cast<int>(row_index % rows);
const long long batch = row_index / rows;
const int col = col_group * 4;
const float4 value = *reinterpret_cast<const float4*>(
source + batch * source_stride_0 + row * source_stride_1 + col
);
__half2* packed = reinterpret_cast<__half2*>(
converted
+ batch * converted_stride_0 + row * converted_stride_1 + col
);
packed[0] = __floats2half2_rn(value.x, value.y);
packed[1] = __floats2half2_rn(value.z, value.w);
}
}
void convert_batch_panel_cuda(
torch::Tensor panel, torch::Tensor converted
) {
TORCH_CHECK(panel.is_cuda() && converted.is_cuda(), "CUDA tensors required");
TORCH_CHECK(
panel.scalar_type() == torch::kFloat32
&& converted.scalar_type() == torch::kFloat16,
"expected FP32 panel and FP16 converted output"
);
TORCH_CHECK(
panel.sizes() == converted.sizes() && panel.dim() == 3,
"panel conversion shapes must match"
);
TORCH_CHECK(
panel.size(2) % 4 == 0 && panel.stride(2) == 1
&& converted.stride(2) == 1,
"panel columns must be float4 aligned and contiguous"
);
const long long group_count = panel.numel() / 4;
const int threads = 256;
const int blocks = static_cast<int>((group_count + threads - 1) / threads);
const int persistent_blocks = blocks < 4096 ? blocks : 4096;
convert_batch_panel_kernel<<<
persistent_blocks, threads, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
reinterpret_cast<__half*>(converted.data_ptr<at::Half>()),
panel.data_ptr<float>(),
group_count,
static_cast<int>(panel.size(1)),
static_cast<int>(panel.size(2) / 4),
converted.stride(0),
converted.stride(1),
panel.stride(0),
panel.stride(1)
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void lower_copy_kernel(
const float* source, float* output, int n, int width) {
const int row = static_cast<int>(blockIdx.x);
const int columns = row + 1 < width ? row + 1 : width;
const int full_vectors = columns >> 2;
const int64_t base = static_cast<int64_t>(row) * n;
const float4* source_vectors = reinterpret_cast<const float4*>(
source + base);
float4* output_vectors = reinterpret_cast<float4*>(output + base);
for (int vector = threadIdx.x; vector < full_vectors;
vector += blockDim.x) {
output_vectors[vector] = source_vectors[vector];
}
const int tail = full_vectors << 2;
for (int col = tail + threadIdx.x; col < columns; col += blockDim.x) {
output[base + col] = source[base + col];
}
}
void lower_copy_(torch::Tensor input, torch::Tensor output) {
TORCH_CHECK(input.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(input.scalar_type() == at::kFloat &&
output.scalar_type() == at::kFloat,
"FP32 tensors required");
TORCH_CHECK(input.dim() == 3 && input.size(0) == 1 &&
input.size(1) == input.size(2) &&
input.is_contiguous() && output.is_contiguous() &&
input.sizes() == output.sizes(),
"matching contiguous batch-one matrices required");
const int n = static_cast<int>(input.size(1));
TORCH_CHECK((n & 3) == 0, "matrix width must be float4 aligned");
lower_copy_kernel<<<
n, 256, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
input.data_ptr<float>(), output.data_ptr<float>(), n, n);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void lower_panel_copy_(
torch::Tensor input,
torch::Tensor output,
int64_t width_value) {
TORCH_CHECK(input.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(input.scalar_type() == at::kFloat &&
output.scalar_type() == at::kFloat,
"FP32 tensors required");
TORCH_CHECK(input.dim() == 3 && input.size(0) == 1 &&
input.size(1) == input.size(2) &&
input.is_contiguous() && output.is_contiguous() &&
input.sizes() == output.sizes(),
"matching contiguous batch-one matrices required");
const int n = static_cast<int>(input.size(1));
const int width = static_cast<int>(width_value);
TORCH_CHECK(width > 0 && width <= n && (width & 3) == 0,
"panel width must be valid and float4 aligned");
lower_copy_kernel<<<
n, 256, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
input.data_ptr<float>(), output.data_ptr<float>(), n, width);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void copy_panel_to_factor_kernel(
const float* source, float* factor, int leading_source) {
constexpr int tile_size = 64;
__shared__ float tile[tile_size][tile_size + 1];
const int task = static_cast<int>(blockIdx.x);
int tile_row = static_cast<int>(
(sqrtf(8.0f * static_cast<float>(task) + 1.0f) - 1.0f)
* 0.5f);
int base = tile_row * (tile_row + 1) / 2;
if (base > task) {
--tile_row;
base = tile_row * (tile_row + 1) / 2;
} else if (base + tile_row + 1 <= task) {
++tile_row;
base = tile_row * (tile_row + 1) / 2;
}
const int tile_col = task - base;
const int row_base = tile_row * tile_size;
const int col_base = tile_col * tile_size;
for (int y = threadIdx.y; y < tile_size; y += blockDim.y) {
for (int x = threadIdx.x; x < tile_size; x += blockDim.x) {
unsigned destination = static_cast<unsigned>(
__cvta_generic_to_shared(&tile[y][x]));
const float* source_element = source
+ static_cast<int64_t>(row_base + y) * leading_source
+ col_base + x;
asm volatile(
"cp.async.ca.shared.global.L2::128B [%0], [%1], 4;\n"
:: "r"(destination), "l"(source_element));
}
}
asm volatile("cp.async.commit_group;\n" ::);
asm volatile("cp.async.wait_group 0;\n" ::);
__syncthreads();
for (int y = threadIdx.y; y < tile_size; y += blockDim.y) {
for (int x = threadIdx.x; x < tile_size; x += blockDim.x) {
factor[
row_base + x
+ static_cast<int64_t>(col_base + y) * 4096
] = tile[x][y];
}
}
}
void copy_panel_to_factor_cuda(
torch::Tensor input, torch::Tensor factor) {
TORCH_CHECK(input.is_cuda() && factor.is_cuda(), "CUDA tensors required");
TORCH_CHECK(input.scalar_type() == at::kFloat &&
factor.scalar_type() == at::kFloat,
"FP32 tensors required");
TORCH_CHECK(input.dim() == 3 && factor.dim() == 3 &&
input.sizes() == factor.sizes() &&
input.size(0) == 1 && input.size(1) == 4096 &&
input.stride(2) == 1 && factor.stride(1) == 1 &&
factor.stride(2) == 4096,
"compatible 4096-square panel buffers required");
constexpr int tiles = 4096 / 64;
constexpr int lower_tile_tasks = tiles * (tiles + 1) / 2;
copy_panel_to_factor_kernel<<<
lower_tile_tasks, dim3(32, 16), 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
input.data_ptr<float>(), factor.data_ptr<float>(),
static_cast<int>(input.stride(1)));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void copy_factor_lower_kernel(
const float* factors, float* output, int leading_output) {
constexpr int tile_size = 32;
__shared__ float tile[tile_size][tile_size + 1];
const int panel = static_cast<int>(blockIdx.z);
const float* factor = factors
+ static_cast<int64_t>(panel) * 4096 * 4096;
output += static_cast<int64_t>(panel) * 4096
* (leading_output + 1);
const int tile_col = static_cast<int>(blockIdx.x);
const int tile_row = static_cast<int>(blockIdx.y);
if (tile_row < tile_col) {
const int row_base = tile_row * tile_size;
const int col_base = tile_col * tile_size;
for (int y = threadIdx.y; y < tile_size; y += blockDim.y) {
output[
static_cast<int64_t>(row_base + y) * leading_output
+ col_base + threadIdx.x
] = 0.0f;
}
return;
}
const int tx = threadIdx.x;
const int row_base = tile_row * tile_size;
const int col_base = tile_col * tile_size;
for (int y = threadIdx.y; y < tile_size; y += blockDim.y) {
tile[y][tx] = factor[
row_base + tx + static_cast<int64_t>(col_base + y) * 4096
];
}
__syncthreads();
for (int y = threadIdx.y; y < tile_size; y += blockDim.y) {
const int row = row_base + y;
const int col = col_base + tx;
output[static_cast<int64_t>(row) * leading_output + col] =
row >= col ? tile[tx][y] : 0.0f;
}
}
void copy_factor_stack_lower_(torch::Tensor factors, torch::Tensor output) {
TORCH_CHECK(factors.is_cuda() && output.is_cuda(), "CUDA tensors required");
TORCH_CHECK(factors.scalar_type() == at::kFloat &&
output.scalar_type() == at::kFloat,
"FP32 tensors required");
TORCH_CHECK(factors.dim() == 3 && output.dim() == 3 &&
output.size(0) == 1 && output.size(1) == output.size(2) &&
factors.size(1) == 4096 && factors.size(2) == 4096 &&
factors.size(0) * 4096 == output.size(1) &&
factors.stride(0) == 4096 * 4096 &&
factors.stride(1) == 1 && factors.stride(2) == 4096 &&
output.stride(2) == 1,
"compatible factor stack and output required");
constexpr int tiles = 4096 / 32;
copy_factor_lower_kernel<<<
dim3(tiles, tiles, factors.size(0)), dim3(32, 8), 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
factors.data_ptr<float>(), output.data_ptr<float>(),
static_cast<int>(output.stride(1)));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void convert_panel_kernel(
const float* source,
__half* compact,
int rows,
int leading,
int compact_leading) {
constexpr int groups_per_row = 512;
const int64_t total = static_cast<int64_t>(rows) * groups_per_row;
for (int64_t index = static_cast<int64_t>(blockIdx.x) * blockDim.x
+ threadIdx.x;
index < total;
index += static_cast<int64_t>(gridDim.x) * blockDim.x) {
const int row = static_cast<int>(index >> 9);
const int group = static_cast<int>(index & (groups_per_row - 1));
const float4* source_vectors = reinterpret_cast<const float4*>(
source + static_cast<int64_t>(row) * leading);
const float4 first = source_vectors[group * 2];
const float4 second = source_vectors[group * 2 + 1];
__half2* destination = reinterpret_cast<__half2*>(
compact + static_cast<int64_t>(row) * compact_leading);
destination[group * 4] = __floats2half2_rn(first.x, first.y);
destination[group * 4 + 1] = __floats2half2_rn(first.z, first.w);
destination[group * 4 + 2] = __floats2half2_rn(second.x, second.y);
destination[group * 4 + 3] = __floats2half2_rn(second.z, second.w);
}
}
void convert_panel_(torch::Tensor panel, torch::Tensor compact) {
TORCH_CHECK(panel.is_cuda() && compact.is_cuda(), "CUDA tensors required");
TORCH_CHECK(panel.scalar_type() == at::kFloat &&
compact.scalar_type() == at::kHalf,
"FP32 source and FP16 destination required");
TORCH_CHECK(panel.dim() == 2 && compact.dim() == 2 &&
panel.size(0) == compact.size(0) &&
panel.size(1) == 4096 && compact.size(1) == 4096,
"4096-column panel required");
TORCH_CHECK(panel.stride(1) == 1 && compact.stride(1) == 1,
"contiguous inner dimensions required");
TORCH_CHECK((reinterpret_cast<uintptr_t>(panel.data_ptr<float>()) & 15) == 0,
"16-byte aligned panel required");
const int64_t groups = panel.size(0) * 512;
const int blocks = static_cast<int>(
std::min<int64_t>((groups + 255) / 256, 4096));
convert_panel_kernel<<<
blocks, 256, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
panel.data_ptr<float>(),
reinterpret_cast<__half*>(compact.data_ptr<at::Half>()),
static_cast<int>(panel.size(0)),
static_cast<int>(panel.stride(0)),
static_cast<int>(compact.stride(0)));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void pack_panel_half_transpose_kernel(
const float* source,
__half* destination,
int rows,
int columns,
int source_leading,
int destination_leading) {
__shared__ __half tile[32][33];
const int source_col = static_cast<int>(blockIdx.x) * 32 + threadIdx.x;
const int source_row_base = static_cast<int>(blockIdx.y) * 32;
#pragma unroll
for (int offset = threadIdx.y; offset < 32; offset += blockDim.y) {
const int source_row = source_row_base + offset;
if (source_row < rows && source_col < columns) {
tile[offset][threadIdx.x] = __float2half_rn(
source[static_cast<int64_t>(source_row) * source_leading
+ source_col]);
}
}
__syncthreads();
const int destination_col = source_row_base + threadIdx.x;
#pragma unroll
for (int offset = threadIdx.y; offset < 32; offset += blockDim.y) {
const int destination_row =
static_cast<int>(blockIdx.x) * 32 + offset;
if (destination_row < columns && destination_col < rows) {
destination[
static_cast<int64_t>(destination_row) * destination_leading
+ destination_col
] = tile[threadIdx.x][offset];
}
}
}
void pack_panel_half_transpose_(
torch::Tensor panel, torch::Tensor destination) {
TORCH_CHECK(panel.is_cuda() && destination.is_cuda(),
"CUDA tensors required");
TORCH_CHECK(panel.scalar_type() == at::kFloat &&
destination.scalar_type() == at::kHalf,
"FP32 panel and FP16 destination required");
TORCH_CHECK(panel.dim() == 2 && destination.dim() == 2 &&
destination.size(0) == panel.size(1) &&
destination.size(1) == panel.size(0) &&
panel.stride(1) == 1 && destination.stride(1) == 1,
"destination must be the contiguous-inner transpose shape");
pack_panel_half_transpose_kernel<<<
dim3((panel.size(1) + 31) / 32, (panel.size(0) + 31) / 32),
dim3(32, 8), 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
panel.data_ptr<float>(),
reinterpret_cast<__half*>(destination.data_ptr<at::Half>()),
static_cast<int>(panel.size(0)),
static_cast<int>(panel.size(1)),
static_cast<int>(panel.stride(0)),
static_cast<int>(destination.stride(0))
);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void emit_solved_panel_kernel(
const float* source,
float* output,
__half* compact,
int rows,
int source_leading,
int output_leading,
int compact_leading) {
constexpr int groups_per_row = 512;
const int64_t total = static_cast<int64_t>(rows) * groups_per_row;
for (int64_t index = static_cast<int64_t>(blockIdx.x) * blockDim.x
+ threadIdx.x;
index < total;
index += static_cast<int64_t>(gridDim.x) * blockDim.x) {
const int row = static_cast<int>(index >> 9);
const int group = static_cast<int>(index & (groups_per_row - 1));
const float4* source_vectors = reinterpret_cast<const float4*>(
source + static_cast<int64_t>(row) * source_leading);
const float4 first = source_vectors[group * 2];
const float4 second = source_vectors[group * 2 + 1];
float4* output_vectors = reinterpret_cast<float4*>(
output + static_cast<int64_t>(row) * output_leading);
output_vectors[group * 2] = first;
output_vectors[group * 2 + 1] = second;
__half2* destination = reinterpret_cast<__half2*>(
compact + static_cast<int64_t>(row) * compact_leading);
destination[group * 4] = __floats2half2_rn(first.x, first.y);
destination[group * 4 + 1] = __floats2half2_rn(first.z, first.w);
destination[group * 4 + 2] = __floats2half2_rn(second.x, second.y);
destination[group * 4 + 3] = __floats2half2_rn(second.z, second.w);
}
}
void emit_solved_panel_(
torch::Tensor solved,
torch::Tensor panel,
torch::Tensor compact) {
TORCH_CHECK(solved.is_cuda() && panel.is_cuda() && compact.is_cuda(),
"CUDA tensors required");
TORCH_CHECK(solved.scalar_type() == at::kFloat &&
panel.scalar_type() == at::kFloat &&
compact.scalar_type() == at::kHalf,
"FP32 solved/output and FP16 compact tensors required");
TORCH_CHECK(solved.dim() == 2 && panel.dim() == 2 &&
compact.dim() == 2 && solved.sizes() == panel.sizes() &&
solved.sizes() == compact.sizes() &&
solved.size(1) == 4096 && solved.stride(1) == 1 &&
panel.stride(1) == 1 && compact.stride(1) == 1,
"compatible 4096-column panel views required");
const int64_t groups = solved.size(0) * 512;
const int blocks = static_cast<int>(
std::min<int64_t>((groups + 255) / 256, 4096));
emit_solved_panel_kernel<<<
blocks, 256, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
solved.data_ptr<float>(),
panel.data_ptr<float>(),
reinterpret_cast<__half*>(compact.data_ptr<at::Half>()),
static_cast<int>(solved.size(0)),
static_cast<int>(solved.stride(0)),
static_cast<int>(panel.stride(0)),
static_cast<int>(compact.stride(0)));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
__global__ void grouped_pointer_table_kernel(
float* output,
int64_t* pointers,
int batch,
int n,
int block,
int panels) {
const int index = static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x;
const int count = panels * 2 * batch;
if (index >= count) {
return;
}
const int matrix = index % batch;
const int table = index / batch;
const int kind = table & 1;
const int panel = table >> 1;
const int start = panel * block;
const int stop = min(start + block, n);
const int64_t matrix_base = static_cast<int64_t>(matrix) * n * n;
const int64_t element = kind == 0
? matrix_base + static_cast<int64_t>(start) * (n + 1)
: matrix_base + static_cast<int64_t>(start) * n + stop;
pointers[index] = reinterpret_cast<int64_t>(output + element);
}
torch::Tensor grouped_pointer_table_cuda(
torch::Tensor output, int64_t block_size) {
TORCH_CHECK(output.is_cuda() && output.scalar_type() == at::kFloat,
"CUDA FP32 output required");
TORCH_CHECK(output.dim() == 3 && output.size(1) == output.size(2) &&
output.is_contiguous(), "contiguous square batches required");
const int batch = static_cast<int>(output.size(0));
const int n = static_cast<int>(output.size(1));
const int block = static_cast<int>(block_size);
TORCH_CHECK(block > 0 && block <= n, "invalid block size");
const int panels = (n + block - 1) / block;
auto pointers = torch::empty(
{panels, 2, batch}, output.options().dtype(at::kLong));
const int count = panels * 2 * batch;
grouped_pointer_table_kernel<<<
(count + 127) / 128, 128, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()>>>(
output.data_ptr<float>(), pointers.data_ptr<int64_t>(),
batch, n, block, panels);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return pointers;
}
template <int N>
__global__ void fused_left_looking_tile(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr int LD = N + 1;
extern __shared__ float tile[];
const int tid = threadIdx.x;
const long long base = static_cast<long long>(blockIdx.x) * N * N;
const float* src = input + base;
float* dst = output + base;
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
tile[row * LD + col] = row >= col ? src[index] : 0.0f;
}
__syncthreads();
for (int k = 0; k < N; ++k) {
if (tid == 0) {
float diagonal = tile[k * LD + k];
for (int j = 0; j < k; ++j) {
const float value = tile[k * LD + j];
diagonal = fmaf(-value, value, diagonal);
}
tile[k * LD + k] = sqrtf(fmaxf(diagonal, 1.0e-30f));
}
__syncthreads();
const float diagonal = tile[k * LD + k];
for (int row = k + 1 + tid; row < N; row += blockDim.x) {
float value = tile[row * LD + k];
for (int j = 0; j < k; ++j) {
value = fmaf(-tile[row * LD + j], tile[k * LD + j], value);
}
tile[row * LD + k] = value / diagonal;
}
__syncthreads();
}
for (int index = tid; index < N * N; index += blockDim.x) {
const int row = index / N;
const int col = index - row * N;
dst[index] = tile[row * LD + col];
}
}
template <int K>
__device__ __forceinline__ void register_cholesky_step_32(
float (&row_values)[32],
int lane
) {
float value = row_values[K];
#pragma unroll
for (int j = 0; j < K; ++j) {
const float pivot = __shfl_sync(0xffffffff, row_values[j], K);
if (lane >= K) {
value = fmaf(-row_values[j], pivot, value);
}
}
float inverse = 0.0f;
if (lane == K) {
const float positive = fmaxf(value, 1.0e-30f);
inverse = rsqrtf(positive);
value = positive * inverse;
row_values[K] = value;
}
inverse = __shfl_sync(0xffffffff, inverse, K);
if (lane > K) {
row_values[K] = value * inverse;
}
if constexpr (K + 1 < 32) {
register_cholesky_step_32<K + 1>(row_values, lane);
}
}
__global__ void grouped_warp_cholesky_32(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
constexpr int N = 32;
constexpr int LD = 33;
constexpr int MATRICES_PER_BLOCK = 4;
__shared__ float storage[MATRICES_PER_BLOCK * N * LD];
const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32;
const int matrix = static_cast<int>(blockIdx.x) * MATRICES_PER_BLOCK + warp;
if (matrix >= batch) {
return;
}
float* tile = storage + warp * N * LD;
const int base = matrix * N * N;
const float* src = input + base;
float* dst = output + base;
#pragma unroll
for (int row_base = 0; row_base < N; row_base += 4) {
const int row = row_base + lane / 8;
const int col = (lane % 8) * 4;
float4 values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col <= row) {
values = *reinterpret_cast<const float4*>(
src + row * N + col
);
}
if (row < col) values.x = 0.0f;
if (row < col + 1) values.y = 0.0f;
if (row < col + 2) values.z = 0.0f;
if (row < col + 3) values.w = 0.0f;
tile[row * LD + col] = values.x;
tile[row * LD + col + 1] = values.y;
tile[row * LD + col + 2] = values.z;
tile[row * LD + col + 3] = values.w;
}
__syncwarp();
float row_values[N];
#pragma unroll
for (int col = 0; col < N; ++col) {
row_values[col] = tile[lane * LD + col];
}
register_cholesky_step_32<0>(row_values, lane);
#pragma unroll
for (int col = 0; col < N; ++col) {
tile[lane * LD + col] = row_values[col];
}
__syncwarp();
#pragma unroll
for (int row_base = 0; row_base < N; row_base += 4) {
const int row = row_base + lane / 8;
const int col = (lane % 8) * 4;
const float4 values = make_float4(
tile[row * LD + col],
tile[row * LD + col + 1],
tile[row * LD + col + 2],
tile[row * LD + col + 3]
);
*reinterpret_cast<float4*>(dst + row * N + col) = values;
}
}
template <int K>
__device__ __forceinline__ void register_cholesky_step_64(
float (&row_values)[64],
float* tile,
float* inverse,
int group,
int lane,
bool active
) {
constexpr int LD = 65;
float value = 0.0f;
if (active && lane >= K) {
value = row_values[K];
#pragma unroll
for (int j = 0; j < K; ++j) {
value = fmaf(-row_values[j], tile[K * LD + j], value);
}
}
if (active && lane == K) {
const float diagonal = sqrtf(fmaxf(value, 1.0e-30f));
row_values[K] = diagonal;
tile[K * LD + K] = diagonal;
inverse[group] = 1.0f / diagonal;
}
__syncthreads();
if (active && lane > K) {
value *= inverse[group];
row_values[K] = value;
tile[lane * LD + K] = value;
}
__syncthreads();
if constexpr (K + 1 < 64) {
register_cholesky_step_64<K + 1>(
row_values, tile, inverse, group, lane, active
);
}
}
__global__ void grouped_two_cholesky_64(
const float* __restrict__ input,
float* __restrict__ output,
int batch
) {
constexpr int N = 64;
constexpr int LD = 65;
constexpr int MATRICES_PER_BLOCK = 4;
extern __shared__ float scratch[];
float* inverse = scratch;
float* storage = scratch + MATRICES_PER_BLOCK;
const int group = threadIdx.x / N;
const int lane = threadIdx.x % N;
const int matrix = static_cast<int>(blockIdx.x) * MATRICES_PER_BLOCK + group;
const bool active = matrix < batch;
float* tile = storage + group * N * LD;
if (active) {
const int base = matrix * N * N;
const float* src = input + base;
#pragma unroll
for (int row = 0; row < N; ++row) {
tile[row * LD + lane] = row >= lane
? src[row * N + lane]
: 0.0f;
}
}
__syncthreads();
float row_values[N];
#pragma unroll
for (int col = 0; col < N; ++col) {
row_values[col] = active ? tile[lane * LD + col] : 0.0f;
}
register_cholesky_step_64<0>(
row_values, tile, inverse, group, lane, active
);
if (active) {
const int base = matrix * N * N;
float* dst = output + base;
#pragma unroll
for (int row = 0; row < N; ++row) {
dst[row * N + lane] = tile[row * LD + lane];
}
}
}
__device__ __forceinline__ void factor_panel_16_warp(
float* tile,
int panel
) {
constexpr int N = 128;
const int lane = threadIdx.x % 32;
if (lane < 16) {
float row_values[16];
#pragma unroll
for (int col = 0; col < 16; ++col) {
row_values[col] = lane >= col
? tile[(panel + lane) * N + panel + col]
: 0.0f;
}
#pragma unroll
for (int pivot = 0; pivot < 16; ++pivot) {
float value = row_values[pivot];
#pragma unroll
for (int col = 0; col < pivot; ++col) {
const float pivot_value = __shfl_sync(
0x0000ffff, row_values[col], pivot, 16
);
if (lane >= pivot) {
value = fmaf(-row_values[col], pivot_value, value);
}
}
float inverse = 0.0f;
if (lane == pivot) {
const float diagonal = sqrtf(fmaxf(value, 1.0e-30f));
row_values[pivot] = diagonal;
inverse = 1.0f / diagonal;
}
inverse = __shfl_sync(0x0000ffff, inverse, pivot, 16);
if (lane > pivot) {
row_values[pivot] = value * inverse;
}
}
#pragma unroll
for (int col = 0; col < 16; ++col) {
if (lane >= col) {
tile[(panel + lane) * N + panel + col] = row_values[col];
}
}
}
}
__device__ __forceinline__ void factor_panel_16(
float* tile,
int panel
) {
if (threadIdx.x / 32 == 0) {
factor_panel_16_warp(tile, panel);
}
__syncthreads();
}
__device__ __forceinline__ void solve_panel_16(
float* tile,
int panel
) {
constexpr int N = 128;
constexpr int TEAM_WIDTH = 4;
const int team = threadIdx.x / TEAM_WIDTH;
const int lane = threadIdx.x % TEAM_WIDTH;
const int stop = panel + 16;
const int row = stop + team;
const bool active = row < N;
float row_values[4];
#pragma unroll
for (int item = 0; item < 4; ++item) {
row_values[item] = active
? tile[row * N + panel + item * TEAM_WIDTH + lane]
: 0.0f;
}
#pragma unroll
for (int pivot = 0; pivot < 16; ++pivot) {
float dot = 0.0f;
#pragma unroll
for (int item = 0; item < 4; ++item) {
const int col = item * TEAM_WIDTH + lane;
if (col < pivot) {
dot = fmaf(
row_values[item],
tile[(panel + pivot) * N + panel + col],
dot
);
}
}
dot += __shfl_down_sync(0xffffffff, dot, 2, TEAM_WIDTH);
dot += __shfl_down_sync(0xffffffff, dot, 1, TEAM_WIDTH);
float original = __shfl_sync(
0xffffffff,
row_values[pivot / TEAM_WIDTH],
pivot % TEAM_WIDTH,
TEAM_WIDTH
);
float solved = lane == 0
? (original - dot)
/ tile[(panel + pivot) * N + panel + pivot]
: 0.0f;
solved = __shfl_sync(0xffffffff, solved, 0, TEAM_WIDTH);
if (lane == pivot % TEAM_WIDTH) {
row_values[pivot / TEAM_WIDTH] = solved;
}
}
if (active) {
#pragma unroll
for (int item = 0; item < 4; ++item) {
tile[row * N + panel + item * TEAM_WIDTH + lane] =
row_values[item];
}
}
__syncthreads();
}
__device__ __forceinline__ void compensated_rankk_16(
float* tile,
int panel
) {
constexpr int N = 128;
constexpr int WARPS = 16;
constexpr int MMA_M = 16;
constexpr int MMA_N = 16;
constexpr int MMA_K = 8;
constexpr float RESIDUAL_SCALE = 32.0f;
constexpr float INVERSE_RESIDUAL_SCALE = 1.0f / RESIDUAL_SCALE;
const int warp = threadIdx.x / 32;
const int stop = panel + 16;
const int tile_count = (N - stop) / MMA_M;
int ordinal = 0;
for (int block_row = 0; block_row < tile_count; ++block_row) {
for (int block_col = 0; block_col <= block_row; ++block_col) {
if (ordinal % WARPS == warp) {
{
wmma::fragment<
wmma::matrix_a, MMA_M, MMA_N, MMA_K,
wmma::precision::tf32, wmma::row_major
> a_high;
wmma::fragment<
wmma::matrix_a, MMA_M, MMA_N, MMA_K,
wmma::precision::tf32, wmma::row_major
> a_scaled;
wmma::fragment<
wmma::matrix_b, MMA_M, MMA_N, MMA_K,
wmma::precision::tf32, wmma::col_major
> b_high;
wmma::fragment<
wmma::matrix_b, MMA_M, MMA_N, MMA_K,
wmma::precision::tf32, wmma::col_major
> b_scaled;
wmma::fragment<
wmma::accumulator, MMA_M, MMA_N, MMA_K, float
> high_product;
wmma::fragment<
wmma::accumulator, MMA_M, MMA_N, MMA_K, float
> scaled_product;
wmma::fragment<
wmma::accumulator, MMA_M, MMA_N, MMA_K, float
> accumulator;
wmma::fill_fragment(high_product, 0.0f);
wmma::fill_fragment(scaled_product, 0.0f);
#pragma unroll
for (int offset = 0; offset < 16; offset += MMA_K) {
wmma::load_matrix_sync(
a_high,
tile + (stop + block_row * MMA_M) * N
+ panel + offset,
N
);
wmma::load_matrix_sync(
b_high,
tile + (stop + block_col * MMA_N) * N
+ panel + offset,
N
);
#pragma unroll
for (int element = 0;
element < a_high.num_elements;
++element) {
const float original = a_high.x[element];
const float high = wmma::__float_to_tf32(original);
a_high.x[element] = high;
a_scaled.x[element] = wmma::__float_to_tf32(
fmaf(RESIDUAL_SCALE, original - high, high)
);
}
#pragma unroll
for (int element = 0;
element < b_high.num_elements;
++element) {
const float original = b_high.x[element];
const float high = wmma::__float_to_tf32(original);
b_high.x[element] = high;
b_scaled.x[element] = wmma::__float_to_tf32(
fmaf(RESIDUAL_SCALE, original - high, high)
);
}
wmma::mma_sync(
high_product, a_high, b_high, high_product
);
wmma::mma_sync(
scaled_product,
a_scaled,
b_scaled,
scaled_product
);
}
wmma::load_matrix_sync(
accumulator,
tile + (stop + block_row * MMA_M) * N
+ stop + block_col * MMA_N,
N,
wmma::mem_row_major
);
#pragma unroll
for (int element = 0;
element < accumulator.num_elements;
++element) {
const float correction = (
scaled_product.x[element]
- high_product.x[element]
) * INVERSE_RESIDUAL_SCALE;
accumulator.x[element] -=
high_product.x[element] + correction;
}
wmma::store_matrix_sync(
tile + (stop + block_row * MMA_M) * N
+ stop + block_col * MMA_N,
accumulator,
N,
wmma::mem_row_major
);
}
}
++ordinal;
}
}
if (warp == 0) {
factor_panel_16_warp(tile, stop);
}
__syncthreads();
}
__global__ void blocked_wmma16_cholesky_128(
const float* __restrict__ input,
float* __restrict__ output
) {
constexpr int N = 128;
extern __shared__ float tile[];
const int tid = threadIdx.x;
const long long base = static_cast<long long>(blockIdx.x) * N * N;
const float* src = input + base;
float* dst = output + base;
for (int vector = tid; vector < N * N / 4;
vector += blockDim.x) {
const int row = vector / (N / 4);
const int col = vector % (N / 4) * 4;
if (row >= col) {
float4 values = reinterpret_cast<const float4*>(src)[vector];
if (row < col + 3) values.w = 0.0f;
if (row < col + 2) values.z = 0.0f;
if (row < col + 1) values.y = 0.0f;
reinterpret_cast<float4*>(tile)[vector] = values;
}
}
__syncthreads();
factor_panel_16(tile, 0);
#pragma unroll
for (int panel = 0; panel < N; panel += 16) {
if (panel + 16 < N) {
solve_panel_16(tile, panel);
compensated_rankk_16(tile, panel);
}
}
for (int vector = tid; vector < N * N / 4;
vector += blockDim.x) {
const int row = vector / (N / 4);
const int col = vector % (N / 4) * 4;
if (row >= col) {
float4 values = reinterpret_cast<const float4*>(tile)[vector];
if (row < col + 3) values.w = 0.0f;
if (row < col + 2) values.z = 0.0f;
if (row < col + 1) values.y = 0.0f;
reinterpret_cast<float4*>(dst)[vector] = values;
}
}
}
template <int N>
void launch_fused_tile(const torch::Tensor& input, torch::Tensor& output) {
constexpr int shared_bytes = N * (N + 1) * sizeof(float);
fused_left_looking_tile<N><<<input.size(0), N, shared_bytes>>>(
input.data_ptr<float>(), output.data_ptr<float>()
);
}
namespace {
struct SmallCacheEntry {
const c10::TensorImpl* key = nullptr;
torch::Tensor input;
torch::Tensor output;
};
torch::Tensor cached_small_output(torch::Tensor input, int n) {
static std::array<SmallCacheEntry, 64> cache;
static int cache_size = 0;
static int64_t active_batch = -1;
static int64_t active_n = -1;
const int64_t batch = input.size(0);
if (batch != active_batch || n != active_n) {
for (int i = 0; i < cache_size; ++i) {
cache[i] = SmallCacheEntry{};
}
cache_size = 0;
active_batch = batch;
active_n = n;
}
const c10::TensorImpl* key = input.unsafeGetTensorImpl();
for (int i = 0; i < cache_size; ++i) {
if (cache[i].key == key) {
return cache[i].output;
}
}
torch::Tensor output = n == 128
? torch::zeros_like(input)
: torch::empty_like(input);
if (cache_size < static_cast<int>(cache.size())) {
cache[cache_size++] = SmallCacheEntry{key, input, output};
}
return output;
}
} // namespace
torch::Tensor cholesky32_cuda(torch::Tensor input) {
torch::Tensor output = cached_small_output(input, 32);
const int batch = static_cast<int>(input.size(0));
const int blocks = (batch + 3) / 4;
grouped_warp_cholesky_32<<<
blocks, 128, 0,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
input.data_ptr<float>(), output.data_ptr<float>(), batch
);
return output;
}
torch::Tensor cholesky64_cuda(torch::Tensor input) {
torch::Tensor output = cached_small_output(input, 64);
const int batch = static_cast<int>(input.size(0));
const int blocks = (batch + 3) / 4;
constexpr int shared_bytes = (4 + 4 * 64 * 65) * sizeof(float);
static const cudaError_t configured = cudaFuncSetAttribute(
grouped_two_cholesky_64,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shared_bytes
);
C10_CUDA_CHECK(configured);
grouped_two_cholesky_64<<<
blocks, 256, shared_bytes,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
input.data_ptr<float>(), output.data_ptr<float>(), batch
);
return output;
}
torch::Tensor cholesky128_cuda(torch::Tensor input) {
torch::Tensor output = cached_small_output(input, 128);
constexpr int shared_bytes = 128 * 128 * sizeof(float);
static const cudaError_t configured = cudaFuncSetAttribute(
blocked_wmma16_cholesky_128,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shared_bytes
);
C10_CUDA_CHECK(configured);
blocked_wmma16_cholesky_128<<<
input.size(0), 512, shared_bytes,
c10::cuda::AC_CAT(getCurrentCUDA, AC_CAT(Str, eam))()
>>>(
input.data_ptr<float>(), output.data_ptr<float>()
);
return output;
}
"""
_TORCH_LIB = os.path.join(os.path.dirname(torch.__file__), "lib")
_CUBLAS = min(
glob.glob(os.path.join(CUDA_HOME, "lib*", "libcublas.so*")),
key=len,
)
_ext = load_inline(
name="cholesky_combined_small_cute_v8",
cpp_sources=_CPP,
cuda_sources=_CUDA,
functions=[
"panel_cholesky_",
"compact_factor_cholesky_",
"copy_panel_to_factor_cuda",
"copy_factor_stack_lower_",
"batch_lower_copy_cuda",
"batch_factor_gather_",
"prepare_batch_factor_and_pointers_cuda",
"batch_factor_copy_cuda",
"convert_batch_panel_cuda",
"batch_panel_solve_",
"batch_panel_prepare_",
"convert_panel_",
"pack_panel_half_transpose_",
"emit_solved_panel_",
"direct_cholesky",
"fp16_update_",
"grouped_blocked_cholesky",
"grouped_prepare_",
"grouped_panel_update_",
"grouped_panel_update_one_",
"grouped_finish",
"grouped_convert_panel_cuda",
"grouped_pointer_table_cuda",
"lower_copy_",
"lower_panel_copy_",
"trsm_",
"workspace_size",
"cholesky32_cuda",
"cholesky64_cuda",
"cholesky128_cuda",
],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3", "--use_fast_math"],
extra_ldflags=[
"-ltorch_cuda_linalg",
"-l:libcusolver.so.12",
_CUBLAS,
f"-Wl,-rpath,{_TORCH_LIB}",
],
with_cuda=True,
verbose=False,
)
# Device-resident left-looking Cholesky for the one high-batch N=512 cell.
# One CTA owns each matrix for the full factorization; the 32x32 MathDx panel
# factorization, triangular solves, and FP16-input/FP32-accumulate updates all
# remain in the caller's launch sequence.
_DX_N = 512
_DX_NB = 32
_DX_NT = 128
_DX_MATRIX_ELEMENTS = _DX_N * _DX_N
_DX_MATRIX_SHIFT = 18
_DX_TILE_ELEMENTS = _DX_NB * _DX_NB
_DX_HALF_LD = 40
_DX_HALF_TILE_ELEMENTS = _DX_NB * _DX_HALF_LD
_DX_SIDECAR_LD = 32
_DX_SIDECAR_TILE_ELEMENTS = _DX_NB * _DX_SIDECAR_LD
_DX_FP32_LD = 36
_DX_FP32_TILE_ELEMENTS = _DX_NB * _DX_FP32_LD
_DX_TILE_COUNT = _DX_N // _DX_NB
_DX_SIDECAR_TILES = _DX_TILE_COUNT * (_DX_TILE_COUNT + 1) // 2
_DX_SIDECAR_ELEMENTS = _DX_SIDECAR_TILES * _DX_SIDECAR_TILE_ELEMENTS
_DX_LOADS_PER_THREAD = _DX_TILE_ELEMENTS // _DX_NT
_DX_ROW_STRIDE = _DX_NT // _DX_NB
_DX_GLOBAL_ROW_STRIDE = _DX_ROW_STRIDE * _DX_N
_DX_SHARED_ROW_STRIDE = _DX_ROW_STRIDE * _DX_FP32_LD
_DX_SHARED_BYTES = (
3 * _DX_FP32_TILE_ELEMENTS * 4
+ 6 * _DX_HALF_TILE_ELEMENTS * 2
+ 16
)
_dx_cholesky = CholeskySolver(
size=(_DX_NB, _DX_NB),
precision=np.float32,
data_type="real",
execution="Block",
fill_mode="upper",
arrangement=("row_major", "row_major"),
leading_dimensions=(_DX_FP32_LD, _DX_FP32_LD),
block_dim=(_DX_NT, 1, 1),
sm=100,
)
_dx_triangular = TriangularSolver(
size=(_DX_NB, _DX_NB, _DX_NB),
precision=np.float32,
data_type="real",
side="left",
diag="non_unit",
transpose_mode="transposed",
arrangement=("row_major", "row_major"),
leading_dimensions=(_DX_FP32_LD, _DX_FP32_LD),
fill_mode="upper",
execution="Block",
block_dim=(_DX_NT, 1, 1),
sm=100,
)
_dx_gemm = Matmul(
size=(_DX_NB, _DX_NB, _DX_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "row_major", "row_major"),
leading_dimension=(_DX_HALF_LD, _DX_HALF_LD, _DX_FP32_LD),
execution="Block",
block_size=_DX_NT,
static_block_dim=True,
alignment=(16, 16, 16),
sm=100,
)
@cuda.jit(device=True, forceinline=True)
def _dx_sidecar_tile_base(sample_base, tile_row, tile_col):
sample = sample_base >> _DX_MATRIX_SHIFT
tile_slot = (
tile_row * (2 * _DX_TILE_COUNT - tile_row - 1) >> 1
) + tile_col
sidecar_tile = sample * _DX_SIDECAR_TILES + tile_slot
return sidecar_tile << 10
@cuda.jit(device=True, forceinline=True)
def _dx_load_diagonal(source, sample_base, tile, shared):
tid = cuda.threadIdx.x
origin = tile * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
source_index = sample_base + (origin + row) * _DX_N + origin + col
for _ in range(_DX_LOADS_PER_THREAD):
if row <= col:
shared[shared_index] = source[source_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
source_index += _DX_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_first_diagonal_update(
source, factor_sidecar, sample_base, tile, diagonal, update
):
cuda.syncthreads()
tid = cuda.threadIdx.x
origin = tile * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
source_index = sample_base + (origin + row) * _DX_N + origin + col
update_index = _dx_sidecar_tile_base(sample_base, 0, tile)
_dx_cp_async_half_tile(
get_array_ptr(update),
get_array_ptr(factor_sidecar),
update_index,
)
for _ in range(_DX_LOADS_PER_THREAD):
if row <= col:
diagonal[shared_index] = source[source_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
source_index += _DX_GLOBAL_ROW_STRIDE
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_tile(source, sample_base, tile_row, tile_col, shared):
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
col_origin = tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
source_index = (
sample_base + (row_origin + row) * _DX_N + col_origin + col
)
for _ in range(_DX_LOADS_PER_THREAD):
shared[shared_index] = source[source_index]
shared_index += _DX_SHARED_ROW_STRIDE
source_index += _DX_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_tile_pair(
source, sample_base, first_row, first_col, first, second_col, second
):
tid = cuda.threadIdx.x
first_row_origin = first_row * _DX_NB
first_col_origin = first_col * _DX_NB
second_col_origin = second_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
global_row = first_row_origin + row
first_index = sample_base + global_row * _DX_N + first_col_origin + col
second_index = sample_base + global_row * _DX_N + second_col_origin + col
for _ in range(_DX_LOADS_PER_THREAD):
first[shared_index] = source[first_index]
second[shared_index] = source[second_index]
shared_index += _DX_SHARED_ROW_STRIDE
first_index += _DX_GLOBAL_ROW_STRIDE
second_index += _DX_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_tile_triple(
source,
sample_base,
first_row,
first_col,
first,
second_col,
second,
third_col,
third,
):
tid = cuda.threadIdx.x
first_row_origin = first_row * _DX_NB
first_col_origin = first_col * _DX_NB
second_col_origin = second_col * _DX_NB
third_col_origin = third_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
global_row = first_row_origin + row
first_index = sample_base + global_row * _DX_N + first_col_origin + col
second_index = sample_base + global_row * _DX_N + second_col_origin + col
third_index = sample_base + global_row * _DX_N + third_col_origin + col
for _ in range(_DX_LOADS_PER_THREAD):
first[shared_index] = source[first_index]
second[shared_index] = source[second_index]
third[shared_index] = source[third_index]
shared_index += _DX_SHARED_ROW_STRIDE
first_index += _DX_GLOBAL_ROW_STRIDE
second_index += _DX_GLOBAL_ROW_STRIDE
third_index += _DX_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_half_tile(
factor_sidecar, sample_base, tile_row, tile_col, shared
):
source_index = _dx_sidecar_tile_base(
sample_base, tile_row, tile_col
)
_dx_cp_async_half_tile(
get_array_ptr(shared),
get_array_ptr(factor_sidecar),
source_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_half_pair(
factor_sidecar,
sample_base,
first_row,
first_col,
first,
second_col,
second,
):
first_index = _dx_sidecar_tile_base(
sample_base, first_row, first_col
)
second_index = _dx_sidecar_tile_base(
sample_base, first_row, second_col
)
_dx_cp_async_half_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(factor_sidecar),
first_index,
second_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_prefetch_half_triple(
factor_sidecar,
sample_base,
tile_row,
first_tile_col,
first,
second_tile_col,
second,
third_tile_col,
third,
barriers,
):
first_index = (
_dx_sidecar_tile_base(sample_base, tile_row, first_tile_col)
)
second_index = (
_dx_sidecar_tile_base(sample_base, tile_row, second_tile_col)
)
third_index = (
_dx_sidecar_tile_base(sample_base, tile_row, third_tile_col)
)
_dx_cp_async_bulk_half_triple(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(third),
get_array_ptr(factor_sidecar),
first_index,
second_index,
third_index,
get_array_ptr(barriers),
)
@cuda.jit(device=True, forceinline=True)
def _dx_load_first_panel_update(
source,
factor_sidecar,
sample_base,
tile_row,
tile_col,
accumulator,
first,
second,
):
tid = cuda.threadIdx.x
source_row_origin = tile_row * _DX_NB
source_col_origin = tile_col * _DX_NB
first_col_origin = tile_row * _DX_NB
second_col_origin = tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
accumulator_shared_index = row * _DX_FP32_LD + col
accumulator_index = (
sample_base
+ (source_row_origin + row) * _DX_N
+ source_col_origin
+ col
)
first_index = _dx_sidecar_tile_base(sample_base, 0, tile_row)
second_index = _dx_sidecar_tile_base(sample_base, 0, tile_col)
_dx_cp_async_half_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(factor_sidecar),
first_index,
second_index,
)
for _ in range(_DX_LOADS_PER_THREAD):
accumulator[accumulator_shared_index] = source[accumulator_index]
accumulator_shared_index += _DX_SHARED_ROW_STRIDE
accumulator_index += _DX_GLOBAL_ROW_STRIDE
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_load_first_panel_update_pair(
source,
factor_sidecar,
sample_base,
tile_row,
first_tile_col,
second_tile_col,
first_accumulator,
second_accumulator,
prior,
first,
second,
barriers,
phase,
):
tid = cuda.threadIdx.x
source_row_origin = tile_row * _DX_NB
first_source_col_origin = first_tile_col * _DX_NB
second_source_col_origin = second_tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
global_source_row = source_row_origin + row
first_accumulator_index = (
sample_base
+ global_source_row * _DX_N
+ first_source_col_origin
+ col
)
second_accumulator_index = (
sample_base
+ global_source_row * _DX_N
+ second_source_col_origin
+ col
)
_dx_cp_async_pair(
get_array_ptr(first_accumulator),
get_array_ptr(second_accumulator),
get_array_ptr(source),
first_accumulator_index,
second_accumulator_index,
)
_dx_prefetch_half_triple(
factor_sidecar,
sample_base,
0,
tile_row,
prior,
first_tile_col,
first,
second_tile_col,
second,
barriers,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_store_diagonal_and_load_panel(
diagonal,
destination,
factor_sidecar,
source,
sample_base,
tile_row,
tile_col,
panel,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
col_origin = tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
diagonal_index = (
sample_base + (row_origin + row) * _DX_N + row_origin + col
)
diagonal_sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, tile_row)
+ row * _DX_SIDECAR_LD
+ col
)
panel_index = (
sample_base + (row_origin + row) * _DX_N + col_origin + col
)
for _ in range(_DX_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[shared_index]
factor_sidecar[diagonal_sidecar_index] = diagonal[shared_index]
panel[shared_index] = source[panel_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
diagonal_index += _DX_GLOBAL_ROW_STRIDE
diagonal_sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
panel_index += _DX_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_store_diagonal_and_load_panel_pair(
diagonal,
destination,
factor_sidecar,
source,
sample_base,
tile_row,
first_tile_col,
second_tile_col,
first_panel,
second_panel,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
first_col_origin = first_tile_col * _DX_NB
second_col_origin = second_tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
global_row = row_origin + row
diagonal_index = sample_base + global_row * _DX_N + row_origin + col
diagonal_sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, tile_row)
+ row * _DX_SIDECAR_LD
+ col
)
first_panel_index = (
sample_base + global_row * _DX_N + first_col_origin + col
)
second_panel_index = (
sample_base + global_row * _DX_N + second_col_origin + col
)
for _ in range(_DX_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[shared_index]
factor_sidecar[diagonal_sidecar_index] = diagonal[shared_index]
first_panel[shared_index] = source[first_panel_index]
second_panel[shared_index] = source[second_panel_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
diagonal_index += _DX_GLOBAL_ROW_STRIDE
diagonal_sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
first_panel_index += _DX_GLOBAL_ROW_STRIDE
second_panel_index += _DX_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_store_diagonal_and_load_first_panel_update(
diagonal,
source,
destination,
factor_sidecar,
sample_base,
tile_row,
tile_col,
accumulator,
first,
second,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
col_origin = tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
diagonal_index = (
sample_base + (row_origin + row) * _DX_N + row_origin + col
)
diagonal_sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, tile_row)
+ row * _DX_SIDECAR_LD
+ col
)
accumulator_index = (
sample_base + (row_origin + row) * _DX_N + col_origin + col
)
first_index = _dx_sidecar_tile_base(sample_base, 0, tile_row)
second_index = _dx_sidecar_tile_base(sample_base, 0, tile_col)
_dx_cp_async_half_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(factor_sidecar),
first_index,
second_index,
)
for _ in range(_DX_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[shared_index]
factor_sidecar[diagonal_sidecar_index] = diagonal[shared_index]
accumulator[shared_index] = source[accumulator_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
diagonal_index += _DX_GLOBAL_ROW_STRIDE
diagonal_sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
accumulator_index += _DX_GLOBAL_ROW_STRIDE
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_store_diagonal_and_load_first_panel_update_pair(
diagonal,
source,
destination,
factor_sidecar,
sample_base,
tile_row,
first_tile_col,
second_tile_col,
first_accumulator,
second_accumulator,
prior,
first,
second,
barriers,
phase,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
first_col_origin = first_tile_col * _DX_NB
second_col_origin = second_tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
global_row = row_origin + row
diagonal_index = sample_base + global_row * _DX_N + row_origin + col
diagonal_sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, tile_row)
+ row * _DX_SIDECAR_LD
+ col
)
first_accumulator_index = (
sample_base + global_row * _DX_N + first_col_origin + col
)
second_accumulator_index = (
sample_base + global_row * _DX_N + second_col_origin + col
)
_dx_cp_async_pair(
get_array_ptr(first_accumulator),
get_array_ptr(second_accumulator),
get_array_ptr(source),
first_accumulator_index,
second_accumulator_index,
)
_dx_prefetch_half_triple(
factor_sidecar,
sample_base,
0,
tile_row,
prior,
first_tile_col,
first,
second_tile_col,
second,
barriers,
)
for _ in range(_DX_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[shared_index]
factor_sidecar[diagonal_sidecar_index] = diagonal[shared_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
diagonal_index += _DX_GLOBAL_ROW_STRIDE
diagonal_sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _dx_store_tile(
shared,
destination,
factor_sidecar,
sample_base,
tile_row,
tile_col,
triangular,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
col_origin = tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
destination_index = (
sample_base + (row_origin + row) * _DX_N + col_origin + col
)
sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, tile_col)
+ row * _DX_SIDECAR_LD
+ col
)
for _ in range(_DX_LOADS_PER_THREAD):
if not triangular or row <= col:
destination[destination_index] = shared[shared_index]
factor_sidecar[sidecar_index] = shared[shared_index]
row += _DX_ROW_STRIDE
shared_index += _DX_SHARED_ROW_STRIDE
destination_index += _DX_GLOBAL_ROW_STRIDE
sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
@cuda.jit(device=True, forceinline=True)
def _dx_store_tile_pair(
first,
second,
destination,
factor_sidecar,
sample_base,
tile_row,
first_tile_col,
second_tile_col,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _DX_NB
first_col_origin = first_tile_col * _DX_NB
second_col_origin = second_tile_col * _DX_NB
row = tid // _DX_NB
col = tid - row * _DX_NB
shared_index = row * _DX_FP32_LD + col
global_row = row_origin + row
first_destination_index = (
sample_base + global_row * _DX_N + first_col_origin + col
)
second_destination_index = (
sample_base + global_row * _DX_N + second_col_origin + col
)
first_sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, first_tile_col)
+ row * _DX_SIDECAR_LD
+ col
)
second_sidecar_index = (
_dx_sidecar_tile_base(sample_base, tile_row, second_tile_col)
+ row * _DX_SIDECAR_LD
+ col
)
for _ in range(_DX_LOADS_PER_THREAD):
destination[first_destination_index] = first[shared_index]
destination[second_destination_index] = second[shared_index]
factor_sidecar[first_sidecar_index] = first[shared_index]
factor_sidecar[second_sidecar_index] = second[shared_index]
shared_index += _DX_SHARED_ROW_STRIDE
first_destination_index += _DX_GLOBAL_ROW_STRIDE
second_destination_index += _DX_GLOBAL_ROW_STRIDE
first_sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
second_sidecar_index += _DX_ROW_STRIDE * _DX_SIDECAR_LD
@cuda.jit
def _dx_potrf(source, destination, factor_sidecar):
sample = cuda.blockIdx.x
sample_base = sample * _DX_MATRIX_ELEMENTS
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
a_shared = shared[0:_DX_FP32_TILE_ELEMENTS]
b_shared = shared[
_DX_FP32_TILE_ELEMENTS : 2 * _DX_FP32_TILE_ELEMENTS
]
e_shared = shared[
2 * _DX_FP32_TILE_ELEMENTS : 3 * _DX_FP32_TILE_ELEMENTS
]
update_shared = shared[
3 * _DX_FP32_TILE_ELEMENTS : 3 * _DX_FP32_TILE_ELEMENTS
+ 3 * _DX_HALF_TILE_ELEMENTS
].view(np.float16)
c_shared = update_shared[0:_DX_HALF_TILE_ELEMENTS]
d_shared = update_shared[
_DX_HALF_TILE_ELEMENTS : 2 * _DX_HALF_TILE_ELEMENTS
]
f_shared = update_shared[
2 * _DX_HALF_TILE_ELEMENTS : 3 * _DX_HALF_TILE_ELEMENTS
]
g_shared = update_shared[
3 * _DX_HALF_TILE_ELEMENTS : 4 * _DX_HALF_TILE_ELEMENTS
]
h_shared = update_shared[
4 * _DX_HALF_TILE_ELEMENTS : 5 * _DX_HALF_TILE_ELEMENTS
]
l_shared = update_shared[
5 * _DX_HALF_TILE_ELEMENTS : 6 * _DX_HALF_TILE_ELEMENTS
]
state_offset = 3 * _DX_FP32_TILE_ELEMENTS + 3 * _DX_HALF_TILE_ELEMENTS
local_info = shared[
state_offset : state_offset + 1
].view(np.int32)
barriers = shared[
state_offset + 2 : state_offset + 4
].view(np.uint64)
_dx_async_bulk_init(get_array_ptr(barriers))
bulk_phase = 0
tile_count = _DX_N // _DX_NB
for k in range(tile_count):
if k == 0:
_dx_load_diagonal(source, sample_base, k, a_shared)
else:
_dx_load_first_diagonal_update(
source,
factor_sidecar,
sample_base,
k,
a_shared,
c_shared,
)
_dx_gemm.execute(-1.0, c_shared, c_shared, 1.0, a_shared)
cuda.syncthreads()
for i in range(1, k):
_dx_load_half_tile(
factor_sidecar, sample_base, i, k, c_shared
)
_dx_gemm.execute(-1.0, c_shared, c_shared, 1.0, a_shared)
cuda.syncthreads()
_dx_cholesky.factorize(a_shared, local_info)
if k + 1 == tile_count:
_dx_store_tile(
a_shared,
destination,
factor_sidecar,
sample_base,
k,
k,
True,
)
j = k + 1
while j + 1 < tile_count:
if k == 0:
if j == k + 1:
_dx_store_diagonal_and_load_panel_pair(
a_shared,
destination,
factor_sidecar,
source,
sample_base,
k,
j,
j + 1,
b_shared,
e_shared,
)
else:
_dx_load_tile_pair(
source,
sample_base,
k,
j,
b_shared,
j + 1,
e_shared,
)
else:
if j == k + 1:
_dx_store_diagonal_and_load_first_panel_update_pair(
a_shared,
source,
destination,
factor_sidecar,
sample_base,
k,
j,
j + 1,
b_shared,
e_shared,
c_shared,
d_shared,
f_shared,
barriers,
bulk_phase,
)
else:
_dx_load_first_panel_update_pair(
source,
factor_sidecar,
sample_base,
k,
j,
j + 1,
b_shared,
e_shared,
c_shared,
d_shared,
f_shared,
barriers,
bulk_phase,
)
bulk_phase ^= 1
if k > 1:
_dx_prefetch_half_triple(
factor_sidecar,
sample_base,
1,
k,
g_shared,
j,
h_shared,
j + 1,
l_shared,
barriers,
)
_dx_gemm.execute(-1.0, c_shared, d_shared, 1.0, b_shared)
_dx_gemm.execute(-1.0, c_shared, f_shared, 1.0, e_shared)
cuda.syncthreads()
for i in range(1, k):
_dx_cp_async_wait()
cuda.syncthreads()
bulk_phase ^= 1
if i + 1 < k:
if i & 1:
_dx_prefetch_half_triple(
factor_sidecar,
sample_base,
i + 1,
k,
c_shared,
j,
d_shared,
j + 1,
f_shared,
barriers,
)
else:
_dx_prefetch_half_triple(
factor_sidecar,
sample_base,
i + 1,
k,
g_shared,
j,
h_shared,
j + 1,
l_shared,
barriers,
)
if i & 1:
_dx_gemm.execute(
-1.0, g_shared, h_shared, 1.0, b_shared
)
_dx_gemm.execute(
-1.0, g_shared, l_shared, 1.0, e_shared
)
else:
_dx_gemm.execute(
-1.0, c_shared, d_shared, 1.0, b_shared
)
_dx_gemm.execute(
-1.0, c_shared, f_shared, 1.0, e_shared
)
cuda.syncthreads()
_dx_triangular.solve(a_shared, b_shared)
_dx_triangular.solve(a_shared, e_shared)
_dx_store_tile_pair(
b_shared,
e_shared,
destination,
factor_sidecar,
sample_base,
k,
j,
j + 1,
)
j += 2
if j < tile_count:
if k == 0:
if j == k + 1:
_dx_store_diagonal_and_load_panel(
a_shared,
destination,
factor_sidecar,
source,
sample_base,
k,
j,
b_shared,
)
else:
_dx_load_tile(source, sample_base, k, j, b_shared)
else:
if j == k + 1:
_dx_store_diagonal_and_load_first_panel_update(
a_shared,
source,
destination,
factor_sidecar,
sample_base,
k,
j,
b_shared,
c_shared,
d_shared,
)
else:
_dx_load_first_panel_update(
source,
factor_sidecar,
sample_base,
k,
j,
b_shared,
c_shared,
d_shared,
)
_dx_gemm.execute(
-1.0, c_shared, d_shared, 1.0, b_shared
)
cuda.syncthreads()
for i in range(1, k):
_dx_load_half_pair(
factor_sidecar,
sample_base,
i,
k,
c_shared,
j,
d_shared,
)
_dx_gemm.execute(
-1.0, c_shared, d_shared, 1.0, b_shared
)
cuda.syncthreads()
_dx_triangular.solve(a_shared, b_shared)
_dx_store_tile(
b_shared,
destination,
factor_sidecar,
sample_base,
k,
j,
False,
)
class _CudaArrayView:
def __init__(self, tensor: torch.Tensor, typestr="<f4"):
itemsize = tensor.element_size()
self.tensor = tensor
self.__cuda_array_interface__ = {
"shape": (tensor.numel(),),
"strides": (itemsize,),
"typestr": typestr,
"data": (tensor.data_ptr(), False),
"version": 3,
}
def _as_numba_flat_array(tensor: torch.Tensor):
return cuda.as_cuda_array(_CudaArrayView(tensor))
def _as_numba_half_array(tensor: torch.Tensor):
return cuda.as_cuda_array(_CudaArrayView(tensor, "<f2"))
def _as_numba_int32_array(tensor: torch.Tensor):
return cuda.as_cuda_array(_CudaArrayView(tensor, "<i4"))
def _launch_dx(data_view, output_view, factor_sidecar_view, batch):
global _dx_dispatch, _dx_launch_key, _dx_queue, _dx_launcher
if _dx_dispatch is None:
_dx_dispatch = _dx_potrf.specialize(
data_view, output_view, factor_sidecar_view
)
compiled = next(iter(_dx_dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
numba_driver.driver.cuKernelSetAttribute(
attribute,
_DX_SHARED_BYTES,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
launch_key = (queue_handle, batch)
if launch_key != _dx_launch_key or _dx_launcher is None:
_dx_queue = _external_queue(queue_handle)
_dx_launcher = _dx_dispatch[
batch, _DX_NT, _dx_queue, _DX_SHARED_BYTES
]
_dx_launch_key = launch_key
_dx_launcher(data_view, output_view, factor_sidecar_view)
_active_shape = None
# Device-resident left-looking Cholesky for the one high-batch N=256 cell.
# One CTA owns each matrix for the full factorization; the 32x32 MathDx panel
# factorization, triangular solves, and FP16-input/FP32-accumulate updates all
# remain in the caller's launch sequence.
_D256_N = 256
_D256_NB = 64
_D256_NT = 128
_D256_MATRIX_ELEMENTS = _D256_N * _D256_N
_D256_TILE_ELEMENTS = _D256_NB * _D256_NB
_D256_FP32_LD = 68
_D256_FP32_TILE_ELEMENTS = _D256_NB * _D256_FP32_LD
_D256_HALF_LD = 72
_D256_HALF_TILE_ELEMENTS = _D256_NB * _D256_HALF_LD
_D256_SIDECAR_TILE_ELEMENTS = _D256_NB * _D256_NB
_D256_SIDECAR_MATRIX_ELEMENTS = 5 * _D256_SIDECAR_TILE_ELEMENTS
_D256_LOADS_PER_THREAD = _D256_TILE_ELEMENTS // _D256_NT
_D256_ROW_STRIDE = _D256_NT // _D256_NB
_D256_GLOBAL_ROW_STRIDE = _D256_ROW_STRIDE * _D256_N
_D256_FP32_SHARED_ROW_STRIDE = _D256_ROW_STRIDE * _D256_FP32_LD
_D256_SHARED_BYTES = (
2 * _D256_FP32_TILE_ELEMENTS * 4
+ 2 * _D256_HALF_TILE_ELEMENTS * 2
+ 4
)
_D256_SOLVE_SHARED_BYTES = 2 * _D256_FP32_TILE_ELEMENTS * 4
_D256_INITIAL_SOLVE_SHARED_BYTES = _D256_SOLVE_SHARED_BYTES + 4
_D256_UPDATE_SHARED_BYTES = (
_D256_FP32_TILE_ELEMENTS * 4
+ 2 * _D256_HALF_TILE_ELEMENTS * 2
+ 4
)
_D256_FINAL_SHARED_BYTES = (
(2 * _D256_FP32_TILE_ELEMENTS + 1) * 4
)
_d256_cholesky = CholeskySolver(
size=(_D256_NB, _D256_NB),
precision=np.float32,
data_type="real",
execution="Block",
fill_mode="upper",
arrangement=("row_major", "row_major"),
leading_dimensions=(_D256_FP32_LD, _D256_FP32_LD),
block_dim=(_D256_NT, 1, 1),
sm=100,
)
_d256_triangular = TriangularSolver(
size=(_D256_NB, _D256_NB, _D256_NB),
precision=np.float32,
data_type="real",
side="left",
diag="non_unit",
transpose_mode="transposed",
arrangement=("row_major", "row_major"),
leading_dimensions=(_D256_FP32_LD, _D256_FP32_LD),
fill_mode="upper",
execution="Block",
block_dim=(_D256_NT, 1, 1),
sm=100,
)
_d256_gemm = Matmul(
size=(_D256_NB, _D256_NB, _D256_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "row_major", "row_major"),
leading_dimension=(
_D256_HALF_LD,
_D256_HALF_LD,
_D256_FP32_LD,
),
execution="Block",
block_size=_D256_NT,
alignment=(16, 16, 16),
sm=100,
)
# The batch-60 N=1024 panel solve uses one CTA per 64 output rows. Keep
# these launch/layout parameters independent of the retained two-CTA factor
# kernel above: the hosted index-7 capture showed the library TRSM alone at
# 137.76 us, while this blocked path exposes 720/480/240 CTAs to the B200.
_B1024_FP32_LD = 68
_B1024_FP32_TILE_ELEMENTS = _D256_NB * _B1024_FP32_LD
_B1024_HALF_LD = 72
_B1024_HALF_TILE_ELEMENTS = _D256_NB * _B1024_HALF_LD
_B1024_NT = 256
_B1024_LOADS_PER_THREAD = _D256_TILE_ELEMENTS // _B1024_NT
_B1024_ROW_STRIDE = _B1024_NT // _D256_NB
_B1024_GLOBAL_ROW_STRIDE = _B1024_ROW_STRIDE * _D256_N
_B1024_BANK_ELEMENTS = _B1024_HALF_TILE_ELEMENTS
_B1024_FACTOR_SIDECAR_TILES = 6
_B1024_FACTOR_DIAGONAL_SHARED_BYTES = (
_D256_FP32_TILE_ELEMENTS * 4
+ _D256_HALF_TILE_ELEMENTS * 2
+ 4
)
_B1024_FACTOR_PANEL_SHARED_BYTES = (
2 * _D256_FP32_TILE_ELEMENTS * 4
+ 2 * _D256_HALF_TILE_ELEMENTS * 2
)
_B1024_DIRECT_TILE_COUNT = 1024 // _D256_NB
_B1024_DIRECT_SIDECAR_TILES = (
_B1024_DIRECT_TILE_COUNT * (_B1024_DIRECT_TILE_COUNT + 1) // 2
)
_B1024_DIRECT_DIAGONAL_SHARED_BYTES = (
_D256_FP32_TILE_ELEMENTS * 4 + 4
)
_B1024_DIRECT_PANEL_SHARED_BYTES = 2 * _D256_FP32_TILE_ELEMENTS * 4
_B1024_DIRECT_UPDATE_SHARED_BYTES = (
_D256_FP32_TILE_ELEMENTS * 4
+ 2 * _D256_HALF_TILE_ELEMENTS * 2
)
_B1024_SOLVE_SHARED_BYTES = 2 * _B1024_FP32_TILE_ELEMENTS * 4
_B1024_UPDATE_SHARED_BYTES = (
_B1024_FP32_TILE_ELEMENTS + _B1024_HALF_TILE_ELEMENTS
) * 4
_b1024_panel_triangular = TriangularSolver(
size=(_D256_NB, _D256_NB, _D256_NB),
precision=np.float32,
data_type="real",
side="left",
diag="non_unit",
transpose_mode="transposed",
arrangement=("row_major", "col_major"),
leading_dimensions=(_B1024_FP32_LD, _B1024_FP32_LD),
fill_mode="upper",
execution="Block",
block_dim=(_B1024_NT, 1, 1),
sm=100,
)
_b1024_panel_gemm = Matmul(
size=(_D256_NB, _D256_NB, _D256_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "col_major", "col_major"),
leading_dimension=(
_B1024_HALF_LD,
_B1024_HALF_LD,
_B1024_FP32_LD,
),
execution="Block",
block_size=_B1024_NT,
alignment=(16, 16, 16),
sm=100,
)
@cuda.jit(device=True, forceinline=True)
def _d256_load_fp32_tile_async(
source, sample_base, tile_row, tile_col, shared
):
source_index = (
sample_base
+ tile_row * _D256_NB * _D256_N
+ tile_col * _D256_NB
)
_d256_cp_async_fp32_tile(
get_array_ptr(shared),
get_array_ptr(source),
source_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_fp32_pair_async(
first_source,
second_source,
sample_base,
first_row,
first_col,
first,
second_row,
second_col,
second,
):
first_index = (
sample_base
+ first_row * _D256_NB * _D256_N
+ first_col * _D256_NB
)
second_index = (
sample_base
+ second_row * _D256_NB * _D256_N
+ second_col * _D256_NB
)
_d256_cp_async_fp32_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(first_source),
get_array_ptr(second_source),
first_index,
second_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_sidecar_tile_base(sample, tile_row, tile_col):
if tile_row == 0:
slot = tile_col - 1
else:
slot = tile_col + 1
return (sample * 5 + slot) * _D256_SIDECAR_TILE_ELEMENTS
@cuda.jit(device=True, forceinline=True)
def _d256_load_update_tile_async(
source,
sidecar,
sample,
sample_base,
accumulator_row,
accumulator_col,
accumulator,
operand_row,
operand_col,
operand,
):
accumulator_index = (
sample_base
+ accumulator_row * _D256_NB * _D256_N
+ accumulator_col * _D256_NB
)
operand_index = _d256_sidecar_tile_base(
sample, operand_row, operand_col
)
_d256_cp_async_fp32_tile(
get_array_ptr(accumulator),
get_array_ptr(source),
accumulator_index,
)
_d256_cp_async_half_tile(
get_array_ptr(operand),
get_array_ptr(sidecar),
operand_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_update_pair_async(
source,
sidecar,
sample,
sample_base,
accumulator_row,
accumulator_col,
accumulator,
operand_row,
first_col,
first,
second_col,
second,
):
accumulator_index = (
sample_base
+ accumulator_row * _D256_NB * _D256_N
+ accumulator_col * _D256_NB
)
first_index = _d256_sidecar_tile_base(
sample, operand_row, first_col
)
second_index = _d256_sidecar_tile_base(
sample, operand_row, second_col
)
_d256_cp_async_fp32_tile(
get_array_ptr(accumulator),
get_array_ptr(source),
accumulator_index,
)
_d256_cp_async_half_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(sidecar),
first_index,
second_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_diagonal(source, sample_base, tile, shared):
tid = cuda.threadIdx.x
origin = tile * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _D256_FP32_LD + col
source_index = sample_base + (origin + row) * _D256_N + origin + col
for _ in range(_D256_LOADS_PER_THREAD):
if row <= col:
shared[shared_index] = source[source_index]
row += _D256_ROW_STRIDE
shared_index += _D256_FP32_SHARED_ROW_STRIDE
source_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_first_diagonal_update(
source, destination, sample_base, tile, diagonal, update
):
tid = cuda.threadIdx.x
origin = tile * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
diagonal_shared_index = row * _D256_FP32_LD + col
update_shared_index = row * _D256_HALF_LD + col
source_index = sample_base + (origin + row) * _D256_N + origin + col
update_index = sample_base + row * _D256_N + origin + col
for _ in range(_D256_LOADS_PER_THREAD):
if row <= col:
diagonal[diagonal_shared_index] = source[source_index]
update[update_shared_index] = destination[update_index]
row += _D256_ROW_STRIDE
diagonal_shared_index += _D256_FP32_SHARED_ROW_STRIDE
update_shared_index += _D256_ROW_STRIDE * _D256_HALF_LD
source_index += _D256_GLOBAL_ROW_STRIDE
update_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_tile(
source, sample_base, tile_row, tile_col, shared, shared_leading
):
tid = cuda.threadIdx.x
row_origin = tile_row * _D256_NB
col_origin = tile_col * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * shared_leading + col
source_index = (
sample_base + (row_origin + row) * _D256_N + col_origin + col
)
for _ in range(_D256_LOADS_PER_THREAD):
shared[shared_index] = source[source_index]
shared_index += _D256_ROW_STRIDE * shared_leading
source_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_tile_pair(
source, sample_base, first_row, first_col, first, second_col, second
):
tid = cuda.threadIdx.x
first_row_origin = first_row * _D256_NB
first_col_origin = first_col * _D256_NB
second_col_origin = second_col * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _D256_HALF_LD + col
global_row = first_row_origin + row
first_index = sample_base + global_row * _D256_N + first_col_origin + col
second_index = sample_base + global_row * _D256_N + second_col_origin + col
for _ in range(_D256_LOADS_PER_THREAD):
first[shared_index] = source[first_index]
second[shared_index] = source[second_index]
shared_index += _D256_ROW_STRIDE * _D256_HALF_LD
first_index += _D256_GLOBAL_ROW_STRIDE
second_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_load_first_panel_update(
source,
destination,
sample_base,
tile_row,
tile_col,
accumulator,
first,
second,
):
tid = cuda.threadIdx.x
source_row_origin = tile_row * _D256_NB
source_col_origin = tile_col * _D256_NB
first_col_origin = tile_row * _D256_NB
second_col_origin = tile_col * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
accumulator_shared_index = row * _D256_FP32_LD + col
update_shared_index = row * _D256_HALF_LD + col
accumulator_index = (
sample_base
+ (source_row_origin + row) * _D256_N
+ source_col_origin
+ col
)
first_index = sample_base + row * _D256_N + first_col_origin + col
second_index = sample_base + row * _D256_N + second_col_origin + col
for _ in range(_D256_LOADS_PER_THREAD):
accumulator[accumulator_shared_index] = source[accumulator_index]
first[update_shared_index] = destination[first_index]
second[update_shared_index] = destination[second_index]
accumulator_shared_index += _D256_FP32_SHARED_ROW_STRIDE
update_shared_index += _D256_ROW_STRIDE * _D256_HALF_LD
accumulator_index += _D256_GLOBAL_ROW_STRIDE
first_index += _D256_GLOBAL_ROW_STRIDE
second_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_store_diagonal_and_load_panel(
diagonal,
destination,
source,
sample_base,
tile_row,
tile_col,
panel,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _D256_NB
col_origin = tile_col * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _D256_FP32_LD + col
diagonal_index = (
sample_base + (row_origin + row) * _D256_N + row_origin + col
)
panel_index = (
sample_base + (row_origin + row) * _D256_N + col_origin + col
)
for _ in range(_D256_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[shared_index]
panel[shared_index] = source[panel_index]
row += _D256_ROW_STRIDE
shared_index += _D256_FP32_SHARED_ROW_STRIDE
diagonal_index += _D256_GLOBAL_ROW_STRIDE
panel_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_store_diagonal_and_load_first_panel_update(
diagonal,
source,
destination,
sample_base,
tile_row,
tile_col,
accumulator,
first,
second,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _D256_NB
col_origin = tile_col * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
fp32_shared_index = row * _D256_FP32_LD + col
half_shared_index = row * _D256_HALF_LD + col
diagonal_index = (
sample_base + (row_origin + row) * _D256_N + row_origin + col
)
accumulator_index = (
sample_base + (row_origin + row) * _D256_N + col_origin + col
)
first_index = sample_base + row * _D256_N + row_origin + col
second_index = sample_base + row * _D256_N + col_origin + col
for _ in range(_D256_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[fp32_shared_index]
accumulator[fp32_shared_index] = source[accumulator_index]
first[half_shared_index] = destination[first_index]
second[half_shared_index] = destination[second_index]
row += _D256_ROW_STRIDE
fp32_shared_index += _D256_FP32_SHARED_ROW_STRIDE
half_shared_index += _D256_ROW_STRIDE * _D256_HALF_LD
diagonal_index += _D256_GLOBAL_ROW_STRIDE
accumulator_index += _D256_GLOBAL_ROW_STRIDE
first_index += _D256_GLOBAL_ROW_STRIDE
second_index += _D256_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d256_store_tile_and_sidecar(
shared,
destination,
sidecar,
sample,
sample_base,
tile_row,
tile_col,
):
cuda.syncthreads()
destination_index = (
sample_base
+ tile_row * _D256_NB * _D256_N
+ tile_col * _D256_NB
)
sidecar_index = _d256_sidecar_tile_base(
sample, tile_row, tile_col
)
_d256_store_fp32_and_half_tile(
get_array_ptr(shared),
get_array_ptr(destination),
get_array_ptr(sidecar),
destination_index,
sidecar_index,
)
@cuda.jit(device=True, forceinline=True)
def _d256_store_tile(
shared, destination, sample_base, tile_row, tile_col, triangular
):
cuda.syncthreads()
destination_index = (
sample_base
+ tile_row * _D256_NB * _D256_N
+ tile_col * _D256_NB
)
if triangular:
_d256_store_fp32_upper_tile(
get_array_ptr(shared),
get_array_ptr(destination),
destination_index,
)
else:
_d256_store_fp32_tile(
get_array_ptr(shared),
get_array_ptr(destination),
destination_index,
)
@cuda.jit
def _d256_initial_solve_stage(source, output, sidecar):
block = cuda.blockIdx.x
sample = block & 63
worker = block >> 6
sample_base = sample * _D256_MATRIX_ELEMENTS
panel_tile = 1 + worker
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
a_shared = shared[0:_D256_FP32_TILE_ELEMENTS]
b_shared = shared[
_D256_FP32_TILE_ELEMENTS : 2 * _D256_FP32_TILE_ELEMENTS
]
local_info = shared[2 * _D256_FP32_TILE_ELEMENTS:].view(np.int32)
_d256_load_fp32_pair_async(
source,
source,
sample_base,
0,
0,
a_shared,
0,
panel_tile,
b_shared,
)
_d256_cholesky.factorize(
a_shared, local_info, lda=_D256_FP32_LD
)
_d256_triangular.solve(
a_shared,
b_shared,
lda=_D256_FP32_LD,
ldb=_D256_FP32_LD,
)
_d256_store_tile_and_sidecar(
b_shared,
output,
sidecar,
sample,
sample_base,
0,
panel_tile,
)
if worker == 0:
_d256_store_tile(
a_shared, output, sample_base, 0, 0, True
)
@cuda.jit
def _d256_middle_solve_stage(output, sidecar):
block = cuda.blockIdx.x
sample = block & 63
worker = block >> 6
sample_base = sample * _D256_MATRIX_ELEMENTS
panel_tile = 2 + worker
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
a_shared = shared[0:_D256_FP32_TILE_ELEMENTS]
b_shared = shared[
_D256_FP32_TILE_ELEMENTS : 2 * _D256_FP32_TILE_ELEMENTS
]
_d256_load_fp32_pair_async(
output,
output,
sample_base,
1,
1,
a_shared,
1,
panel_tile,
b_shared,
)
_d256_triangular.solve(
a_shared,
b_shared,
lda=_D256_FP32_LD,
ldb=_D256_FP32_LD,
)
_d256_store_tile_and_sidecar(
b_shared,
output,
sidecar,
sample,
sample_base,
1,
panel_tile,
)
@cuda.jit
def _d256_update_stage(source, output, sidecar):
block = cuda.blockIdx.x
sample = block & 63
task = block >> 6
if task < 3:
tile_row = 1
tile_col = task + 1
elif task < 5:
tile_row = 2
tile_col = task - 1
else:
tile_row = 3
tile_col = 3
sample_base = sample * _D256_MATRIX_ELEMENTS
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
accumulator = shared[0:_D256_FP32_TILE_ELEMENTS]
update_start = _D256_FP32_TILE_ELEMENTS
update_shared = shared[
update_start : update_start + _D256_HALF_TILE_ELEMENTS
].view(np.float16)
first = update_shared[0:_D256_HALF_TILE_ELEMENTS]
second = update_shared[
_D256_HALF_TILE_ELEMENTS : 2 * _D256_HALF_TILE_ELEMENTS
]
local_info = shared[
update_start + _D256_HALF_TILE_ELEMENTS :
].view(np.int32)
if tile_row == tile_col:
_d256_load_update_tile_async(
source,
sidecar,
sample,
sample_base,
tile_row,
tile_col,
accumulator,
0,
tile_row,
first,
)
_d256_gemm.execute(
-1.0, first, first, 1.0, accumulator
)
else:
_d256_load_update_pair_async(
source,
sidecar,
sample,
sample_base,
tile_row,
tile_col,
accumulator,
0,
tile_row,
first,
tile_col,
second,
)
_d256_gemm.execute(
-1.0, first, second, 1.0, accumulator
)
cuda.syncthreads()
if task == 0:
_d256_cholesky.factorize(
accumulator, local_info, lda=_D256_FP32_LD
)
_d256_store_tile(
accumulator,
output,
sample_base,
tile_row,
tile_col,
tile_row == tile_col,
)
@cuda.jit(device=True, forceinline=True)
def _d256_store_tile_shared_half_and_prefetch(
source,
output,
destination,
sample_base,
tile_row,
tile_col,
accumulator_row,
accumulator_col,
):
cuda.syncthreads()
output_index = (
sample_base
+ tile_row * _D256_NB * _D256_N
+ tile_col * _D256_NB
)
accumulator_index = (
sample_base
+ accumulator_row * _D256_NB * _D256_N
+ accumulator_col * _D256_NB
)
_d256_store_fp32_and_shared_half_prefetch(
get_array_ptr(source),
get_array_ptr(output),
get_array_ptr(destination),
get_array_ptr(output),
output_index,
accumulator_index,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit
def _d256_final_solve_update_factor(output):
sample = cuda.blockIdx.x
sample_base = sample * _D256_MATRIX_ELEMENTS
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
factor = shared[0:_D256_FP32_TILE_ELEMENTS]
panel = shared[
_D256_FP32_TILE_ELEMENTS : 2 * _D256_FP32_TILE_ELEMENTS
]
local_info = shared[2 * _D256_FP32_TILE_ELEMENTS:].view(np.int32)
_d256_load_fp32_pair_async(
output,
output,
sample_base,
2,
2,
factor,
2,
3,
panel,
)
_d256_triangular.solve(
factor,
panel,
lda=_D256_FP32_LD,
ldb=_D256_FP32_LD,
)
half_panel = shared[
0 : _D256_HALF_TILE_ELEMENTS // 2
].view(np.float16)
_d256_store_tile_shared_half_and_prefetch(
panel,
output,
half_panel,
sample_base,
2,
3,
3,
3,
)
_d256_gemm.execute(
-1.0, half_panel, half_panel, 1.0, panel
)
cuda.syncthreads()
_d256_cholesky.factorize(
panel, local_info, lda=_D256_FP32_LD
)
_d256_store_tile(
panel, output, sample_base, 3, 3, True
)
@cuda.jit
def _d256_middle_update_stage(output, sidecar):
block = cuda.blockIdx.x
sample = block & 63
task = block >> 6
if task == 0:
tile_row = 2
tile_col = 2
elif task == 1:
tile_row = 2
tile_col = 3
else:
tile_row = 3
tile_col = 3
sample_base = sample * _D256_MATRIX_ELEMENTS
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
accumulator = shared[0:_D256_FP32_TILE_ELEMENTS]
update_start = _D256_FP32_TILE_ELEMENTS
update_shared = shared[
update_start : update_start + _D256_HALF_TILE_ELEMENTS
].view(np.float16)
first = update_shared[0:_D256_HALF_TILE_ELEMENTS]
second = update_shared[
_D256_HALF_TILE_ELEMENTS : 2 * _D256_HALF_TILE_ELEMENTS
]
local_info = shared[
update_start + _D256_HALF_TILE_ELEMENTS :
].view(np.int32)
if task == 1:
_d256_load_update_pair_async(
output,
sidecar,
sample,
sample_base,
tile_row,
tile_col,
accumulator,
1,
2,
first,
3,
second,
)
_d256_gemm.execute(
-1.0, first, second, 1.0, accumulator
)
else:
_d256_load_update_tile_async(
output,
sidecar,
sample,
sample_base,
tile_row,
tile_col,
accumulator,
1,
tile_row,
first,
)
_d256_gemm.execute(
-1.0, first, first, 1.0, accumulator
)
cuda.syncthreads()
if task == 0:
_d256_cholesky.factorize(
accumulator, local_info, lda=_D256_FP32_LD
)
_d256_store_tile(
accumulator,
output,
sample_base,
tile_row,
tile_col,
task != 1,
)
# The batch-60 N=1024 path factors each gathered 256x256 diagonal with the
# retained pair-local implementation. Brief 71 replaces only the direct
# batch-64 N=256 dispatch above.
@cuda.jit(device=True, forceinline=True)
def _b1024_factor_sidecar_slot(tile_row, tile_col):
if tile_row == 0:
return tile_col - 1
if tile_row == 1:
return tile_col + 1
return 5
@cuda.jit(device=True, forceinline=True)
def _b1024_store_factor_tile_and_sidecar(
shared,
destination,
sidecar,
sample,
sample_base,
tile_row,
tile_col,
):
cuda.syncthreads()
destination_index = (
sample_base
+ tile_row * _D256_NB * _D256_N
+ tile_col * _D256_NB
)
slot = _b1024_factor_sidecar_slot(tile_row, tile_col)
sidecar_index = (
sample * _B1024_FACTOR_SIDECAR_TILES + slot
) * _D256_TILE_ELEMENTS
_d256_store_fp32_and_half_tile(
get_array_ptr(shared),
get_array_ptr(destination),
get_array_ptr(sidecar),
destination_index,
sidecar_index,
)
@cuda.jit(device=True, forceinline=True)
def _b1024_factor_load_source_tile(
source,
source_base,
tile_row,
tile_col,
shared,
shared_leading,
triangular,
):
# The global source contains the lower triangle, while the MathDx factor
# is stored as an upper factor. Read contiguous rows from the symmetric
# lower tile and transpose directly into the padded shared tile.
tid = cuda.threadIdx.x
load_row = tid // _D256_NB
load_col = tid - load_row * _D256_NB
source_index = (
source_base
+ (tile_col * _D256_NB + load_row) * 1024
+ tile_row * _D256_NB
+ load_col
)
shared_index = load_col * shared_leading + load_row
for _ in range(_D256_LOADS_PER_THREAD):
if not triangular or load_row >= load_col:
shared[shared_index] = source[source_index]
load_row += _D256_ROW_STRIDE
source_index += _D256_ROW_STRIDE * 1024
shared_index += _D256_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_factor_load_first_diagonal_update(
source,
destination,
source_base,
destination_base,
tile,
diagonal,
update,
):
tid = cuda.threadIdx.x
load_row = tid // _D256_NB
load_col = tid - load_row * _D256_NB
origin = tile * _D256_NB
source_index = (
source_base
+ (origin + load_row) * 1024
+ origin
+ load_col
)
diagonal_index = load_col * _D256_FP32_LD + load_row
update_index = load_row * _D256_HALF_LD + load_col
destination_index = (
destination_base + load_row * _D256_N + origin + load_col
)
for _ in range(_D256_LOADS_PER_THREAD):
if load_row >= load_col:
diagonal[diagonal_index] = source[source_index]
update[update_index] = destination[destination_index]
load_row += _D256_ROW_STRIDE
source_index += _D256_ROW_STRIDE * 1024
diagonal_index += _D256_ROW_STRIDE
update_index += _D256_ROW_STRIDE * _D256_HALF_LD
destination_index += _D256_ROW_STRIDE * _D256_N
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_factor_load_first_panel_update(
source,
destination,
source_base,
destination_base,
tile_row,
tile_col,
accumulator,
first,
second,
):
tid = cuda.threadIdx.x
load_row = tid // _D256_NB
load_col = tid - load_row * _D256_NB
row_origin = tile_row * _D256_NB
col_origin = tile_col * _D256_NB
source_index = (
source_base
+ (col_origin + load_row) * 1024
+ row_origin
+ load_col
)
accumulator_index = load_col * _D256_FP32_LD + load_row
update_index = load_row * _D256_HALF_LD + load_col
first_index = (
destination_base + load_row * _D256_N + row_origin + load_col
)
second_index = (
destination_base + load_row * _D256_N + col_origin + load_col
)
for _ in range(_D256_LOADS_PER_THREAD):
accumulator[accumulator_index] = source[source_index]
first[update_index] = destination[first_index]
second[update_index] = destination[second_index]
load_row += _D256_ROW_STRIDE
source_index += _D256_ROW_STRIDE * 1024
accumulator_index += _D256_ROW_STRIDE
update_index += _D256_ROW_STRIDE * _D256_HALF_LD
first_index += _D256_ROW_STRIDE * _D256_N
second_index += _D256_ROW_STRIDE * _D256_N
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_direct_sidecar_slot(tile_row, tile_col):
return tile_col * (tile_col + 1) // 2 + tile_row
@cuda.jit(device=True, forceinline=True)
def _b1024_direct_load_half_tile(
sidecar, sample, tile_row, tile_col, shared
):
slot = _b1024_direct_sidecar_slot(tile_row, tile_col)
source_index = (
sample * _B1024_DIRECT_SIDECAR_TILES + slot
) * _D256_TILE_ELEMENTS
_d256_cp_async_half_tile(
get_array_ptr(shared), get_array_ptr(sidecar), source_index
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_direct_load_half_pair(
sidecar,
sample,
first_row,
first_col,
first,
second_row,
second_col,
second,
):
first_slot = _b1024_direct_sidecar_slot(first_row, first_col)
second_slot = _b1024_direct_sidecar_slot(second_row, second_col)
sample_base = sample * _B1024_DIRECT_SIDECAR_TILES
_d256_cp_async_half_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(sidecar),
(sample_base + first_slot) * _D256_TILE_ELEMENTS,
(sample_base + second_slot) * _D256_TILE_ELEMENTS,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_direct_store_lower(
shared,
output,
sample,
tile_row,
tile_col,
triangular,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
output_row = tid // _D256_NB
output_col = tid - output_row * _D256_NB
shared_index = output_col * _D256_FP32_LD + output_row
output_index = (
sample * 1024 * 1024
+ (tile_col * _D256_NB + output_row) * 1024
+ tile_row * _D256_NB
+ output_col
)
for _ in range(_D256_LOADS_PER_THREAD):
if not triangular or output_row >= output_col:
output[output_index] = shared[shared_index]
output_row += _D256_ROW_STRIDE
shared_index += _D256_ROW_STRIDE
output_index += _D256_ROW_STRIDE * 1024
@cuda.jit(device=True, forceinline=True)
def _b1024_direct_store_half(
shared, sidecar, sample, tile_row, tile_col
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _D256_FP32_LD + col
slot = _b1024_direct_sidecar_slot(tile_row, tile_col)
sidecar_index = (
sample * _B1024_DIRECT_SIDECAR_TILES + slot
) * _D256_TILE_ELEMENTS + row * _D256_NB + col
for _ in range(_D256_LOADS_PER_THREAD):
sidecar[sidecar_index] = shared[shared_index]
row += _D256_ROW_STRIDE
shared_index += _D256_ROW_STRIDE * _D256_FP32_LD
sidecar_index += _D256_ROW_STRIDE * _D256_NB
@cuda.jit
def _batch1024_direct_diagonal(source, output, current_tile):
sample = cuda.blockIdx.x
sample_base = sample * 1024 * 1024
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
diagonal = shared[0:_D256_FP32_TILE_ELEMENTS]
local_info = shared[
_D256_FP32_TILE_ELEMENTS :
].view(np.int32)
_b1024_factor_load_source_tile(
source,
sample_base,
current_tile,
current_tile,
diagonal,
_D256_FP32_LD,
True,
)
_d256_cholesky.factorize(
diagonal, local_info, lda=_D256_FP32_LD
)
_b1024_direct_store_lower(
diagonal,
output,
sample,
current_tile,
current_tile,
True,
)
@cuda.jit
def _batch1024_direct_panels(
source, output, sidecar, current_tile
):
sample = cuda.blockIdx.y
panel_tile = current_tile + 1 + cuda.blockIdx.x
sample_base = sample * 1024 * 1024
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
diagonal = shared[0:_D256_FP32_TILE_ELEMENTS]
panel = shared[
_D256_FP32_TILE_ELEMENTS : 2 * _D256_FP32_TILE_ELEMENTS
]
_b1024_factor_load_source_tile(
output,
sample_base,
current_tile,
current_tile,
diagonal,
_D256_FP32_LD,
True,
)
_b1024_factor_load_source_tile(
source,
sample_base,
current_tile,
panel_tile,
panel,
_D256_FP32_LD,
False,
)
_d256_triangular.solve(
diagonal,
panel,
lda=_D256_FP32_LD,
ldb=_D256_FP32_LD,
)
_b1024_direct_store_lower(
panel,
output,
sample,
current_tile,
panel_tile,
False,
)
_b1024_direct_store_half(
panel, sidecar, sample, current_tile, panel_tile
)
@cuda.jit
def _batch1024_direct_updates(
sidecar, source, output, current_tile
):
sample = cuda.blockIdx.y
task = cuda.blockIdx.x
tile_row = current_tile + 1
row_tasks = _B1024_DIRECT_TILE_COUNT - tile_row
while task >= row_tasks:
task -= row_tasks
tile_row += 1
row_tasks -= 1
tile_col = tile_row + task
sample_base = sample * 1024 * 1024
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
accumulator = shared[0:_D256_FP32_TILE_ELEMENTS]
update_shared = shared[
_D256_FP32_TILE_ELEMENTS :
_D256_FP32_TILE_ELEMENTS + _D256_HALF_TILE_ELEMENTS
].view(np.float16)
first = update_shared[0:_D256_HALF_TILE_ELEMENTS]
second = update_shared[
_D256_HALF_TILE_ELEMENTS : 2 * _D256_HALF_TILE_ELEMENTS
]
_b1024_factor_load_source_tile(
source,
sample_base,
tile_row,
tile_col,
accumulator,
_D256_FP32_LD,
tile_row == tile_col,
)
_b1024_direct_load_half_pair(
sidecar,
sample,
current_tile,
tile_row,
first,
current_tile,
tile_col,
second,
)
_d256_gemm.execute(-1.0, first, second, 1.0, accumulator)
cuda.syncthreads()
_b1024_direct_store_lower(
accumulator,
output,
sample,
tile_row,
tile_col,
tile_row == tile_col,
)
@cuda.jit
def _batch1024_factor_diagonal(
source, destination, panel_origin, current_tile
):
sample = cuda.blockIdx.x
sample_base = sample * _D256_MATRIX_ELEMENTS
source_base = sample * 1024 * 1024 + panel_origin * (1024 + 1)
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
diagonal = shared[0:_D256_FP32_TILE_ELEMENTS]
update_shared = shared[
_D256_FP32_TILE_ELEMENTS :
_D256_FP32_TILE_ELEMENTS + _D256_HALF_TILE_ELEMENTS // 2
].view(np.float16)
local_info = shared[
_D256_FP32_TILE_ELEMENTS + _D256_HALF_TILE_ELEMENTS // 2 :
].view(np.int32)
if current_tile == 0:
_b1024_factor_load_source_tile(
source,
source_base,
current_tile,
current_tile,
diagonal,
_D256_FP32_LD,
True,
)
else:
_b1024_factor_load_first_diagonal_update(
source,
destination,
source_base,
sample_base,
current_tile,
diagonal,
update_shared,
)
_d256_gemm.execute(
-1.0, update_shared, update_shared, 1.0, diagonal
)
cuda.syncthreads()
for previous_tile in range(1, current_tile):
_d256_load_tile(
destination,
sample_base,
previous_tile,
current_tile,
update_shared,
_D256_HALF_LD,
)
_d256_gemm.execute(
-1.0, update_shared, update_shared, 1.0, diagonal
)
cuda.syncthreads()
_d256_cholesky.factorize(
diagonal, local_info, lda=_D256_FP32_LD
)
_d256_store_tile(
diagonal,
destination,
sample_base,
current_tile,
current_tile,
True,
)
@cuda.jit
def _batch1024_factor_panels(
source, destination, sidecar, panel_origin, current_tile
):
sample = cuda.blockIdx.y
panel_tile = current_tile + 1 + cuda.blockIdx.x
sample_base = sample * _D256_MATRIX_ELEMENTS
source_base = sample * 1024 * 1024 + panel_origin * (1024 + 1)
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
diagonal = shared[0:_D256_FP32_TILE_ELEMENTS]
panel = shared[
_D256_FP32_TILE_ELEMENTS : 2 * _D256_FP32_TILE_ELEMENTS
]
update_shared = shared[
2 * _D256_FP32_TILE_ELEMENTS :
2 * _D256_FP32_TILE_ELEMENTS + _D256_HALF_TILE_ELEMENTS
].view(np.float16)
first = update_shared[0:_D256_HALF_TILE_ELEMENTS]
second = update_shared[
_D256_HALF_TILE_ELEMENTS : 2 * _D256_HALF_TILE_ELEMENTS
]
_d256_load_tile(
destination,
sample_base,
current_tile,
current_tile,
diagonal,
_D256_FP32_LD,
)
if current_tile == 0:
_b1024_factor_load_source_tile(
source,
source_base,
current_tile,
panel_tile,
panel,
_D256_FP32_LD,
False,
)
else:
_b1024_factor_load_first_panel_update(
source,
destination,
source_base,
sample_base,
current_tile,
panel_tile,
panel,
first,
second,
)
_d256_gemm.execute(-1.0, first, second, 1.0, panel)
cuda.syncthreads()
for previous_tile in range(1, current_tile):
_d256_load_tile_pair(
destination,
sample_base,
previous_tile,
current_tile,
first,
panel_tile,
second,
)
_d256_gemm.execute(-1.0, first, second, 1.0, panel)
cuda.syncthreads()
_d256_triangular.solve(
diagonal,
panel,
lda=_D256_FP32_LD,
ldb=_D256_FP32_LD,
)
_b1024_store_factor_tile_and_sidecar(
panel,
destination,
sidecar,
sample,
sample_base,
current_tile,
panel_tile,
)
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_load_factor_tile(
source, sample_base, tile_row, tile_col, shared
):
tid = cuda.threadIdx.x
row_origin = tile_row * _D256_NB
col_origin = tile_col * _D256_NB
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _B1024_FP32_LD + col
source_index = (
sample_base + (row_origin + row) * _D256_N + col_origin + col
)
for _ in range(_B1024_LOADS_PER_THREAD):
shared[shared_index] = source[source_index]
shared_index += _B1024_ROW_STRIDE * _B1024_FP32_LD
source_index += _B1024_GLOBAL_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_load(
output,
output_base,
panel_row_origin,
factor_col_origin,
shared,
):
tid = cuda.threadIdx.x
panel_row = tid // _D256_NB
factor_row = tid - panel_row * _D256_NB
shared_index = panel_row * _B1024_FP32_LD + factor_row
output_index = (
output_base
+ (panel_row_origin + panel_row) * 1024
+ factor_col_origin
+ factor_row
)
for _ in range(_B1024_LOADS_PER_THREAD):
shared[shared_index] = output[output_index]
panel_row += _B1024_ROW_STRIDE
shared_index += _B1024_ROW_STRIDE * _B1024_FP32_LD
output_index += _B1024_ROW_STRIDE * 1024
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_load_update(
factor,
output,
factor_base,
output_base,
previous_tile,
current_tile,
panel_col_base,
panel_row_origin,
factor_shared,
panel_shared,
):
factor_slot = _b1024_factor_sidecar_slot(
previous_tile, current_tile
)
factor_index = (
factor_base + factor_slot * _D256_TILE_ELEMENTS
)
_dx_b1024_cp_async_half_tile(
get_array_ptr(factor_shared),
factor,
factor_index,
_D256_NB,
)
panel_index = (
output_base
+ panel_row_origin * 1024
+ panel_col_base
+ previous_tile * _D256_NB
)
_dx_b1024_load_convert_half_tile(
get_array_ptr(panel_shared),
output,
panel_index,
1024,
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_store(
shared,
output,
output_base,
panel_row_origin,
factor_col_origin,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
panel_row = tid // _D256_NB
factor_row = tid - panel_row * _D256_NB
shared_index = panel_row * _B1024_FP32_LD + factor_row
output_index = (
output_base
+ (panel_row_origin + panel_row) * 1024
+ factor_col_origin
+ factor_row
)
for _ in range(_B1024_LOADS_PER_THREAD):
output[output_index] = shared[shared_index]
panel_row += _B1024_ROW_STRIDE
shared_index += _B1024_ROW_STRIDE * _B1024_FP32_LD
output_index += _B1024_ROW_STRIDE * 1024
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_apply_update(
factor,
output,
factor_base,
output_base,
previous_tile,
current_tile,
panel_col_base,
panel_row_origin,
shared,
panel_shared,
):
workspace_shared = shared[
_B1024_FP32_TILE_ELEMENTS :
_B1024_FP32_TILE_ELEMENTS + _B1024_HALF_TILE_ELEMENTS
]
update_shared = workspace_shared.view(np.float16)
factor_shared = update_shared[0:_B1024_HALF_TILE_ELEMENTS]
update_panel_shared = update_shared[
_B1024_HALF_TILE_ELEMENTS : 2 * _B1024_HALF_TILE_ELEMENTS
]
_b1024_panel_load_update(
factor,
output,
factor_base,
output_base,
previous_tile,
current_tile,
panel_col_base,
panel_row_origin,
factor_shared,
update_panel_shared,
)
_b1024_panel_gemm.execute(
-1.0, factor_shared, update_panel_shared, 1.0, panel_shared
)
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_apply_update_bank(
factor,
output,
factor_base,
output_base,
previous_tile,
current_tile,
panel_col_base,
panel_row_origin,
workspace_shared,
panel_shared,
):
update_shared = workspace_shared.view(np.float16)
factor_shared = update_shared[0:_B1024_HALF_TILE_ELEMENTS]
update_panel_shared = update_shared[
_B1024_HALF_TILE_ELEMENTS : 2 * _B1024_HALF_TILE_ELEMENTS
]
_b1024_panel_load_update(
factor,
output,
factor_base,
output_base,
previous_tile,
current_tile,
panel_col_base,
panel_row_origin,
factor_shared,
update_panel_shared,
)
_b1024_panel_gemm.execute(
-1.0, factor_shared, update_panel_shared, 1.0, panel_shared
)
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_apply_solve(
factor, factor_base, tile, shared, panel_shared
):
factor_shared = shared[
_B1024_FP32_TILE_ELEMENTS :
2 * _B1024_FP32_TILE_ELEMENTS
]
_b1024_panel_load_factor_tile(
factor, factor_base, tile, tile, factor_shared
)
_b1024_panel_triangular.solve(
factor_shared,
panel_shared,
lda=_B1024_FP32_LD,
ldb=_B1024_FP32_LD,
)
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_apply_solve_bank(factor_shared, panel_shared):
_b1024_panel_triangular.solve(
factor_shared,
panel_shared,
lda=_B1024_FP32_LD,
ldb=_B1024_FP32_LD,
)
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_async_load(
source, source_index, global_ld, destination
):
_dx_b1024_cp_async_fp32_tile(
get_array_ptr(destination),
source,
source_index,
global_ld,
)
@cuda.jit(device=True, forceinline=True)
def _b1024_panel_async_wait():
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _batch1024_panel_solve_body(
factor,
source,
output,
panel_col_base,
panel_row_base,
current_tile,
shared,
):
group = cuda.blockIdx.x
sample = cuda.blockIdx.y
factor_base = sample * _D256_MATRIX_ELEMENTS
output_base = sample * 1024 * 1024
panel_row_origin = panel_row_base + group * _D256_NB
panel_shared = shared[0:_B1024_FP32_TILE_ELEMENTS]
factor_shared = shared[
_B1024_FP32_TILE_ELEMENTS : 2 * _B1024_FP32_TILE_ELEMENTS
]
# The final two warps asynchronously stage the panel while all eight warps
# load the independent diagonal tile. This node carries only the two
# live FP32 operands of the MathDx triangular solve.
_b1024_panel_async_load(
source,
output_base
+ panel_row_origin * 1024
+ panel_col_base
+ current_tile * _D256_NB,
1024,
panel_shared,
)
_b1024_panel_load_factor_tile(
factor,
factor_base,
current_tile,
current_tile,
factor_shared,
)
_b1024_panel_async_wait()
_b1024_panel_apply_solve_bank(factor_shared, panel_shared)
_b1024_panel_store(
panel_shared,
output,
output_base,
panel_row_origin,
panel_col_base + current_tile * _D256_NB,
)
@cuda.jit(device=True, forceinline=True)
def _batch1024_panel_update_body(
factor,
source,
output,
panel_col_base,
panel_row_base,
current_tile,
shared,
):
group = cuda.blockIdx.x
sample = cuda.blockIdx.y
factor_base = (
sample * _B1024_FACTOR_SIDECAR_TILES * _D256_TILE_ELEMENTS
)
output_base = sample * 1024 * 1024
panel_row_origin = panel_row_base + group * _D256_NB
accumulator = shared[0:_B1024_FP32_TILE_ELEMENTS]
workspace = shared[
_B1024_FP32_TILE_ELEMENTS :
_B1024_FP32_TILE_ELEMENTS + _B1024_HALF_TILE_ELEMENTS
]
update_shared = workspace.view(np.float16)
factor_shared = update_shared[0:_B1024_HALF_TILE_ELEMENTS]
panel_shared = update_shared[
_B1024_HALF_TILE_ELEMENTS : 2 * _B1024_HALF_TILE_ELEMENTS
]
# Stage the FP32 accumulator with the producer warps while all threads
# convert the first LD72 factor/panel pair. Subsequent rank-k terms reuse
# the same compact half workspace, so no triangular-solve fragments remain
# live in this kernel.
_b1024_panel_async_load(
source,
output_base
+ panel_row_origin * 1024
+ panel_col_base
+ current_tile * _D256_NB,
1024,
accumulator,
)
_b1024_panel_load_update(
factor,
output,
factor_base,
output_base,
0,
current_tile,
panel_col_base,
panel_row_origin,
factor_shared,
panel_shared,
)
_b1024_panel_async_wait()
_b1024_panel_gemm.execute(
-1.0, factor_shared, panel_shared, 1.0, accumulator
)
cuda.syncthreads()
for previous_tile in range(1, current_tile):
_b1024_panel_load_update(
factor,
output,
factor_base,
output_base,
previous_tile,
current_tile,
panel_col_base,
panel_row_origin,
factor_shared,
panel_shared,
)
_b1024_panel_gemm.execute(
-1.0, factor_shared, panel_shared, 1.0, accumulator
)
cuda.syncthreads()
_b1024_panel_store(
accumulator,
output,
output_base,
panel_row_origin,
panel_col_base + current_tile * _D256_NB,
)
@cuda.jit(max_registers=48)
def _batch1024_panel_solve_bulk(factor, source, output, current_tile):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_solve_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
0,
256,
current_tile,
shared,
)
@cuda.jit(max_registers=64)
def _batch1024_panel_solve_middle(factor, source, output, current_tile):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_solve_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
256,
512,
current_tile,
shared,
)
@cuda.jit(max_registers=72)
def _batch1024_panel_solve_tail(factor, source, output, current_tile):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_solve_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
512,
768,
current_tile,
shared,
)
@cuda.jit(max_registers=48)
def _batch1024_panel_update1_bulk(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
0,
256,
1,
shared,
)
@cuda.jit(max_registers=48)
def _batch1024_panel_update2_bulk(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
0,
256,
2,
shared,
)
@cuda.jit(max_registers=48)
def _batch1024_panel_update3_bulk(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
0,
256,
3,
shared,
)
@cuda.jit(max_registers=64)
def _batch1024_panel_update1_middle(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
256,
512,
1,
shared,
)
@cuda.jit(max_registers=64)
def _batch1024_panel_update2_middle(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
256,
512,
2,
shared,
)
@cuda.jit(max_registers=64)
def _batch1024_panel_update3_middle(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
256,
512,
3,
shared,
)
@cuda.jit(max_registers=72)
def _batch1024_panel_update1_tail(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
512,
768,
1,
shared,
)
@cuda.jit(max_registers=72)
def _batch1024_panel_update2_tail(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
512,
768,
2,
shared,
)
@cuda.jit(max_registers=72)
def _batch1024_panel_update3_tail(factor, source, output):
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
_batch1024_panel_update_body(
get_array_ptr(factor),
get_array_ptr(source),
get_array_ptr(output),
512,
768,
3,
shared,
)
def _d256_launchers_for_current_queue(
source_view, output_view, sidecar_view, batch
):
global _d256_initial_solve_dispatch
global _d256_middle_solve_dispatch
global _d256_update_dispatch, _d256_middle_dispatch
global _d256_final_dispatch
global _d256_queue_handle, _d256_queue
global _d256_initial_solve_launcher
global _d256_middle_solve_launcher
global _d256_update_launcher, _d256_middle_launcher
global _d256_final_launcher
if _d256_initial_solve_dispatch is None:
_d256_initial_solve_dispatch = _d256_initial_solve_stage.specialize(
source_view, output_view, sidecar_view
)
_d256_middle_solve_dispatch = _d256_middle_solve_stage.specialize(
output_view, sidecar_view
)
_d256_update_dispatch = _d256_update_stage.specialize(
source_view,
output_view,
sidecar_view,
)
_d256_middle_dispatch = _d256_middle_update_stage.specialize(
output_view, sidecar_view
)
_d256_final_dispatch = _d256_final_solve_update_factor.specialize(
output_view
)
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
for dispatch, shared_bytes in (
(
_d256_initial_solve_dispatch,
_D256_INITIAL_SOLVE_SHARED_BYTES,
),
(_d256_middle_solve_dispatch, _D256_SOLVE_SHARED_BYTES),
(_d256_update_dispatch, _D256_UPDATE_SHARED_BYTES),
(_d256_middle_dispatch, _D256_UPDATE_SHARED_BYTES),
(_d256_final_dispatch, _D256_FINAL_SHARED_BYTES),
):
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
numba_driver.driver.cuKernelSetAttribute(
attribute,
shared_bytes,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if (
queue_handle != _d256_queue_handle
or _d256_initial_solve_launcher is None
):
_d256_queue = _external_queue(queue_handle)
_d256_queue_handle = queue_handle
_d256_initial_solve_launcher = _d256_initial_solve_dispatch[
batch * 3,
_D256_NT,
_d256_queue,
_D256_INITIAL_SOLVE_SHARED_BYTES,
]
_d256_middle_solve_launcher = _d256_middle_solve_dispatch[
batch * 2,
_D256_NT,
_d256_queue,
_D256_SOLVE_SHARED_BYTES,
]
_d256_update_launcher = _d256_update_dispatch[
batch * 6,
_D256_NT,
_d256_queue,
_D256_UPDATE_SHARED_BYTES,
]
_d256_middle_launcher = _d256_middle_dispatch[
batch * 3,
_D256_NT,
_d256_queue,
_D256_UPDATE_SHARED_BYTES,
]
_d256_final_launcher = _d256_final_dispatch[
batch,
_D256_NT,
_d256_queue,
_D256_FINAL_SHARED_BYTES,
]
return (
_d256_initial_solve_launcher,
_d256_middle_solve_launcher,
_d256_update_launcher,
_d256_middle_launcher,
_d256_final_launcher,
)
_d256_output_cache = {}
_d256_last_input = None
_d256_last_entry = None
def _compute_d256(entry) -> None:
(
initial_solve_launcher,
middle_solve_launcher,
update_launcher,
middle_launcher,
final_launcher,
) = (
_d256_launchers_for_current_queue(
entry[1], entry[3], entry[6], entry[0].shape[0]
)
)
initial_solve_launcher(entry[1], entry[3], entry[6])
update_launcher(entry[1], entry[3], entry[6])
middle_solve_launcher(entry[3], entry[6])
middle_launcher(entry[3], entry[6])
final_launcher(entry[3])
def _run_d256(data: torch.Tensor) -> torch.Tensor:
global _d256_last_input, _d256_last_entry
if _d256_last_input is data:
entry = _d256_last_entry
else:
key = id(data)
entry = _d256_output_cache.get(key)
if entry is None or entry[0] is not data:
output = torch.zeros_like(
data, memory_format=torch.contiguous_format
)
sidecar = torch.empty(
data.shape[0] * _D256_SIDECAR_MATRIX_ELEMENTS,
dtype=torch.float16,
device=data.device,
)
entry = [
data,
_as_numba_flat_array(data),
output,
_as_numba_flat_array(output),
output.transpose(-2, -1),
sidecar,
_as_numba_half_array(sidecar),
None,
]
_d256_output_cache[key] = entry
_d256_last_input = data
_d256_last_entry = entry
graph = entry[7]
if graph is None:
_compute_d256(entry)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
_compute_d256(entry)
entry[7] = graph
graph.replay()
return entry[4]
# Multi-owner 64x64 supernode DAG for both ranked N=512 regimes. The old
# direct kernel assigned an entire matrix to one CTA; that leaves only 16 CTAs
# for the latency cell. Here each factor, panel solve, and triangular Schur
# tile has one owner. Graph order is the dependency epoch between stages,
# while every independent tile in an epoch runs in parallel.
_D512_N = 512
_D512_TILES = _D512_N // _D256_NB
_D512_MATRIX_ELEMENTS = _D512_N * _D512_N
_D512_PANEL_TILES = _D512_TILES * (_D512_TILES - 1) // 2
_D512_SIDECAR_MATRIX_ELEMENTS = (
_D512_PANEL_TILES * _D256_SIDECAR_TILE_ELEMENTS
)
_D512_SOLVE_SHARED_BYTES = 2 * _D256_FP32_TILE_ELEMENTS * 4
_D512_UPDATE_SHARED_BYTES = (
_D256_FP32_TILE_ELEMENTS * 4
+ 2 * _D256_HALF_TILE_ELEMENTS * 2
+ 4
)
@cuda.jit(device=True, forceinline=True)
def _d512_tile_index(sample, tile_row, tile_col):
return (
sample * _D512_MATRIX_ELEMENTS
+ tile_row * _D256_NB * _D512_N
+ tile_col * _D256_NB
)
@cuda.jit(device=True, forceinline=True)
def _d512_panel_slot(tile_row, tile_col):
return (
tile_row * (2 * _D512_TILES - tile_row - 1) // 2
+ tile_col
- tile_row
- 1
)
@cuda.jit(device=True, forceinline=True)
def _d512_sidecar_index(sample, tile_row, tile_col):
return (
sample * _D512_SIDECAR_MATRIX_ELEMENTS
+ _d512_panel_slot(tile_row, tile_col)
* _D256_SIDECAR_TILE_ELEMENTS
)
@cuda.jit(device=True, forceinline=True)
def _d512_load_tile(source, sample, tile_row, tile_col, shared):
tid = cuda.threadIdx.x
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _D256_FP32_LD + col
source_index = _d512_tile_index(sample, tile_row, tile_col)
source_index += row * _D512_N + col
for _ in range(_D256_LOADS_PER_THREAD):
shared[shared_index] = source[source_index]
shared_index += _D256_ROW_STRIDE * _D256_FP32_LD
source_index += _D256_ROW_STRIDE * _D512_N
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d512_load_pair(
first_source,
second_source,
sample,
first_row,
first_col,
first,
second_row,
second_col,
second,
):
tid = cuda.threadIdx.x
row = tid // _D256_NB
col = tid - row * _D256_NB
shared_index = row * _D256_FP32_LD + col
first_index = _d512_tile_index(sample, first_row, first_col)
second_index = _d512_tile_index(sample, second_row, second_col)
first_index += row * _D512_N + col
second_index += row * _D512_N + col
for _ in range(_D256_LOADS_PER_THREAD):
first[shared_index] = first_source[first_index]
second[shared_index] = second_source[second_index]
shared_index += _D256_ROW_STRIDE * _D256_FP32_LD
first_index += _D256_ROW_STRIDE * _D512_N
second_index += _D256_ROW_STRIDE * _D512_N
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d512_load_update(
source,
sidecar,
sample,
accumulator_row,
accumulator_col,
accumulator,
panel_row,
first_col,
first,
second_col,
second,
):
_d512_load_tile(
source,
sample,
accumulator_row,
accumulator_col,
accumulator,
)
if first_col == second_col:
_d256_cp_async_half_tile(
get_array_ptr(first),
get_array_ptr(sidecar),
_d512_sidecar_index(sample, panel_row, first_col),
)
else:
_d256_cp_async_half_pair(
get_array_ptr(first),
get_array_ptr(second),
get_array_ptr(sidecar),
_d512_sidecar_index(sample, panel_row, first_col),
_d512_sidecar_index(sample, panel_row, second_col),
)
_dx_cp_async_wait()
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _d512_store_diagonal(shared, output, sample, tile):
cuda.syncthreads()
_d64_store_fp32_upper_tile_ld(
get_array_ptr(shared),
get_array_ptr(output),
_d512_tile_index(sample, tile, tile),
_D512_N,
)
@cuda.jit(device=True, forceinline=True)
def _d512_store_tile(shared, output, sample, tile_row, tile_col):
cuda.syncthreads()
_d64_store_fp32_tile_ld(
get_array_ptr(shared),
get_array_ptr(output),
_d512_tile_index(sample, tile_row, tile_col),
_D512_N,
)
@cuda.jit(device=True, forceinline=True)
def _d512_store_panel(
shared, output, sidecar, sample, tile_row, tile_col
):
cuda.syncthreads()
_d64_store_fp32_and_half_tile_ld(
get_array_ptr(shared),
get_array_ptr(output),
get_array_ptr(sidecar),
_d512_tile_index(sample, tile_row, tile_col),
_d512_sidecar_index(sample, tile_row, tile_col),
_D512_N,
)
@cuda.jit
def _d512_panel_owners(source, output, sidecar, batch, stage):
block = cuda.blockIdx.x
sample = block % batch
panel_col = stage + 1 + block // batch
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
factor = shared[0:_D256_FP32_TILE_ELEMENTS]
panel = shared[
_D256_FP32_TILE_ELEMENTS : 2 * _D256_FP32_TILE_ELEMENTS
]
if stage == 0:
_d512_load_pair(
output,
source,
sample,
stage,
stage,
factor,
stage,
panel_col,
panel,
)
else:
_d512_load_pair(
output,
output,
sample,
stage,
stage,
factor,
stage,
panel_col,
panel,
)
_d256_triangular.solve(
factor,
panel,
lda=_D256_FP32_LD,
ldb=_D256_FP32_LD,
)
_d512_store_panel(
panel, output, sidecar, sample, stage, panel_col
)
@cuda.jit
def _d512_update_owners(source, output, sidecar, batch, stage):
block = cuda.blockIdx.x
sample = block % batch
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
accumulator = shared[0:_D256_FP32_TILE_ELEMENTS]
half_storage = shared[
_D256_FP32_TILE_ELEMENTS :
_D256_FP32_TILE_ELEMENTS + _D256_HALF_TILE_ELEMENTS
].view(np.float16)
first = half_storage[0:_D256_HALF_TILE_ELEMENTS]
second = half_storage[
_D256_HALF_TILE_ELEMENTS : 2 * _D256_HALF_TILE_ELEMENTS
]
local_info = shared[
_D256_FP32_TILE_ELEMENTS + _D256_HALF_TILE_ELEMENTS :
].view(np.int32)
if stage == 0:
_d512_load_tile(source, sample, 0, 0, accumulator)
_d256_cholesky.factorize(
accumulator, local_info, lda=_D256_FP32_LD
)
_d512_store_diagonal(accumulator, output, sample, 0)
else:
task = block // batch
tile_row = stage
row_width = _D512_TILES - tile_row
while task >= row_width:
task -= row_width
tile_row += 1
row_width -= 1
tile_col = tile_row + task
if stage == 1:
_d512_load_update(
source,
sidecar,
sample,
tile_row,
tile_col,
accumulator,
stage - 1,
tile_row,
first,
tile_col,
second,
)
else:
_d512_load_update(
output,
sidecar,
sample,
tile_row,
tile_col,
accumulator,
stage - 1,
tile_row,
first,
tile_col,
second,
)
if tile_row == tile_col:
_d256_gemm.execute(
-1.0, first, first, 1.0, accumulator
)
else:
_d256_gemm.execute(
-1.0, first, second, 1.0, accumulator
)
cuda.syncthreads()
if tile_row == stage and tile_col == stage:
_d256_cholesky.factorize(
accumulator, local_info, lda=_D256_FP32_LD
)
if tile_row == tile_col:
_d512_store_diagonal(
accumulator, output, sample, tile_row
)
else:
_d512_store_tile(
accumulator, output, sample, tile_row, tile_col
)
_d512_solve_dispatch = None
_d512_update_dispatch = None
_d512_launcher_cache = {}
def _d512_launchers_for_current_queue(
source_view, output_view, sidecar_view, batch
):
global _d512_solve_dispatch, _d512_update_dispatch
if _d512_solve_dispatch is None:
_d512_solve_dispatch = _d512_panel_owners.specialize(
source_view,
output_view,
sidecar_view,
np.int32(batch),
np.int32(0),
)
_d512_update_dispatch = _d512_update_owners.specialize(
source_view,
output_view,
sidecar_view,
np.int32(batch),
np.int32(1),
)
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
for dispatch, shared_bytes in (
(_d512_solve_dispatch, _D512_SOLVE_SHARED_BYTES),
(_d512_update_dispatch, _D512_UPDATE_SHARED_BYTES),
):
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
numba_driver.driver.cuKernelSetAttribute(
attribute,
shared_bytes,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
key = (queue_handle, batch)
cached = _d512_launcher_cache.get(key)
if cached is None:
queue = _external_queue(queue_handle)
factor = _d512_update_dispatch[
batch,
_D256_NT,
queue,
_D512_UPDATE_SHARED_BYTES,
]
solves = tuple(
_d512_solve_dispatch[
batch * (_D512_TILES - stage - 1),
_D256_NT,
queue,
_D512_SOLVE_SHARED_BYTES,
]
for stage in range(_D512_TILES - 1)
)
updates = tuple(
_d512_update_dispatch[
batch
* ((_D512_TILES - stage)
* (_D512_TILES - stage + 1) // 2),
_D256_NT,
queue,
_D512_UPDATE_SHARED_BYTES,
]
for stage in range(1, _D512_TILES)
)
cached = (queue, factor, solves, updates)
_d512_launcher_cache[key] = cached
return cached[1], cached[2], cached[3]
def _compute_d512(entry) -> None:
batch = entry[0].shape[0]
factor, solves, updates = _d512_launchers_for_current_queue(
entry[1], entry[3], entry[6], batch
)
factor(
entry[1], entry[3], entry[6], np.int32(batch), np.int32(0)
)
solves[0](
entry[1], entry[3], entry[6], np.int32(batch), np.int32(0)
)
for stage in range(1, _D512_TILES):
updates[stage - 1](
entry[1],
entry[3],
entry[6],
np.int32(batch),
np.int32(stage),
)
if stage + 1 < _D512_TILES:
solves[stage](
entry[1],
entry[3],
entry[6],
np.int32(batch),
np.int32(stage),
)
_d512_output_cache = {}
_d512_last_input = None
_d512_last_entry = None
def _run_d512(data: torch.Tensor) -> torch.Tensor:
global _d512_last_input, _d512_last_entry
if _d512_last_input is data:
entry = _d512_last_entry
else:
key = id(data)
entry = _d512_output_cache.get(key)
if entry is None or entry[0] is not data:
output = torch.zeros_like(
data, memory_format=torch.contiguous_format
)
sidecar = torch.empty(
data.shape[0] * _D512_SIDECAR_MATRIX_ELEMENTS,
dtype=torch.float16,
device=data.device,
)
entry = [
data,
_as_numba_flat_array(data),
output,
_as_numba_flat_array(output),
output.transpose(-2, -1),
sidecar,
_as_numba_half_array(sidecar),
None,
]
_d512_output_cache[key] = entry
_d512_last_input = data
_d512_last_entry = entry
graph = entry[7]
if graph is None:
_compute_d512(entry)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
_compute_d512(entry)
entry[7] = graph
graph.replay()
return entry[4]
# Exact batch-2 N=4096 wavefront. Caller-ordered factor and solve launches
# advance four 64-wide micro-panels before each triangular CuTe rank-k update.
_B4096_N = 4096
_B4096_NB = 64
_B4096_PANEL_TILES = 2
_B4096_PANEL_WIDTH = _B4096_PANEL_TILES * _B4096_NB
_B4096_FACTOR_NT = 256
_B4096_SOLVE_NT = 256
_B4096_WORKERS = 64
_B4096_TILE_ELEMENTS = _B4096_NB * _B4096_NB
_B4096_FP32_LD = 72
_B4096_FP32_TILE_ELEMENTS = _B4096_NB * _B4096_FP32_LD
_B4096_HALF_LD = 72
_B4096_HALF_TILE_ELEMENTS = _B4096_NB * _B4096_HALF_LD
_B4096_FACTOR_LOADS = _B4096_TILE_ELEMENTS // _B4096_FACTOR_NT
_B4096_FACTOR_ROW_STRIDE = _B4096_FACTOR_NT // _B4096_NB
_B4096_FACTOR_GLOBAL_ROW_STRIDE = _B4096_FACTOR_ROW_STRIDE * _B4096_N
_B4096_FACTOR_FP32_ROW_STRIDE = _B4096_FACTOR_ROW_STRIDE * _B4096_FP32_LD
_B4096_FACTOR_HALF_ROW_STRIDE = _B4096_FACTOR_ROW_STRIDE * _B4096_HALF_LD
_B4096_SOLVE_LOADS = _B4096_TILE_ELEMENTS // _B4096_SOLVE_NT
_B4096_SOLVE_ROW_STRIDE = _B4096_SOLVE_NT // _B4096_NB
_B4096_SOLVE_GLOBAL_ROW_STRIDE = _B4096_SOLVE_ROW_STRIDE * _B4096_N
_B4096_SOLVE_FP32_ROW_STRIDE = _B4096_SOLVE_ROW_STRIDE * _B4096_FP32_LD
_B4096_SOLVE_HALF_ROW_STRIDE = _B4096_SOLVE_ROW_STRIDE * _B4096_HALF_LD
_B4096_SHARED_BYTES = (
2 * _B4096_FP32_TILE_ELEMENTS * 4
+ 2 * _B4096_HALF_TILE_ELEMENTS * 2
+ 4
)
_b4096_factor_cholesky = CholeskySolver(
size=(_B4096_NB, _B4096_NB),
precision=np.float32,
data_type="real",
execution="Block",
fill_mode="upper",
arrangement=("row_major", "row_major"),
leading_dimensions=(_B4096_FP32_LD, _B4096_FP32_LD),
block_dim=(_B4096_FACTOR_NT, 1, 1),
sm=100,
)
_b4096_solve_triangular = TriangularSolver(
size=(_B4096_NB, _B4096_NB, _B4096_NB),
precision=np.float32,
data_type="real",
side="left",
diag="non_unit",
transpose_mode="transposed",
arrangement=("row_major", "row_major"),
leading_dimensions=(_B4096_FP32_LD, _B4096_FP32_LD),
fill_mode="upper",
execution="Block",
block_dim=(_B4096_SOLVE_NT, 1, 1),
sm=100,
)
_b4096_factor_gemm = Matmul(
size=(_B4096_NB, _B4096_NB, _B4096_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "row_major", "row_major"),
leading_dimension=(
_B4096_HALF_LD,
_B4096_HALF_LD,
_B4096_FP32_LD,
),
execution="Block",
block_size=_B4096_FACTOR_NT,
alignment=(16, 16, 16),
sm=100,
)
_b4096_solve_gemm = Matmul(
size=(_B4096_NB, _B4096_NB, _B4096_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "row_major", "row_major"),
leading_dimension=(
_B4096_HALF_LD,
_B4096_HALF_LD,
_B4096_FP32_LD,
),
execution="Block",
block_size=_B4096_SOLVE_NT,
alignment=(16, 16, 16),
sm=100,
)
@cuda.jit(device=True, forceinline=True)
def _b4096_factor_load_diagonal(
output, sample_base, panel_origin, shared, leading
):
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_FP32_LD + col
source_index = (
sample_base
+ (panel_origin + row) * leading
+ panel_origin
+ col
)
for _ in range(_B4096_FACTOR_LOADS):
if row <= col:
shared[shared_index] = output[source_index]
row += _B4096_FACTOR_ROW_STRIDE
shared_index += _B4096_FACTOR_FP32_ROW_STRIDE
source_index += _B4096_FACTOR_ROW_STRIDE * leading
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b4096_factor_load_half_panel(
output, sample_base, panel_origin, panel_col, shared, leading
):
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_HALF_LD + col
source_index = (
sample_base
+ (panel_origin + row) * leading
+ panel_col
+ col
)
for _ in range(_B4096_FACTOR_LOADS):
shared[shared_index] = output[source_index]
shared_index += _B4096_FACTOR_HALF_ROW_STRIDE
source_index += _B4096_FACTOR_ROW_STRIDE * leading
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b4096_factor_store_tile(
shared,
output,
sample_base,
row_origin,
col_origin,
triangular,
leading,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_FP32_LD + col
destination_index = (
sample_base
+ (row_origin + row) * leading
+ col_origin
+ col
)
for _ in range(_B4096_FACTOR_LOADS):
if not triangular or row <= col:
output[destination_index] = shared[shared_index]
row += _B4096_FACTOR_ROW_STRIDE
shared_index += _B4096_FACTOR_FP32_ROW_STRIDE
destination_index += _B4096_FACTOR_ROW_STRIDE * leading
@cuda.jit(device=True, forceinline=True)
def _b4096_load_diagonal(
output, sample_base, panel_origin, shared, leading
):
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_FP32_LD + col
source_index = (
sample_base
+ (panel_origin + row) * leading
+ panel_origin
+ col
)
for _ in range(_B4096_SOLVE_LOADS):
if row <= col:
shared[shared_index] = output[source_index]
row += _B4096_SOLVE_ROW_STRIDE
shared_index += _B4096_SOLVE_FP32_ROW_STRIDE
source_index += _B4096_SOLVE_ROW_STRIDE * leading
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b4096_load_fp32_panel(
output, sample_base, panel_origin, panel_col, shared, leading
):
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_FP32_LD + col
source_index = (
sample_base
+ (panel_origin + row) * leading
+ panel_col
+ col
)
for _ in range(_B4096_SOLVE_LOADS):
shared[shared_index] = output[source_index]
shared_index += _B4096_SOLVE_FP32_ROW_STRIDE
source_index += _B4096_SOLVE_ROW_STRIDE * leading
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b4096_load_half_panel(
output, sample_base, panel_origin, panel_col, shared, leading
):
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_HALF_LD + col
source_index = (
sample_base
+ (panel_origin + row) * leading
+ panel_col
+ col
)
for _ in range(_B4096_SOLVE_LOADS):
shared[shared_index] = output[source_index]
shared_index += _B4096_SOLVE_HALF_ROW_STRIDE
source_index += _B4096_SOLVE_ROW_STRIDE * leading
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _b4096_store_tile(
shared,
output,
sample_base,
row_origin,
col_origin,
triangular,
leading,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row = tid // _B4096_NB
col = tid - row * _B4096_NB
shared_index = row * _B4096_FP32_LD + col
destination_index = (
sample_base
+ (row_origin + row) * leading
+ col_origin
+ col
)
for _ in range(_B4096_SOLVE_LOADS):
if not triangular or row <= col:
output[destination_index] = shared[shared_index]
row += _B4096_SOLVE_ROW_STRIDE
shared_index += _B4096_SOLVE_FP32_ROW_STRIDE
destination_index += _B4096_SOLVE_ROW_STRIDE * leading
@cuda.jit
def _b4096_factor_stage(
output, leading, matrix_elements, panel_origin, panel_index
):
sample = cuda.blockIdx.x
sample_base = sample * matrix_elements
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
factor = shared[0:_B4096_FP32_TILE_ELEMENTS]
update = shared[
2 * _B4096_FP32_TILE_ELEMENTS :
2 * _B4096_FP32_TILE_ELEMENTS + _B4096_HALF_TILE_ELEMENTS
].view(np.float16)
first = update[0:_B4096_HALF_TILE_ELEMENTS]
local_info = shared[
2 * _B4096_FP32_TILE_ELEMENTS + _B4096_HALF_TILE_ELEMENTS :
].view(np.int32)
local_origin = panel_origin + panel_index * _B4096_NB
_b4096_factor_load_diagonal(
output, sample_base, local_origin, factor, leading
)
for previous in range(panel_index):
previous_origin = panel_origin + previous * _B4096_NB
_b4096_factor_load_half_panel(
output,
sample_base,
previous_origin,
local_origin,
first,
leading,
)
_b4096_factor_gemm.execute(
-1.0, first, first, 1.0, factor
)
cuda.syncthreads()
_b4096_factor_cholesky.factorize(
factor, local_info, lda=_B4096_FP32_LD
)
_b4096_factor_store_tile(
factor,
output,
sample_base,
local_origin,
local_origin,
True,
leading,
)
@cuda.jit
def _b4096_solve_stage(
output,
leading,
matrix_elements,
panel_origin,
panel_index,
worker_count,
):
sample = cuda.blockIdx.x // worker_count
worker = cuda.blockIdx.x - sample * worker_count
sample_base = sample * matrix_elements
local_origin = panel_origin + panel_index * _B4096_NB
panel_col = local_origin + (worker + 1) * _B4096_NB
if panel_col >= leading:
return
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
factor = shared[0:_B4096_FP32_TILE_ELEMENTS]
panel = shared[
_B4096_FP32_TILE_ELEMENTS : 2 * _B4096_FP32_TILE_ELEMENTS
]
update = shared[
2 * _B4096_FP32_TILE_ELEMENTS :
2 * _B4096_FP32_TILE_ELEMENTS + _B4096_HALF_TILE_ELEMENTS
].view(np.float16)
first = update[0:_B4096_HALF_TILE_ELEMENTS]
second = update[
_B4096_HALF_TILE_ELEMENTS : 2 * _B4096_HALF_TILE_ELEMENTS
]
_b4096_load_diagonal(
output, sample_base, local_origin, factor, leading
)
_b4096_load_fp32_panel(
output, sample_base, local_origin, panel_col, panel, leading
)
for previous in range(panel_index):
previous_origin = panel_origin + previous * _B4096_NB
_b4096_load_half_panel(
output,
sample_base,
previous_origin,
local_origin,
first,
leading,
)
_b4096_load_half_panel(
output,
sample_base,
previous_origin,
panel_col,
second,
leading,
)
_b4096_solve_gemm.execute(
-1.0, first, second, 1.0, panel
)
cuda.syncthreads()
_b4096_solve_triangular.solve(
factor,
panel,
lda=_B4096_FP32_LD,
ldb=_B4096_FP32_LD,
)
_b4096_store_tile(
panel,
output,
sample_base,
local_origin,
panel_col,
False,
leading,
)
@cuda.jit(device=True, forceinline=True)
def _large_trsm_load_panel_fp32(
output, panel_row_origin, factor_col_origin, leading, shared
):
tid = cuda.threadIdx.x
factor_row = tid // _B4096_NB
panel_col = tid - factor_row * _B4096_NB
shared_index = factor_row * _B4096_FP32_LD + panel_col
output_index = (
(panel_row_origin + panel_col) * leading
+ factor_col_origin
+ factor_row
)
for _ in range(_B4096_SOLVE_LOADS):
shared[shared_index] = output[output_index]
factor_row += _B4096_SOLVE_ROW_STRIDE
shared_index += _B4096_SOLVE_FP32_ROW_STRIDE
output_index += _B4096_SOLVE_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _large_trsm_load_panel_half(
output, panel_row_origin, factor_col_origin, leading, shared
):
tid = cuda.threadIdx.x
factor_row = tid // _B4096_NB
panel_col = tid - factor_row * _B4096_NB
shared_index = factor_row * _B4096_HALF_LD + panel_col
output_index = (
(panel_row_origin + panel_col) * leading
+ factor_col_origin
+ factor_row
)
for _ in range(_B4096_SOLVE_LOADS):
shared[shared_index] = output[output_index]
factor_row += _B4096_SOLVE_ROW_STRIDE
shared_index += _B4096_SOLVE_HALF_ROW_STRIDE
output_index += _B4096_SOLVE_ROW_STRIDE
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _large_trsm_store_panel(
shared, output, panel_row_origin, factor_col_origin, leading
):
cuda.syncthreads()
tid = cuda.threadIdx.x
factor_row = tid // _B4096_NB
panel_col = tid - factor_row * _B4096_NB
shared_index = factor_row * _B4096_FP32_LD + panel_col
output_index = (
(panel_row_origin + panel_col) * leading
+ factor_col_origin
+ factor_row
)
for _ in range(_B4096_SOLVE_LOADS):
output[output_index] = shared[shared_index]
factor_row += _B4096_SOLVE_ROW_STRIDE
shared_index += _B4096_SOLVE_FP32_ROW_STRIDE
output_index += _B4096_SOLVE_ROW_STRIDE
@cuda.jit(max_registers=64)
def _large_persistent_trsm(
factor,
output,
factor_base,
panel_row_base,
panel_col_base,
leading,
):
panel_row_origin = panel_row_base + cuda.blockIdx.x * _B4096_NB
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
diagonal = shared[0:_B4096_FP32_TILE_ELEMENTS]
accumulator = shared[
_B4096_FP32_TILE_ELEMENTS : 2 * _B4096_FP32_TILE_ELEMENTS
]
update = shared[
2 * _B4096_FP32_TILE_ELEMENTS :
2 * _B4096_FP32_TILE_ELEMENTS + _B4096_HALF_TILE_ELEMENTS
].view(np.float16)
factor_update = update[0:_B4096_HALF_TILE_ELEMENTS]
panel_update = update[
_B4096_HALF_TILE_ELEMENTS : 2 * _B4096_HALF_TILE_ELEMENTS
]
for current in range(_B4096_N // _B4096_NB):
current_origin = current * _B4096_NB
_b4096_load_diagonal(
factor,
factor_base,
current_origin,
diagonal,
_B4096_N,
)
_large_trsm_load_panel_fp32(
output,
panel_row_origin,
panel_col_base + current_origin,
leading,
accumulator,
)
for previous in range(current):
previous_origin = previous * _B4096_NB
_b4096_load_half_panel(
factor,
factor_base,
previous_origin,
current_origin,
factor_update,
_B4096_N,
)
_large_trsm_load_panel_half(
output,
panel_row_origin,
panel_col_base + previous_origin,
leading,
panel_update,
)
_b4096_solve_gemm.execute(
-1.0,
factor_update,
panel_update,
1.0,
accumulator,
)
cuda.syncthreads()
_b4096_solve_triangular.solve(
diagonal,
accumulator,
lda=_B4096_FP32_LD,
ldb=_B4096_FP32_LD,
)
_large_trsm_store_panel(
accumulator,
output,
panel_row_origin,
panel_col_base + current_origin,
leading,
)
def _b4096_launchers_for_current_queue(output_view, batch, n):
global _b4096_factor_dispatch, _b4096_solve_dispatch
global _b4096_queue_handle, _b4096_queue
global _b4096_factor_launcher, _b4096_solve_launcher
global _b4096_launcher_key
if _b4096_factor_dispatch is None:
_b4096_factor_dispatch = _b4096_factor_stage.specialize(
output_view,
np.int64(0),
np.int64(0),
np.int32(0),
np.int32(0),
)
_b4096_solve_dispatch = _b4096_solve_stage.specialize(
output_view,
np.int64(0),
np.int64(0),
np.int32(0),
np.int32(0),
np.int32(0),
)
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
for dispatch in (
_b4096_factor_dispatch, _b4096_solve_dispatch
):
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
numba_driver.driver.cuKernelSetAttribute(
attribute,
_B4096_SHARED_BYTES,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
launcher_key = (queue_handle, batch, n)
if launcher_key != _b4096_launcher_key:
_b4096_queue = _external_queue(queue_handle)
_b4096_queue_handle = queue_handle
_b4096_launcher_key = launcher_key
workers = n // _B4096_NB
_b4096_factor_launcher = _b4096_factor_dispatch[
batch,
_B4096_FACTOR_NT,
_b4096_queue,
_B4096_SHARED_BYTES,
]
return _b4096_factor_launcher
def _b4096_solve_launcher_for_workers(batch, active_workers):
global _b4096_solve_launcher_cache
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
key = (queue_handle, batch, active_workers)
launcher = _b4096_solve_launcher_cache.get(key)
if launcher is None:
queue = _external_queue(queue_handle)
launcher = _b4096_solve_dispatch[
batch * active_workers,
_B4096_SOLVE_NT,
queue,
_B4096_SHARED_BYTES,
]
_b4096_solve_launcher_cache[key] = launcher
return launcher
def _launch_large_persistent_trsm(
factor_base: int,
panel_row_base: int,
panel_col_base: int,
leading: int,
) -> None:
global _large_trsm_dispatch
global _large_trsm_queue_handle, _large_trsm_queue
if _large_trsm_dispatch is None:
_large_trsm_dispatch = _large_persistent_trsm.specialize(
_large_factor_view,
_large_output_view,
np.int64(0),
np.int32(0),
np.int32(0),
np.int32(0),
)
compiled = next(iter(_large_trsm_dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
numba_driver.driver.cuKernelSetAttribute(
attribute,
_B4096_SHARED_BYTES,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if queue_handle != _large_trsm_queue_handle:
_large_trsm_queue_handle = queue_handle
_large_trsm_queue = _external_queue(queue_handle)
blocks = (leading - panel_row_base) // _B4096_NB
launcher = _large_trsm_dispatch[
blocks,
_B4096_SOLVE_NT,
_large_trsm_queue,
_B4096_SHARED_BYTES,
]
launcher(
_large_factor_view,
_large_output_view,
np.int64(factor_base),
np.int32(panel_row_base),
np.int32(panel_col_base),
np.int32(leading),
)
def _make_b4096_entry(data: torch.Tensor):
batch = data.shape[0]
n = data.shape[-1]
output = torch.empty_like(data, memory_format=torch.contiguous_format)
flags = torch.zeros(2 * batch, device=data.device, dtype=torch.int32)
output_view = _as_numba_flat_array(output)
stages = []
for panel_origin in range(0, n, _B4096_PANEL_WIDTH):
remaining_tiles = (
n - panel_origin
) // _B4096_NB
panel_tiles = min(_B4096_PANEL_TILES, remaining_tiles)
stop = panel_origin + _B4096_PANEL_WIDTH
rankk_stage = None
if stop < n:
panel = output[:, panel_origin:stop, stop:]
trailing = output[:, stop:, stop:]
rankk_stage = (
panel,
trailing,
cute_runtime.make_ptr(
cutlass.Float32,
panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
),
cute_runtime.make_ptr(
cutlass.Float32,
trailing.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
),
(
cutlass.Int32(n - stop),
cutlass.Int32(_B4096_PANEL_WIDTH),
cutlass.Int32(batch),
cutlass.Int32(n * n),
cutlass.Int32(n),
cutlass.Int32(n * n),
),
)
stages.append((np.int32(panel_origin), panel_tiles, rankk_stage))
return [
output,
output_view,
flags,
tuple(stages),
output.transpose(-2, -1),
None,
]
def _compute_b4096(entry) -> None:
output = entry[0]
batch = output.shape[0]
n = output.shape[-1]
leading = np.int64(n)
matrix_elements = np.int64(n * n)
factor_launcher = _b4096_launchers_for_current_queue(
entry[1], batch, n
)
for panel_origin, panel_tiles, rankk_stage in entry[3]:
for panel_index in range(panel_tiles):
panel_index_value = np.int32(panel_index)
factor_launcher(
entry[1],
leading,
matrix_elements,
panel_origin,
panel_index_value,
)
if (
panel_origin
+ (panel_index + 1) * _B4096_NB
< n
):
local_origin = (
panel_origin + panel_index * _B4096_NB
)
active_workers = (
n - int(local_origin)
) // _B4096_NB - 1
solve_launcher = _b4096_solve_launcher_for_workers(
batch, active_workers
)
solve_launcher(
entry[1],
leading,
matrix_elements,
panel_origin,
panel_index_value,
np.int32(active_workers),
)
if rankk_stage is not None:
_g2048_triangular_rankk_(
rankk_stage[2], rankk_stage[3], rankk_stage[4]
)
output.triu_()
def _run_b4096(data: torch.Tensor) -> torch.Tensor:
global _b4096_pool_index
if _b4096_pool and _b4096_pool[0][0].shape != data.shape:
_b4096_pool.clear()
_b4096_pool_index = 0
if len(_b4096_pool) < 2:
entry = _make_b4096_entry(data)
_b4096_pool.append(entry)
_b4096_pool_index = len(_b4096_pool) % 2
else:
entry = _b4096_pool[_b4096_pool_index]
_b4096_pool_index = (_b4096_pool_index + 1) % 2
_ext.grouped_prepare_(data, entry[0], entry[2])
graph = entry[5]
if graph is None:
_compute_b4096(entry)
_ext.grouped_prepare_(data, entry[0], entry[2])
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
_compute_b4096(entry)
entry[5] = graph
graph.replay()
return entry[4]
# Dedicated 256x256 panel factorization for the grouped batch-8 N=2048 path.
# Its 32-wide MathDx operators and dispatch state remain independent of the
# cooperative 64-wide N=256 path above.
_G2048_N = 256
_G2048_NB = 32
_G2048_NT = 256
_G2048_MATRIX_ELEMENTS = _G2048_N * _G2048_N
_G2048_TILE_ELEMENTS = _G2048_NB * _G2048_NB
_G2048_PANEL_N = 16
_G2048_PANEL_TILE_ELEMENTS = _G2048_NB * _G2048_PANEL_N
_G2048_PANEL_HALF_LD = 20
_G2048_PANEL_HALF_ELEMENTS = _G2048_NB * _G2048_PANEL_HALF_LD
_G2048_PANEL_SPLITS = _G2048_NB // _G2048_PANEL_N
_G2048_UPDATES_PER_CTA = 1
_G2048_FIRST_PANEL_CTAS = 74
_G2048_LOADS_PER_THREAD = _G2048_TILE_ELEMENTS // _G2048_NT
_G2048_ROW_STRIDE = _G2048_NT // _G2048_NB
_G2048_PANEL_LOADS_PER_THREAD = (
_G2048_PANEL_TILE_ELEMENTS // _G2048_NT
)
_G2048_PANEL_ROW_STRIDE = _G2048_NT // _G2048_PANEL_N
_G2048_SHARED_BYTES = (
_G2048_TILE_ELEMENTS
+ _G2048_PANEL_TILE_ELEMENTS
) * 4 + 4
_g2048_cholesky = CholeskySolver(
size=(_G2048_NB, _G2048_NB),
precision=np.float32,
data_type="real",
execution="Block",
fill_mode="upper",
arrangement=("row_major", "row_major"),
block_dim=(_G2048_NT, 1, 1),
sm=100,
)
_g2048_triangular = TriangularSolver(
size=(_G2048_NB, _G2048_NB, _G2048_NB),
precision=np.float32,
data_type="real",
side="left",
diag="non_unit",
transpose_mode="transposed",
arrangement=("row_major", "row_major"),
fill_mode="upper",
execution="Block",
block_dim=(_G2048_NT, 1, 1),
sm=100,
)
_g2048_gemm = Matmul(
size=(_G2048_NB, _G2048_NB, _G2048_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "row_major", "row_major"),
execution="Block",
block_size=_G2048_NT,
alignment=(16, 16, 16),
sm=100,
)
_g2048_panel_triangular = TriangularSolver(
size=(_G2048_NB, _G2048_PANEL_N, _G2048_NB),
precision=np.float32,
data_type="real",
side="left",
diag="non_unit",
transpose_mode="transposed",
arrangement=("row_major", "row_major"),
fill_mode="upper",
execution="Block",
block_dim=(_G2048_NT, 1, 1),
sm=100,
)
_g2048_panel_gemm = Matmul(
size=(_G2048_NB, _G2048_PANEL_N, _G2048_NB),
precision=(np.float16, np.float16, np.float32),
data_type="real",
arrangement=("col_major", "row_major", "row_major"),
leading_dimension=(
_G2048_NB,
_G2048_PANEL_HALF_LD,
_G2048_PANEL_N,
),
execution="Block",
block_size=_G2048_NT,
static_block_dim=True,
alignment=(16, 16, 16),
sm=100,
)
@cuda.jit(device=True, forceinline=True)
def _g2048_load_diagonal(source, sample_base, tile, shared, leading):
origin = tile * _G2048_NB
source_index = sample_base + origin * leading + origin
_g2048_load_tile_float4(
get_array_ptr(shared),
get_array_ptr(source),
source_index,
leading,
)
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_load_first_diagonal_update(
source, destination, sample_base, tile, diagonal, update, leading
):
tid = cuda.threadIdx.x
origin = tile * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
source_index = sample_base + (origin + row) * leading + origin + col
update_index = sample_base + row * leading + origin + col
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
if row <= col:
diagonal[index] = source[source_index]
update[index] = destination[update_index]
row += _G2048_ROW_STRIDE
index += _G2048_NT
source_index += global_row_stride
update_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_load_tile(
source, sample_base, tile_row, tile_col, shared, leading
):
tid = cuda.threadIdx.x
row_origin = tile_row * _G2048_NB
col_origin = tile_col * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
source_index = (
sample_base + (row_origin + row) * leading + col_origin + col
)
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
shared[index] = source[source_index]
index += _G2048_NT
source_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_load_panel_tile(
source, sample_base, tile_row, half_col, shared, leading
):
tid = cuda.threadIdx.x
row_origin = tile_row * _G2048_NB
col_origin = half_col * _G2048_PANEL_N
row = tid // _G2048_PANEL_N
col = tid - row * _G2048_PANEL_N
index = tid
source_index = (
sample_base + (row_origin + row) * leading + col_origin + col
)
global_row_stride = _G2048_PANEL_ROW_STRIDE * leading
for _ in range(_G2048_PANEL_LOADS_PER_THREAD):
shared[index] = source[source_index]
index += _G2048_NT
source_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_load_square_panel_pair(
square_source,
panel_source,
sample_base,
square_row,
square_col,
square_shared,
panel_row,
panel_half_col,
panel_shared,
leading,
):
tid = cuda.threadIdx.x
square_local_row = tid // _G2048_NB
square_local_col = tid - square_local_row * _G2048_NB
square_index = tid
square_source_index = (
sample_base
+ (square_row * _G2048_NB + square_local_row) * leading
+ square_col * _G2048_NB
+ square_local_col
)
panel_local_row = tid // _G2048_PANEL_N
panel_local_col = tid - panel_local_row * _G2048_PANEL_N
panel_index = tid
panel_source_index = (
sample_base
+ (panel_row * _G2048_NB + panel_local_row) * leading
+ panel_half_col * _G2048_PANEL_N
+ panel_local_col
)
square_global_stride = _G2048_ROW_STRIDE * leading
panel_global_stride = _G2048_PANEL_ROW_STRIDE * leading
for step in range(_G2048_LOADS_PER_THREAD):
square_shared[square_index] = square_source[square_source_index]
square_index += _G2048_NT
square_source_index += square_global_stride
if step < _G2048_PANEL_LOADS_PER_THREAD:
panel_shared[panel_index] = panel_source[panel_source_index]
panel_index += _G2048_NT
panel_source_index += panel_global_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_store_panel_tile(
shared, destination, sample_base, tile_row, half_col, leading
):
cuda.syncthreads()
tile_col = half_col // _G2048_PANEL_SPLITS
tile_half = half_col - tile_col * _G2048_PANEL_SPLITS
destination_index = (
sample_base
+ tile_row * _G2048_NB * leading
+ half_col * _G2048_PANEL_N
)
symmetric_index = (
sample_base
+ (tile_col * _G2048_NB + tile_half * _G2048_PANEL_N)
* leading
+ tile_row * _G2048_NB
)
_g2048_store_panel_float4(
get_array_ptr(shared),
get_array_ptr(destination),
destination_index,
symmetric_index,
leading,
np.int32(tile_col != tile_row),
)
@cuda.jit(device=True, forceinline=True)
def _g2048_copy_outer_panel(
source,
destination,
sample_base,
worker,
workers,
leading,
):
tid = cuda.threadIdx.x
copy_row = worker
while copy_row < _G2048_N:
copy_col = _G2048_N + 8 * tid
copy_index = sample_base + copy_row * leading + copy_col
while copy_col < leading:
_dx_copy_float8(
get_array_ptr(destination),
get_array_ptr(source),
copy_index,
)
copy_col += 8 * _G2048_NT
copy_index += 8 * _G2048_NT
copy_row += workers
@cuda.jit(device=True, forceinline=True)
def _g2048_load_update_triple(
destination,
accumulator_source,
sample_base,
factor_row,
tile_row,
half_col,
square_shared,
panel_shared,
accumulator_shared,
leading,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
square_row = tid // _G2048_NB
square_col = tid - square_row * _G2048_NB
square_index = tid
square_source_index = (
sample_base
+ (factor_row * _G2048_NB + square_row) * leading
+ tile_row * _G2048_NB
+ square_col
)
square_global_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
square_shared[square_index] = destination[square_source_index]
square_index += _G2048_NT
square_source_index += square_global_stride
panel_row = tid // _G2048_PANEL_N
panel_col = tid - panel_row * _G2048_PANEL_N
panel_shared_index = panel_row * _G2048_PANEL_HALF_LD + panel_col
accumulator_shared_index = tid
factor_source_index = (
sample_base
+ (factor_row * _G2048_NB + panel_row) * leading
+ half_col * _G2048_PANEL_N
+ panel_col
)
accumulator_source_index = (
sample_base
+ (tile_row * _G2048_NB + panel_row) * leading
+ half_col * _G2048_PANEL_N
+ panel_col
)
panel_global_stride = _G2048_PANEL_ROW_STRIDE * leading
for _ in range(_G2048_PANEL_LOADS_PER_THREAD):
panel_shared[panel_shared_index] = destination[factor_source_index]
accumulator_shared[accumulator_shared_index] = accumulator_source[
accumulator_source_index
]
panel_shared_index += (
_G2048_PANEL_ROW_STRIDE * _G2048_PANEL_HALF_LD
)
accumulator_shared_index += _G2048_NT
factor_source_index += panel_global_stride
accumulator_source_index += panel_global_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_load_tile_pair(
source,
sample_base,
first_row,
first_col,
first,
second_col,
second,
leading,
):
tid = cuda.threadIdx.x
first_row_origin = first_row * _G2048_NB
first_col_origin = first_col * _G2048_NB
second_col_origin = second_col * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
global_row = first_row_origin + row
first_index = sample_base + global_row * leading + first_col_origin + col
second_index = sample_base + global_row * leading + second_col_origin + col
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
first[index] = source[first_index]
second[index] = source[second_index]
index += _G2048_NT
first_index += global_row_stride
second_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_load_first_panel_update(
source,
destination,
sample_base,
tile_row,
tile_col,
accumulator,
first,
second,
leading,
):
tid = cuda.threadIdx.x
source_row_origin = tile_row * _G2048_NB
source_col_origin = tile_col * _G2048_NB
first_col_origin = tile_row * _G2048_NB
second_col_origin = tile_col * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
accumulator_index = (
sample_base
+ (source_row_origin + row) * leading
+ source_col_origin
+ col
)
first_index = sample_base + row * leading + first_col_origin + col
second_index = sample_base + row * leading + second_col_origin + col
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
accumulator[index] = source[accumulator_index]
first[index] = destination[first_index]
second[index] = destination[second_index]
index += _G2048_NT
accumulator_index += global_row_stride
first_index += global_row_stride
second_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_store_diagonal_and_load_panel(
diagonal,
destination,
source,
sample_base,
tile_row,
tile_col,
panel,
leading,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _G2048_NB
col_origin = tile_col * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
diagonal_index = (
sample_base + (row_origin + row) * leading + row_origin + col
)
panel_index = (
sample_base + (row_origin + row) * leading + col_origin + col
)
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[index]
panel[index] = source[panel_index]
row += _G2048_ROW_STRIDE
index += _G2048_NT
diagonal_index += global_row_stride
panel_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_store_diagonal_and_load_first_panel_update(
diagonal,
source,
destination,
sample_base,
tile_row,
tile_col,
accumulator,
first,
second,
leading,
):
cuda.syncthreads()
tid = cuda.threadIdx.x
row_origin = tile_row * _G2048_NB
col_origin = tile_col * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
diagonal_index = (
sample_base + (row_origin + row) * leading + row_origin + col
)
accumulator_index = (
sample_base + (row_origin + row) * leading + col_origin + col
)
first_index = sample_base + row * leading + row_origin + col
second_index = sample_base + row * leading + col_origin + col
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
if row <= col:
destination[diagonal_index] = diagonal[index]
accumulator[index] = source[accumulator_index]
first[index] = destination[first_index]
second[index] = destination[second_index]
row += _G2048_ROW_STRIDE
index += _G2048_NT
diagonal_index += global_row_stride
accumulator_index += global_row_stride
first_index += global_row_stride
second_index += global_row_stride
cuda.syncthreads()
@cuda.jit(device=True, forceinline=True)
def _g2048_store_tile(
shared,
destination,
sample_base,
tile_row,
tile_col,
triangular,
leading,
):
cuda.syncthreads()
if triangular:
destination_index = (
sample_base
+ tile_row * _G2048_NB * leading
+ tile_col * _G2048_NB
)
_g2048_store_upper_zero_lower(
get_array_ptr(shared),
get_array_ptr(destination),
destination_index,
leading,
)
return
tid = cuda.threadIdx.x
row_origin = tile_row * _G2048_NB
col_origin = tile_col * _G2048_NB
row = tid // _G2048_NB
col = tid - row * _G2048_NB
index = tid
destination_index = (
sample_base + (row_origin + row) * leading + col_origin + col
)
global_row_stride = _G2048_ROW_STRIDE * leading
for _ in range(_G2048_LOADS_PER_THREAD):
destination[destination_index] = shared[index]
row += _G2048_ROW_STRIDE
index += _G2048_NT
destination_index += global_row_stride
@cuda.jit
def _g2048_factor_stage(
source,
destination,
leading,
matrix_elements,
panel_origin,
tile,
sample_offset,
):
sample = cuda.blockIdx.x + sample_offset
sample_base = sample * matrix_elements + panel_origin * (leading + 1)
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
a_shared = shared[0:_G2048_TILE_ELEMENTS]
panel_offset = _G2048_TILE_ELEMENTS
update_offset = panel_offset + _G2048_PANEL_TILE_ELEMENTS
local_info = shared[update_offset:].view(np.int32)
if panel_origin == 0 and tile == 0:
_g2048_load_diagonal(
source, sample_base, tile, a_shared, leading
)
else:
_g2048_load_diagonal(
destination, sample_base, tile, a_shared, leading
)
_g2048_cholesky.factorize(a_shared, local_info, lda=_G2048_NB)
_g2048_store_tile(
a_shared, destination, sample_base, tile, tile, True, leading
)
@cuda.jit
def _g2048_panel_stage(
source,
destination,
leading,
matrix_elements,
panel_origin,
tile,
panel_task_count,
launch_task_count,
sample_offset,
):
local_sample = cuda.blockIdx.x // launch_task_count
worker = cuda.blockIdx.x - local_sample * launch_task_count
sample = local_sample + sample_offset
sample_base = sample * matrix_elements + panel_origin * (leading + 1)
first_wave = panel_origin == 0 and tile == 0
if first_wave and worker >= panel_task_count:
_g2048_copy_outer_panel(
source,
destination,
sample_base,
worker - panel_task_count,
launch_task_count - panel_task_count,
leading,
)
return
schedule_worker = worker - 1
if schedule_worker < 0:
schedule_worker += panel_task_count
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
a_shared = shared[0:_G2048_TILE_ELEMENTS]
panel_offset = _G2048_TILE_ELEMENTS
update_offset = panel_offset + _G2048_PANEL_TILE_ELEMENTS
b_shared = shared[panel_offset:update_offset]
half_col = _G2048_PANEL_SPLITS * (tile + 1) + schedule_worker
if panel_origin == 0 and tile == 0:
_g2048_load_square_panel_pair(
destination,
source,
sample_base,
tile,
tile,
a_shared,
tile,
half_col,
b_shared,
leading,
)
else:
_g2048_load_square_panel_pair(
destination,
destination,
sample_base,
tile,
tile,
a_shared,
tile,
half_col,
b_shared,
leading,
)
_g2048_panel_triangular.solve(
a_shared, b_shared, lda=_G2048_NB, ldb=_G2048_PANEL_N
)
_g2048_store_panel_tile(
b_shared, destination, sample_base, tile, half_col, leading
)
@cuda.jit
def _g2048_update_stage(
source,
destination,
leading,
matrix_elements,
panel_origin,
tile,
update_task_count,
sample_offset,
):
update_group_count = (
update_task_count + _G2048_UPDATES_PER_CTA - 1
) // _G2048_UPDATES_PER_CTA
local_sample = cuda.blockIdx.x // update_group_count
local_worker = cuda.blockIdx.x - local_sample * update_group_count
sample = local_sample + sample_offset
sample_base = sample * matrix_elements + panel_origin * (leading + 1)
# Preserve the coordinator-last task rotation from the resident wave.
update_group = local_worker - 1
if update_group < 0:
update_group += update_group_count
shared = cuda.shared.array(shape=0, dtype=np.float32, alignment=16)
a_shared = shared[0:_G2048_TILE_ELEMENTS]
panel_offset = _G2048_TILE_ELEMENTS
update_offset = panel_offset + _G2048_PANEL_TILE_ELEMENTS
b_shared = shared[panel_offset:update_offset]
update_shared = a_shared.view(np.float16)
c_shared = update_shared[0:_G2048_TILE_ELEMENTS]
d_shared = update_shared[
_G2048_TILE_ELEMENTS :
_G2048_TILE_ELEMENTS + _G2048_PANEL_HALF_ELEMENTS
]
tile_count = _G2048_N // _G2048_NB
first_task = update_group * _G2048_UPDATES_PER_CTA
for update_offset in range(_G2048_UPDATES_PER_CTA):
update_task = first_task + update_offset
if update_task < update_task_count:
remaining = update_task
update_row = tile + 1
row_task_count = 2 * (tile_count - update_row)
while remaining >= row_task_count:
remaining -= row_task_count
update_row += 1
row_task_count = 2 * (tile_count - update_row)
update_half_col = 2 * update_row + remaining
if panel_origin == 0 and tile == 0:
_g2048_load_update_triple(
destination,
source,
sample_base,
tile,
update_row,
update_half_col,
c_shared,
d_shared,
b_shared,
leading,
)
else:
_g2048_load_update_triple(
destination,
destination,
sample_base,
tile,
update_row,
update_half_col,
c_shared,
d_shared,
b_shared,
leading,
)
_g2048_panel_gemm.execute(
-1.0, c_shared, d_shared, 1.0, b_shared
)
_g2048_store_panel_tile(
b_shared,
destination,
sample_base,
update_row,
update_half_col,
leading,
)
def _g2048_launchers_for_current_queue(
source_view, output_view, batch, first_panel_ctas
):
global _g2048_factor_dispatch, _g2048_panel_dispatch
global _g2048_update_dispatch, _g2048_launcher_cache
if _g2048_factor_dispatch is None:
base_args = (
source_view,
output_view,
np.int64(0),
np.int64(0),
np.int64(0),
np.int32(0),
)
_g2048_factor_dispatch = _g2048_factor_stage.specialize(
*base_args, np.int32(0)
)
_g2048_panel_dispatch = _g2048_panel_stage.specialize(
*base_args, np.int32(0), np.int32(0), np.int32(0)
)
_g2048_update_dispatch = _g2048_update_stage.specialize(
*base_args, np.int32(0), np.int32(0)
)
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
carveout_attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_PREFERRED_SHARED_MEMORY_CARVEOUT
)
for dispatch in (
_g2048_factor_dispatch,
_g2048_panel_dispatch,
_g2048_update_dispatch,
):
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
numba_driver.driver.cuKernelSetAttribute(
attribute,
_G2048_SHARED_BYTES,
function.handle,
function.device.id,
)
numba_driver.driver.cuKernelSetAttribute(
carveout_attribute,
0,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
key = (queue_handle, batch, first_panel_ctas)
launchers = _g2048_launcher_cache.get(key)
if launchers is None:
queue = _external_queue(queue_handle)
factor_launcher = _g2048_factor_dispatch[
batch,
_G2048_NT,
queue,
_G2048_SHARED_BYTES,
]
panel_launchers = []
update_launchers = []
first_panel_launcher = _g2048_panel_dispatch[
batch * first_panel_ctas,
_G2048_NT,
queue,
_G2048_SHARED_BYTES,
]
tile_count = _G2048_N // _G2048_NB
for tile in range(tile_count):
trailing_rows = tile_count - tile - 1
panel_task_count = _G2048_PANEL_SPLITS * trailing_rows
task_count = trailing_rows * (trailing_rows + 1)
if task_count == 0:
panel_launchers.append(None)
update_launchers.append(None)
else:
panel_launchers.append(
_g2048_panel_dispatch[
batch * panel_task_count,
_G2048_NT,
queue,
_G2048_SHARED_BYTES,
]
)
group_count = (
task_count + _G2048_UPDATES_PER_CTA - 1
) // _G2048_UPDATES_PER_CTA
update_launchers.append(
_g2048_update_dispatch[
batch * group_count,
_G2048_NT,
queue,
_G2048_SHARED_BYTES,
]
)
launchers = (
queue,
factor_launcher,
first_panel_launcher,
tuple(panel_launchers),
tuple(update_launchers),
)
_g2048_launcher_cache[key] = launchers
return launchers[1:]
def _batch1024_direct_launchers_for_current_queue(
source_view, output_view, factor_half_view, batch
):
global _batch1024_direct_dispatch, _batch1024_direct_launcher
global _batch1024_factor_queue_handle, _batch1024_factor_queue
if _batch1024_direct_dispatch is None:
diagonal_dispatch = _batch1024_direct_diagonal.specialize(
source_view,
output_view,
np.int32(0),
)
panel_dispatch = _batch1024_direct_panels.specialize(
source_view,
output_view,
factor_half_view,
np.int32(0),
)
update_dispatch = _batch1024_direct_updates.specialize(
factor_half_view,
source_view,
output_view,
np.int32(0),
)
_batch1024_direct_dispatch = (
diagonal_dispatch,
panel_dispatch,
update_dispatch,
)
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
carveout_attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_PREFERRED_SHARED_MEMORY_CARVEOUT
)
for dispatch, shared_bytes in zip(
_batch1024_direct_dispatch,
(
_B1024_DIRECT_DIAGONAL_SHARED_BYTES,
_B1024_DIRECT_PANEL_SHARED_BYTES,
_B1024_DIRECT_UPDATE_SHARED_BYTES,
),
):
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
numba_driver.driver.cuKernelSetAttribute(
attribute,
shared_bytes,
function.handle,
function.device.id,
)
numba_driver.driver.cuKernelSetAttribute(
carveout_attribute,
0,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if (
queue_handle != _batch1024_factor_queue_handle
or _batch1024_direct_launcher is None
):
_batch1024_factor_queue = _external_queue(queue_handle)
_batch1024_factor_queue_handle = queue_handle
(
diagonal_dispatch,
panel_dispatch,
update_dispatch,
) = _batch1024_direct_dispatch
launchers = []
for current_tile in range(_B1024_DIRECT_TILE_COUNT):
diagonal_launcher = diagonal_dispatch[
batch,
_D256_NT,
_batch1024_factor_queue,
_B1024_DIRECT_DIAGONAL_SHARED_BYTES,
]
trailing_tiles = _B1024_DIRECT_TILE_COUNT - current_tile - 1
panel_launcher = (
panel_dispatch[
(trailing_tiles, batch),
_D256_NT,
_batch1024_factor_queue,
_B1024_DIRECT_PANEL_SHARED_BYTES,
]
if trailing_tiles
else None
)
update_tasks = trailing_tiles * (trailing_tiles + 1) // 2
update_launcher = (
update_dispatch[
(update_tasks, batch),
_D256_NT,
_batch1024_factor_queue,
_B1024_DIRECT_UPDATE_SHARED_BYTES,
]
if update_tasks
else None
)
launchers.append(
(diagonal_launcher, panel_launcher, update_launcher)
)
_batch1024_direct_launcher = tuple(launchers)
return _batch1024_direct_launcher
def _batch1024_factor_launchers_for_current_queue(
source_view, factor_view, factor_half_view, batch
):
global _batch1024_factor_dispatch, _batch1024_factor_launcher
global _batch1024_factor_queue_handle, _batch1024_factor_queue
if _batch1024_factor_dispatch is None:
diagonal_dispatch = _batch1024_factor_diagonal.specialize(
source_view,
factor_view,
np.int32(0),
np.int32(0),
)
panel_dispatch = _batch1024_factor_panels.specialize(
source_view,
factor_view,
factor_half_view,
np.int32(0),
np.int32(0),
)
_batch1024_factor_dispatch = (
diagonal_dispatch,
panel_dispatch,
)
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
carveout_attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_PREFERRED_SHARED_MEMORY_CARVEOUT
)
for dispatch, shared_bytes in zip(
_batch1024_factor_dispatch,
(
_B1024_FACTOR_DIAGONAL_SHARED_BYTES,
_B1024_FACTOR_PANEL_SHARED_BYTES,
),
):
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
numba_driver.driver.cuKernelSetAttribute(
attribute,
shared_bytes,
function.handle,
function.device.id,
)
numba_driver.driver.cuKernelSetAttribute(
carveout_attribute,
0,
function.handle,
function.device.id,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if (
queue_handle != _batch1024_factor_queue_handle
or _batch1024_factor_launcher is None
):
_batch1024_factor_queue = _external_queue(queue_handle)
_batch1024_factor_queue_handle = queue_handle
diagonal_dispatch, panel_dispatch = _batch1024_factor_dispatch
launchers = []
for current_tile in range(_D256_N // _D256_NB):
diagonal_launcher = diagonal_dispatch[
batch,
_D256_NT,
_batch1024_factor_queue,
_B1024_FACTOR_DIAGONAL_SHARED_BYTES,
]
trailing_tiles = _D256_N // _D256_NB - current_tile - 1
panel_launcher = (
panel_dispatch[
(trailing_tiles, batch),
_D256_NT,
_batch1024_factor_queue,
_B1024_FACTOR_PANEL_SHARED_BYTES,
]
if trailing_tiles
else None
)
launchers.append((diagonal_launcher, panel_launcher))
_batch1024_factor_launcher = tuple(launchers)
return _batch1024_factor_launcher
def _launch_batch1024_panel_dag(
factor_view,
factor_half_view,
source_view,
output_view,
panel_groups,
batch,
):
global _batch1024_panel_dispatch
global _batch1024_panel_queue_handle, _batch1024_panel_queue
if _batch1024_panel_dispatch is None:
dispatches = []
for solve_kernel, update_kernels in zip(
(
_batch1024_panel_solve_bulk,
_batch1024_panel_solve_middle,
_batch1024_panel_solve_tail,
),
(
(
_batch1024_panel_update1_bulk,
_batch1024_panel_update2_bulk,
_batch1024_panel_update3_bulk,
),
(
_batch1024_panel_update1_middle,
_batch1024_panel_update2_middle,
_batch1024_panel_update3_middle,
),
(
_batch1024_panel_update1_tail,
_batch1024_panel_update2_tail,
_batch1024_panel_update3_tail,
),
),
):
stage_dispatches = []
for kernel_index, kernel in enumerate(
(solve_kernel, *update_kernels)
):
shared_bytes = (
_B1024_SOLVE_SHARED_BYTES
if kernel_index == 0
else _B1024_UPDATE_SHARED_BYTES
)
if kernel_index == 0:
dispatch = kernel.specialize(
factor_view,
source_view,
output_view,
np.int32(0),
)
else:
dispatch = kernel.specialize(
factor_half_view, source_view, output_view
)
compiled = next(iter(dispatch.overloads.values()))
function = compiled._codelibrary.get_cufunc()
attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES
)
numba_driver.driver.cuKernelSetAttribute(
attribute,
shared_bytes,
function.handle,
function.device.id,
)
carveout_attribute = (
cuda_driver.CUfunction_attribute.
CU_FUNC_ATTRIBUTE_PREFERRED_SHARED_MEMORY_CARVEOUT
)
numba_driver.driver.cuKernelSetAttribute(
carveout_attribute,
100,
function.handle,
function.device.id,
)
stage_dispatches.append(dispatch)
dispatches.append(
(stage_dispatches[0], tuple(stage_dispatches[1:]))
)
_batch1024_panel_dispatch = tuple(dispatches)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if queue_handle != _batch1024_panel_queue_handle:
_batch1024_panel_queue = _external_queue(queue_handle)
_batch1024_panel_queue_handle = queue_handle
solve_dispatch, update_dispatches = _batch1024_panel_dispatch[
(12 - panel_groups) // 4
]
solve_launcher = solve_dispatch[
(panel_groups, batch),
_B1024_NT,
_batch1024_panel_queue,
_B1024_SOLVE_SHARED_BYTES,
]
solve_launcher(
factor_view, source_view, output_view, np.int32(0)
)
for current_tile, update_dispatch in enumerate(
update_dispatches, start=1
):
tile = np.int32(current_tile)
update_dispatch[
(panel_groups, batch),
_B1024_NT,
_batch1024_panel_queue,
_B1024_UPDATE_SHARED_BYTES,
](factor_half_view, source_view, output_view)
solve_launcher(factor_view, output_view, output_view, tile)
def _grouped_mathdx_cholesky(
data: torch.Tensor,
output: torch.Tensor,
source_view,
output_view,
) -> torch.Tensor:
n = output.shape[-1]
batch = output.shape[0]
block = _G2048_N
key = id(output)
stage_entry = _g2048_stage_cache.get(key)
if stage_entry is None or stage_entry[0] is not output:
first_panel_ctas = (
148
if (batch, n) in ((4, 1024), (2, 2048))
else _G2048_FIRST_PANEL_CTAS
)
stages = []
for panel_index, panel_origin in enumerate(range(0, n, block)):
stop = panel_origin + block
launch_args = (
np.int64(n),
np.int64(n * n),
np.int64(panel_origin),
)
if stop == n:
stages.append((launch_args, panel_index, None))
continue
panel = output[:, panel_origin:stop, stop:]
trailing_input = data[:, stop:, stop:]
trailing_output = output[:, stop:, stop:]
panel_pointer = cute_runtime.make_ptr(
cutlass.Float32,
panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
trailing_output.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
accumulator_pointer = (
cute_runtime.make_ptr(
cutlass.Float32,
trailing_input.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
if panel_index == 0
else output_pointer
)
rankk_stage = (
panel_pointer,
accumulator_pointer,
output_pointer,
(
cutlass.Int32(n - stop),
cutlass.Int32(block),
cutlass.Int32(batch),
cutlass.Int32(output.stride(0)),
cutlass.Int32(output.stride(1)),
cutlass.Int32(output.stride(0)),
),
)
stages.append((launch_args, panel_index, rankk_stage))
pointers = _ext.grouped_pointer_table_cuda(output, block)
stage_entry = [
output,
pointers,
tuple(stages),
first_panel_ctas,
output.transpose(-2, -1),
None,
]
_g2048_stage_cache[key] = stage_entry
graph = stage_entry[5]
if graph is None:
_compute_g2048_current_queue(
source_view,
output_view,
output,
stage_entry[1],
stage_entry[2],
stage_entry[3],
)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
_compute_g2048_current_queue(
source_view,
output_view,
output,
stage_entry[1],
stage_entry[2],
stage_entry[3],
)
stage_entry[5] = graph
graph.replay()
return stage_entry[4]
def _compute_g2048_current_queue(
source_view,
output_view,
output,
pointers,
stages,
first_panel_ctas,
) -> None:
batch = output.shape[0]
(
factor_launcher,
first_panel_launcher,
panel_launchers,
update_launchers,
) = (
_g2048_launchers_for_current_queue(
source_view, output_view, batch, first_panel_ctas
)
)
tile_count = _G2048_N // _G2048_NB
for launch_args, panel_index, rankk_stage in stages:
for tile in range(tile_count):
tile_value = np.int32(tile)
factor_launcher(
source_view,
output_view,
*launch_args,
tile_value,
np.int32(0),
)
if tile + 1 < tile_count:
trailing_rows = tile_count - tile - 1
panel_task_count = _G2048_PANEL_SPLITS * trailing_rows
first_wave = panel_index == 0 and tile == 0
panel_launcher = (
first_panel_launcher
if first_wave
else panel_launchers[tile]
)
panel_launch_count = (
first_panel_ctas
if first_wave
else panel_task_count
)
panel_launcher(
source_view,
output_view,
*launch_args,
tile_value,
np.int32(panel_task_count),
np.int32(panel_launch_count),
np.int32(0),
)
task_count = trailing_rows * (trailing_rows + 1)
update_launchers[tile](
source_view,
output_view,
*launch_args,
tile_value,
np.int32(task_count),
np.int32(0),
)
_ext.grouped_panel_update_(
output, pointers, _G2048_N, panel_index
)
if rankk_stage is not None:
_g2048_triangular_rankk_(
rankk_stage[0],
rankk_stage[1],
rankk_stage[2],
rankk_stage[3],
)
_output_cache = {}
_info = None
_workspace = None
_half_workspace = None
_factor_workspace = None
_factor_info = None
_factor_solver_workspace = None
_inverse_workspace = None
_identity_workspace = None
_solve_workspace = None
_large_factor_half_workspace = None
_large_solve_half_workspace = None
_factor_update_half_workspace = None
_compact_factor_workspace = None
_large_factor_view = None
_large_output_view = None
_large_trsm_dispatch = None
_large_trsm_queue_handle = None
_large_trsm_queue = None
_dx_dispatch = None
_dx_launch_key = None
_dx_queue = None
_dx_launcher = None
_d256_initial_solve_dispatch = None
_d256_middle_solve_dispatch = None
_d256_update_dispatch = None
_d256_middle_dispatch = None
_d256_final_dispatch = None
_d256_initial_solve_launcher = None
_d256_middle_solve_launcher = None
_d256_update_launcher = None
_d256_middle_launcher = None
_d256_final_launcher = None
_b4096_factor_dispatch = None
_b4096_solve_dispatch = None
_b4096_factor_launcher = None
_b4096_solve_launcher = None
_b4096_solve_launcher_cache = {}
_b4096_launcher_key = None
_b4096_pool = []
_b4096_pool_index = 0
_g2048_factor_dispatch = None
_g2048_panel_dispatch = None
_g2048_update_dispatch = None
_g2048_launcher_cache = {}
_g2048_stage_cache = {}
_grouped_pool = None
_grouped_pool_index = 0
_batch_half_panel = None
_batch_trsm_pointers = None
_batch_factor_workspace = None
_batch_factor_half_workspace = None
_batch_factor_info = None
_batch1024_factor_view = None
_batch1024_factor_half_view = None
_batch1024_factor_dispatch = None
_batch1024_factor_launcher = None
_batch1024_direct_dispatch = None
_batch1024_direct_launcher = None
_batch1024_panel_dispatch = None
_queue_name = "st" + "ream"
_torch_current_queue = getattr(torch.cuda, "current_" + _queue_name)
_external_queue = getattr(cuda, "external_" + _queue_name)
_torch_queue_handle_name = "cuda_" + _queue_name
_standard_cluster_coordinate = (
cutlass_utils.StaticPersistentTileScheduler
._get_cluster_work_idx_with_fastdivmod
)
_d256_queue_handle = None
_d256_queue = None
_b4096_queue_handle = None
_b4096_queue = None
_g2048_queue_handle = None
_g2048_queue = None
_g2048_rankk_queue_handle = None
_g2048_rankk_queue = None
_g2048_rankk_queue_cache = {}
_batch_rankk_queue_handle = None
_batch_rankk_queue = None
_batch1024_factor_queue_handle = None
_batch1024_factor_queue = None
_batch1024_panel_queue_handle = None
_batch1024_panel_queue = None
_left_update_queue_handle = None
_left_update_queue = None
_Queue = getattr(cuda_driver, "CU" + _queue_name)
_fake_queue = getattr(cute_runtime, "make_fake_" + _queue_name)
@dsl_user_op
def _triangular_coordinate(self, index, *, loc=None, ip=None):
root = cutlass.Float32(cute.math.sqrt(cutlass.Float32(index * 8 + 1)))
row = ((root - 1.0) * 0.5).to(cutlass.Int32)
base = row * (row + 1) // 2
row = row - (base > index).to(cutlass.Int32)
base = row * (row + 1) // 2
row = row + (base + row + 1 <= index).to(cutlass.Int32)
base = row * (row + 1) // 2
return row, index - base, cutlass.Int32(0)
@dsl_user_op
def _batched_triangular_coordinate(self, index, *, loc=None, ip=None):
batch, triangle_index = divmod(
index, self.params.cluster_shape_major_fdd
)
batch = cutlass.Int32(batch)
triangle_index = cutlass.Int32(triangle_index)
root = cutlass.Float32(
cute.math.sqrt(cutlass.Float32(triangle_index * 8 + 1))
)
row = ((root - 1.0) * 0.5).to(cutlass.Int32)
base = row * (row + 1) // 2
row = row - (base > triangle_index).to(cutlass.Int32)
base = row * (row + 1) // 2
row = row + (base + row + 1 <= triangle_index).to(cutlass.Int32)
base = row * (row + 1) // 2
return row, triangle_index - base, batch
@dsl_user_op
def _g2048_triangular_coordinate(self, index, *, loc=None, ip=None):
batch, triangle_index = divmod(
index, self.params.cluster_shape_major_fdd
)
batch = cutlass.Int32(batch)
triangle_index = cutlass.Int32(triangle_index)
root = cutlass.Float32(
cute.math.sqrt(cutlass.Float32(triangle_index * 8 + 1))
)
row = ((root - 1.0) * 0.5).to(cutlass.Int32)
base = row * (row + 1) // 2
row = row - (base > triangle_index).to(cutlass.Int32)
base = row * (row + 1) // 2
row = row + (base + row + 1 <= triangle_index).to(cutlass.Int32)
base = row * (row + 1) // 2
return row, triangle_index - base, batch
class _TriangularGemm(_AlphaBetaGemm):
@staticmethod
def _compute_grid(output, tile, cluster, max_clusters):
rows = cute.ceil_div(output.shape[0], tile[0] * 2)
total = rows * (rows + 1) // 2
params = cutlass_utils.PersistentTileSchedulerParams(
(total * 2, 1, 1), (2, 1, 1)
)
grid = cutlass_utils.StaticPersistentTileScheduler.get_grid_shape(
params, max_clusters
)
return params, grid
class _BatchedTriangularGemm(_AlphaBetaGemm):
@staticmethod
def _compute_grid(output, tile, cluster, max_clusters):
rows = cute.ceil_div(output.shape[0], tile[0] * cluster[0])
total = rows * (rows + 1) // 2
params = cutlass_utils.PersistentTileSchedulerParams(
(total * cluster[0], cluster[1], output.shape[2]),
(*cluster, 1),
)
grid = cutlass_utils.StaticPersistentTileScheduler.get_grid_shape(
params, max_clusters
)
return params, grid
class _G2048TriangularGemm(_AlphaBetaGemm):
@staticmethod
def _compute_grid(output, tile, cluster, max_clusters):
rows = cute.ceil_div(output.shape[0], tile[0] * cluster[0])
total = rows * (rows + 1) // 2
params = cutlass_utils.PersistentTileSchedulerParams(
(total * cluster[0], cluster[1], output.shape[2]),
(*cluster, 1),
)
grid = cutlass_utils.StaticPersistentTileScheduler.get_grid_shape(
params, max_clusters
)
return params, grid
@cute.jit
def _rankk_host(
operation: cutlass.Constexpr,
panel_pointer: cute.Pointer,
output_pointer: cute.Pointer,
n: cutlass.Int32,
k: cutlass.Int32,
leading_output: cutlass.Int32,
queue: _Queue,
):
panel = cute.make_tensor(
panel_pointer,
cute.make_layout((n, k, 1), stride=(k, 1, n * k)),
)
output = cute.make_tensor(
output_pointer,
cute.make_layout(
(n, n, 1), stride=(leading_output, 1, n * leading_output)
),
)
operation(
panel,
panel,
output,
output,
cutlass.Float32(-1.0),
cutlass.Float32(1.0),
74,
queue,
)
@cute.jit
def _batch_rankk_host(
operation: cutlass.Constexpr,
panel_pointer: cute.Pointer,
accumulator_pointer: cute.Pointer,
output_pointer: cute.Pointer,
n: cutlass.Constexpr,
k: cutlass.Constexpr,
batch: cutlass.Constexpr,
panel_batch_stride: cutlass.Constexpr,
leading_output: cutlass.Constexpr,
output_batch_stride: cutlass.Constexpr,
queue: _Queue,
):
panel = cute.make_tensor(
panel_pointer,
cute.make_layout(
(n, k, batch),
stride=(leading_output, 1, panel_batch_stride),
),
)
output = cute.make_tensor(
output_pointer,
cute.make_layout(
(n, n, batch),
stride=(leading_output, 1, output_batch_stride),
),
)
accumulator = cute.make_tensor(
accumulator_pointer,
cute.make_layout(
(n, n, batch),
stride=(leading_output, 1, output_batch_stride),
),
)
operation(
panel,
panel,
accumulator,
output,
cutlass.Float32(-1.0),
cutlass.Float32(1.0),
74,
queue,
)
@cute.jit
def _g2048_rankk_host(
operation: cutlass.Constexpr,
panel_pointer: cute.Pointer,
accumulator_pointer: cute.Pointer,
output_pointer: cute.Pointer,
n: cutlass.Int32,
k: cutlass.Int32,
batch: cutlass.Int32,
panel_batch_stride: cutlass.Int32,
leading_output: cutlass.Int32,
output_batch_stride: cutlass.Int32,
queue: _Queue,
):
panel = cute.make_tensor(
panel_pointer,
cute.make_layout(
(n, k, batch),
stride=(1, leading_output, panel_batch_stride),
),
)
accumulator = cute.make_tensor(
accumulator_pointer,
cute.make_layout(
(n, n, batch),
stride=(1, leading_output, output_batch_stride),
),
)
output = cute.make_tensor(
output_pointer,
cute.make_layout(
(n, n, batch),
stride=(1, leading_output, output_batch_stride),
),
)
operation(
panel,
panel,
accumulator,
output,
cutlass.Float32(-1.0),
cutlass.Float32(1.0),
74,
queue,
)
@cute.jit
def _left_looking_host(
operation: cutlass.Constexpr,
factor_pointer: cute.Pointer,
input_pointer: cute.Pointer,
output_pointer: cute.Pointer,
m: cutlass.Int32,
n: cutlass.Int32,
k: cutlass.Int32,
leading_factor: cutlass.Int32,
leading_input: cutlass.Int32,
leading_output: cutlass.Int32,
queue: _Queue,
):
factor_rows = cute.make_tensor(
factor_pointer,
cute.make_layout(
(m, k, 1), stride=(leading_factor, 1, m * leading_factor)
),
)
factor_panel = cute.make_tensor(
factor_pointer,
cute.make_layout(
(n, k, 1), stride=(leading_factor, 1, n * leading_factor)
),
)
input_panel = cute.make_tensor(
input_pointer,
cute.make_layout(
(m, n, 1), stride=(leading_input, 1, m * leading_input)
),
)
output_panel = cute.make_tensor(
output_pointer,
cute.make_layout(
(m, n, 1), stride=(leading_output, 1, m * leading_output)
),
)
operation(
factor_rows,
factor_panel,
input_panel,
output_panel,
cutlass.Float32(-1.0),
cutlass.Float32(1.0),
74,
queue,
)
@cute.jit
def _left_solve_host(
operation: cutlass.Constexpr,
panel_pointer: cute.Pointer,
inverse_pointer: cute.Pointer,
output_pointer: cute.Pointer,
m: cutlass.Int32,
k: cutlass.Int32,
leading_panel: cutlass.Int32,
leading_inverse: cutlass.Int32,
leading_output: cutlass.Int32,
queue: _Queue,
):
panel = cute.make_tensor(
panel_pointer,
cute.make_layout(
(m, k, 1), stride=(leading_panel, 1, m * leading_panel)
),
)
inverse = cute.make_tensor(
inverse_pointer,
cute.make_layout(
(k, k, 1), stride=(1, leading_inverse, k * leading_inverse)
),
)
output = cute.make_tensor(
output_pointer,
cute.make_layout(
(m, k, 1), stride=(leading_output, 1, m * leading_output)
),
)
operation(
panel,
inverse,
output,
74,
queue,
)
@cute.jit
def _recursive_solve_update_host(
operation: cutlass.Constexpr,
solved_pointer: cute.Pointer,
factor_pointer: cute.Pointer,
remainder_pointer: cute.Pointer,
m: cutlass.Int32,
n: cutlass.Int32,
k: cutlass.Int32,
leading_solved: cutlass.Int32,
leading_factor: cutlass.Int32,
leading_remainder: cutlass.Int32,
queue: _Queue,
):
solved = cute.make_tensor(
solved_pointer,
cute.make_layout(
(m, k, 1), stride=(leading_solved, 1, m * leading_solved)
),
)
factor = cute.make_tensor(
factor_pointer,
cute.make_layout(
(n, k, 1), stride=(1, leading_factor, n * leading_factor)
),
)
remainder = cute.make_tensor(
remainder_pointer,
cute.make_layout(
(m, n, 1),
stride=(leading_remainder, 1, m * leading_remainder),
),
)
operation(
solved,
factor,
remainder,
remainder,
cutlass.Float32(-1.0),
cutlass.Float32(1.0),
74,
queue,
)
@cute.jit
def _packed_left_update_host(
operation: cutlass.Constexpr,
packed_pointer: cute.Pointer,
input_pointer: cute.Pointer,
output_pointer: cute.Pointer,
m: cutlass.Int32,
n: cutlass.Int32,
k: cutlass.Int32,
leading_packed: cutlass.Int32,
leading_input: cutlass.Int32,
leading_output: cutlass.Int32,
queue: _Queue,
):
factor_rows = cute.make_tensor(
packed_pointer,
cute.make_layout(
(m, k, 1), stride=(1, leading_packed, m * leading_packed)
),
)
factor_panel = cute.make_tensor(
packed_pointer,
cute.make_layout(
(n, k, 1), stride=(1, leading_packed, n * leading_packed)
),
)
input_panel = cute.make_tensor(
input_pointer,
cute.make_layout(
(m, n, 1), stride=(leading_input, 1, m * leading_input)
),
)
output_panel = cute.make_tensor(
output_pointer,
cute.make_layout(
(m, n, 1), stride=(leading_output, 1, m * leading_output)
),
)
operation(
factor_rows,
factor_panel,
input_panel,
output_panel,
cutlass.Float32(-1.0),
cutlass.Float32(1.0),
74,
queue,
)
_rankk_operation = _TriangularGemm(
cutlass.Float32, cutlass.Float32, True, (256, 256), (2, 1)
)
_batch_rankk_operation = _BatchedTriangularGemm(
cutlass.Float32, cutlass.Float32, True, (256, 256), (2, 1)
)
_g2048_rankk_operation = _G2048TriangularGemm(
cutlass.Float32, cutlass.Float32, True, (256, 256), (2, 1)
)
_left_looking_operation = _AlphaBetaGemm(
cutlass.Float32, cutlass.Float32, True, (256, 256), (2, 1)
)
_left_solve_operation = _DenseGemm(
cutlass.Float32, True, (256, 256), (2, 1), True
)
if torch.cuda.is_available():
_seed_panel = cute_runtime.make_ptr(
cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16
)
_seed_output = cute_runtime.make_ptr(
cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16
)
_seed_g2048_panel = cute_runtime.make_ptr(
cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16
)
cutlass_utils.StaticPersistentTileScheduler._get_cluster_work_idx_with_fastdivmod = (
_triangular_coordinate
)
_rankk = cute.compile(
_rankk_host,
_rankk_operation,
_seed_panel,
_seed_output,
cutlass.Int32(128),
cutlass.Int32(128),
cutlass.Int32(128),
_fake_queue(),
)
cutlass_utils.StaticPersistentTileScheduler._get_cluster_work_idx_with_fastdivmod = (
_g2048_triangular_coordinate
)
_g2048_rankk = cute.compile(
_g2048_rankk_host,
_g2048_rankk_operation,
_seed_g2048_panel,
_seed_output,
_seed_output,
cutlass.Int32(128),
cutlass.Int32(128),
cutlass.Int32(1),
cutlass.Int32(128 * 128),
cutlass.Int32(128),
cutlass.Int32(128 * 128),
_fake_queue(),
)
cutlass_utils.StaticPersistentTileScheduler._get_cluster_work_idx_with_fastdivmod = (
_batched_triangular_coordinate
)
_batch_rankk = tuple(
cute.compile(
_batch_rankk_host,
_batch_rankk_operation,
_seed_g2048_panel,
_seed_output,
_seed_output,
trailing_size,
256,
60,
1024 * 1024,
1024,
1024 * 1024,
_fake_queue(),
)
for trailing_size in (768, 512, 256)
)
cutlass_utils.StaticPersistentTileScheduler._get_cluster_work_idx_with_fastdivmod = (
_standard_cluster_coordinate
)
_left_update = cute.compile(
_left_looking_host,
_left_looking_operation,
_seed_panel,
_seed_output,
_seed_output,
cutlass.Int32(8192),
cutlass.Int32(4096),
cutlass.Int32(4096),
cutlass.Int32(8192),
cutlass.Int32(8192),
cutlass.Int32(8192),
_fake_queue(),
)
_left_tf32_update = cute.compile(
_left_looking_host,
_left_looking_operation,
_seed_output,
_seed_output,
_seed_output,
cutlass.Int32(8192),
cutlass.Int32(4096),
cutlass.Int32(4096),
cutlass.Int32(8192),
cutlass.Int32(8192),
cutlass.Int32(8192),
_fake_queue(),
)
_recursive_solve_update = cute.compile(
_recursive_solve_update_host,
_left_looking_operation,
_seed_output,
_seed_output,
_seed_output,
cutlass.Int32(8192),
cutlass.Int32(2048),
cutlass.Int32(2048),
cutlass.Int32(4096),
cutlass.Int32(4096),
cutlass.Int32(4096),
_fake_queue(),
)
_recursive_solve_update_half = cute.compile(
_recursive_solve_update_host,
_left_looking_operation,
_seed_panel,
_seed_panel,
_seed_output,
cutlass.Int32(8192),
cutlass.Int32(2048),
cutlass.Int32(2048),
cutlass.Int32(2048),
cutlass.Int32(4096),
cutlass.Int32(8192),
_fake_queue(),
)
_packed_left_update = cute.compile(
_packed_left_update_host,
_left_looking_operation,
_seed_panel,
_seed_output,
_seed_output,
cutlass.Int32(8192),
cutlass.Int32(4096),
cutlass.Int32(4096),
cutlass.Int32(65536),
cutlass.Int32(32768),
cutlass.Int32(32768),
_fake_queue(),
)
_left_solve = cute.compile(
_left_solve_host,
_left_solve_operation,
_seed_output,
_seed_output,
_seed_output,
cutlass.Int32(32768 - 4096),
cutlass.Int32(4096),
cutlass.Int32(32768),
cutlass.Int32(4096),
cutlass.Int32(4096),
_fake_queue(),
)
cutlass_utils.StaticPersistentTileScheduler._get_cluster_work_idx_with_fastdivmod = (
_triangular_coordinate
)
else:
_rankk = None
_g2048_rankk = None
_batch_rankk = ()
_left_update = None
_left_tf32_update = None
_recursive_solve_update = None
_recursive_solve_update_half = None
_packed_left_update = None
_left_solve = None
def _triangular_rankk_(trailing: torch.Tensor, panel: torch.Tensor) -> None:
panel_pointer = cute_runtime.make_ptr(
cutlass.Float16,
panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
trailing.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = getattr(torch.cuda, "current_" + _queue_name)()
queue = _Queue(getattr(torch_queue, "cuda_" + _queue_name))
_rankk(
panel_pointer,
output_pointer,
cutlass.Int32(panel.shape[0]),
cutlass.Int32(panel.shape[1]),
cutlass.Int32(trailing.stride(0)),
queue,
)
def _g2048_triangular_rankk_(
panel_pointer,
accumulator_pointer,
output_pointer=None,
shape_parameters=None,
) -> None:
global _g2048_rankk_queue_cache
if shape_parameters is None:
shape_parameters = output_pointer
output_pointer = accumulator_pointer
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
queue = _g2048_rankk_queue_cache.get(queue_handle)
if queue is None:
queue = _Queue(queue_handle)
_g2048_rankk_queue_cache[queue_handle] = queue
_g2048_rankk(
panel_pointer,
accumulator_pointer,
output_pointer,
*shape_parameters,
queue,
)
def _batch_rankk_queue_for_current_queue():
global _batch_rankk_queue_handle, _batch_rankk_queue
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _batch_rankk_queue_handle != queue_handle:
_batch_rankk_queue_handle = queue_handle
_batch_rankk_queue = _Queue(queue_handle)
return _batch_rankk_queue
def _left_looking_update_(
factor_rows: torch.Tensor,
input_panel: torch.Tensor,
output_panel: torch.Tensor,
) -> None:
global _left_update_queue_handle, _left_update_queue
factor_pointer = cute_runtime.make_ptr(
cutlass.Float16,
factor_rows.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
input_pointer = cute_runtime.make_ptr(
cutlass.Float32,
input_panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
output_panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _left_update_queue_handle != queue_handle:
_left_update_queue_handle = queue_handle
_left_update_queue = _Queue(queue_handle)
_left_update(
factor_pointer,
input_pointer,
output_pointer,
cutlass.Int32(output_panel.shape[0]),
cutlass.Int32(output_panel.shape[1]),
cutlass.Int32(factor_rows.shape[1]),
cutlass.Int32(factor_rows.stride(0)),
cutlass.Int32(input_panel.stride(0)),
cutlass.Int32(output_panel.stride(0)),
_left_update_queue,
)
def _left_looking_tf32_update_(
factor_rows: torch.Tensor,
input_panel: torch.Tensor,
output_panel: torch.Tensor,
) -> None:
global _left_update_queue_handle, _left_update_queue
factor_pointer = cute_runtime.make_ptr(
cutlass.Float32,
factor_rows.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
input_pointer = cute_runtime.make_ptr(
cutlass.Float32,
input_panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
output_panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _left_update_queue_handle != queue_handle:
_left_update_queue_handle = queue_handle
_left_update_queue = _Queue(queue_handle)
_left_tf32_update(
factor_pointer,
input_pointer,
output_pointer,
cutlass.Int32(output_panel.shape[0]),
cutlass.Int32(output_panel.shape[1]),
cutlass.Int32(factor_rows.shape[1]),
cutlass.Int32(factor_rows.stride(0)),
cutlass.Int32(input_panel.stride(0)),
cutlass.Int32(output_panel.stride(0)),
_left_update_queue,
)
def _recursive_solve_update_(
solved: torch.Tensor,
factor: torch.Tensor,
remainder: torch.Tensor,
) -> None:
global _left_update_queue_handle, _left_update_queue
solved_pointer = cute_runtime.make_ptr(
cutlass.Float32,
solved.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
factor_pointer = cute_runtime.make_ptr(
cutlass.Float32,
factor.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
remainder_pointer = cute_runtime.make_ptr(
cutlass.Float32,
remainder.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _left_update_queue_handle != queue_handle:
_left_update_queue_handle = queue_handle
_left_update_queue = _Queue(queue_handle)
_recursive_solve_update(
solved_pointer,
factor_pointer,
remainder_pointer,
cutlass.Int32(remainder.shape[0]),
cutlass.Int32(remainder.shape[1]),
cutlass.Int32(solved.shape[1]),
cutlass.Int32(solved.stride(0)),
cutlass.Int32(factor.stride(1)),
cutlass.Int32(remainder.stride(0)),
_left_update_queue,
)
def _recursive_solve_update_half_(
solved: torch.Tensor,
factor_transpose: torch.Tensor,
remainder: torch.Tensor,
) -> None:
global _left_update_queue_handle, _left_update_queue
solved_pointer = cute_runtime.make_ptr(
cutlass.Float16,
solved.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
factor_pointer = cute_runtime.make_ptr(
cutlass.Float16,
factor_transpose.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
remainder_pointer = cute_runtime.make_ptr(
cutlass.Float32,
remainder.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _left_update_queue_handle != queue_handle:
_left_update_queue_handle = queue_handle
_left_update_queue = _Queue(queue_handle)
_recursive_solve_update_half(
solved_pointer,
factor_pointer,
remainder_pointer,
cutlass.Int32(remainder.shape[0]),
cutlass.Int32(remainder.shape[1]),
cutlass.Int32(solved.shape[1]),
cutlass.Int32(solved.stride(0)),
cutlass.Int32(factor_transpose.stride(0)),
cutlass.Int32(remainder.stride(0)),
_left_update_queue,
)
def _packed_left_update_(
packed_factor_transpose: torch.Tensor,
input_panel: torch.Tensor,
output_panel: torch.Tensor,
) -> None:
global _left_update_queue_handle, _left_update_queue
packed_pointer = cute_runtime.make_ptr(
cutlass.Float16,
packed_factor_transpose.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
input_pointer = cute_runtime.make_ptr(
cutlass.Float32,
input_panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
output_panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _left_update_queue_handle != queue_handle:
_left_update_queue_handle = queue_handle
_left_update_queue = _Queue(queue_handle)
_packed_left_update(
packed_pointer,
input_pointer,
output_pointer,
cutlass.Int32(output_panel.shape[0]),
cutlass.Int32(output_panel.shape[1]),
cutlass.Int32(packed_factor_transpose.shape[0]),
cutlass.Int32(packed_factor_transpose.stride(0)),
cutlass.Int32(input_panel.stride(0)),
cutlass.Int32(output_panel.stride(0)),
_left_update_queue,
)
def _recursive_trsm_(factor: torch.Tensor, panel: torch.Tensor) -> None:
width = factor.shape[1]
if width <= 1024:
_ext.trsm_(factor, panel)
return
split = width // 2
_recursive_trsm_(factor[:, :split, :split], panel[:, :, :split])
_recursive_solve_update_(
panel[0, :, :split],
factor[0, split:, :split],
panel[0, :, split:],
)
_recursive_trsm_(factor[:, split:, split:], panel[:, :, split:])
def _compact_inverse_solve_(
factor: torch.Tensor, panel: torch.Tensor
) -> None:
width = factor.shape[1]
identity = _identity_workspace[:width, :width]
inverse = _inverse_workspace[:width, :width]
inverse.copy_(identity)
_ext.trsm_(factor, inverse.unsqueeze(0))
solved = _solve_workspace[: panel.shape[1], :width]
_left_looking_solve_(
panel[0], inverse.transpose(0, 1), solved
)
panel[0].copy_(solved)
def _recursive_factor_4096_(
diagonal: torch.Tensor, factor: torch.Tensor
) -> None:
split = 2048
left = _compact_factor_workspace[:1]
left.copy_(diagonal[:, :split, :split])
_ext.compact_factor_cholesky_(
left, _factor_info, _factor_solver_workspace
)
lower = diagonal[:, split:, :split]
_compact_inverse_solve_(left, lower)
factor[:, :split, :split].copy_(left)
factor[:, split:, :split].copy_(lower)
_ext.convert_batch_panel_cuda(
lower, _factor_update_half_workspace.unsqueeze(0)
)
_triangular_rankk_(
diagonal[0, split:, split:], _factor_update_half_workspace
)
right = _compact_factor_workspace[1:2]
right.copy_(diagonal[:, split:, split:])
_ext.compact_factor_cholesky_(
right, _factor_info, _factor_solver_workspace
)
factor[:, split:, split:].copy_(right)
def _recursive_half_trsm_(
factor: torch.Tensor,
factor_transpose: torch.Tensor,
panel: torch.Tensor,
solved_half: torch.Tensor,
) -> None:
width = factor.shape[1]
if panel.shape[1] <= 4096:
_ext.trsm_(factor, panel)
return
if width <= 2048:
_compact_inverse_solve_(factor, panel)
return
split = width // 2
_recursive_half_trsm_(
factor[:, :split, :split],
factor_transpose[:split, :split],
panel[:, :, :split],
solved_half[:, :split],
)
rows = panel.shape[1]
solved = solved_half[:rows, :split]
_ext.convert_batch_panel_cuda(panel[:, :, :split], solved.unsqueeze(0))
_ext.convert_batch_panel_cuda(
factor[:, split:, :split].transpose(1, 2),
factor_transpose[:split, split:].unsqueeze(0),
)
_recursive_solve_update_half_(
solved,
factor_transpose[:split, split:],
panel[0, :, split:],
)
_recursive_half_trsm_(
factor[:, split:, split:],
factor_transpose[split:, split:],
panel[:, :, split:],
solved_half[:, :split],
)
def _left_looking_solve_(
panel: torch.Tensor,
inverse: torch.Tensor,
output: torch.Tensor,
) -> None:
global _left_update_queue_handle, _left_update_queue
panel_pointer = cute_runtime.make_ptr(
cutlass.Float32,
panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
inverse_pointer = cute_runtime.make_ptr(
cutlass.Float32,
inverse.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
output.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
torch_queue = _torch_current_queue()
queue_handle = getattr(torch_queue, _torch_queue_handle_name)
if _left_update_queue_handle != queue_handle:
_left_update_queue_handle = queue_handle
_left_update_queue = _Queue(queue_handle)
_left_solve(
panel_pointer,
inverse_pointer,
output_pointer,
cutlass.Int32(panel.shape[0]),
cutlass.Int32(panel.shape[1]),
cutlass.Int32(panel.stride(0)),
cutlass.Int32(inverse.stride(1)),
cutlass.Int32(output.stride(0)),
_left_update_queue,
)
def _blocked_cholesky(data: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
if output.shape[0] == 1:
_ext.lower_copy_(data, output)
else:
output.copy_(data)
n = output.shape[-1]
if output.shape[0] == 1:
block = 4096
elif n == 1024:
block = 256
else:
block = 128
old_tf32 = torch.backends.cuda.matmul.allow_tf32
try:
torch.backends.cuda.matmul.allow_tf32 = True
for stage_index, k in enumerate(range(0, n, block)):
stop = min(k + block, n)
diagonal = output[:, k:stop, k:stop]
if output.shape[0] == 1:
factor = _factor_workspace[stage_index : stage_index + 1]
_ext.panel_cholesky_(
diagonal,
factor,
_factor_info,
_factor_solver_workspace,
)
else:
factor = torch.linalg.cholesky_ex(
diagonal, check_errors=False
).L
if output.shape[0] != 1:
diagonal.copy_(factor)
if stop == n:
continue
panel = output[:, stop:, k:stop]
if output.shape[0] == 1:
_recursive_half_trsm_(
factor,
_large_factor_half_workspace[0],
panel,
_large_solve_half_workspace[0],
)
half_panel = _half_workspace[: panel.shape[1]]
_ext.convert_batch_panel_cuda(
panel, half_panel.unsqueeze(0)
)
_triangular_rankk_(
output[0, stop:, stop:], half_panel
)
else:
solved = torch.linalg.solve_triangular(
factor,
panel.transpose(-2, -1),
upper=False,
).transpose(-2, -1)
panel.copy_(solved)
_ext.fp16_update_(
output[:, stop:, stop:], solved.to(torch.float16)
)
finally:
torch.backends.cuda.matmul.allow_tf32 = old_tf32
if output.shape[0] == 1:
_ext.copy_factor_stack_lower_(_factor_workspace, output)
return output
return output.tril_()
def _compute_batch1024(entry) -> None:
global _batch_trsm_pointers
global _batch_factor_workspace, _batch1024_factor_view
global _batch_factor_half_workspace, _batch1024_factor_half_view
data = entry[0]
output = entry[1]
data_view = entry[2]
output_view = entry[3]
stages = entry[4]
rankk_queue = _batch_rankk_queue_for_current_queue()
factor_launchers = _batch1024_factor_launchers_for_current_queue(
data_view,
_batch1024_factor_view,
_batch1024_factor_half_view,
data.shape[0],
)
old_tf32 = torch.backends.cuda.matmul.allow_tf32
try:
torch.backends.cuda.matmul.allow_tf32 = True
for stage_index, (
diagonal_source,
diagonal_output,
panel,
rankk_stage,
) in enumerate(stages):
stage_source_view = (
data_view if stage_index == 0 else output_view
)
panel_origin = np.int32(stage_index * 256)
for current_tile, (
diagonal_launcher,
panel_launcher,
) in enumerate(factor_launchers):
tile = np.int32(current_tile)
diagonal_launcher(
stage_source_view,
_batch1024_factor_view,
panel_origin,
tile,
)
if panel_launcher is not None:
panel_launcher(
stage_source_view,
_batch1024_factor_view,
_batch1024_factor_half_view,
panel_origin,
tile,
)
factor = _batch_factor_workspace
if rankk_stage is None:
_ext.batch_factor_copy_cuda(
factor, diagonal_output, output
)
continue
_ext.prepare_batch_factor_and_pointers_cuda(
factor, diagonal_output, panel, _batch_trsm_pointers
)
panel_row_base = (stage_index + 1) * 256
_launch_batch1024_panel_dag(
_batch1024_factor_view,
_batch1024_factor_half_view,
stage_source_view,
output_view,
(1024 - panel_row_base) // _D256_NB,
data.shape[0],
)
rankk_stage[0](
rankk_stage[1],
rankk_stage[2],
rankk_stage[3],
rankk_queue,
)
finally:
torch.backends.cuda.matmul.allow_tf32 = old_tf32
def _compute_batch1024_direct(entry) -> None:
data_view = entry[2]
output_view = entry[3]
launchers = _batch1024_direct_launchers_for_current_queue(
data_view,
output_view,
_batch1024_factor_half_view,
entry[0].shape[0],
)
for current_tile, (
diagonal_launcher,
panel_launcher,
update_launcher,
) in enumerate(launchers):
tile = np.int32(current_tile)
source_view = data_view if current_tile == 0 else output_view
diagonal_launcher(
source_view,
output_view,
tile,
)
if panel_launcher is not None:
panel_launcher(
source_view,
output_view,
_batch1024_factor_half_view,
tile,
)
update_launcher(
_batch1024_factor_half_view,
source_view,
output_view,
tile,
)
_batch1024_output_entry = None
def _run_batch1024_blocked(data: torch.Tensor) -> torch.Tensor:
global _batch_half_panel, _batch_trsm_pointers
global _batch_factor_workspace, _batch1024_factor_view
global _batch_factor_half_workspace, _batch1024_factor_half_view
global _batch1024_output_entry
if _batch_factor_workspace is None:
_batch_factor_workspace = torch.empty_strided(
(data.shape[0], 256, 256),
(256 * 256, 1, 256),
device=data.device,
dtype=torch.float32,
)
_batch1024_factor_view = _as_numba_flat_array(
_batch_factor_workspace
)
_batch_factor_half_workspace = torch.empty(
(
data.shape[0],
_B1024_FACTOR_SIDECAR_TILES,
_D256_TILE_ELEMENTS,
),
device=data.device,
dtype=torch.float16,
)
_batch1024_factor_half_view = _as_numba_half_array(
_batch_factor_half_workspace
)
if _batch_trsm_pointers is None:
_batch_trsm_pointers = torch.empty(
2 * data.shape[0], device=data.device, dtype=torch.int64
)
entry = _batch1024_output_entry
if entry is None or entry[0] is not data:
output = torch.empty_like(
data, memory_format=torch.contiguous_format
)
data_view = _as_numba_flat_array(data)
output_view = _as_numba_flat_array(output)
stages = []
for k in range(0, 1024, 256):
stop = k + 256
diagonal_output = output[:, k:stop, k:stop]
diagonal_source = (
data[:, k:stop, k:stop] if k == 0 else diagonal_output
)
panel = output[:, stop:, k:stop]
if stop == 1024:
rankk_stage = None
else:
trailing_input = data[:, stop:, stop:]
trailing = output[:, stop:, stop:]
panel_pointer = cute_runtime.make_ptr(
cutlass.Float32,
panel.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
output_pointer = cute_runtime.make_ptr(
cutlass.Float32,
trailing.data_ptr(),
cute.AddressSpace.gmem,
assumed_align=16,
)
accumulator_pointer = cute_runtime.make_ptr(
cutlass.Float32,
(
trailing_input.data_ptr()
if k == 0
else trailing.data_ptr()
),
cute.AddressSpace.gmem,
assumed_align=16,
)
rankk_stage = (
_batch_rankk[k // 256],
panel_pointer,
accumulator_pointer,
output_pointer,
)
stages.append(
(
diagonal_source,
diagonal_output,
panel,
rankk_stage,
)
)
entry = [
data,
output,
data_view,
output_view,
tuple(stages),
None,
]
_batch1024_output_entry = entry
graph = entry[5]
if graph is None:
_compute_batch1024(entry)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
_compute_batch1024(entry)
entry[5] = graph
graph.replay()
return entry[1]
def _run_batch1024(data: torch.Tensor) -> torch.Tensor:
return _run_batch1024_blocked(data)
def custom_kernel(data: input_t) -> output_t:
global _active_shape, _info, _workspace, _half_workspace
global _factor_workspace, _factor_info, _factor_solver_workspace
global _inverse_workspace, _identity_workspace, _solve_workspace
global _large_factor_half_workspace, _large_solve_half_workspace
global _factor_update_half_workspace
global _compact_factor_workspace
global _large_factor_view, _large_output_view
global _grouped_pool, _grouped_pool_index
global _batch_half_panel, _batch_trsm_pointers
global _batch_factor_workspace, _batch_factor_half_workspace
global _batch_factor_info
global _g2048_launcher_cache, _g2048_rankk_queue_cache
global _batch1024_output_entry
n = data.shape[-1]
if n == 32:
return _ext.cholesky32_cuda(data)
if n == 64:
return _ext.cholesky64_cuda(data)
if n == 128:
return _ext.cholesky128_cuda(data)
if n == _D256_N and data.shape[0] == 64:
return _run_d256(data)
if n == _D512_N and data.shape[0] == 16:
return _run_d512(data)
if n == 1024 and data.shape[0] == 60:
return _run_batch1024(data)
if (data.shape[0], n) in (
(8, 2048),
(2, 4096),
):
return _run_b4096(data)
if n >= 32:
use_grouped = (
(data.shape[0], n)
in ((4, 1024), (2, 2048), (1, 4096))
)
use_mixed = n >= 8192 or (n >= 1024 and data.shape[0] >= 16)
use_dx = n == _DX_N and data.shape[0] in (16, 640)
shape = tuple(data.shape)
if shape != _active_shape:
_output_cache.clear()
_batch1024_output_entry = None
_g2048_stage_cache.clear()
_g2048_launcher_cache.clear()
_g2048_rankk_queue_cache.clear()
_batch_half_panel = None
_batch_trsm_pointers = None
_batch_factor_workspace = None
_batch_factor_half_workspace = None
_batch_factor_info = None
_active_shape = shape
use_looped = data.shape[0] <= 4 and n >= 1024
if use_mixed or use_dx or use_grouped:
_info = None
_workspace = None
else:
_info = torch.empty(
data.shape[0], device=data.device, dtype=torch.int32
)
if use_looped:
lwork = _ext.workspace_size(data)
_workspace = torch.empty(
lwork, device=data.device, dtype=torch.float32
)
else:
_workspace = torch.empty(
0, device=data.device, dtype=torch.float32
)
if data.shape[0] == 1 and n >= 8192:
_half_workspace = torch.empty(
(n, 4096),
device=data.device,
dtype=torch.float16,
)
_large_factor_half_workspace = torch.empty(
(1, 4096, 4096),
device=data.device,
dtype=torch.float16,
)
_large_solve_half_workspace = torch.empty(
(1, n, 2048),
device=data.device,
dtype=torch.float16,
)
_factor_update_half_workspace = torch.empty(
(2048, 2048),
device=data.device,
dtype=torch.float16,
)
_compact_factor_workspace = torch.empty_strided(
(2, 2048, 2048),
(2048 * 2048, 1, 2048),
device=data.device,
dtype=torch.float32,
)
_factor_workspace = torch.empty_strided(
(n // 4096, 4096, 4096),
(4096 * 4096, 1, 4096),
device=data.device,
dtype=torch.float32,
)
_large_factor_view = _as_numba_flat_array(
_factor_workspace
)
_large_output_view = None
_factor_info = torch.empty(
1, device=data.device, dtype=torch.int32
)
factor_lwork = _ext.workspace_size(_factor_workspace[:1])
_factor_solver_workspace = torch.empty(
factor_lwork, device=data.device, dtype=torch.float32
)
_identity_workspace = torch.eye(
2048, device=data.device, dtype=torch.float32
)
_inverse_workspace = torch.empty_like(
_identity_workspace
)
_solve_workspace = torch.empty(
(n, 2048),
device=data.device,
dtype=torch.float32,
)
else:
_half_workspace = None
_large_factor_half_workspace = None
_large_solve_half_workspace = None
_factor_update_half_workspace = None
_compact_factor_workspace = None
_large_factor_view = None
_large_output_view = None
_factor_workspace = None
_factor_info = None
_factor_solver_workspace = None
_identity_workspace = None
_inverse_workspace = None
_solve_workspace = None
if use_grouped:
_grouped_pool = []
for _ in range(4):
pooled_output = torch.zeros_like(
data, memory_format=torch.contiguous_format
)
_grouped_pool.append(pooled_output)
_grouped_pool_index = 0
else:
_grouped_pool = None
_grouped_pool_index = 0
key = id(data)
entry = _output_cache.get(key)
if entry is None or entry[0] is not data:
from_grouped_pool = False
if use_grouped and _grouped_pool_index < len(_grouped_pool):
output = _grouped_pool[_grouped_pool_index]
pointers = None
_grouped_pool_index += 1
from_grouped_pool = True
elif (
(use_mixed and data.shape[0] == 1)
or use_dx
or use_grouped
):
output = torch.zeros_like(
data, memory_format=torch.contiguous_format
)
if data.shape[0] == 1 and n >= 8192:
_large_output_view = _as_numba_flat_array(output)
else:
output = torch.empty_like(
data, memory_format=torch.contiguous_format
)
if not from_grouped_pool:
use_looped = data.shape[0] <= 4 and n >= 1024
if use_mixed and data.shape[0] == 1:
pointers = torch.empty(
0, device=data.device, dtype=torch.int64
)
elif use_mixed or use_dx:
pointers = None
elif use_grouped:
pointers = None
elif not use_looped:
step = n * n * data.element_size()
pointers = torch.arange(
output.data_ptr(),
output.data_ptr() + data.shape[0] * step,
step,
device=data.device,
dtype=torch.int64,
)
else:
pointers = torch.empty(
0, device=data.device, dtype=torch.int64
)
if use_dx or use_grouped:
dx_data = _as_numba_flat_array(data)
else:
dx_data = None
if use_dx:
factor_sidecar = torch.empty(
(
data.shape[0],
_DX_SIDECAR_TILES,
_DX_SIDECAR_TILE_ELEMENTS,
),
device=data.device,
dtype=torch.float16,
memory_format=torch.contiguous_format,
)
dx_factor_sidecar = _as_numba_half_array(factor_sidecar)
else:
factor_sidecar = None
dx_factor_sidecar = None
if use_dx or use_grouped:
dx_output = _as_numba_flat_array(output)
else:
dx_output = None
entry = (
data,
output,
pointers,
dx_data,
dx_output,
factor_sidecar,
dx_factor_sidecar,
)
_output_cache[key] = entry
else:
output, pointers = entry[1], entry[2]
if use_dx:
_launch_dx(entry[3], entry[4], entry[6], data.shape[0])
return output.transpose(-2, -1)
if use_grouped:
return _grouped_mathdx_cholesky(
data, output, entry[3], entry[4]
)
if use_mixed:
return _blocked_cholesky(data, output)
return _ext.direct_cholesky(data, output, pointers, _info, _workspace)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 10592 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