submission 890168
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 6472 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-890168?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:e01a3ca70fddc1b83886b6a3264f594758eaa6570b386364864e8dcc28afc34f
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
cluster.sync();mbarrier
asm volatile("bar.sync 1, 256;" ::: "memory");mma
namespace sp_wmma = nvcuda::wmma;num-warps = 8
num_warps=8,shared-memory
__shared__ float tile[32][33];stages = 3
num_stages=3,tile-m = 16
BLOCK_M=16,tile-n = 256
BLOCK_N=256,vector-width = float4
const float4* input4 = reinterpret_cast<const float4*>(input);Kernel source
submission.py6472 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import re
import torch
import triton
import triton.language as tl
import torch.utils.cpp_extension as cpp_extension
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
_CPP_SRC = r"""
torch::Tensor direct_potrf(torch::Tensor input);
torch::Tensor direct_potrf_split4(torch::Tensor input);
torch::Tensor xpotrf_bf16x9(torch::Tensor input);
torch::Tensor xpotrf_bf16x9_2(torch::Tensor input);
torch::Tensor xpotrf_bf16x9_4(torch::Tensor input);
torch::Tensor xpotrf_bf16x9_8(torch::Tensor input);
torch::Tensor xpotrf_bf16x9_16(torch::Tensor input);
void fp8_rankk_update_lt(
torch::Tensor trailing,
torch::Tensor panel,
torch::Tensor inverse_scale,
torch::Tensor workspace);
void fp8_lower_rankk_update_lt(
torch::Tensor trailing,
torch::Tensor output,
torch::Tensor panel,
torch::Tensor inverse_scale,
torch::Tensor workspace,
int64_t strip_size,
int64_t lane_count,
int64_t first_row,
bool triangular_k);
void fp16_lower_rankk_update_lt(
torch::Tensor trailing,
torch::Tensor output,
torch::Tensor panel,
torch::Tensor workspace,
int64_t strip_size);
void tf32_lower_rankk_update_4096_lt(
torch::Tensor target,
torch::Tensor output,
torch::Tensor panel,
torch::Tensor workspace);
torch::Tensor init_lower(torch::Tensor input);
void init_lower_into(torch::Tensor input, torch::Tensor output);
void reuse_lower_into(torch::Tensor input, torch::Tensor output);
void clear_upper(torch::Tensor matrix);
void clear_upper_view(torch::Tensor matrix);
void diagonal_tail(torch::Tensor matrix);
void jacobi_tail(torch::Tensor matrix, double alpha);
void jacobi_refine(
torch::Tensor matrix,
torch::Tensor residual,
torch::Tensor target,
double beta,
int64_t first_row,
bool reciprocal_diagonal);
torch::Tensor cluster_potrf1024(torch::Tensor input);
torch::Tensor cluster_potrf512_gemm(torch::Tensor input);
torch::Tensor cluster_potrf1024_gemm(torch::Tensor input);
torch::Tensor grid_potrf2048(torch::Tensor input);
torch::Tensor grid_potrf2048_gemm(torch::Tensor input);
torch::Tensor cluster_potrf256_b64(torch::Tensor input);
torch::Tensor cluster_potrf128_b256(torch::Tensor input);
torch::Tensor shared_potrf128_b256(torch::Tensor input);
"""
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublasLt.h>
#include <cooperative_groups.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cusolverDn.h>
#include <mma.h>
#include <algorithm>
#include <vector>
static cusolverDnHandle_t solver = nullptr;
static cusolverDnHandle_t solver_batch4[5] = {};
static cusolverDnHandle_t xsolver = nullptr;
static cusolverDnParams_t xparams = nullptr;
static size_t xdevice_bytes = 0;
static size_t xhost_bytes = 0;
static torch::Tensor xworkspace;
static std::vector<unsigned char> xhost_workspace;
static cusolverDnHandle_t xsolver2[2] = {nullptr, nullptr};
static cusolverDnParams_t xparams2[2] = {nullptr, nullptr};
static size_t xdevice_bytes2[2] = {0, 0};
static size_t xhost_bytes2[2] = {0, 0};
static int xworkspace_n2[2] = {0, 0};
static torch::Tensor xworkspace2[2];
static std::vector<unsigned char> xhost_workspace2[2];
static cusolverDnHandle_t xsolver4[4] = {};
static cusolverDnParams_t xparams4[4] = {};
static size_t xdevice_bytes4[4] = {};
static size_t xhost_bytes4[4] = {};
static torch::Tensor xworkspace4[4];
static std::vector<unsigned char> xhost_workspace4[4];
static cusolverDnHandle_t xsolver8[8] = {};
static cusolverDnParams_t xparams8[8] = {};
static size_t xdevice_bytes8[8] = {};
static size_t xhost_bytes8[8] = {};
static torch::Tensor xworkspace8[8];
static std::vector<unsigned char> xhost_workspace8[8];
static cusolverDnHandle_t xsolver16[16] = {};
static cusolverDnParams_t xparams16[16] = {};
static size_t xdevice_bytes16[16] = {};
static size_t xhost_bytes16[16] = {};
static torch::Tensor xworkspace16[16];
static std::vector<unsigned char> xhost_workspace16[16];
static cublasLtHandle_t lt = nullptr;
static cublasHandle_t sg_blas = nullptr;
#define PC_CAT2_(x, y) x##y
#define PC_CAT_(x, y) PC_CAT2_(x, y)
static inline auto current_q() {
auto q = at::cuda::PC_CAT_(getCurrentCUDAStr, eam)();
return q.PC_CAT_(str, eam)();
}
static inline void check_lt(cublasStatus_t status, const char* where) {
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS,
where, " failed: ", static_cast<int>(status));
}
struct Fp8LowerPlan {
int batch = 1;
int rows = 0;
int cols = 0;
int k = 0;
int panel_ld = 0;
int ldc = 0;
int ldd = 0;
int64_t panel_stride = 0;
int64_t trailing_stride = 0;
int64_t output_stride = 0;
cublasLtMatmulDesc_t operation = nullptr;
cublasLtMatrixLayout_t a_layout = nullptr;
cublasLtMatrixLayout_t b_layout = nullptr;
cublasLtMatrixLayout_t c_layout = nullptr;
cublasLtMatrixLayout_t d_layout = nullptr;
cublasLtMatmulAlgo_t algorithm = {};
bool ready = false;
};
static std::vector<Fp8LowerPlan> lower_plans;
static constexpr int kLowerLanes = 4;
static decltype(current_q()) lower_queues[kLowerLanes] = {};
static cudaEvent_t lower_ready = nullptr;
static cudaEvent_t lower_done[kLowerLanes] = {};
static bool lower_queues_initialized = false;
struct Fp16LowerPlan {
int batch = 0;
int rows = 0;
int cols = 0;
int k = 0;
int panel_ld = 0;
int ldc = 0;
int ldd = 0;
int64_t panel_stride = 0;
int64_t trailing_stride = 0;
int64_t output_stride = 0;
size_t workspace_bytes = 0;
cublasLtMatmulDesc_t operation = nullptr;
cublasLtMatrixLayout_t a_layout = nullptr;
cublasLtMatrixLayout_t b_layout = nullptr;
cublasLtMatrixLayout_t c_layout = nullptr;
cublasLtMatrixLayout_t d_layout = nullptr;
cublasLtMatmulAlgo_t algorithm = {};
bool ready = false;
};
static std::vector<Fp16LowerPlan> half_lower_plans;
struct Tf32LowerPlan {
int batch = 1;
int k = 0;
int64_t panel_stride = 0;
int64_t target_stride = 0;
int64_t output_stride = 0;
size_t workspace_bytes = 0;
cublasLtMatmulDesc_t operation = nullptr;
cublasLtMatrixLayout_t a_layout = nullptr;
cublasLtMatrixLayout_t b_layout = nullptr;
cublasLtMatrixLayout_t c_layout = nullptr;
cublasLtMatrixLayout_t d_layout = nullptr;
cublasLtMatmulAlgo_t algorithm = {};
bool ready = false;
};
static std::vector<Tf32LowerPlan> tf32_lower_plans;
void fp8_rankk_update_lt(
torch::Tensor trailing,
torch::Tensor panel,
torch::Tensor inverse_scale,
torch::Tensor workspace) {
TORCH_CHECK(trailing.is_cuda() && trailing.scalar_type() == torch::kFloat32,
"trailing must be CUDA FP32");
TORCH_CHECK(trailing.dim() == 2 && trailing.size(0) == trailing.size(1) &&
trailing.stride(1) == 1,
"trailing must be a row-major square view");
TORCH_CHECK(panel.is_cuda() && panel.dim() == 2 && panel.is_contiguous() &&
panel.element_size() == 1 && panel.size(0) == trailing.size(0),
"panel must be contiguous E4M3 and row aligned");
TORCH_CHECK(inverse_scale.is_cuda() &&
inverse_scale.scalar_type() == torch::kFloat32 &&
inverse_scale.numel() == 1,
"inverse_scale must be one CUDA FP32 value");
TORCH_CHECK(workspace.is_cuda() && workspace.element_size() == 1 &&
workspace.is_contiguous(),
"workspace must be contiguous CUDA bytes");
c10::cuda::CUDAGuard guard(trailing.device());
if (lt == nullptr) check_lt(cublasLtCreate(<), "cublasLtCreate");
const int m = static_cast<int>(panel.size(0));
const int k = static_cast<int>(panel.size(1));
const int ldc = static_cast<int>(trailing.stride(0));
const size_t workspace_bytes = static_cast<size_t>(workspace.numel());
const float alpha = -1.0f;
const float beta = 1.0f;
const float* scale_ptr = inverse_scale.data_ptr<float>();
const cublasOperation_t transa = CUBLAS_OP_T;
const cublasOperation_t transb = CUBLAS_OP_N;
cublasLtMatmulDesc_t operation = nullptr;
cublasLtMatrixLayout_t a_layout = nullptr;
cublasLtMatrixLayout_t b_layout = nullptr;
cublasLtMatrixLayout_t c_layout = nullptr;
cublasLtMatrixLayout_t d_layout = nullptr;
cublasLtMatmulPreference_t preference = nullptr;
check_lt(cublasLtMatmulDescCreate(
&operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
"cublasLtMatmulDescCreate");
check_lt(cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_TRANSA,
&transa, sizeof(transa)),
"cublasLt TRANSA");
check_lt(cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_TRANSB,
&transb, sizeof(transb)),
"cublasLt TRANSB");
check_lt(cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_A_SCALE_POINTER,
&scale_ptr, sizeof(scale_ptr)),
"cublasLt A scale");
check_lt(cublasLtMatmulDescSetAttribute(
operation, CUBLASLT_MATMUL_DESC_B_SCALE_POINTER,
&scale_ptr, sizeof(scale_ptr)),
"cublasLt B scale");
check_lt(cublasLtMatrixLayoutCreate(
&a_layout, CUDA_R_8F_E4M3, k, m, k),
"cublasLt A layout");
check_lt(cublasLtMatrixLayoutCreate(
&b_layout, CUDA_R_8F_E4M3, k, m, k),
"cublasLt B layout");
check_lt(cublasLtMatrixLayoutCreate(
&c_layout, CUDA_R_32F, m, m, ldc),
"cublasLt C layout");
check_lt(cublasLtMatrixLayoutCreate(
&d_layout, CUDA_R_32F, m, m, ldc),
"cublasLt D layout");
check_lt(cublasLtMatmulPreferenceCreate(&preference),
"cublasLt preference create");
check_lt(cublasLtMatmulPreferenceSetAttribute(
preference, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_bytes, sizeof(workspace_bytes)),
"cublasLt workspace preference");
int returned = 0;
cublasLtMatmulHeuristicResult_t heuristic = {};
check_lt(cublasLtMatmulAlgoGetHeuristic(
lt, operation, a_layout, b_layout, c_layout, d_layout,
preference, 1, &heuristic, &returned),
"cublasLt heuristic");
TORCH_CHECK(returned > 0, "cublasLt found no FP8 algorithm");
check_lt(cublasLtMatmul(
lt, operation,
&alpha, panel.data_ptr(), a_layout,
panel.data_ptr(), b_layout,
&beta, trailing.data_ptr<float>(), c_layout,
trailing.data_ptr<float>(), d_layout,
&heuristic.algo, workspace.data_ptr(), workspace_bytes,
current_q()),
"cublasLtMatmul");
cublasLtMatmulPreferenceDestroy(preference);
cublasLtMatrixLayoutDestroy(d_layout);
cublasLtMatrixLayoutDestroy(c_layout);
cublasLtMatrixLayoutDestroy(b_layout);
cublasLtMatrixLayoutDestroy(a_layout);
cublasLtMatmulDescDestroy(operation);
}
void fp8_lower_rankk_update_lt(
torch::Tensor trailing,
torch::Tensor output,
torch::Tensor panel,
torch::Tensor inverse_scale,
torch::Tensor workspace,
int64_t strip_size,
int64_t lane_count,
int64_t first_row,
bool triangular_k) {
TORCH_CHECK(trailing.is_cuda() && trailing.scalar_type() == torch::kFloat32,
"trailing must be CUDA FP32");
const bool unbatched =
trailing.dim() == 2 && trailing.size(0) == trailing.size(1) &&
trailing.stride(1) == 1;
const bool batched =
trailing.dim() == 3 && trailing.size(1) == trailing.size(2) &&
trailing.stride(2) == 1;
TORCH_CHECK(unbatched || batched,
"trailing must be row-major square views");
TORCH_CHECK(output.is_cuda() &&
output.scalar_type() == torch::kFloat32 &&
output.sizes() == trailing.sizes() &&
output.stride(output.dim() - 1) == 1,
"output must match trailing");
const bool panel_matches = unbatched
? panel.dim() == 2 && panel.size(0) == trailing.size(0)
: panel.dim() == 3 && panel.size(0) == trailing.size(0) &&
panel.size(1) == trailing.size(1);
TORCH_CHECK(panel.is_cuda() && panel_matches && panel.is_contiguous() &&
panel.element_size() == 1,
"panel must be contiguous E4M3 and row aligned");
TORCH_CHECK(!triangular_k ||
panel.size(panel.dim() - 2) == panel.size(panel.dim() - 1),
"triangular panel must be square");
TORCH_CHECK(inverse_scale.is_cuda() &&
inverse_scale.scalar_type() == torch::kFloat32 &&
inverse_scale.numel() == 1,
"inverse_scale must be one CUDA FP32 value");
TORCH_CHECK(workspace.is_cuda() && workspace.element_size() == 1 &&
workspace.is_contiguous(),
"workspace must be contiguous CUDA bytes");
const int matrix_dim = trailing.dim() - 2;
TORCH_CHECK(strip_size > 0 && strip_size <= trailing.size(matrix_dim),
"invalid lower-update strip size");
TORCH_CHECK(first_row >= 0 && first_row <= trailing.size(matrix_dim) &&
first_row % strip_size == 0,
"lower-update first row must be strip aligned");
TORCH_CHECK(lane_count == 1 || lane_count == kLowerLanes,
"lower-update lanes must be 1 or 4");
c10::cuda::CUDAGuard guard(trailing.device());
if (lt == nullptr) check_lt(cublasLtCreate(<), "cublasLtCreate");
const int batch = unbatched ? 1 : static_cast<int>(panel.size(0));
const int m = static_cast<int>(panel.size(panel.dim() - 2));
const int panel_ld = static_cast<int>(panel.size(panel.dim() - 1));
const int ldc = static_cast<int>(trailing.stride(matrix_dim));
const int ldd = static_cast<int>(output.stride(matrix_dim));
const int64_t panel_stride = unbatched ? 0 : panel.stride(0);
const int64_t trailing_stride = unbatched ? 0 : trailing.stride(0);
const int64_t output_stride = unbatched ? 0 : output.stride(0);
const int strip = static_cast<int>(strip_size);
const int lanes = static_cast<int>(lane_count);
TORCH_CHECK(workspace.numel() % lanes == 0,
"lower-update workspace must divide across lanes");
const size_t workspace_bytes =
static_cast<size_t>(workspace.numel() / lanes);
const float alpha = -1.0f;
const float beta = 1.0f;
const float* scale_ptr = inverse_scale.data_ptr<float>();
const cublasOperation_t transa = CUBLAS_OP_T;
const cublasOperation_t transb = CUBLAS_OP_N;
const unsigned char* panel_ptr =
reinterpret_cast<const unsigned char*>(panel.data_ptr());
float* trailing_ptr = trailing.data_ptr<float>();
float* output_ptr = output.data_ptr<float>();
auto caller = current_q();
if (lanes == kLowerLanes) {
if (!lower_queues_initialized) {
TORCH_CHECK(cudaEventCreateWithFlags(
&lower_ready, cudaEventDisableTiming) ==
cudaSuccess,
"lower ready event creation failed");
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(PC_CAT_(cudaStr, eamCreateWithFlags)(
&lower_queues[lane], 1) == cudaSuccess,
"lower queue creation failed");
TORCH_CHECK(cudaEventCreateWithFlags(
&lower_done[lane], cudaEventDisableTiming) ==
cudaSuccess,
"lower done event creation failed");
}
lower_queues_initialized = true;
}
TORCH_CHECK(cudaEventRecord(lower_ready, caller) == cudaSuccess,
"lower ready event record failed");
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
lower_queues[lane], lower_ready, 0) == cudaSuccess,
"lower queue ready wait failed");
}
}
const int call_count = (m + strip - 1) / strip;
const int first_call = static_cast<int>(first_row) / strip;
const int active_calls = call_count - first_call;
const bool balanced_b2 = batch == 2 && m == 4096 && strip == 512 &&
triangular_k && first_row == 0 &&
lanes == kLowerLanes;
const bool tiled_triangular = triangular_k &&
(m == 20480 || m == 24576 || m == 26624 || m == 28672);
constexpr int kTriangularTile = 4096;
static constexpr int kB2Start[11] = {
3072, 2560, 2048, 1536, 3584, 3584, 1024, 3584, 512, 3584, 0};
static constexpr int kB2RowStart[11] = {
0, 0, 0, 0, 3072, 2048, 0, 1024, 0, 0, 0};
static constexpr int kB2Rows[11] = {
3584, 3072, 2560, 2048, 1024, 1024, 1536, 1024, 1024, 1024, 512};
static constexpr int kB2K[11] = {
3584, 3072, 2560, 2048, 4096, 3072, 1536, 2048, 1024, 1024, 512};
const int sequence_count = balanced_b2 ? 11 : active_calls;
int64_t lane_work[kLowerLanes] = {};
for (int sequence = 0; sequence < sequence_count; ++sequence) {
const int call = lanes == 1
? first_call + sequence
: call_count - 1 - sequence;
const int start = balanced_b2 ? kB2Start[sequence] : call * strip;
const int end = std::min(start + strip, m);
const int cols = end - start;
const int off_diagonal_tiles = tiled_triangular
? (start + kTriangularTile - 1) / kTriangularTile
: 0;
const int horizontal_jobs = tiled_triangular && !balanced_b2
? off_diagonal_tiles + 1
: 1;
for (int horizontal = 0; horizontal < horizontal_jobs; ++horizontal) {
int row_start = balanced_b2 ? kB2RowStart[sequence] : 0;
int row_end = balanced_b2
? row_start + kB2Rows[sequence]
: end;
if (tiled_triangular) {
if (horizontal == 0) {
row_start = start;
} else {
const int tile = off_diagonal_tiles - horizontal;
row_start = tile * kTriangularTile;
row_end = std::min(row_start + kTriangularTile, start);
}
}
const int rows = row_end - row_start;
const int k = balanced_b2
? kB2K[sequence]
: (triangular_k ? row_end : panel_ld);
const int64_t job_work = triangular_k
? static_cast<int64_t>(rows) * cols * k
: end;
int lane = 0;
for (int candidate = 1; candidate < lanes; ++candidate) {
if (lane_work[candidate] < lane_work[lane]) lane = candidate;
}
lane_work[lane] += job_work;
auto queue = lanes == 1 ? caller : lower_queues[lane];
Fp8LowerPlan* plan = nullptr;
for (auto& candidate : lower_plans) {
if (candidate.batch == batch &&
candidate.rows == rows && candidate.cols == cols &&
candidate.k == k && candidate.panel_ld == panel_ld &&
candidate.ldc == ldc &&
candidate.ldd == ldd &&
candidate.panel_stride == panel_stride &&
candidate.trailing_stride == trailing_stride &&
candidate.output_stride == output_stride) {
plan = &candidate;
break;
}
}
if (plan == nullptr) {
lower_plans.emplace_back();
plan = &lower_plans.back();
plan->batch = batch;
plan->rows = rows;
plan->cols = cols;
plan->k = k;
plan->panel_ld = panel_ld;
plan->ldc = ldc;
plan->ldd = ldd;
plan->panel_stride = panel_stride;
plan->trailing_stride = trailing_stride;
plan->output_stride = output_stride;
check_lt(cublasLtMatmulDescCreate(
&plan->operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
"cublasLt lower desc create");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation, CUBLASLT_MATMUL_DESC_TRANSA,
&transa, sizeof(transa)),
"cublasLt lower TRANSA");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation, CUBLASLT_MATMUL_DESC_TRANSB,
&transb, sizeof(transb)),
"cublasLt lower TRANSB");
check_lt(cublasLtMatrixLayoutCreate(
&plan->a_layout, CUDA_R_8F_E4M3,
k, rows, panel_ld),
"cublasLt lower A layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->b_layout, CUDA_R_8F_E4M3,
k, cols, panel_ld),
"cublasLt lower B layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->c_layout, CUDA_R_32F, rows, cols, ldc),
"cublasLt lower C layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->d_layout, CUDA_R_32F, rows, cols, ldd),
"cublasLt lower D layout");
if (batch > 1) {
for (cublasLtMatrixLayout_t layout : {
plan->a_layout, plan->b_layout,
plan->c_layout, plan->d_layout}) {
check_lt(cublasLtMatrixLayoutSetAttribute(
layout,
CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT,
&batch, sizeof(batch)),
"cublasLt lower batch count");
}
for (cublasLtMatrixLayout_t layout : {
plan->a_layout, plan->b_layout}) {
check_lt(cublasLtMatrixLayoutSetAttribute(
layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&panel_stride, sizeof(panel_stride)),
"cublasLt lower panel stride");
}
check_lt(cublasLtMatrixLayoutSetAttribute(
plan->c_layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&trailing_stride, sizeof(trailing_stride)),
"cublasLt lower target stride");
check_lt(cublasLtMatrixLayoutSetAttribute(
plan->d_layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&output_stride, sizeof(output_stride)),
"cublasLt lower output stride");
}
}
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation,
CUBLASLT_MATMUL_DESC_A_SCALE_POINTER,
&scale_ptr, sizeof(scale_ptr)),
"cublasLt lower A scale");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation,
CUBLASLT_MATMUL_DESC_B_SCALE_POINTER,
&scale_ptr, sizeof(scale_ptr)),
"cublasLt lower B scale");
if (!plan->ready) {
cublasLtMatmulPreference_t preference = nullptr;
check_lt(cublasLtMatmulPreferenceCreate(&preference),
"cublasLt lower preference create");
check_lt(cublasLtMatmulPreferenceSetAttribute(
preference,
CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_bytes, sizeof(workspace_bytes)),
"cublasLt lower workspace preference");
int returned = 0;
cublasLtMatmulHeuristicResult_t heuristic = {};
check_lt(cublasLtMatmulAlgoGetHeuristic(
lt, plan->operation, plan->a_layout,
plan->b_layout, plan->c_layout, plan->d_layout,
preference, 1, &heuristic, &returned),
"cublasLt lower heuristic");
TORCH_CHECK(returned > 0,
"cublasLt found no lower FP8 algorithm");
plan->algorithm = heuristic.algo;
plan->ready = true;
cublasLtMatmulPreferenceDestroy(preference);
}
const void* a = panel_ptr +
static_cast<size_t>(row_start) * panel_ld;
const void* b = panel_ptr +
static_cast<size_t>(start) * panel_ld;
float* c = trailing_ptr + static_cast<size_t>(start) * ldc + row_start;
float* d = output_ptr + static_cast<size_t>(start) * ldd + row_start;
unsigned char* workspace_ptr =
static_cast<unsigned char*>(workspace.data_ptr()) +
static_cast<size_t>(lane) * workspace_bytes;
check_lt(cublasLtMatmul(
lt, plan->operation,
&alpha, a, plan->a_layout,
b, plan->b_layout,
&beta, c, plan->c_layout,
d, plan->d_layout,
&plan->algorithm, workspace_ptr,
workspace_bytes, queue),
"cublasLt lower matmul");
}
}
if (lanes == kLowerLanes) {
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(cudaEventRecord(
lower_done[lane], lower_queues[lane]) ==
cudaSuccess,
"lower done event record failed");
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
caller, lower_done[lane], 0) == cudaSuccess,
"caller lower wait failed");
}
}
}
void fp16_lower_rankk_update_lt(
torch::Tensor trailing,
torch::Tensor output,
torch::Tensor panel,
torch::Tensor workspace,
int64_t strip_size) {
TORCH_CHECK(trailing.is_cuda() && trailing.scalar_type() == torch::kFloat32,
"trailing must be CUDA FP32");
TORCH_CHECK(trailing.dim() == 3 && trailing.size(1) == trailing.size(2) &&
trailing.stride(2) == 1,
"trailing must be batched row-major square views");
TORCH_CHECK(output.is_cuda() && output.scalar_type() == torch::kFloat32 &&
output.sizes() == trailing.sizes() && output.stride(2) == 1,
"output must match trailing");
TORCH_CHECK(panel.is_cuda() && panel.scalar_type() == torch::kFloat16 &&
panel.dim() == 3 && panel.stride(2) == 1 &&
panel.stride(1) == panel.size(2) &&
panel.size(0) == trailing.size(0) &&
panel.size(1) == trailing.size(1),
"panel must be row-contiguous batched FP16 and row aligned");
TORCH_CHECK(workspace.is_cuda() && workspace.element_size() == 1 &&
workspace.is_contiguous(),
"workspace must be contiguous CUDA bytes");
TORCH_CHECK(strip_size > 0 && strip_size <= trailing.size(1),
"invalid FP16 lower-update strip size");
c10::cuda::CUDAGuard guard(trailing.device());
if (lt == nullptr) check_lt(cublasLtCreate(<), "cublasLtCreate");
const int batch = static_cast<int>(panel.size(0));
const int m = static_cast<int>(panel.size(1));
const int panel_ld = static_cast<int>(panel.size(2));
const int ldc = static_cast<int>(trailing.stride(1));
const int ldd = static_cast<int>(output.stride(1));
const int strip = static_cast<int>(strip_size);
const int64_t panel_stride = panel.stride(0);
const int64_t trailing_stride = trailing.stride(0);
const int64_t output_stride = output.stride(0);
const bool tiled_tail = (batch == 1 || batch == 2) && m == 4096 &&
panel_ld == 4096 &&
strip == 1024;
const int lanes = tiled_tail ? kLowerLanes : 1;
TORCH_CHECK(workspace.numel() % lanes == 0,
"FP16 lower-update workspace must divide across queues");
const size_t workspace_bytes =
static_cast<size_t>(workspace.numel() / lanes);
const float alpha = -1.0f;
const float beta = 1.0f;
const cublasOperation_t transa = CUBLAS_OP_T;
const cublasOperation_t transb = CUBLAS_OP_N;
const at::Half* panel_ptr = panel.data_ptr<at::Half>();
float* trailing_ptr = trailing.data_ptr<float>();
float* output_ptr = output.data_ptr<float>();
auto caller = current_q();
if (tiled_tail) {
if (!lower_queues_initialized) {
TORCH_CHECK(cudaEventCreateWithFlags(
&lower_ready, cudaEventDisableTiming) ==
cudaSuccess,
"lower ready event creation failed");
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(PC_CAT_(cudaStr, eamCreateWithFlags)(
&lower_queues[lane], 1) == cudaSuccess,
"lower queue creation failed");
TORCH_CHECK(cudaEventCreateWithFlags(
&lower_done[lane], cudaEventDisableTiming) ==
cudaSuccess,
"lower done event creation failed");
}
lower_queues_initialized = true;
}
TORCH_CHECK(cudaEventRecord(lower_ready, caller) == cudaSuccess,
"lower ready event record failed");
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
lower_queues[lane], lower_ready, 0) ==
cudaSuccess,
"lower queue ready wait failed");
}
}
static constexpr int kTailRowBlock[10] = {3, 2, 2, 1, 1, 1, 0, 0, 0, 0};
static constexpr int kTailColBlock[10] = {3, 3, 2, 3, 2, 1, 3, 2, 1, 0};
const int job_count = tiled_tail ? 10 : (m + strip - 1) / strip;
int64_t lane_work[kLowerLanes] = {};
for (int job = 0; job < job_count; ++job) {
const int row_start = tiled_tail ? kTailRowBlock[job] * 1024 : 0;
const int row_end = tiled_tail
? row_start + 1024
: std::min((job + 1) * strip, m);
const int start = tiled_tail ? kTailColBlock[job] * 1024 : job * strip;
const int end = tiled_tail ? start + 1024 : row_end;
const int cols = end - start;
const int rows = row_end - row_start;
const int active_k = tiled_tail ? row_end : panel_ld;
const int64_t job_work =
static_cast<int64_t>(rows) * cols * active_k;
int lane = 0;
for (int candidate = 1; candidate < lanes; ++candidate) {
if (lane_work[candidate] < lane_work[lane]) lane = candidate;
}
lane_work[lane] += job_work;
auto queue = tiled_tail ? lower_queues[lane] : caller;
Fp16LowerPlan* plan = nullptr;
for (auto& candidate : half_lower_plans) {
if (candidate.batch == batch && candidate.rows == rows &&
candidate.cols == cols && candidate.k == active_k &&
candidate.panel_ld == panel_ld &&
candidate.ldc == ldc &&
candidate.ldd == ldd &&
candidate.panel_stride == panel_stride &&
candidate.trailing_stride == trailing_stride &&
candidate.output_stride == output_stride &&
candidate.workspace_bytes == workspace_bytes) {
plan = &candidate;
break;
}
}
if (plan == nullptr) {
half_lower_plans.emplace_back();
plan = &half_lower_plans.back();
plan->batch = batch;
plan->rows = rows;
plan->cols = cols;
plan->k = active_k;
plan->panel_ld = panel_ld;
plan->ldc = ldc;
plan->ldd = ldd;
plan->panel_stride = panel_stride;
plan->trailing_stride = trailing_stride;
plan->output_stride = output_stride;
plan->workspace_bytes = workspace_bytes;
check_lt(cublasLtMatmulDescCreate(
&plan->operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
"cublasLt half lower desc create");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation, CUBLASLT_MATMUL_DESC_TRANSA,
&transa, sizeof(transa)),
"cublasLt half lower TRANSA");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation, CUBLASLT_MATMUL_DESC_TRANSB,
&transb, sizeof(transb)),
"cublasLt half lower TRANSB");
check_lt(cublasLtMatrixLayoutCreate(
&plan->a_layout, CUDA_R_16F,
active_k, rows, panel_ld),
"cublasLt half lower A layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->b_layout, CUDA_R_16F,
active_k, cols, panel_ld),
"cublasLt half lower B layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->c_layout, CUDA_R_32F, rows, cols, ldc),
"cublasLt half lower C layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->d_layout, CUDA_R_32F, rows, cols, ldd),
"cublasLt half lower D layout");
for (cublasLtMatrixLayout_t layout : {
plan->a_layout, plan->b_layout,
plan->c_layout, plan->d_layout}) {
check_lt(cublasLtMatrixLayoutSetAttribute(
layout, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT,
&batch, sizeof(batch)),
"cublasLt half lower batch count");
}
for (cublasLtMatrixLayout_t layout : {
plan->a_layout, plan->b_layout}) {
check_lt(cublasLtMatrixLayoutSetAttribute(
layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&panel_stride, sizeof(panel_stride)),
"cublasLt half lower panel stride");
}
check_lt(cublasLtMatrixLayoutSetAttribute(
plan->c_layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&trailing_stride, sizeof(trailing_stride)),
"cublasLt half lower trailing stride");
check_lt(cublasLtMatrixLayoutSetAttribute(
plan->d_layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&output_stride, sizeof(output_stride)),
"cublasLt half lower output stride");
}
if (!plan->ready) {
cublasLtMatmulPreference_t preference = nullptr;
check_lt(cublasLtMatmulPreferenceCreate(&preference),
"cublasLt half lower preference create");
check_lt(cublasLtMatmulPreferenceSetAttribute(
preference,
CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_bytes, sizeof(workspace_bytes)),
"cublasLt half lower workspace preference");
int returned = 0;
cublasLtMatmulHeuristicResult_t heuristic = {};
check_lt(cublasLtMatmulAlgoGetHeuristic(
lt, plan->operation, plan->a_layout,
plan->b_layout, plan->c_layout, plan->d_layout,
preference, 1, &heuristic, &returned),
"cublasLt half lower heuristic");
TORCH_CHECK(returned > 0,
"cublasLt found no batched FP16 lower algorithm");
plan->algorithm = heuristic.algo;
plan->ready = true;
cublasLtMatmulPreferenceDestroy(preference);
}
const void* a = panel_ptr +
static_cast<size_t>(row_start) * panel_ld;
const void* b = panel_ptr +
static_cast<size_t>(start) * panel_ld;
float* c = trailing_ptr +
static_cast<size_t>(start) * ldc + row_start;
float* d = output_ptr +
static_cast<size_t>(start) * ldd + row_start;
unsigned char* workspace_ptr =
static_cast<unsigned char*>(workspace.data_ptr()) +
static_cast<size_t>(lane) * workspace_bytes;
check_lt(cublasLtMatmul(
lt, plan->operation,
&alpha, a, plan->a_layout,
b, plan->b_layout,
&beta, c, plan->c_layout,
d, plan->d_layout,
&plan->algorithm, workspace_ptr,
workspace_bytes, queue),
"cublasLt half lower matmul");
}
if (tiled_tail) {
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(cudaEventRecord(
lower_done[lane], lower_queues[lane]) ==
cudaSuccess,
"lower done event record failed");
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
caller, lower_done[lane], 0) == cudaSuccess,
"caller lower wait failed");
}
}
}
void tf32_lower_rankk_update_4096_lt(
torch::Tensor target,
torch::Tensor output,
torch::Tensor panel,
torch::Tensor workspace) {
const bool unbatched =
target.dim() == 2 && target.size(0) == 4096 &&
target.size(1) == 4096;
const bool batched =
target.dim() == 3 && target.size(0) == 2 &&
target.size(1) == 4096 && target.size(2) == 4096;
TORCH_CHECK(target.is_cuda() && target.scalar_type() == torch::kFloat32 &&
(unbatched || batched) && target.is_contiguous(),
"target must be contiguous 4096-square CUDA FP32");
TORCH_CHECK(output.is_cuda() && output.scalar_type() == torch::kFloat32 &&
output.sizes() == target.sizes() && output.is_contiguous(),
"output must match target");
TORCH_CHECK(panel.is_cuda() && panel.scalar_type() == torch::kFloat32 &&
panel.sizes() == target.sizes() && panel.is_contiguous(),
"panel must match target");
TORCH_CHECK(workspace.is_cuda() && workspace.element_size() == 1 &&
workspace.is_contiguous() &&
workspace.numel() % kLowerLanes == 0,
"workspace must be divisible across four queues");
c10::cuda::CUDAGuard guard(target.device());
if (lt == nullptr) check_lt(cublasLtCreate(<), "cublasLtCreate");
constexpr int kSize = 4096;
constexpr int kTile = 1024;
static constexpr int kSmallBlock[10] = {
3, 2, 2, 1, 1, 1, 0, 0, 0, 0};
static constexpr int kLargeBlock[10] = {
3, 3, 2, 3, 2, 1, 3, 2, 1, 0};
static constexpr int kJobLane[10] = {
0, 1, 2, 3, 3, 1, 2, 0, 2, 3};
const int batch = unbatched ? 1 : 2;
const int64_t panel_stride = unbatched ? 0 : panel.stride(0);
const int64_t target_stride = unbatched ? 0 : target.stride(0);
const int64_t output_stride = unbatched ? 0 : output.stride(0);
const size_t workspace_bytes =
static_cast<size_t>(workspace.numel() / kLowerLanes);
const float alpha = -1.0f;
const float beta = 1.0f;
const cublasOperation_t transa = CUBLAS_OP_T;
const cublasOperation_t transb = CUBLAS_OP_N;
const float* panel_ptr = panel.data_ptr<float>();
const float* target_ptr = target.data_ptr<float>();
float* output_ptr = output.data_ptr<float>();
auto caller = current_q();
if (!lower_queues_initialized) {
TORCH_CHECK(cudaEventCreateWithFlags(
&lower_ready, cudaEventDisableTiming) == cudaSuccess,
"lower ready event creation failed");
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(PC_CAT_(cudaStr, eamCreateWithFlags)(
&lower_queues[lane], 1) == cudaSuccess,
"lower queue creation failed");
TORCH_CHECK(cudaEventCreateWithFlags(
&lower_done[lane], cudaEventDisableTiming) ==
cudaSuccess,
"lower done event creation failed");
}
lower_queues_initialized = true;
}
TORCH_CHECK(cudaEventRecord(lower_ready, caller) == cudaSuccess,
"lower ready event record failed");
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
lower_queues[lane], lower_ready, 0) == cudaSuccess,
"lower queue ready wait failed");
}
for (int job = 0; job < 10; ++job) {
const int small_start = kSmallBlock[job] * kTile;
const int large_start = kLargeBlock[job] * kTile;
const int active_k = small_start + kTile;
const int lane = kJobLane[job];
Tf32LowerPlan* plan = nullptr;
for (auto& candidate : tf32_lower_plans) {
if (candidate.batch == batch && candidate.k == active_k &&
candidate.panel_stride == panel_stride &&
candidate.target_stride == target_stride &&
candidate.output_stride == output_stride &&
candidate.workspace_bytes == workspace_bytes) {
plan = &candidate;
break;
}
}
if (plan == nullptr) {
tf32_lower_plans.emplace_back();
plan = &tf32_lower_plans.back();
plan->batch = batch;
plan->k = active_k;
plan->panel_stride = panel_stride;
plan->target_stride = target_stride;
plan->output_stride = output_stride;
plan->workspace_bytes = workspace_bytes;
check_lt(cublasLtMatmulDescCreate(
&plan->operation,
CUBLAS_COMPUTE_32F_FAST_TF32, CUDA_R_32F),
"cublasLt TF32 lower desc create");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation, CUBLASLT_MATMUL_DESC_TRANSA,
&transa, sizeof(transa)),
"cublasLt TF32 lower TRANSA");
check_lt(cublasLtMatmulDescSetAttribute(
plan->operation, CUBLASLT_MATMUL_DESC_TRANSB,
&transb, sizeof(transb)),
"cublasLt TF32 lower TRANSB");
check_lt(cublasLtMatrixLayoutCreate(
&plan->a_layout, CUDA_R_32F,
active_k, kTile, kSize),
"cublasLt TF32 lower A layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->b_layout, CUDA_R_32F,
active_k, kTile, kSize),
"cublasLt TF32 lower B layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->c_layout, CUDA_R_32F,
kTile, kTile, kSize),
"cublasLt TF32 lower C layout");
check_lt(cublasLtMatrixLayoutCreate(
&plan->d_layout, CUDA_R_32F,
kTile, kTile, kSize),
"cublasLt TF32 lower D layout");
if (batch > 1) {
for (cublasLtMatrixLayout_t layout : {
plan->a_layout, plan->b_layout,
plan->c_layout, plan->d_layout}) {
check_lt(cublasLtMatrixLayoutSetAttribute(
layout,
CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT,
&batch, sizeof(batch)),
"cublasLt TF32 lower batch count");
}
for (cublasLtMatrixLayout_t layout : {
plan->a_layout, plan->b_layout}) {
check_lt(cublasLtMatrixLayoutSetAttribute(
layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&panel_stride, sizeof(panel_stride)),
"cublasLt TF32 lower panel stride");
}
check_lt(cublasLtMatrixLayoutSetAttribute(
plan->c_layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&target_stride, sizeof(target_stride)),
"cublasLt TF32 lower target stride");
check_lt(cublasLtMatrixLayoutSetAttribute(
plan->d_layout,
CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
&output_stride, sizeof(output_stride)),
"cublasLt TF32 lower output stride");
}
}
if (!plan->ready) {
cublasLtMatmulPreference_t preference = nullptr;
check_lt(cublasLtMatmulPreferenceCreate(&preference),
"cublasLt TF32 lower preference create");
check_lt(cublasLtMatmulPreferenceSetAttribute(
preference,
CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_bytes, sizeof(workspace_bytes)),
"cublasLt TF32 lower workspace preference");
int returned = 0;
cublasLtMatmulHeuristicResult_t heuristic = {};
check_lt(cublasLtMatmulAlgoGetHeuristic(
lt, plan->operation, plan->a_layout,
plan->b_layout, plan->c_layout, plan->d_layout,
preference, 1, &heuristic, &returned),
"cublasLt TF32 lower heuristic");
TORCH_CHECK(returned > 0,
"cublasLt found no TF32 lower algorithm");
plan->algorithm = heuristic.algo;
plan->ready = true;
cublasLtMatmulPreferenceDestroy(preference);
}
const float* a = panel_ptr +
static_cast<size_t>(small_start) * kSize;
const float* b = panel_ptr +
static_cast<size_t>(large_start) * kSize;
const float* c = target_ptr +
static_cast<size_t>(large_start) * kSize + small_start;
float* d = output_ptr +
static_cast<size_t>(large_start) * kSize + small_start;
unsigned char* workspace_ptr =
static_cast<unsigned char*>(workspace.data_ptr()) +
static_cast<size_t>(lane) * workspace_bytes;
check_lt(cublasLtMatmul(
lt, plan->operation,
&alpha, a, plan->a_layout,
b, plan->b_layout,
&beta, c, plan->c_layout,
d, plan->d_layout,
&plan->algorithm, workspace_ptr,
workspace_bytes, lower_queues[lane]),
"cublasLt TF32 lower matmul");
}
for (int lane = 0; lane < kLowerLanes; ++lane) {
TORCH_CHECK(cudaEventRecord(
lower_done[lane], lower_queues[lane]) == cudaSuccess,
"lower done event record failed");
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
caller, lower_done[lane], 0) == cudaSuccess,
"caller lower wait failed");
}
}
__global__ void clear_upper_kernel(
float* __restrict__ matrix,
int n,
size_t matrix_stride) {
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int row = blockIdx.x * 8 + warp;
if (row >= n) return;
float* matrix_row = matrix +
static_cast<size_t>(blockIdx.y) * matrix_stride +
static_cast<size_t>(row) * n;
for (int col = row + 1 + lane; col < n; col += 32) {
matrix_row[col] = 0.0f;
}
}
__global__ void clear_upper_view_kernel(
float* __restrict__ matrix,
int n,
int64_t batch_stride,
int64_t row_stride) {
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int row = blockIdx.x * 8 + warp;
if (row >= n) return;
float* matrix_row = matrix +
static_cast<int64_t>(blockIdx.y) * batch_stride +
static_cast<int64_t>(row) * row_stride;
for (int col = row + 1 + lane; col < n; col += 32) {
matrix_row[col] = 0.0f;
}
}
__global__ void clear_lower_kernel(
float* __restrict__ matrix,
int n,
size_t matrix_stride) {
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int row = blockIdx.x * 8 + warp;
if (row >= n) return;
float* matrix_row = matrix +
static_cast<size_t>(blockIdx.y) * matrix_stride +
static_cast<size_t>(row) * n;
for (int col = lane; col < row; col += 32) {
matrix_row[col] = 0.0f;
}
}
void clear_upper(torch::Tensor matrix) {
TORCH_CHECK(matrix.is_cuda() && matrix.scalar_type() == torch::kFloat32 &&
matrix.is_contiguous() && matrix.dim() == 3 &&
matrix.size(1) == matrix.size(2),
"matrix must be contiguous square CUDA FP32");
c10::cuda::CUDAGuard guard(matrix.device());
const int batch = static_cast<int>(matrix.size(0));
const int n = static_cast<int>(matrix.size(1));
const size_t stride = static_cast<size_t>(n) * n;
const dim3 grid((n + 7) / 8, batch);
clear_upper_kernel<<<grid, 256, 0, current_q()>>>(
matrix.data_ptr<float>(), n, stride);
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "clear_upper failed: ",
cudaGetErrorString(error));
}
void clear_upper_view(torch::Tensor matrix) {
TORCH_CHECK(matrix.is_cuda() && matrix.scalar_type() == torch::kFloat32 &&
matrix.dim() == 3 && matrix.size(1) == matrix.size(2) &&
matrix.stride(2) == 1,
"matrix must be a row-contiguous square CUDA FP32 view");
c10::cuda::CUDAGuard guard(matrix.device());
const int batch = static_cast<int>(matrix.size(0));
const int n = static_cast<int>(matrix.size(1));
const dim3 grid((n + 7) / 8, batch);
clear_upper_view_kernel<<<grid, 256, 0, current_q()>>>(
matrix.data_ptr<float>(), n, matrix.stride(0), matrix.stride(1));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "clear_upper_view failed: ",
cudaGetErrorString(error));
}
__global__ void diagonal_tail_kernel(
float* __restrict__ matrix,
int n,
int64_t batch_stride,
int64_t row_stride) {
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int row = blockIdx.x * 8 + warp;
if (row >= n) return;
float* matrix_row = matrix +
static_cast<int64_t>(blockIdx.y) * batch_stride +
static_cast<int64_t>(row) * row_stride;
for (int col = lane; col < n; col += 32) {
matrix_row[col] = col == row
? sqrtf(fmaxf(matrix_row[col], 1.0e-30f))
: 0.0f;
}
}
void diagonal_tail(torch::Tensor matrix) {
TORCH_CHECK(matrix.is_cuda() && matrix.scalar_type() == torch::kFloat32 &&
matrix.dim() == 3 && matrix.size(1) == matrix.size(2) &&
matrix.stride(2) == 1,
"matrix must be a row-contiguous square CUDA FP32 view");
c10::cuda::CUDAGuard guard(matrix.device());
const int batch = static_cast<int>(matrix.size(0));
const int n = static_cast<int>(matrix.size(1));
const dim3 grid((n + 7) / 8, batch);
diagonal_tail_kernel<<<grid, 256, 0, current_q()>>>(
matrix.data_ptr<float>(), n, matrix.stride(0), matrix.stride(1));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "diagonal_tail failed: ",
cudaGetErrorString(error));
}
__global__ void jacobi_tail_diagonal_kernel(
const float* __restrict__ matrix,
float* __restrict__ inverse_diagonal,
int n,
int64_t batch_stride,
int64_t row_stride) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index >= n) return;
const float diagonal = matrix[
static_cast<int64_t>(blockIdx.y) * batch_stride +
static_cast<int64_t>(index) * row_stride + index];
inverse_diagonal[static_cast<int64_t>(blockIdx.y) * n + index] =
rsqrtf(fmaxf(diagonal, 1.0e-30f));
}
__global__ void jacobi_refine_reciprocal_diagonal_kernel(
const float* __restrict__ matrix,
float* __restrict__ inverse_diagonal,
int n,
int64_t batch_stride,
int64_t row_stride) {
const int index = blockIdx.x * blockDim.x + threadIdx.x;
if (index >= n) return;
const float diagonal = matrix[
static_cast<int64_t>(blockIdx.y) * batch_stride +
static_cast<int64_t>(index) * row_stride + index];
inverse_diagonal[static_cast<int64_t>(blockIdx.y) * n + index] =
1.0f / fmaxf(diagonal, 1.0e-30f);
}
__global__ void jacobi_tail_factor_kernel(
float* __restrict__ matrix,
const float* __restrict__ inverse_diagonal,
int n,
float alpha,
int64_t batch_stride,
int64_t row_stride) {
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int row = blockIdx.x * 8 + warp;
if (row >= n) return;
float* matrix_row = matrix +
static_cast<int64_t>(blockIdx.y) * batch_stride +
static_cast<int64_t>(row) * row_stride;
const float* inverse = inverse_diagonal +
static_cast<int64_t>(blockIdx.y) * n;
const float diagonal = matrix_row[row];
float energy = 0.0f;
for (int col = lane; col < row; col += 32) {
const float value = alpha * matrix_row[col] * inverse[col];
matrix_row[col] = value;
energy = fmaf(value, value, energy);
}
#pragma unroll
for (int offset = 16; offset; offset >>= 1) {
energy += __shfl_down_sync(0xffffffffu, energy, offset);
}
if (lane == 0) {
matrix_row[row] = sqrtf(fmaxf(diagonal - energy, 1.0e-30f));
}
for (int col = row + 1 + lane; col < n; col += 32) {
matrix_row[col] = 0.0f;
}
}
void jacobi_tail(torch::Tensor matrix, double alpha) {
TORCH_CHECK(matrix.is_cuda() && matrix.scalar_type() == torch::kFloat32 &&
matrix.dim() == 3 && matrix.size(1) == matrix.size(2) &&
matrix.stride(2) == 1,
"matrix must be a row-contiguous square CUDA FP32 view");
TORCH_CHECK(alpha > 0.0 && alpha <= 1.0,
"alpha must be in (0, 1]");
c10::cuda::CUDAGuard guard(matrix.device());
const int batch = static_cast<int>(matrix.size(0));
const int n = static_cast<int>(matrix.size(1));
auto inverse_diagonal = torch::empty({batch, n}, matrix.options());
jacobi_tail_diagonal_kernel<<<dim3((n + 255) / 256, batch), 256, 0,
current_q()>>>(
matrix.data_ptr<float>(), inverse_diagonal.data_ptr<float>(), n,
matrix.stride(0), matrix.stride(1));
jacobi_tail_factor_kernel<<<dim3((n + 7) / 8, batch), 256, 0,
current_q()>>>(
matrix.data_ptr<float>(), inverse_diagonal.data_ptr<float>(), n,
static_cast<float>(alpha), matrix.stride(0), matrix.stride(1));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "jacobi_tail failed: ",
cudaGetErrorString(error));
}
__global__ void jacobi_tail_refine_kernel(
float* __restrict__ matrix,
const float* __restrict__ residual,
const float* __restrict__ target_diagonal,
const float* __restrict__ inverse_diagonal,
int n,
float beta,
int64_t matrix_batch_stride,
int64_t matrix_row_stride,
int64_t residual_batch_stride,
int64_t residual_row_stride,
int64_t target_batch_stride,
int64_t target_diagonal_stride,
int first_row) {
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int row = first_row + blockIdx.x * 8 + warp;
if (row >= n) return;
float* matrix_row = matrix +
static_cast<int64_t>(blockIdx.y) * matrix_batch_stride +
static_cast<int64_t>(row) * matrix_row_stride;
const float* residual_row = residual +
static_cast<int64_t>(blockIdx.y) * residual_batch_stride +
static_cast<int64_t>(row) * residual_row_stride;
const float target_value = target_diagonal[
static_cast<int64_t>(blockIdx.y) * target_batch_stride +
static_cast<int64_t>(row) * target_diagonal_stride];
const float* inverse = inverse_diagonal +
static_cast<int64_t>(blockIdx.y) * n;
float energy = 0.0f;
for (int col = lane; col < row; col += 32) {
const float correction = residual_row[col] * inverse[col];
const float value = fmaf(beta, correction, matrix_row[col]);
matrix_row[col] = value;
energy = fmaf(value, value, energy);
}
#pragma unroll
for (int offset = 16; offset; offset >>= 1) {
energy += __shfl_down_sync(0xffffffffu, energy, offset);
}
if (lane == 0) {
matrix_row[row] = sqrtf(
fmaxf(target_value - energy, 1.0e-30f));
}
}
void jacobi_refine(
torch::Tensor matrix,
torch::Tensor residual,
torch::Tensor target,
double beta,
int64_t first_row,
bool reciprocal_diagonal) {
TORCH_CHECK(matrix.is_cuda() && matrix.scalar_type() == torch::kFloat32 &&
matrix.dim() == 3 && matrix.size(1) == matrix.size(2) &&
matrix.stride(2) == 1,
"matrix must be a row-contiguous square CUDA FP32 view");
TORCH_CHECK(residual.is_cuda() && residual.scalar_type() == torch::kFloat32 &&
residual.sizes() == matrix.sizes() && residual.stride(2) == 1,
"residual must match matrix");
const bool target_matrix =
target.is_cuda() && target.scalar_type() == torch::kFloat32 &&
target.dim() == 3 && target.sizes() == matrix.sizes() &&
target.stride(2) == 1;
const bool target_diagonal =
target.is_cuda() && target.scalar_type() == torch::kFloat32 &&
target.dim() == 2 && target.size(0) == matrix.size(0) &&
target.size(1) == matrix.size(1) && target.stride(1) == 1;
TORCH_CHECK(target_matrix || target_diagonal,
"target must be the matrix or its diagonal");
TORCH_CHECK(beta > 0.0 && beta <= 1.0,
"beta must be in (0, 1]");
TORCH_CHECK(first_row >= 0 && first_row <= matrix.size(1),
"first row must be inside the matrix");
c10::cuda::CUDAGuard guard(matrix.device());
const int batch = static_cast<int>(matrix.size(0));
const int n = static_cast<int>(matrix.size(1));
auto inverse_diagonal = torch::empty({batch, n}, matrix.options());
if (reciprocal_diagonal) {
jacobi_refine_reciprocal_diagonal_kernel
<<<dim3((n + 255) / 256, batch), 256, 0, current_q()>>>(
matrix.data_ptr<float>(), inverse_diagonal.data_ptr<float>(),
n, matrix.stride(0), matrix.stride(1));
} else {
jacobi_tail_diagonal_kernel
<<<dim3((n + 255) / 256, batch), 256, 0, current_q()>>>(
matrix.data_ptr<float>(), inverse_diagonal.data_ptr<float>(),
n, matrix.stride(0), matrix.stride(1));
}
const int64_t target_batch_stride = target.stride(0);
const int64_t target_diagonal_stride = target_matrix
? target.stride(1) + target.stride(2)
: target.stride(1);
jacobi_tail_refine_kernel<<<dim3((n - first_row + 7) / 8, batch), 256, 0,
current_q()>>>(
matrix.data_ptr<float>(), residual.data_ptr<float>(),
target.data_ptr<float>(), inverse_diagonal.data_ptr<float>(), n,
static_cast<float>(beta), matrix.stride(0), matrix.stride(1),
residual.stride(0), residual.stride(1), target_batch_stride,
target_diagonal_stride, static_cast<int>(first_row));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "jacobi_refine failed: ",
cudaGetErrorString(error));
}
template <int N, int ROWS>
__global__ void init_lower_kernel(
const float* __restrict__ input,
float* __restrict__ output) {
constexpr int VECTORS_PER_ROW = N / 4;
const int matrix = blockIdx.x;
const int row_base = blockIdx.y * ROWS;
const size_t matrix_base = static_cast<size_t>(matrix) * N * N;
const float4* input4 = reinterpret_cast<const float4*>(input);
float4* output4 = reinterpret_cast<float4*>(output);
const size_t matrix_base4 = matrix_base / 4;
for (int index = threadIdx.x; index < ROWS * VECTORS_PER_ROW;
index += blockDim.x) {
const int row_offset = index / VECTORS_PER_ROW;
const int row = row_base + row_offset;
const int vector_col = index - row_offset * VECTORS_PER_ROW;
const int col = 4 * vector_col;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col <= row) {
value = input4[matrix_base4 +
static_cast<size_t>(row) * VECTORS_PER_ROW +
vector_col];
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
output4[matrix_base4 +
static_cast<size_t>(row) * VECTORS_PER_ROW + vector_col] =
value;
}
}
torch::Tensor init_lower(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 &&
input.is_contiguous() && input.dim() == 3 &&
input.size(1) == input.size(2),
"input must be contiguous square CUDA FP32");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
auto output = torch::empty_like(input);
#define LAUNCH_INIT_LOWER(N) \
init_lower_kernel<N, 8><<<dim3(batch, N / 8), 256, 0, current_q()>>>( \
input.data_ptr<float>(), output.data_ptr<float>())
if (n == 1024) {
LAUNCH_INIT_LOWER(1024);
} else if (n == 2048) {
LAUNCH_INIT_LOWER(2048);
} else if (n == 4096) {
LAUNCH_INIT_LOWER(4096);
} else if (n == 8192) {
LAUNCH_INIT_LOWER(8192);
} else if (n == 16384) {
LAUNCH_INIT_LOWER(16384);
} else if (n == 32768) {
LAUNCH_INIT_LOWER(32768);
} else {
TORCH_CHECK(false, "unsupported init_lower size");
}
#undef LAUNCH_INIT_LOWER
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "init_lower failed: ",
cudaGetErrorString(error));
return output;
}
void init_lower_into(torch::Tensor input, torch::Tensor output) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 &&
input.is_contiguous() && input.dim() == 3 &&
input.size(1) == input.size(2),
"input must be contiguous square CUDA FP32");
TORCH_CHECK(output.is_cuda() && output.scalar_type() == torch::kFloat32 &&
output.is_contiguous() && output.sizes() == input.sizes(),
"output must be a matching contiguous CUDA FP32 tensor");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
#define LAUNCH_INIT_LOWER_INTO(N) \
init_lower_kernel<N, 8><<<dim3(batch, N / 8), 256, 0, current_q()>>>( \
input.data_ptr<float>(), output.data_ptr<float>())
if (n == 8192) {
LAUNCH_INIT_LOWER_INTO(8192);
} else if (n == 16384) {
LAUNCH_INIT_LOWER_INTO(16384);
} else if (n == 32768) {
LAUNCH_INIT_LOWER_INTO(32768);
} else {
TORCH_CHECK(false, "unsupported init_lower_into size");
}
#undef LAUNCH_INIT_LOWER_INTO
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "init_lower_into failed: ",
cudaGetErrorString(error));
}
template <int N, int ROWS>
__global__ void reuse_lower_kernel(
const float* __restrict__ input,
float* __restrict__ output) {
constexpr int VECTORS_PER_ROW = N / 4;
const int matrix = blockIdx.x;
const int row_base = blockIdx.y * ROWS;
const size_t matrix_base = static_cast<size_t>(matrix) * N * N;
const float4* input4 = reinterpret_cast<const float4*>(input);
float4* output4 = reinterpret_cast<float4*>(output);
const size_t matrix_base4 = matrix_base / 4;
#pragma unroll
for (int row_offset = 0; row_offset < ROWS; ++row_offset) {
const int row = row_base + row_offset;
const int row_vectors = row / 4 + 1;
for (int vector_col = threadIdx.x; vector_col < row_vectors;
vector_col += blockDim.x) {
const int col = 4 * vector_col;
float4 value = input4[
matrix_base4 + static_cast<size_t>(row) * VECTORS_PER_ROW +
vector_col];
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
output4[matrix_base4 +
static_cast<size_t>(row) * VECTORS_PER_ROW +
vector_col] = value;
}
}
}
void reuse_lower_into(torch::Tensor input, torch::Tensor output) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 &&
input.is_contiguous() && input.dim() == 3 &&
input.size(1) == input.size(2),
"input must be contiguous square CUDA FP32");
TORCH_CHECK(output.is_cuda() && output.scalar_type() == torch::kFloat32 &&
output.is_contiguous() && output.sizes() == input.sizes(),
"output must be a matching contiguous CUDA FP32 tensor");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
#define LAUNCH_REUSE_LOWER(N) \
reuse_lower_kernel<N, 8><<<dim3(batch, N / 8), 256, 0, current_q()>>>(\
input.data_ptr<float>(), output.data_ptr<float>())
if (n == 8192) {
LAUNCH_REUSE_LOWER(8192);
} else if (n == 16384) {
LAUNCH_REUSE_LOWER(16384);
} else if (n == 32768) {
LAUNCH_REUSE_LOWER(32768);
} else {
TORCH_CHECK(false, "unsupported reuse_lower_into size");
}
#undef LAUNCH_REUSE_LOWER
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "reuse_lower_into failed: ",
cudaGetErrorString(error));
}
template <int N>
__global__ void prepare_column_major(
const float* __restrict__ input,
float* __restrict__ raw,
float** __restrict__ pointers) {
constexpr int STRIDE = N * N;
const int matrix = blockIdx.x;
const size_t base = static_cast<size_t>(matrix) * STRIDE;
if (threadIdx.x == 0) pointers[matrix] = raw + base;
const float4* input4 = reinterpret_cast<const float4*>(input);
float4* raw4 = reinterpret_cast<float4*>(raw);
const size_t base4 = base / 4;
for (int index = threadIdx.x; index < STRIDE / 4;
index += blockDim.x) {
raw4[base4 + index] = input4[base4 + index];
}
}
__global__ void prepare_pointer_array(
float* __restrict__ raw,
float** __restrict__ pointers,
size_t stride,
int batch) {
const int matrix = blockIdx.x * blockDim.x + threadIdx.x;
if (matrix < batch) {
pointers[matrix] = raw + static_cast<size_t>(matrix) * stride;
}
}
template <int N, int ROWS>
__global__ void prepare_view(
const float* __restrict__ input,
float* __restrict__ raw) {
constexpr int VECTORS_PER_ROW = N / 4;
const int matrix = blockIdx.x;
const int row_base = blockIdx.y * ROWS;
const size_t matrix_base = static_cast<size_t>(matrix) * N * N;
const float4* input4 = reinterpret_cast<const float4*>(input);
float4* raw4 = reinterpret_cast<float4*>(raw);
const size_t matrix_base4 = matrix_base / 4;
for (int index = threadIdx.x; index < ROWS * VECTORS_PER_ROW;
index += blockDim.x) {
const int row = row_base + index / VECTORS_PER_ROW;
const int vector_col = index - (index / VECTORS_PER_ROW) * VECTORS_PER_ROW;
const int col = 4 * vector_col;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col + 3 >= row) {
value = input4[matrix_base4 +
static_cast<size_t>(row) * VECTORS_PER_ROW +
vector_col];
if (col < row) value.x = 0.0f;
if (col + 1 < row) value.y = 0.0f;
if (col + 2 < row) value.z = 0.0f;
}
raw4[matrix_base4 +
static_cast<size_t>(row) * VECTORS_PER_ROW + vector_col] = value;
}
}
template <int N>
__global__ void transpose_upper_to_lower(
const float* __restrict__ raw,
float* __restrict__ output) {
__shared__ float tile[32][33];
constexpr int STRIDE = N * N;
constexpr int TILES = N / 32;
const int matrix = blockIdx.x / TILES;
const int tile_row = blockIdx.x - matrix * TILES;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const size_t base = static_cast<size_t>(matrix) * STRIDE;
const int row_base = tile_row * 32;
for (int tile_col = 0; tile_col <= tile_row; ++tile_col) {
const int col_base = tile_col * 32;
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int row = warp + 8 * q;
tile[row][lane] = raw[
base + static_cast<size_t>(col_base + row) * N +
row_base + lane];
}
__syncthreads();
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int row = warp + 8 * q;
const float value = tile[lane][row];
output[base + static_cast<size_t>(row_base + row) * N +
col_base + lane] =
(tile_row > tile_col || row >= lane) ? value : 0.0f;
if (tile_row > tile_col) {
output[base + static_cast<size_t>(col_base + row) * N +
row_base + lane] = 0.0f;
}
}
__syncthreads();
}
}
torch::Tensor direct_potrf(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(1) == input.size(2),
"input must be contiguous (batch,n,n)");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
TORCH_CHECK(n == 64 || n == 128 || n == 256 || n == 512,
"unsupported direct Cholesky size");
auto raw = torch::empty_like(input);
auto pointers = torch::empty({batch}, input.options().dtype(torch::kInt64));
auto info = torch::empty({batch}, input.options().dtype(torch::kInt32));
if (n == 64) {
prepare_column_major<64><<<batch, 256>>>(
input.data_ptr<float>(), raw.data_ptr<float>(),
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()));
} else if (n == 128) {
prepare_column_major<128><<<batch, 256>>>(
input.data_ptr<float>(), raw.data_ptr<float>(),
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()));
} else {
const size_t stride = static_cast<size_t>(n) * n;
if (n == 256) {
prepare_view<256, 16><<<dim3(batch, 16), 256, 0, current_q()>>>(
input.data_ptr<float>(), raw.data_ptr<float>());
} else {
prepare_view<512, 32><<<dim3(batch, 16), 256, 0, current_q()>>>(
input.data_ptr<float>(), raw.data_ptr<float>());
}
prepare_pointer_array<<<(batch + 255) / 256, 256, 0, current_q()>>>(
raw.data_ptr<float>(),
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()),
stride, batch);
}
if (solver == nullptr) {
const cusolverStatus_t create_status = cusolverDnCreate(&solver);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnCreate failed: ", static_cast<int>(create_status));
}
const cusolverStatus_t status = cusolverDnSpotrfBatched(
solver, CUBLAS_FILL_MODE_LOWER, n,
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()), n,
info.data_ptr<int>(), batch);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnSpotrfBatched failed: ", static_cast<int>(status));
if (n == 64 || n == 128) {
const size_t stride = static_cast<size_t>(n) * n;
const dim3 clear_grid((n + 7) / 8, batch);
clear_lower_kernel<<<clear_grid, 256>>>(raw.data_ptr<float>(), n, stride);
}
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "direct_potrf launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
torch::Tensor direct_potrf_split4(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 640 && input.size(1) == 512 &&
input.size(2) == 512,
"input must be contiguous batch640 n512");
c10::cuda::CUDAGuard guard(input.device());
constexpr int batch = 640;
constexpr int groups = 5;
constexpr int group_batch = batch / groups;
constexpr int n = 512;
constexpr size_t stride = static_cast<size_t>(n) * n;
auto raw = torch::empty_like(input);
auto pointers = torch::empty({batch}, input.options().dtype(torch::kInt64));
auto info = torch::empty({batch}, input.options().dtype(torch::kInt32));
prepare_view<n, 32><<<dim3(batch, n / 32), 256, 0, current_q()>>>(
input.data_ptr<float>(), raw.data_ptr<float>());
prepare_pointer_array<<<(batch + 255) / 256, 256, 0, current_q()>>>(
raw.data_ptr<float>(),
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()),
stride, batch);
using Queue = decltype(current_q());
static Queue queues[groups] = {};
static cudaEvent_t ready = nullptr;
static cudaEvent_t done[groups] = {};
if (ready == nullptr) {
TORCH_CHECK(cudaEventCreateWithFlags(
&ready, cudaEventDisableTiming) == cudaSuccess,
"batch4 ready event creation failed");
}
for (int index = 0; index < groups; ++index) {
if (queues[index] == nullptr) {
const cudaError_t create_error =
PC_CAT_(cudaStr, eamCreateWithFlags)(&queues[index], 1);
TORCH_CHECK(create_error == cudaSuccess,
"batch4 queue creation failed: ",
cudaGetErrorString(create_error));
}
if (done[index] == nullptr) {
TORCH_CHECK(cudaEventCreateWithFlags(
&done[index], cudaEventDisableTiming) == cudaSuccess,
"batch4 done event creation failed");
}
if (solver_batch4[index] == nullptr) {
const cusolverStatus_t create_status =
cusolverDnCreate(&solver_batch4[index]);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"batch4 solver creation failed: ",
static_cast<int>(create_status));
}
const cusolverStatus_t queue_status =
PC_CAT_(cusolverDnSetStr, eam)(solver_batch4[index], queues[index]);
TORCH_CHECK(queue_status == CUSOLVER_STATUS_SUCCESS,
"batch4 queue binding failed: ",
static_cast<int>(queue_status));
}
TORCH_CHECK(cudaEventRecord(ready, current_q()) == cudaSuccess,
"batch4 ready event record failed");
for (int index = 0; index < groups; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
queues[index], ready, 0) == cudaSuccess,
"batch4 queue wait failed");
const int offset = index * group_batch;
const cusolverStatus_t status = cusolverDnSpotrfBatched(
solver_batch4[index], CUBLAS_FILL_MODE_LOWER, n,
reinterpret_cast<float**>(pointers.data_ptr<int64_t>()) + offset,
n, info.data_ptr<int>() + offset, group_batch);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"batch4 POTRF failed: ", static_cast<int>(status));
TORCH_CHECK(cudaEventRecord(done[index], queues[index]) == cudaSuccess,
"batch4 done event record failed");
}
for (int index = 0; index < groups; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
current_q(), done[index], 0) == cudaSuccess,
"batch4 caller wait failed");
}
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "direct split2 launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
torch::Tensor xpotrf_bf16x9(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 1 && input.size(1) == 4096 &&
input.size(2) == 4096,
"input must be contiguous (1,4096,4096)");
c10::cuda::CUDAGuard guard(input.device());
constexpr int n = 4096;
auto raw = torch::empty_like(input);
auto info = torch::empty({1}, input.options().dtype(torch::kInt32));
prepare_view<n, 16><<<dim3(1, n / 16), 256, 0, current_q()>>>(
input.data_ptr<float>(), raw.data_ptr<float>());
if (xsolver == nullptr) {
const cusolverStatus_t create_status = cusolverDnCreate(&xsolver);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"xsolver creation failed: ",
static_cast<int>(create_status));
const cusolverStatus_t params_status = cusolverDnCreateParams(&xparams);
TORCH_CHECK(params_status == CUSOLVER_STATUS_SUCCESS,
"xsolver params creation failed: ",
static_cast<int>(params_status));
const cusolverStatus_t math_status = cusolverDnSetMathMode(
xsolver, CUSOLVER_FP32_EMULATED_BF16X9_MATH);
TORCH_CHECK(math_status == CUSOLVER_STATUS_SUCCESS,
"xsolver math mode failed: ",
static_cast<int>(math_status));
}
const cusolverStatus_t queue_status =
PC_CAT_(cusolverDnSetStr, eam)(xsolver, current_q());
TORCH_CHECK(queue_status == CUSOLVER_STATUS_SUCCESS,
"xsolver queue binding failed: ",
static_cast<int>(queue_status));
if (xdevice_bytes == 0) {
const cusolverStatus_t size_status = cusolverDnXpotrf_bufferSize(
xsolver, xparams, CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F, raw.data_ptr<float>(), n, CUDA_R_32F,
&xdevice_bytes, &xhost_bytes);
TORCH_CHECK(size_status == CUSOLVER_STATUS_SUCCESS,
"xpotrf workspace query failed: ",
static_cast<int>(size_status));
xworkspace = torch::empty(
{static_cast<int64_t>(xdevice_bytes)},
input.options().dtype(torch::kUInt8));
xhost_workspace.resize(xhost_bytes);
}
const cusolverStatus_t status = cusolverDnXpotrf(
xsolver, xparams, CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F, raw.data_ptr<float>(), n, CUDA_R_32F,
xworkspace.data_ptr(), xdevice_bytes,
xhost_bytes ? xhost_workspace.data() : nullptr, xhost_bytes,
info.data_ptr<int>());
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnXpotrf failed: ", static_cast<int>(status));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "xpotrf launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
torch::Tensor xpotrf_bf16x9_2(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 2 && input.size(1) == input.size(2) &&
(input.size(1) == 2048 || input.size(1) == 4096),
"input must be contiguous batch2 n2048 or n4096");
c10::cuda::CUDAGuard guard(input.device());
const int n = static_cast<int>(input.size(1));
const size_t matrix_elements = static_cast<size_t>(n) * n;
auto raw = torch::empty_like(input);
auto info = torch::empty({2}, input.options().dtype(torch::kInt32));
using Queue = decltype(current_q());
static Queue queues[2] = {nullptr, nullptr};
for (int index = 0; index < 2; ++index) {
if (queues[index] == nullptr) {
const cudaError_t create_error =
PC_CAT_(cudaStr, eamCreateWithFlags)(
&queues[index], 1);
TORCH_CHECK(create_error == cudaSuccess,
"queue creation failed: ",
cudaGetErrorString(create_error));
}
if (xsolver2[index] == nullptr) {
const cusolverStatus_t create_status =
cusolverDnCreate(&xsolver2[index]);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"xsolver2 creation failed: ",
static_cast<int>(create_status));
const cusolverStatus_t params_status =
cusolverDnCreateParams(&xparams2[index]);
TORCH_CHECK(params_status == CUSOLVER_STATUS_SUCCESS,
"xsolver2 params creation failed: ",
static_cast<int>(params_status));
const cusolverStatus_t math_status = cusolverDnSetMathMode(
xsolver2[index], CUSOLVER_FP32_EMULATED_BF16X9_MATH);
TORCH_CHECK(math_status == CUSOLVER_STATUS_SUCCESS,
"xsolver2 math mode failed: ",
static_cast<int>(math_status));
}
const cusolverStatus_t queue_status =
PC_CAT_(cusolverDnSetStr, eam)(xsolver2[index], queues[index]);
TORCH_CHECK(queue_status == CUSOLVER_STATUS_SUCCESS,
"xsolver2 queue binding failed: ",
static_cast<int>(queue_status));
if (xworkspace_n2[index] != n) {
const cusolverStatus_t size_status = cusolverDnXpotrf_bufferSize(
xsolver2[index], xparams2[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n,
CUDA_R_32F, &xdevice_bytes2[index], &xhost_bytes2[index]);
TORCH_CHECK(size_status == CUSOLVER_STATUS_SUCCESS,
"xpotrf2 workspace query failed: ",
static_cast<int>(size_status));
xworkspace2[index] = torch::empty(
{static_cast<int64_t>(xdevice_bytes2[index])},
input.options().dtype(torch::kUInt8));
xhost_workspace2[index].resize(xhost_bytes2[index]);
xworkspace_n2[index] = n;
}
}
cudaEvent_t ready = nullptr;
cudaEvent_t done[2] = {nullptr, nullptr};
TORCH_CHECK(cudaEventCreateWithFlags(&ready, cudaEventDisableTiming) ==
cudaSuccess,
"ready event creation failed");
TORCH_CHECK(cudaEventRecord(ready, current_q()) == cudaSuccess,
"ready event record failed");
for (int index = 0; index < 2; ++index) {
TORCH_CHECK(cudaEventCreateWithFlags(
&done[index], cudaEventDisableTiming) == cudaSuccess,
"done event creation failed");
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
queues[index], ready, 0) == cudaSuccess,
"queue wait failed");
if (n == 2048) {
prepare_view<2048, 4>
<<<dim3(1, 2048 / 4), 256, 0, queues[index]>>>(
input.data_ptr<float>() + index * matrix_elements,
raw.data_ptr<float>() + index * matrix_elements);
} else {
prepare_view<4096, 4>
<<<dim3(1, 4096 / 4), 256, 0, queues[index]>>>(
input.data_ptr<float>() + index * matrix_elements,
raw.data_ptr<float>() + index * matrix_elements);
}
const cusolverStatus_t status = cusolverDnXpotrf(
xsolver2[index], xparams2[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n, CUDA_R_32F,
xworkspace2[index].data_ptr(), xdevice_bytes2[index],
xhost_bytes2[index] ? xhost_workspace2[index].data() : nullptr,
xhost_bytes2[index], info.data_ptr<int>() + index);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnXpotrf2 failed: ",
static_cast<int>(status));
TORCH_CHECK(cudaEventRecord(done[index], queues[index]) == cudaSuccess,
"done event record failed");
}
for (int index = 0; index < 2; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
current_q(), done[index], 0) == cudaSuccess,
"caller wait failed");
cudaEventDestroy(done[index]);
}
cudaEventDestroy(ready);
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "xpotrf2 launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
torch::Tensor xpotrf_bf16x9_8(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 8 && input.size(1) == 2048 &&
input.size(2) == 2048,
"input must be contiguous batch8 n2048");
c10::cuda::CUDAGuard guard(input.device());
constexpr int n = 2048;
constexpr size_t matrix_elements = static_cast<size_t>(n) * n;
auto raw = torch::empty_like(input);
auto info = torch::empty({8}, input.options().dtype(torch::kInt32));
using Queue = decltype(current_q());
static Queue queues[8] = {};
for (int index = 0; index < 8; ++index) {
if (queues[index] == nullptr) {
const cudaError_t create_error =
PC_CAT_(cudaStr, eamCreateWithFlags)(&queues[index], 1);
TORCH_CHECK(create_error == cudaSuccess,
"queue8 creation failed: ",
cudaGetErrorString(create_error));
}
if (xsolver8[index] == nullptr) {
const cusolverStatus_t create_status =
cusolverDnCreate(&xsolver8[index]);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"xsolver8 creation failed: ",
static_cast<int>(create_status));
const cusolverStatus_t params_status =
cusolverDnCreateParams(&xparams8[index]);
TORCH_CHECK(params_status == CUSOLVER_STATUS_SUCCESS,
"xsolver8 params creation failed: ",
static_cast<int>(params_status));
const cusolverStatus_t math_status = cusolverDnSetMathMode(
xsolver8[index], CUSOLVER_FP32_EMULATED_BF16X9_MATH);
TORCH_CHECK(math_status == CUSOLVER_STATUS_SUCCESS,
"xsolver8 math mode failed: ",
static_cast<int>(math_status));
}
const cusolverStatus_t queue_status =
PC_CAT_(cusolverDnSetStr, eam)(xsolver8[index], queues[index]);
TORCH_CHECK(queue_status == CUSOLVER_STATUS_SUCCESS,
"xsolver8 queue binding failed: ",
static_cast<int>(queue_status));
if (xdevice_bytes8[index] == 0) {
const cusolverStatus_t size_status = cusolverDnXpotrf_bufferSize(
xsolver8[index], xparams8[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n,
CUDA_R_32F, &xdevice_bytes8[index], &xhost_bytes8[index]);
TORCH_CHECK(size_status == CUSOLVER_STATUS_SUCCESS,
"xpotrf8 workspace query failed: ",
static_cast<int>(size_status));
xworkspace8[index] = torch::empty(
{static_cast<int64_t>(xdevice_bytes8[index])},
input.options().dtype(torch::kUInt8));
xhost_workspace8[index].resize(xhost_bytes8[index]);
}
}
cudaEvent_t ready = nullptr;
cudaEvent_t done[8] = {};
TORCH_CHECK(cudaEventCreateWithFlags(&ready, cudaEventDisableTiming) ==
cudaSuccess,
"ready8 event creation failed");
TORCH_CHECK(cudaEventRecord(ready, current_q()) == cudaSuccess,
"ready8 event record failed");
for (int index = 0; index < 8; ++index) {
TORCH_CHECK(cudaEventCreateWithFlags(
&done[index], cudaEventDisableTiming) == cudaSuccess,
"done8 event creation failed");
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
queues[index], ready, 0) == cudaSuccess,
"queue8 wait failed");
prepare_view<2048, 4>
<<<dim3(1, 2048 / 4), 256, 0, queues[index]>>>(
input.data_ptr<float>() + index * matrix_elements,
raw.data_ptr<float>() + index * matrix_elements);
const cusolverStatus_t status = cusolverDnXpotrf(
xsolver8[index], xparams8[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n, CUDA_R_32F,
xworkspace8[index].data_ptr(), xdevice_bytes8[index],
xhost_bytes8[index] ? xhost_workspace8[index].data() : nullptr,
xhost_bytes8[index], info.data_ptr<int>() + index);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnXpotrf8 failed: ", static_cast<int>(status));
TORCH_CHECK(cudaEventRecord(done[index], queues[index]) == cudaSuccess,
"done8 event record failed");
}
for (int index = 0; index < 8; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
current_q(), done[index], 0) == cudaSuccess,
"caller8 wait failed");
cudaEventDestroy(done[index]);
}
cudaEventDestroy(ready);
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "xpotrf8 launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
torch::Tensor xpotrf_bf16x9_4(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 4 && input.size(1) == 1024 &&
input.size(2) == 1024,
"input must be contiguous batch4 n1024");
c10::cuda::CUDAGuard guard(input.device());
constexpr int n = 1024;
constexpr size_t matrix_elements = static_cast<size_t>(n) * n;
auto raw = torch::empty_like(input);
auto info = torch::empty({4}, input.options().dtype(torch::kInt32));
using Queue = decltype(current_q());
static Queue queues[4] = {};
for (int index = 0; index < 4; ++index) {
if (queues[index] == nullptr) {
const cudaError_t create_error =
PC_CAT_(cudaStr, eamCreateWithFlags)(&queues[index], 1);
TORCH_CHECK(create_error == cudaSuccess,
"queue4 creation failed: ",
cudaGetErrorString(create_error));
}
if (xsolver4[index] == nullptr) {
const cusolverStatus_t create_status =
cusolverDnCreate(&xsolver4[index]);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"xsolver4 creation failed: ",
static_cast<int>(create_status));
const cusolverStatus_t params_status =
cusolverDnCreateParams(&xparams4[index]);
TORCH_CHECK(params_status == CUSOLVER_STATUS_SUCCESS,
"xsolver4 params creation failed: ",
static_cast<int>(params_status));
const cusolverStatus_t math_status = cusolverDnSetMathMode(
xsolver4[index], CUSOLVER_FP32_EMULATED_BF16X9_MATH);
TORCH_CHECK(math_status == CUSOLVER_STATUS_SUCCESS,
"xsolver4 math mode failed: ",
static_cast<int>(math_status));
}
const cusolverStatus_t queue_status =
PC_CAT_(cusolverDnSetStr, eam)(xsolver4[index], queues[index]);
TORCH_CHECK(queue_status == CUSOLVER_STATUS_SUCCESS,
"xsolver4 queue binding failed: ",
static_cast<int>(queue_status));
if (xdevice_bytes4[index] == 0) {
const cusolverStatus_t size_status = cusolverDnXpotrf_bufferSize(
xsolver4[index], xparams4[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n,
CUDA_R_32F, &xdevice_bytes4[index], &xhost_bytes4[index]);
TORCH_CHECK(size_status == CUSOLVER_STATUS_SUCCESS,
"xpotrf4 workspace query failed: ",
static_cast<int>(size_status));
xworkspace4[index] = torch::empty(
{static_cast<int64_t>(xdevice_bytes4[index])},
input.options().dtype(torch::kUInt8));
xhost_workspace4[index].resize(xhost_bytes4[index]);
}
}
cudaEvent_t ready = nullptr;
cudaEvent_t done[4] = {};
TORCH_CHECK(cudaEventCreateWithFlags(&ready, cudaEventDisableTiming) ==
cudaSuccess,
"ready4 event creation failed");
TORCH_CHECK(cudaEventRecord(ready, current_q()) == cudaSuccess,
"ready4 event record failed");
for (int index = 0; index < 4; ++index) {
TORCH_CHECK(cudaEventCreateWithFlags(
&done[index], cudaEventDisableTiming) == cudaSuccess,
"done4 event creation failed");
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
queues[index], ready, 0) == cudaSuccess,
"queue4 wait failed");
prepare_view<1024, 4>
<<<dim3(1, 1024 / 4), 256, 0, queues[index]>>>(
input.data_ptr<float>() + index * matrix_elements,
raw.data_ptr<float>() + index * matrix_elements);
const cusolverStatus_t status = cusolverDnXpotrf(
xsolver4[index], xparams4[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n, CUDA_R_32F,
xworkspace4[index].data_ptr(), xdevice_bytes4[index],
xhost_bytes4[index] ? xhost_workspace4[index].data() : nullptr,
xhost_bytes4[index], info.data_ptr<int>() + index);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnXpotrf4 failed: ", static_cast<int>(status));
TORCH_CHECK(cudaEventRecord(done[index], queues[index]) == cudaSuccess,
"done4 event record failed");
}
for (int index = 0; index < 4; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
current_q(), done[index], 0) == cudaSuccess,
"caller4 wait failed");
cudaEventDestroy(done[index]);
}
cudaEventDestroy(ready);
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "xpotrf4 launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
torch::Tensor xpotrf_bf16x9_16(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 16 && input.size(1) == 512 &&
input.size(2) == 512,
"input must be contiguous batch16 n512");
c10::cuda::CUDAGuard guard(input.device());
constexpr int batch = 16;
constexpr int n = 512;
constexpr size_t matrix_elements = static_cast<size_t>(n) * n;
auto raw = torch::empty_like(input);
auto info = torch::empty({batch}, input.options().dtype(torch::kInt32));
using Queue = decltype(current_q());
static Queue queues[batch] = {};
static cudaEvent_t ready = nullptr;
static cudaEvent_t done[batch] = {};
if (ready == nullptr) {
TORCH_CHECK(cudaEventCreateWithFlags(
&ready, cudaEventDisableTiming) == cudaSuccess,
"ready16 event creation failed");
}
for (int index = 0; index < batch; ++index) {
if (queues[index] == nullptr) {
const cudaError_t create_error =
PC_CAT_(cudaStr, eamCreateWithFlags)(&queues[index], 1);
TORCH_CHECK(create_error == cudaSuccess,
"queue16 creation failed: ",
cudaGetErrorString(create_error));
}
if (done[index] == nullptr) {
TORCH_CHECK(cudaEventCreateWithFlags(
&done[index], cudaEventDisableTiming) ==
cudaSuccess,
"done16 event creation failed");
}
if (xsolver16[index] == nullptr) {
const cusolverStatus_t create_status =
cusolverDnCreate(&xsolver16[index]);
TORCH_CHECK(create_status == CUSOLVER_STATUS_SUCCESS,
"xsolver16 creation failed: ",
static_cast<int>(create_status));
const cusolverStatus_t params_status =
cusolverDnCreateParams(&xparams16[index]);
TORCH_CHECK(params_status == CUSOLVER_STATUS_SUCCESS,
"xsolver16 params creation failed: ",
static_cast<int>(params_status));
const cusolverStatus_t math_status = cusolverDnSetMathMode(
xsolver16[index], CUSOLVER_FP32_EMULATED_BF16X9_MATH);
TORCH_CHECK(math_status == CUSOLVER_STATUS_SUCCESS,
"xsolver16 math mode failed: ",
static_cast<int>(math_status));
}
const cusolverStatus_t queue_status =
PC_CAT_(cusolverDnSetStr, eam)(xsolver16[index], queues[index]);
TORCH_CHECK(queue_status == CUSOLVER_STATUS_SUCCESS,
"xsolver16 queue binding failed: ",
static_cast<int>(queue_status));
if (xdevice_bytes16[index] == 0) {
const cusolverStatus_t size_status = cusolverDnXpotrf_bufferSize(
xsolver16[index], xparams16[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n,
CUDA_R_32F, &xdevice_bytes16[index], &xhost_bytes16[index]);
TORCH_CHECK(size_status == CUSOLVER_STATUS_SUCCESS,
"xpotrf16 workspace query failed: ",
static_cast<int>(size_status));
xworkspace16[index] = torch::empty(
{static_cast<int64_t>(xdevice_bytes16[index])},
input.options().dtype(torch::kUInt8));
xhost_workspace16[index].resize(xhost_bytes16[index]);
}
}
TORCH_CHECK(cudaEventRecord(ready, current_q()) == cudaSuccess,
"ready16 event record failed");
for (int index = 0; index < batch; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
queues[index], ready, 0) == cudaSuccess,
"queue16 wait failed");
prepare_view<512, 32>
<<<dim3(1, 512 / 32), 256, 0, queues[index]>>>(
input.data_ptr<float>() + index * matrix_elements,
raw.data_ptr<float>() + index * matrix_elements);
const cusolverStatus_t status = cusolverDnXpotrf(
xsolver16[index], xparams16[index], CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F,
raw.data_ptr<float>() + index * matrix_elements, n, CUDA_R_32F,
xworkspace16[index].data_ptr(), xdevice_bytes16[index],
xhost_bytes16[index] ? xhost_workspace16[index].data() : nullptr,
xhost_bytes16[index], info.data_ptr<int>() + index);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
"cusolverDnXpotrf16 failed: ", static_cast<int>(status));
TORCH_CHECK(cudaEventRecord(done[index], queues[index]) == cudaSuccess,
"done16 event record failed");
}
for (int index = 0; index < batch; ++index) {
TORCH_CHECK(PC_CAT_(cudaStr, eamWaitEvent)(
current_q(), done[index], 0) == cudaSuccess,
"caller16 wait failed");
}
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "xpotrf16 launch failed: ",
cudaGetErrorString(error));
return raw.transpose(1, 2);
}
namespace {
namespace sp_cg = cooperative_groups;
namespace sp_wmma = nvcuda::wmma;
constexpr int sp_nb = 32;
using sp_accum_fragment =
sp_wmma::fragment<sp_wmma::accumulator, 16, 16, 16, float>;
using sp_left_fragment = sp_wmma::fragment<
sp_wmma::matrix_a, 16, 16, 16, __half, sp_wmma::row_major>;
using sp_right_fragment = sp_wmma::fragment<
sp_wmma::matrix_b, 16, 16, 16, __half, sp_wmma::col_major>;
using sp_right_row_fragment = sp_wmma::fragment<
sp_wmma::matrix_b, 16, 16, 16, __half, sp_wmma::row_major>;
__device__ __forceinline__ float sp_rsqrt_nr(float value) {
float reciprocal;
asm("rsqrt.approx.f32 %0, %1;" : "=f"(reciprocal) : "f"(value));
return reciprocal;
}
template <int SP_N, int SP_CTAS, int SP_THREADS, int SP_RANKK, int SP_PAIR,
int SP_GROUP, bool SP_EXTERNAL_UPDATE = false,
bool SP_FRONTIER_PIPELINE = false, int SP_MIN_BLOCKS = 1>
__global__ void __launch_bounds__(SP_THREADS, SP_MIN_BLOCKS)
sp_cluster_potrf1024(
const float* __restrict__ input,
float* __restrict__ output,
__half* __restrict__ panel_input,
__half* __restrict__ panel_half,
__half* __restrict__ inverse_half,
int stage_start) {
constexpr int sp_n = SP_N;
constexpr int sp_rankk = SP_RANKK;
constexpr int sp_pair = SP_PAIR;
constexpr int sp_group = SP_GROUP;
constexpr int sp_threads = SP_THREADS;
constexpr int sp_ctas = SP_CTAS;
constexpr int sp_warps = sp_threads / 32;
sp_cg::cluster_group cluster = sp_cg::this_cluster();
const int rank = static_cast<int>(cluster.block_rank());
const int matrix = static_cast<int>(blockIdx.x) / sp_ctas;
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
const int cluster_thread = rank * sp_threads + tid;
const int cluster_warp = rank * sp_warps + warp;
constexpr int matrix_elements = sp_n * sp_n;
constexpr int input_panel_elements = sp_n * sp_nb;
constexpr int pair_panel_elements = sp_n * sp_pair;
const float* src = input + static_cast<size_t>(matrix) * matrix_elements;
float* dst = output + static_cast<size_t>(matrix) * matrix_elements;
__half* hi = panel_input +
static_cast<size_t>(matrix) * input_panel_elements;
__half* hp = panel_half +
static_cast<size_t>(matrix) * pair_panel_elements;
__half* ih = inverse_half + static_cast<size_t>(matrix) * sp_nb * sp_nb;
__shared__ float diagonal[sp_nb * sp_nb];
__shared__ float panel_store[sp_warps * 2 * 16 * 16];
if (!SP_EXTERNAL_UPDATE || stage_start == 0) {
const float4* src4 = reinterpret_cast<const float4*>(src);
float4* dst4 = reinterpret_cast<float4*>(dst);
constexpr bool direct_lower_init =
SP_N == 512 && SP_EXTERNAL_UPDATE && SP_MIN_BLOCKS >= 3;
bool initialize_full = true;
if constexpr (direct_lower_init) {
initialize_full = gridDim.x != 640;
if (!initialize_full) {
constexpr int vectors_per_row = sp_n / 4;
for (int row = cluster_warp; row < sp_n;
row += sp_ctas * sp_warps) {
const int last_vector = row / 4;
for (int vector_col = lane; vector_col <= last_vector;
vector_col += 32) {
const int index = row * vectors_per_row + vector_col;
float4 value = src4[index];
if (vector_col == last_vector) {
const int col = 4 * vector_col;
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
dst4[index] = value;
}
}
}
}
if (initialize_full) {
for (int index = cluster_thread; index < matrix_elements / 4;
index += sp_ctas * sp_threads) {
constexpr int vectors_per_row = sp_n / 4;
const int row = index / vectors_per_row;
const int vector_col = index - row * vectors_per_row;
const int col = 4 * vector_col;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col <= row) {
value = src4[index];
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
dst4[index] = value;
}
}
}
cluster.sync();
const int first_start = SP_EXTERNAL_UPDATE ? stage_start : 0;
const int stop_start = SP_EXTERNAL_UPDATE
? min(stage_start + sp_rankk, sp_n) : sp_n;
for (int start = first_start; start < stop_start; start += sp_nb) {
const int end = start + sp_nb;
const int phase = SP_EXTERNAL_UPDATE
? (start - stage_start) / sp_nb
: (start / sp_nb) % (sp_rankk / sp_nb);
const int phase_offset = phase * sp_nb;
const int panel_count = (sp_n - end) * sp_nb;
if (rank == 0 && warp == 0) {
constexpr bool blocked_factor =
(SP_N == 512 && SP_MIN_BLOCKS >= 3) ||
(SP_N == 1024 && SP_EXTERNAL_UPDATE);
constexpr bool high_only_factor =
(SP_N == 512 && SP_MIN_BLOCKS >= 3) ||
(SP_N == 1024 && SP_EXTERNAL_UPDATE);
if constexpr (blocked_factor) {
constexpr int factor_tile = 16;
volatile float* factor_shared = diagonal;
#pragma unroll
for (int base = 0; base < sp_nb; base += factor_tile) {
if constexpr (blocked_factor) {
if (base == factor_tile) {
__half* factor_high =
reinterpret_cast<__half*>(panel_store);
__half* factor_low = factor_high + 16 * 16;
for (int index = lane; index < 16 * 16;
index += 32) {
const int row = 16 + index / 16;
const int col = index % 16;
const float value =
factor_shared[row * sp_nb + col];
const __half high = __float2half_rn(value);
factor_high[index] = high;
if constexpr (!high_only_factor) {
factor_low[index] = __float2half_rn(
value - __half2float(high));
}
}
__syncwarp();
{
sp_accum_fragment schur;
sp_wmma::load_matrix_sync(
schur,
dst + (start + 16) * sp_n + start + 16,
sp_n, sp_wmma::mem_row_major);
sp_left_fragment left;
sp_right_fragment right;
sp_wmma::load_matrix_sync(
left, factor_high, 16);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
sp_wmma::load_matrix_sync(
right, factor_high, 16);
sp_wmma::mma_sync(schur, left, right, schur);
if constexpr (!high_only_factor) {
sp_wmma::load_matrix_sync(
right, factor_low, 16);
sp_wmma::mma_sync(
schur, left, right, schur);
sp_wmma::load_matrix_sync(
left, factor_low, 16);
#pragma unroll
for (int element = 0;
element < left.num_elements;
++element) {
left.x[element] =
__hneg(left.x[element]);
}
sp_wmma::load_matrix_sync(
right, factor_high, 16);
sp_wmma::mma_sync(
schur, left, right, schur);
}
sp_wmma::store_matrix_sync(
panel_store, schur, 16,
sp_wmma::mem_row_major);
}
__syncwarp();
}
}
float factor_chunk[factor_tile];
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
if constexpr (blocked_factor) {
factor_chunk[q] = base == factor_tile
? ((lane >= 16)
? panel_store[(lane - 16) * 16 + q]
: 0.0f)
: dst[(start + lane) * sp_n + start + q];
} else {
factor_chunk[q] = dst[(start + lane) * sp_n +
start + base + q];
}
}
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
const int k = base + q;
float value = factor_chunk[q];
float diagonal_reciprocal = 0.0f;
if (lane == k) {
if constexpr (!blocked_factor) {
#pragma unroll
for (int p = 0; p < base; ++p) {
const float x =
factor_shared[lane * sp_nb + p];
value = fmaf(-x, x, value);
}
}
#pragma unroll
for (int p = 0; p < q; ++p) {
const float x = factor_chunk[p];
value = fmaf(-x, x, value);
}
value = fmaxf(value, 1.0e-20f);
if constexpr (SP_N >= 512) {
diagonal_reciprocal = sp_rsqrt_nr(value);
value *= diagonal_reciprocal;
} else {
value = sqrtf(value);
diagonal_reciprocal = 1.0f / value;
}
factor_chunk[q] = value;
}
diagonal_reciprocal = __shfl_sync(
0xffffffffu, diagonal_reciprocal, k);
if constexpr (!blocked_factor) {
if (lane > k) {
#pragma unroll
for (int p = 0; p < base; ++p) {
value = fmaf(
-factor_shared[lane * sp_nb + p],
factor_shared[k * sp_nb + p], value);
}
}
}
#pragma unroll
for (int p = 0; p < q; ++p) {
const float pivot = __shfl_sync(
0xffffffffu, factor_chunk[p], k);
if (lane > k) {
value = fmaf(-factor_chunk[p], pivot, value);
}
}
if (lane > k) {
factor_chunk[q] = value * diagonal_reciprocal;
}
}
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
const int col = base + q;
if (lane >= col) {
const float value = factor_chunk[q];
factor_shared[lane * sp_nb + col] = value;
dst[(start + lane) * sp_n + start + col] = value;
}
}
__syncwarp();
}
} else {
float factor_row[sp_nb];
#pragma unroll
for (int col = 0; col < sp_nb; ++col) {
factor_row[col] =
dst[(start + lane) * sp_n + start + col];
}
#pragma unroll
for (int k = 0; k < sp_nb; ++k) {
float value = factor_row[k];
float diagonal_reciprocal = 0.0f;
if (lane == k) {
#pragma unroll
for (int p = 0; p < k; ++p) {
const float x = factor_row[p];
value = fmaf(-x, x, value);
}
value = fmaxf(value, 1.0e-20f);
if constexpr (SP_N >= 512) {
diagonal_reciprocal = sp_rsqrt_nr(value);
value *= diagonal_reciprocal;
} else {
value = sqrtf(value);
diagonal_reciprocal = 1.0f / value;
}
factor_row[k] = value;
}
diagonal_reciprocal = __shfl_sync(
0xffffffffu, diagonal_reciprocal, k);
if (lane > k) {
value = factor_row[k];
}
#pragma unroll
for (int p = 0; p < k; ++p) {
const float pivot = __shfl_sync(
0xffffffffu, factor_row[p], k);
if (lane > k) {
value = fmaf(-factor_row[p], pivot, value);
}
}
if (lane > k) {
factor_row[k] = value * diagonal_reciprocal;
}
}
#pragma unroll
for (int col = 0; col < sp_nb; ++col) {
if (lane >= col) {
dst[(start + lane) * sp_n + start + col] =
factor_row[col];
}
}
}
}
if (end < sp_n && !(rank == 0 && warp == 0)) {
if (SP_FRONTIER_PIPELINE && phase > 0) {
constexpr int worker_warps = sp_ctas * sp_warps - 1;
const int worker_warp = cluster_warp - 1;
const int active_k = phase * sp_nb;
const int row_tiles = (sp_n - end) / 16;
for (int tile_row = worker_warp; tile_row < row_tiles;
tile_row += worker_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_wmma::load_matrix_sync(
accum[q],
dst + row * sp_n + start + 16 * q,
sp_n, sp_wmma::mem_row_major);
}
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sp_pair + k, sp_pair);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
hp + (start + 16 * q) * sp_pair + k,
sp_pair);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
float* tile = panel_store +
(warp * 2 + q) * 16 * 16;
sp_wmma::store_matrix_sync(
tile, accum[q], 16,
sp_wmma::mem_row_major);
__syncwarp();
for (int index = lane; index < 16 * 16;
index += 32) {
const int local_row = index / 16;
const int local_col = index % 16;
const float value = tile[index];
hi[(row + local_row) * sp_nb +
q * 16 + local_col] =
__float2half_rn(value);
}
__syncwarp();
}
}
} else {
constexpr int reserved = 32;
constexpr int workers = sp_ctas * sp_threads - reserved;
const int worker = (rank == 0)
? tid - reserved : sp_threads - reserved + tid;
for (int index = worker; index < panel_count;
index += workers) {
const int row = end + index / sp_nb;
const int col = index % sp_nb;
hi[row * sp_nb + col] =
__float2half_rn(dst[row * sp_n + start + col]);
}
}
}
cluster.sync();
if (end == sp_n) break;
if (tid < sp_nb) {
diagonal[tid] = 1.0f /
dst[(start + tid) * sp_n + start + tid];
}
__syncthreads();
if constexpr (SP_CTAS == 0) {
if (start >= sp_n - 4 * sp_nb) {
const int inverse_row = cluster_thread / sp_nb;
const int inverse_col = cluster_thread % sp_nb;
float inverse_value = 0.0f;
if (inverse_col <= inverse_row) {
const float inverse_row_diagonal = diagonal[inverse_row];
if (inverse_col == inverse_row) {
inverse_value = inverse_row_diagonal;
} else {
const float inverse_col_diagonal = diagonal[inverse_col];
inverse_value = -dst[(start + inverse_row) * sp_n +
start + inverse_col];
#pragma unroll
for (int middle = inverse_col + 1;
middle < inverse_row; ++middle) {
inverse_value = fmaf(
dst[(start + inverse_row) * sp_n +
start + middle] * diagonal[middle],
dst[(start + middle) * sp_n +
start + inverse_col],
inverse_value);
}
inverse_value *=
inverse_row_diagonal * inverse_col_diagonal;
}
}
ih[inverse_row * sp_nb + inverse_col] =
__float2half_rn(inverse_value);
} else {
const int inverse_col = cluster_warp;
const int inverse_index = inverse_col + lane;
float inverse_value = 0.0f;
if (lane < inverse_col) {
ih[lane * sp_nb + inverse_col] = __float2half_rn(0.0f);
}
#pragma unroll
for (int inverse_row = inverse_col;
inverse_row < sp_nb; ++inverse_row) {
float product = 0.0f;
if (inverse_index < inverse_row) {
product =
dst[(start + inverse_row) * sp_n +
start + inverse_index] * inverse_value;
}
#pragma unroll
for (int offset = 16; offset; offset >>= 1) {
product += __shfl_down_sync(
0xffffffffu, product, offset);
}
const float sum = __shfl_sync(0xffffffffu, product, 0);
if (inverse_index == inverse_row) {
inverse_value =
((inverse_row == inverse_col) ? 1.0f : 0.0f) - sum;
inverse_value *= diagonal[inverse_row];
ih[inverse_row * sp_nb + inverse_col] =
__float2half_rn(inverse_value);
}
}
}
} else {
if (rank == 0) {
constexpr int half = sp_nb / 2;
float* inverse_a = diagonal + sp_nb;
float* inverse_d = inverse_a + half * half;
float* middle = inverse_d + half * half;
const int inverse_lane = lane & 15;
const int inverse_offset = (lane >> 4) * half;
float* inverse_block = inverse_a + inverse_offset * half;
#pragma unroll
for (int inverse_base = 0; inverse_base < half;
inverse_base += sp_warps) {
const int inverse_col = inverse_base + warp;
if (inverse_col < half) {
const int inverse_index = inverse_col + inverse_lane;
float inverse_value = 0.0f;
if (inverse_lane < inverse_col) {
inverse_block[inverse_lane * half + inverse_col] =
0.0f;
}
#pragma unroll
for (int inverse_row = inverse_col;
inverse_row < half; ++inverse_row) {
float product = 0.0f;
if (inverse_index < inverse_row) {
product = dst[
(start + inverse_offset + inverse_row) * sp_n +
start + inverse_offset + inverse_index] *
inverse_value;
}
#pragma unroll
for (int offset = 8; offset; offset >>= 1) {
product += __shfl_down_sync(
0xffffffffu, product, offset, 16);
}
const float sum =
__shfl_sync(0xffffffffu, product, 0, 16);
if (inverse_index == inverse_row) {
const float identity =
(inverse_row == inverse_col) ? 1.0f : 0.0f;
inverse_value = (identity - sum) *
diagonal[inverse_offset + inverse_row];
inverse_block[inverse_row * half + inverse_col] =
inverse_value;
}
}
}
}
__syncthreads();
__half* tc_a = reinterpret_cast<__half*>(panel_store);
__half* tc_d = tc_a + half * half;
__half* tc_c = tc_d + half * half;
__half* tc_m = tc_c + half * half;
if (tid < half * half) {
const int row = tid / half;
const int col = tid - row * half;
tc_a[tid] = __float2half_rn(inverse_a[tid]);
tc_d[tid] = __float2half_rn(inverse_d[tid]);
tc_c[tid] = __float2half_rn(
dst[(start + half + row) * sp_n + start + col]);
}
__syncthreads();
if (warp == 0) {
sp_accum_fragment accum;
sp_left_fragment left;
sp_right_row_fragment right;
sp_wmma::fill_fragment(accum, 0.0f);
sp_wmma::load_matrix_sync(left, tc_d, half);
sp_wmma::load_matrix_sync(right, tc_c, half);
sp_wmma::mma_sync(accum, left, right, accum);
sp_wmma::store_matrix_sync(
middle, accum, half, sp_wmma::mem_row_major);
}
__syncthreads();
if (tid < half * half) {
tc_m[tid] = __float2half_rn(middle[tid]);
}
__syncthreads();
if (warp == 0) {
sp_accum_fragment accum;
sp_left_fragment left;
sp_right_row_fragment right;
sp_wmma::fill_fragment(accum, 0.0f);
sp_wmma::load_matrix_sync(left, tc_m, half);
sp_wmma::load_matrix_sync(right, tc_a, half);
sp_wmma::mma_sync(accum, left, right, accum);
sp_wmma::store_matrix_sync(
middle, accum, half, sp_wmma::mem_row_major);
}
__syncthreads();
if (tid < half * half) {
const int row = tid / half;
const int col = tid - row * half;
ih[row * sp_nb + col] = tc_a[tid];
ih[row * sp_nb + half + col] = __float2half_rn(0.0f);
ih[(half + row) * sp_nb + col] =
__float2half_rn(-middle[tid]);
ih[(half + row) * sp_nb + half + col] = tc_d[tid];
}
__syncthreads();
}
}
cluster.sync();
const int panel_row_tiles = (sp_n - end) / 16;
for (int tile_row = cluster_warp; tile_row < panel_row_tiles;
tile_row += sp_ctas * sp_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
sp_wmma::fill_fragment(accum[0], 0.0f);
sp_wmma::fill_fragment(accum[1], 0.0f);
#pragma unroll
for (int k = 0; k < sp_nb; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hi + row * sp_nb + k, sp_nb);
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right, ih + q * 16 * sp_nb + k, sp_nb);
sp_wmma::mma_sync(accum[q], left, right, accum[q]);
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
float* tile = panel_store + (warp * 2 + q) * 16 * 16;
sp_wmma::store_matrix_sync(
tile, accum[q], 16, sp_wmma::mem_row_major);
__syncwarp();
for (int index = lane; index < 16 * 16; index += 32) {
const int local_row = index / 16;
const int local_col = index % 16;
const float value = tile[index];
dst[(row + local_row) * sp_n + start + q * 16 + local_col] =
value;
hp[(row + local_row) * sp_pair + phase_offset +
q * 16 + local_col] =
__float2half_rn(value);
}
__syncwarp();
}
}
cluster.sync();
if (phase < (sp_rankk / sp_nb) - 1) {
const int active_k = (phase + 1) * sp_nb;
const int row_tiles = (sp_n - end) / 16;
if constexpr (SP_FRONTIER_PIPELINE) {
if (rank == 0 && warp == 0) {
#pragma unroll
for (int tile_row = 0; tile_row < 2; ++tile_row) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::load_matrix_sync(
accum[q],
dst + row * sp_n + end + 16 * q,
sp_n, sp_wmma::mem_row_major);
}
}
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sp_pair + k, sp_pair);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
hp + (end + 16 * q) * sp_pair + k,
sp_pair);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::store_matrix_sync(
dst + row * sp_n + end + 16 * q,
accum[q], sp_n,
sp_wmma::mem_row_major);
}
}
}
__syncwarp();
}
continue;
}
for (int tile_row = cluster_warp; tile_row < row_tiles;
tile_row += sp_ctas * sp_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::load_matrix_sync(
accum[q], dst + row * sp_n + end + 16 * q,
sp_n, sp_wmma::mem_row_major);
}
}
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sp_pair + k,
sp_pair);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
hp + (end + 16 * q) * sp_pair +
k,
sp_pair);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::store_matrix_sync(
dst + row * sp_n + end + 16 * q,
accum[q], sp_n, sp_wmma::mem_row_major);
}
}
}
cluster.sync();
continue;
}
if constexpr (SP_EXTERNAL_UPDATE) {
continue;
}
const int tile_count = (sp_n - end) / 16;
int job = 0;
for (int tile_row = 0; tile_row < tile_count; ++tile_row) {
const int groups = tile_row / sp_group + 1;
for (int group = 0; group < groups; ++group, ++job) {
if (job % (sp_ctas * sp_warps) != cluster_warp) continue;
const int row = end + 16 * tile_row;
const int first_col = sp_group * group;
sp_accum_fragment accum[sp_group];
#pragma unroll
for (int q = 0; q < sp_group; ++q) {
if (first_col + q <= tile_row) {
const int col = end + 16 * (first_col + q);
sp_wmma::load_matrix_sync(
accum[q], dst + row * sp_n + col, sp_n,
sp_wmma::mem_row_major);
}
}
#pragma unroll
for (int k = 0; k < sp_rankk; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sp_pair + k, sp_pair);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < sp_group; ++q) {
if (first_col + q <= tile_row) {
const int col = end + 16 * (first_col + q);
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right, hp + col * sp_pair + k, sp_pair);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
}
#pragma unroll
for (int q = 0; q < sp_group; ++q) {
if (first_col + q <= tile_row) {
const int col = end + 16 * (first_col + q);
sp_wmma::store_matrix_sync(
dst + row * sp_n + col, accum[q], sp_n,
sp_wmma::mem_row_major);
}
}
}
}
cluster.sync();
}
if constexpr (!SP_EXTERNAL_UPDATE) {
constexpr int diagonal_tile = 16;
constexpr int diagonal_elements =
(sp_n / diagonal_tile) * diagonal_tile * diagonal_tile;
for (int index = cluster_thread; index < diagonal_elements;
index += sp_ctas * sp_threads) {
const int tile = index / (diagonal_tile * diagonal_tile);
const int local = index - tile * diagonal_tile * diagonal_tile;
const int row = local / diagonal_tile;
const int col = local - row * diagonal_tile;
if (col > row) {
const int base = tile * diagonal_tile;
dst[(base + row) * sp_n + base + col] = 0.0f;
}
}
}
}
constexpr int sg_nb = 32;
template <int SG_N, int SG_RANKK, int SG_STRIDE, int SG_CTAS,
int SG_THREADS, bool SG_FP32_PANEL, bool SG_FP32_UPDATE,
bool SG_EXTERNAL_UPDATE = false>
__global__ void __launch_bounds__(SG_THREADS, 1) sg_grid_potrf2048(
const float* __restrict__ input,
float* __restrict__ output,
__half* __restrict__ panel_input,
__half* __restrict__ panel_half,
__half* __restrict__ inverse_half,
int stage_start) {
constexpr int sg_n = SG_N;
constexpr int sg_rankk = SG_RANKK;
constexpr int sg_stride = SG_STRIDE;
constexpr int sg_ctas = SG_CTAS;
constexpr int sg_threads = SG_THREADS;
constexpr int sg_warps = sg_threads / 32;
sp_cg::cluster_group grid = sp_cg::this_cluster();
const int rank = static_cast<int>(blockIdx.x) % sg_ctas;
const int matrix = static_cast<int>(blockIdx.x) / sg_ctas;
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
const int grid_thread = rank * sg_threads + tid;
const int grid_warp = rank * sg_warps + warp;
constexpr int matrix_elements = sg_n * sg_n;
constexpr int input_panel_elements = sg_n * sg_nb;
constexpr int rankk_panel_elements = sg_n * sg_stride;
const float* src = input + static_cast<size_t>(matrix) * matrix_elements;
float* dst = output + static_cast<size_t>(matrix) * matrix_elements;
__half* hi = panel_input +
static_cast<size_t>(matrix) * input_panel_elements;
__half* hp = panel_half +
static_cast<size_t>(matrix) * rankk_panel_elements;
__half* ih = inverse_half + static_cast<size_t>(matrix) * sg_nb * sg_nb;
__shared__ float panel_store[sg_warps * 2 * 16 * 16];
__shared__ float panel_reciprocal[sg_nb];
constexpr bool sg_factor_frontier_pipeline =
SG_N == 2048 && SG_RANKK == 256 && SG_STRIDE == 288 &&
SG_CTAS == 8 && SG_THREADS == 512 && !SG_FP32_PANEL &&
!SG_FP32_UPDATE && SG_EXTERNAL_UPDATE;
if (!SG_EXTERNAL_UPDATE || stage_start == 0) {
const float4* src4 = reinterpret_cast<const float4*>(src);
float4* dst4 = reinterpret_cast<float4*>(dst);
for (int index = grid_thread; index < matrix_elements / 4;
index += sg_ctas * sg_threads) {
constexpr int vectors_per_row = sg_n / 4;
const int row = index / vectors_per_row;
const int vector_col = index - row * vectors_per_row;
const int col = 4 * vector_col;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col <= row) {
value = src4[index];
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
dst4[index] = value;
}
}
grid.sync();
const int first_start = SG_EXTERNAL_UPDATE ? stage_start : 0;
const int stop_start = SG_EXTERNAL_UPDATE
? min(stage_start + sg_rankk, sg_n) : sg_n;
for (int start = first_start; start < stop_start; start += sg_nb) {
const int end = start + sg_nb;
const int phase = SG_EXTERNAL_UPDATE
? (start - stage_start) / sg_nb
: (start / sg_nb) % (sg_rankk / sg_nb);
const int phase_offset = phase * sg_nb;
const int panel_count = (sg_n - end) * sg_nb;
if (rank == 0 && warp == 0) {
if constexpr (SG_N == 256 || SG_N == 2048) {
constexpr int factor_tile = 16;
__half* factor_high =
reinterpret_cast<__half*>(panel_store);
__half* factor_low =
factor_high + factor_tile * factor_tile;
float* factor_schur = reinterpret_cast<float*>(
factor_low + factor_tile * factor_tile);
#pragma unroll
for (int base = 0; base < sg_nb; base += factor_tile) {
if (base == factor_tile) {
#pragma unroll
for (int group = 0; group < 2; ++group) {
const int vector_index = lane + 32 * group;
const int row = factor_tile + (vector_index >> 2);
const int col = 4 * (vector_index & 3);
const int index = 4 * vector_index;
const float4 value =
*reinterpret_cast<const float4*>(
dst + (start + row) * sg_n + start + col);
const __half2 high01 = __floats2half2_rn(
value.x, value.y);
const __half2 high23 = __floats2half2_rn(
value.z, value.w);
*reinterpret_cast<__half2*>(factor_high + index) =
high01;
*reinterpret_cast<__half2*>(
factor_high + index + 2) = high23;
if constexpr (SG_N == 256) {
const float2 highf01 = __half22float2(high01);
const float2 highf23 = __half22float2(high23);
*reinterpret_cast<__half2*>(
factor_low + index) = __floats2half2_rn(
value.x - highf01.x,
value.y - highf01.y);
*reinterpret_cast<__half2*>(
factor_low + index + 2) =
__floats2half2_rn(
value.z - highf23.x,
value.w - highf23.y);
}
}
__syncwarp();
sp_accum_fragment schur;
sp_wmma::load_matrix_sync(
schur,
dst + (start + factor_tile) * sg_n +
start + factor_tile,
sg_n, sp_wmma::mem_row_major);
sp_left_fragment left;
sp_right_fragment right;
sp_wmma::load_matrix_sync(
left, factor_high, factor_tile);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
sp_wmma::load_matrix_sync(
right, factor_high, factor_tile);
sp_wmma::mma_sync(schur, left, right, schur);
if constexpr (SG_N == 256) {
sp_wmma::load_matrix_sync(
right, factor_low, factor_tile);
sp_wmma::mma_sync(schur, left, right, schur);
sp_wmma::load_matrix_sync(
left, factor_low, factor_tile);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
sp_wmma::load_matrix_sync(
right, factor_high, factor_tile);
sp_wmma::mma_sync(
schur, left, right, schur);
}
sp_wmma::store_matrix_sync(
factor_schur, schur, factor_tile,
sp_wmma::mem_row_major);
__syncwarp();
}
float factor_chunk[factor_tile];
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
factor_chunk[q] = base == factor_tile
? ((lane >= factor_tile)
? factor_schur[
(lane - factor_tile) * factor_tile + q]
: 0.0f)
: dst[(start + lane) * sg_n + start + q];
}
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
const int k = base + q;
float value = factor_chunk[q];
#pragma unroll
for (int p = 0; p < q; ++p) {
const float pivot = __shfl_sync(
0xffffffffu, factor_chunk[p], k);
if (lane >= k) {
value = fmaf(-factor_chunk[p], pivot, value);
}
}
float inverse_diagonal = 0.0f;
if (lane == k) {
value = fmaxf(value, 1.0e-20f);
inverse_diagonal = sp_rsqrt_nr(value);
}
inverse_diagonal = __shfl_sync(
0xffffffffu, inverse_diagonal, k);
if (lane >= k) {
factor_chunk[q] = value * inverse_diagonal;
}
}
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
const int col = base + q;
if (lane >= col) {
dst[(start + lane) * sg_n + start + col] =
factor_chunk[q];
}
}
__syncwarp();
}
} else {
float factor_row[sg_nb];
#pragma unroll
for (int col = 0; col < sg_nb; ++col) {
factor_row[col] =
dst[(start + lane) * sg_n + start + col];
}
#pragma unroll
for (int k = 0; k < sg_nb; ++k) {
float value = factor_row[k];
float diagonal_reciprocal = 0.0f;
if (lane == k) {
#pragma unroll
for (int p = 0; p < k; ++p) {
const float x = factor_row[p];
value = fmaf(-x, x, value);
}
value = fmaxf(value, 1.0e-20f);
if constexpr (SG_N == 2048) {
diagonal_reciprocal = sp_rsqrt_nr(value);
value *= diagonal_reciprocal;
} else {
value = sqrtf(value);
}
factor_row[k] = value;
}
const float diagonal_value =
__shfl_sync(0xffffffffu, value, k);
if constexpr (SG_N == 2048) {
diagonal_reciprocal = __shfl_sync(
0xffffffffu, diagonal_reciprocal, k);
}
if (lane > k) value = factor_row[k];
#pragma unroll
for (int p = 0; p < k; ++p) {
const float pivot = __shfl_sync(
0xffffffffu, factor_row[p], k);
if (lane > k) {
value = fmaf(-factor_row[p], pivot, value);
}
}
if (lane > k) {
if constexpr (SG_N == 2048) {
factor_row[k] = value * diagonal_reciprocal;
} else {
factor_row[k] = value / diagonal_value;
}
}
}
#pragma unroll
for (int col = 0; col < sg_nb; ++col) {
if (lane >= col) {
dst[(start + lane) * sg_n + start + col] =
factor_row[col];
}
}
}
}
const bool factor_scheduler_reserved = sg_factor_frontier_pipeline
? (rank == 0)
: (rank == 0 && warp == 0);
if (!SG_FP32_PANEL && end < sg_n &&
!factor_scheduler_reserved) {
if constexpr (sg_factor_frontier_pipeline) {
constexpr int worker_warps =
(sg_ctas - 1) * sg_warps;
const int worker_warp = grid_warp - sg_warps;
if (phase > 0) {
const int active_k = phase * sg_nb;
const int row_tiles = (sg_n - end) / 16;
for (int tile_row = worker_warp; tile_row < row_tiles;
tile_row += worker_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_wmma::load_matrix_sync(
accum[q],
dst + row * sg_n + start + 16 * q,
sg_n, sp_wmma::mem_row_major);
}
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sg_stride + k, sg_stride);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
hp + (start + 16 * q) * sg_stride + k,
sg_stride);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
float* tile = panel_store +
(warp * 2 + q) * 16 * 16;
sp_wmma::store_matrix_sync(
tile, accum[q], 16,
sp_wmma::mem_row_major);
__syncwarp();
for (int index = lane; index < 16 * 16;
index += 32) {
const int local_row = index / 16;
const int local_col = index % 16;
const float value = tile[index];
dst[(row + local_row) * sg_n + start +
q * 16 + local_col] = value;
hi[(row + local_row) * sg_nb +
q * 16 + local_col] =
__float2half_rn(value);
}
__syncwarp();
}
}
} else {
constexpr int workers = worker_warps * 32;
const int worker = worker_warp * 32 + lane;
for (int index = worker; index < panel_count;
index += workers) {
const int row = end + index / sg_nb;
const int col = index % sg_nb;
hi[row * sg_nb + col] = __float2half_rn(
dst[row * sg_n + start + col]);
}
}
} else {
constexpr int reserved = 32;
constexpr int workers = sg_ctas * sg_threads - reserved;
const int worker = grid_thread - reserved;
for (int index = worker; index < panel_count;
index += workers) {
const int row = end + index / sg_nb;
const int col = index % sg_nb;
hi[row * sg_nb + col] =
__float2half_rn(dst[row * sg_n + start + col]);
}
}
}
grid.sync();
if (end == sg_n) break;
if constexpr (SG_FP32_PANEL) {
if (tid < sg_nb) {
panel_reciprocal[tid] =
1.0f / dst[(start + tid) * sg_n + start + tid];
}
__syncthreads();
for (int row = end + grid_thread; row < sg_n;
row += sg_ctas * sg_threads) {
float solved[sg_nb];
#pragma unroll
for (int col = 0; col < sg_nb; ++col) {
float value = dst[row * sg_n + start + col];
#pragma unroll
for (int p = 0; p < col; ++p) {
value = fmaf(
-dst[(start + col) * sg_n + start + p],
solved[p], value);
}
solved[col] = value * panel_reciprocal[col];
dst[row * sg_n + start + col] = solved[col];
hp[row * sg_stride + phase_offset + col] =
__float2half_rn(solved[col]);
}
}
} else {
if (rank == 0) {
constexpr int half = sg_nb / 2;
float* inverse_a = panel_store;
float* inverse_d = inverse_a + half * half;
float* middle = inverse_d + half * half;
if (tid < sg_nb) {
panel_reciprocal[tid] = 1.0f /
dst[(start + tid) * sg_n + start + tid];
}
__syncthreads();
if constexpr (SG_N == 256 || SG_N == 2048) {
if (tid < 2 * 8 * 8) {
const int block16 = tid >> 6;
const int local = tid & 63;
const int row = local >> 3;
const int col = local & 7;
float* target = inverse_a + block16 * half * half;
target[row * half + 8 + col] = 0.0f;
}
if (warp < 4) {
const int block16 = warp >> 1;
const int block8 = (warp & 1) * 8;
const int inverse4_element = lane & 15;
const int inverse4_block = lane >> 4;
const int inverse4_row = inverse4_element >> 2;
const int inverse4_col = inverse4_element & 3;
const int inverse4_offset = block8 + 4 * inverse4_block;
float* target = inverse_a + block16 * half * half;
const int matrix_offset = block16 * half;
if (lane < 4 * 4) {
const int row = lane >> 2;
const int col = lane & 3;
target[(block8 + row) * half + block8 + 4 + col] =
0.0f;
}
float inverse4_value = 0.0f;
if (inverse4_row == inverse4_col) {
inverse4_value = panel_reciprocal[
matrix_offset + inverse4_offset + inverse4_row];
}
const float inverse4_diag = __shfl_sync(
0xffffffffu, inverse4_value,
4 * inverse4_col + inverse4_col, 16);
if (inverse4_row == inverse4_col + 1) {
const float product = __fmul_rn(
dst[(start + matrix_offset + inverse4_offset +
inverse4_row) * sg_n +
start + matrix_offset + inverse4_offset +
inverse4_col], inverse4_diag);
inverse4_value =
(0.0f - product) * panel_reciprocal[
matrix_offset + inverse4_offset + inverse4_row];
}
const int inverse4_next = min(inverse4_col + 1, 3);
const float inverse4_d2_0 = __shfl_sync(
0xffffffffu, inverse4_value,
4 * inverse4_col + inverse4_col, 16);
const float inverse4_d2_1 = __shfl_sync(
0xffffffffu, inverse4_value,
4 * inverse4_next + inverse4_col, 16);
if (inverse4_row == inverse4_col + 2) {
const float product0 = __fmul_rn(
dst[(start + matrix_offset + inverse4_offset +
inverse4_row) * sg_n +
start + matrix_offset + inverse4_offset +
inverse4_col], inverse4_d2_0);
const float product1 = __fmul_rn(
dst[(start + matrix_offset + inverse4_offset +
inverse4_row) * sg_n +
start + matrix_offset + inverse4_offset +
inverse4_col + 1], inverse4_d2_1);
const float sum = __fadd_rn(product0, product1);
inverse4_value =
(0.0f - sum) * panel_reciprocal[
matrix_offset + inverse4_offset + inverse4_row];
}
const float inverse4_d3_0 = __shfl_sync(
0xffffffffu, inverse4_value, 0, 16);
const float inverse4_d3_1 = __shfl_sync(
0xffffffffu, inverse4_value, 4, 16);
const float inverse4_d3_2 = __shfl_sync(
0xffffffffu, inverse4_value, 8, 16);
if (inverse4_row == 3 && inverse4_col == 0) {
const float product0 = __fmul_rn(
dst[(start + matrix_offset + inverse4_offset + 3) *
sg_n + start + matrix_offset +
inverse4_offset], inverse4_d3_0);
const float product1 = __fmul_rn(
dst[(start + matrix_offset + inverse4_offset + 3) *
sg_n + start + matrix_offset +
inverse4_offset + 1], inverse4_d3_1);
const float product2 = __fmul_rn(
dst[(start + matrix_offset + inverse4_offset + 3) *
sg_n + start + matrix_offset +
inverse4_offset + 2], inverse4_d3_2);
const float sum = __fadd_rn(
__fadd_rn(product0, product2), product1);
inverse4_value =
(0.0f - sum) * panel_reciprocal[
matrix_offset + inverse4_offset + 3];
}
target[(inverse4_offset + inverse4_row) * half +
inverse4_offset + inverse4_col] = inverse4_value;
__syncwarp();
if (lane < 4 * 4) {
const int row = lane >> 2;
const int col = lane & 3;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 4; ++k) {
value = fmaf(
target[(block8 + 4 + row) * half +
block8 + 4 + k],
dst[(start + matrix_offset + block8 + 4 + k) *
sg_n + start + matrix_offset + block8 + col],
value);
}
middle[warp * 16 + row * 4 + col] = value;
}
__syncwarp();
{
const int element = lane & 15;
const int row = element >> 2;
const int col = element & 3;
const int k0 = (lane >> 4) * 2;
float value = 0.0f;
#pragma unroll
for (int local_k = 0; local_k < 2; ++local_k) {
const int k = k0 + local_k;
value = fmaf(
middle[warp * 16 + row * 4 + k],
target[(block8 + k) * half + block8 + col],
value);
}
value += __shfl_xor_sync(
0xffffffffu, value, 16);
if (lane < 16) {
target[(block8 + 4 + row) * half + block8 + col] =
-value;
}
}
}
__syncthreads();
if (tid < 2 * 8 * 8) {
const int block16 = tid >> 6;
const int local = tid & 63;
const int row = local >> 3;
const int col = local & 7;
float* target = inverse_a + block16 * half * half;
const int matrix_offset = block16 * half;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 8; ++k) {
value = fmaf(
target[(8 + row) * half + 8 + k],
dst[(start + matrix_offset + 8 + k) * sg_n +
start + matrix_offset + col], value);
}
middle[block16 * 64 + row * 8 + col] = value;
}
__syncthreads();
if (tid < 2 * 8 * 8) {
const int block16 = tid >> 6;
const int local = tid & 63;
const int row = local >> 3;
const int col = local & 7;
float* target = inverse_a + block16 * half * half;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 8; ++k) {
value = fmaf(
middle[block16 * 64 + row * 8 + k],
target[k * half + col], value);
}
target[(8 + row) * half + col] = -value;
}
} else {
const int inverse_lane = lane & 15;
const int inverse_offset = (lane >> 4) * half;
float* inverse_block = inverse_a + inverse_offset * half;
const int inverse_col = warp;
const int inverse_index = inverse_col + inverse_lane;
float inverse_value = 0.0f;
if (inverse_lane < inverse_col) {
inverse_block[inverse_lane * half + inverse_col] = 0.0f;
}
#pragma unroll
for (int inverse_row = inverse_col;
inverse_row < half; ++inverse_row) {
float product = 0.0f;
if (inverse_index < inverse_row) {
product = dst[
(start + inverse_offset + inverse_row) * sg_n +
start + inverse_offset + inverse_index] *
inverse_value;
}
#pragma unroll
for (int offset = 8; offset; offset >>= 1) {
product += __shfl_down_sync(
0xffffffffu, product, offset, 16);
}
const float sum =
__shfl_sync(0xffffffffu, product, 0, 16);
if (inverse_index == inverse_row) {
const float identity =
(inverse_row == inverse_col) ? 1.0f : 0.0f;
inverse_value = (identity - sum) *
panel_reciprocal[inverse_offset + inverse_row];
inverse_block[inverse_row * half + inverse_col] =
inverse_value;
}
}
}
__syncthreads();
if constexpr (SG_N == 256 || SG_N == 2048) {
if (warp < half) {
const int row = warp;
const int col = lane & (half - 1);
const int k0 = (lane >> 4) * 8;
float value = 0.0f;
#pragma unroll
for (int local_k = 0; local_k < 8; ++local_k) {
const int k = k0 + local_k;
if (k <= row) {
value = fmaf(
inverse_d[row * half + k],
dst[(start + half + k) * sg_n + start + col],
value);
}
}
value += __shfl_xor_sync(0xffffffffu, value, half);
if (lane < half) {
middle[row * half + col] = value;
}
}
} else if (tid < half * half) {
const int row = tid / half;
const int col = tid - row * half;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < half; ++k) {
if (k <= row) {
value = fmaf(
inverse_d[row * half + k],
dst[(start + half + k) * sg_n + start + col],
value);
}
}
middle[row * half + col] = value;
}
__syncthreads();
if (tid < half * half) {
const int row = tid / half;
const int col = tid - row * half;
float lower_left = 0.0f;
#pragma unroll
for (int k = 0; k < half; ++k) {
if (k >= col) {
lower_left = fmaf(
middle[row * half + k],
inverse_a[k * half + col], lower_left);
}
}
ih[row * sg_nb + col] =
__float2half_rn(inverse_a[row * half + col]);
ih[row * sg_nb + half + col] = __float2half_rn(0.0f);
ih[(half + row) * sg_nb + col] =
__float2half_rn(-lower_left);
ih[(half + row) * sg_nb + half + col] =
__float2half_rn(inverse_d[row * half + col]);
}
__syncthreads();
}
grid.sync();
const int panel_row_tiles = (sg_n - end) / 16;
bool split_panel_outputs = false;
if constexpr (sg_factor_frontier_pipeline) {
split_panel_outputs = panel_row_tiles <= 64;
}
if (split_panel_outputs) {
const int panel_jobs = panel_row_tiles * 2;
for (int job = grid_warp; job < panel_jobs;
job += sg_ctas * sg_warps) {
const int tile_row = job >> 1;
const int q = job & 1;
const int row = end + 16 * tile_row;
sp_accum_fragment accum;
sp_wmma::fill_fragment(accum, 0.0f);
#pragma unroll
for (int k = 0; k < sg_nb; k += 16) {
sp_left_fragment left;
sp_right_fragment right;
sp_wmma::load_matrix_sync(
left, hi + row * sg_nb + k, sg_nb);
sp_wmma::load_matrix_sync(
right, ih + q * 16 * sg_nb + k, sg_nb);
sp_wmma::mma_sync(accum, left, right, accum);
}
float* tile = panel_store + (warp * 2 + q) * 16 * 16;
sp_wmma::store_matrix_sync(
tile, accum, 16, sp_wmma::mem_row_major);
__syncwarp();
for (int index = lane; index < 16 * 16; index += 32) {
const int local_row = index / 16;
const int local_col = index % 16;
const float value = tile[index];
dst[(row + local_row) * sg_n + start + q * 16 +
local_col] = value;
hp[(row + local_row) * sg_stride + phase_offset +
q * 16 + local_col] = __float2half_rn(value);
}
__syncwarp();
}
} else {
for (int tile_row = grid_warp; tile_row < panel_row_tiles;
tile_row += sg_ctas * sg_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
sp_wmma::fill_fragment(accum[0], 0.0f);
sp_wmma::fill_fragment(accum[1], 0.0f);
#pragma unroll
for (int k = 0; k < sg_nb; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(left, hi + row * sg_nb + k, sg_nb);
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right, ih + q * 16 * sg_nb + k, sg_nb);
sp_wmma::mma_sync(accum[q], left, right, accum[q]);
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
float* tile = panel_store + (warp * 2 + q) * 16 * 16;
sp_wmma::store_matrix_sync(
tile, accum[q], 16, sp_wmma::mem_row_major);
__syncwarp();
for (int index = lane; index < 16 * 16; index += 32) {
const int local_row = index / 16;
const int local_col = index % 16;
const float value = tile[index];
dst[(row + local_row) * sg_n + start + q * 16 +
local_col] = value;
hp[(row + local_row) * sg_stride + phase_offset +
q * 16 + local_col] = __float2half_rn(value);
}
__syncwarp();
}
}
}
}
grid.sync();
if (phase < (sg_rankk / sg_nb) - 1) {
const int active_k = (phase + 1) * sg_nb;
if constexpr (sg_factor_frontier_pipeline) {
if (rank == 0) {
if (warp < 3) {
const int tile_row = warp == 0 ? 0 : 1;
const int q = warp == 2 ? 1 : 0;
const int row = end + 16 * tile_row;
sp_accum_fragment accum;
sp_wmma::load_matrix_sync(
accum, dst + row * sg_n + end + 16 * q,
sg_n, sp_wmma::mem_row_major);
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sg_stride + k, sg_stride);
#pragma unroll
for (int element = 0;
element < left.num_elements; ++element) {
left.x[element] = __hneg(left.x[element]);
}
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
hp + (end + 16 * q) * sg_stride + k,
sg_stride);
sp_wmma::mma_sync(
accum, left, right, accum);
}
sp_wmma::store_matrix_sync(
dst + row * sg_n + end + 16 * q,
accum, sg_n, sp_wmma::mem_row_major);
}
__syncthreads();
}
continue;
} else if constexpr (SG_FP32_UPDATE) {
const int group_start = start - phase * sg_nb;
const int frontier_elements = (sg_n - end) * sg_nb;
for (int index = grid_thread; index < frontier_elements;
index += sg_ctas * sg_threads) {
const int row = end + index / sg_nb;
const int col = end + index % sg_nb;
if (col <= row) {
float value = dst[row * sg_n + col];
#pragma unroll 4
for (int k = 0; k < active_k; ++k) {
value = fmaf(
-dst[row * sg_n + group_start + k],
dst[col * sg_n + group_start + k], value);
}
dst[row * sg_n + col] = value;
}
}
grid.sync();
continue;
}
const int row_tiles = (sg_n - end) / 16;
for (int tile_row = grid_warp; tile_row < row_tiles;
tile_row += sg_ctas * sg_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::load_matrix_sync(
accum[q], dst + row * sg_n + end + 16 * q,
sg_n, sp_wmma::mem_row_major);
}
}
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sg_stride + k, sg_stride);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
hp + (end + 16 * q) * sg_stride + k,
sg_stride);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::store_matrix_sync(
dst + row * sg_n + end + 16 * q,
accum[q], sg_n, sp_wmma::mem_row_major);
}
}
}
grid.sync();
continue;
}
if constexpr (SG_EXTERNAL_UPDATE) {
continue;
}
const int tile_count = (sg_n - end) / 16;
int job = 0;
for (int tile_row = 0; tile_row < tile_count; ++tile_row) {
const int groups = tile_row / 5 + 1;
for (int group = 0; group < groups; ++group, ++job) {
if (job % (sg_ctas * sg_warps) != grid_warp) continue;
const int row = end + 16 * tile_row;
const int first_col = 5 * group;
sp_accum_fragment accum[5];
#pragma unroll
for (int q = 0; q < 5; ++q) {
if (first_col + q <= tile_row) {
const int col = end + 16 * (first_col + q);
sp_wmma::load_matrix_sync(
accum[q], dst + row * sg_n + col, sg_n,
sp_wmma::mem_row_major);
}
}
#pragma unroll
for (int k = 0; k < sg_rankk; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, hp + row * sg_stride + k, sg_stride);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 5; ++q) {
if (first_col + q <= tile_row) {
const int col = end + 16 * (first_col + q);
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right, hp + col * sg_stride + k, sg_stride);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
}
#pragma unroll
for (int q = 0; q < 5; ++q) {
if (first_col + q <= tile_row) {
const int col = end + 16 * (first_col + q);
sp_wmma::store_matrix_sync(
dst + row * sg_n + col, accum[q], sg_n,
sp_wmma::mem_row_major);
}
}
}
}
grid.sync();
}
if constexpr (!SG_EXTERNAL_UPDATE) {
constexpr int diagonal_tile = 16;
constexpr int diagonal_elements =
(sg_n / diagonal_tile) * diagonal_tile * diagonal_tile;
for (int index = grid_thread; index < diagonal_elements;
index += sg_ctas * sg_threads) {
const int tile = index / (diagonal_tile * diagonal_tile);
const int local = index - tile * diagonal_tile * diagonal_tile;
const int row = local / diagonal_tile;
const int col = local - row * diagonal_tile;
if (col > row) {
const int base = tile * diagonal_tile;
dst[(base + row) * sg_n + base + col] = 0.0f;
}
}
}
}
constexpr int ss_n = 128;
constexpr int ss_nb = 32;
constexpr int ss_stride = 160;
constexpr int ss_threads = 256;
constexpr int ss_warps = ss_threads / 32;
constexpr size_t ss_matrix_elements = ss_n * ss_n;
constexpr size_t ss_half_elements = ss_n * ss_stride;
constexpr size_t ss_shared_bytes =
ss_matrix_elements * sizeof(float) +
ss_half_elements * sizeof(__half) + ss_nb * sizeof(float);
__global__ void __launch_bounds__(ss_threads, 1) ss_potrf128_kernel(
const float* __restrict__ input,
float* __restrict__ output) {
extern __shared__ unsigned char shared_bytes[];
float* matrix = reinterpret_cast<float*>(shared_bytes);
__half* panel = reinterpret_cast<__half*>(matrix + ss_matrix_elements);
float* reciprocal = reinterpret_cast<float*>(panel + ss_half_elements);
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
const int batch_index = static_cast<int>(blockIdx.x);
const float* src = input + static_cast<size_t>(batch_index) *
ss_matrix_elements;
float* dst = output + static_cast<size_t>(batch_index) *
ss_matrix_elements;
float* inverse = reinterpret_cast<float*>(panel);
float* inverse_a = inverse + ss_nb * ss_nb;
float* inverse_d = inverse_a + 16 * 16;
float* middle = inverse_d + 16 * 16;
const float4* src4 = reinterpret_cast<const float4*>(src);
for (int index = tid; index < ss_matrix_elements / 4;
index += ss_threads) {
constexpr int vectors_per_row = ss_n / 4;
const int row = index / vectors_per_row;
const int vector_col = index - row * vectors_per_row;
const int col = 4 * vector_col;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col <= row) {
value = src4[index];
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
*reinterpret_cast<float4*>(matrix + 4 * index) = value;
}
__syncthreads();
for (int start = 0; start < ss_n; start += ss_nb) {
const int end = start + ss_nb;
const int phase = start / ss_nb;
const int phase_offset = phase * ss_nb;
if (warp == 0) {
constexpr int factor_tile = 16;
__half* factor_high = panel;
__half* factor_low = factor_high + factor_tile * factor_tile;
float* factor_schur = reinterpret_cast<float*>(
factor_low + factor_tile * factor_tile);
#pragma unroll
for (int base = 0; base < ss_nb; base += factor_tile) {
if (base == factor_tile) {
for (int index = lane;
index < factor_tile * factor_tile; index += 32) {
const int row = factor_tile + index / factor_tile;
const int col = index % factor_tile;
const float value =
matrix[(start + row) * ss_n + start + col];
const __half high = __float2half_rn(value);
factor_high[index] = high;
factor_low[index] = __float2half_rn(
value - __half2float(high));
}
__syncwarp();
sp_accum_fragment schur;
sp_wmma::load_matrix_sync(
schur,
matrix + (start + factor_tile) * ss_n +
start + factor_tile,
ss_n, sp_wmma::mem_row_major);
sp_left_fragment left;
sp_right_fragment right;
sp_wmma::load_matrix_sync(left, factor_high, factor_tile);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
sp_wmma::load_matrix_sync(right, factor_high, factor_tile);
sp_wmma::mma_sync(schur, left, right, schur);
sp_wmma::load_matrix_sync(right, factor_low, factor_tile);
sp_wmma::mma_sync(schur, left, right, schur);
sp_wmma::load_matrix_sync(left, factor_low, factor_tile);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
sp_wmma::load_matrix_sync(right, factor_high, factor_tile);
sp_wmma::mma_sync(schur, left, right, schur);
sp_wmma::store_matrix_sync(
factor_schur, schur, factor_tile,
sp_wmma::mem_row_major);
__syncwarp();
}
float factor_chunk[factor_tile];
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
factor_chunk[q] = base == factor_tile
? ((lane >= factor_tile)
? factor_schur[
(lane - factor_tile) * factor_tile + q]
: 0.0f)
: matrix[(start + lane) * ss_n + start + q];
}
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
const int k = base + q;
float value = factor_chunk[q];
#pragma unroll
for (int p = 0; p < q; ++p) {
const float pivot = __shfl_sync(
0xffffffffu, factor_chunk[p], k);
if (lane >= k) {
value = fmaf(-factor_chunk[p], pivot, value);
}
}
float inverse_diagonal = 0.0f;
if (lane == k) {
value = fmaxf(value, 1.0e-20f);
inverse_diagonal = sp_rsqrt_nr(value);
}
inverse_diagonal = __shfl_sync(
0xffffffffu, inverse_diagonal, k);
if (lane >= k) {
factor_chunk[q] = value * inverse_diagonal;
}
}
#pragma unroll
for (int q = 0; q < factor_tile; ++q) {
const int col = base + q;
if (lane >= col) {
matrix[(start + lane) * ss_n + start + col] =
factor_chunk[q];
}
}
__syncwarp();
if (base == 0 && end < ss_n) {
if (lane < factor_tile) {
reciprocal[lane] = 1.0f / matrix[
(start + lane) * ss_n + start + lane];
}
__syncwarp();
asm volatile("bar.arrive 1, 256;" ::: "memory");
}
}
}
if (end < ss_n && warp != 0) {
constexpr int workers = ss_threads - 32;
const int worker = tid - 32;
const int panel_count = (ss_n - end) * ss_nb;
for (int index = worker; index < panel_count; index += workers) {
const int row = end + index / ss_nb;
const int col = index % ss_nb;
const float value =
matrix[row * ss_n + start + col];
const __half high = __float2half_rn(value);
panel[row * ss_stride + phase_offset + col] = high;
}
asm volatile("bar.sync 1, 256;" ::: "memory");
const int inverse_lane_early = lane & 15;
const int inverse_group_early = lane >> 4;
const unsigned inverse_mask_early = inverse_group_early == 0
? 0x0000ffffu : 0xffff0000u;
const int inverse_worker_early =
2 * (warp - 1) + inverse_group_early;
#pragma unroll
for (int inverse_col_early = inverse_worker_early;
inverse_col_early < 16; inverse_col_early += 14) {
const int inverse_index_early =
inverse_col_early + inverse_lane_early;
float inverse_value_early = 0.0f;
if (inverse_lane_early < inverse_col_early) {
inverse_a[inverse_lane_early * 16 +
inverse_col_early] = 0.0f;
}
#pragma unroll
for (int inverse_row_early = inverse_col_early;
inverse_row_early < 16; ++inverse_row_early) {
float product_early = 0.0f;
if (inverse_index_early < inverse_row_early) {
product_early = matrix[
(start + inverse_row_early) * ss_n +
start + inverse_index_early] *
inverse_value_early;
}
#pragma unroll
for (int offset = 8; offset; offset >>= 1) {
product_early += __shfl_down_sync(
inverse_mask_early, product_early, offset, 16);
}
const float sum_early = __shfl_sync(
inverse_mask_early, product_early, 0, 16);
if (inverse_index_early == inverse_row_early) {
const float identity =
inverse_row_early == inverse_col_early
? 1.0f : 0.0f;
inverse_value_early = (identity - sum_early) *
reciprocal[inverse_row_early];
inverse_a[inverse_row_early * 16 +
inverse_col_early] = inverse_value_early;
}
}
}
}
__syncthreads();
if (end == ss_n) break;
if (tid >= 16 && tid < ss_nb) {
reciprocal[tid] = 1.0f /
matrix[(start + tid) * ss_n + start + tid];
}
__syncthreads();
if (tid < 8 * 8) {
const int row = tid >> 3;
const int col = tid & 7;
inverse_d[row * 16 + 8 + col] = 0.0f;
}
if (warp < 2) {
const int block8 = warp * 8;
const int inverse4_lane = lane & 3;
const int inverse4_group = lane >> 2;
const int inverse4_block = inverse4_group >> 2;
const int inverse4_col = inverse4_group & 3;
const int inverse4_offset = block8 + 4 * inverse4_block;
const int inverse4_index = inverse4_col + inverse4_lane;
if (lane < 4 * 4) {
const int row = lane >> 2;
const int col = lane & 3;
inverse_d[(block8 + row) * 16 + block8 + 4 + col] = 0.0f;
}
float inverse4_value = 0.0f;
if (inverse4_lane < inverse4_col) {
inverse_d[(inverse4_offset + inverse4_lane) * 16 +
inverse4_offset + inverse4_col] = 0.0f;
}
#pragma unroll
for (int inverse4_row = 0; inverse4_row < 4; ++inverse4_row) {
float inverse4_product = 0.0f;
if (inverse4_row >= inverse4_col &&
inverse4_index < inverse4_row) {
inverse4_product = matrix[
(start + 16 + inverse4_offset + inverse4_row) * ss_n +
start + 16 + inverse4_offset + inverse4_index] *
inverse4_value;
}
#pragma unroll
for (int offset = 2; offset; offset >>= 1) {
inverse4_product += __shfl_down_sync(
0xffffffffu, inverse4_product, offset, 4);
}
const float inverse4_sum = __shfl_sync(
0xffffffffu, inverse4_product, 0, 4);
if (inverse4_row >= inverse4_col &&
inverse4_index == inverse4_row) {
const float identity =
inverse4_row == inverse4_col ? 1.0f : 0.0f;
inverse4_value = (identity - inverse4_sum) * reciprocal[
16 + inverse4_offset + inverse4_row];
inverse_d[(inverse4_offset + inverse4_row) * 16 +
inverse4_offset + inverse4_col] = inverse4_value;
}
}
__syncwarp();
if (lane < 4 * 4) {
const int row = lane >> 2;
const int col = lane & 3;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 4; ++k) {
value = fmaf(
inverse_d[(block8 + 4 + row) * 16 + block8 + 4 + k],
matrix[(start + 16 + block8 + 4 + k) * ss_n +
start + 16 + block8 + col],
value);
}
middle[warp * 16 + row * 4 + col] = value;
}
__syncwarp();
if (lane < 4 * 4) {
const int row = lane >> 2;
const int col = lane & 3;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 4; ++k) {
value = fmaf(
middle[warp * 16 + row * 4 + k],
inverse_d[(block8 + k) * 16 + block8 + col], value);
}
inverse_d[(block8 + 4 + row) * 16 + block8 + col] = -value;
}
}
__syncthreads();
if (tid < 8 * 8) {
const int row = tid >> 3;
const int col = tid & 7;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 8; ++k) {
value = fmaf(
inverse_d[(8 + row) * 16 + 8 + k],
matrix[(start + 24 + k) * ss_n + start + 16 + col],
value);
}
middle[row * 8 + col] = value;
}
__syncthreads();
if (tid < 8 * 8) {
const int row = tid >> 3;
const int col = tid & 7;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 8; ++k) {
value = fmaf(
middle[row * 8 + k], inverse_d[k * 16 + col], value);
}
inverse_d[(8 + row) * 16 + col] = -value;
}
__syncthreads();
if (tid < 16 * 16) {
const int row = tid / 16;
const int col = tid - row * 16;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < 16; ++k) {
if (k <= row) {
value = fmaf(
inverse_d[row * 16 + k],
matrix[(start + 16 + k) * ss_n + start + col],
value);
}
}
middle[row * 16 + col] = value;
}
__syncthreads();
if (tid < 16 * 16) {
const int row = tid / 16;
const int col = tid - row * 16;
float lower_left = 0.0f;
#pragma unroll
for (int k = 0; k < 16; ++k) {
if (k >= col) {
lower_left = fmaf(
middle[row * 16 + k],
inverse_a[k * 16 + col], lower_left);
}
}
inverse[row * ss_nb + col] = inverse_a[row * 16 + col];
inverse[row * ss_nb + 16 + col] = 0.0f;
inverse[(16 + row) * ss_nb + col] = -lower_left;
inverse[(16 + row) * ss_nb + 16 + col] =
inverse_d[row * 16 + col];
}
__syncthreads();
__half* inverse_hi = panel + 2 * ss_nb * ss_nb;
__half* inverse_lo = inverse_hi + ss_nb * ss_nb;
for (int index = tid; index < ss_nb * ss_nb;
index += ss_threads) {
const float value = inverse[index];
const __half high = __float2half_rn(value);
inverse_hi[index] = high;
inverse_lo[index] =
__float2half_rn(value - __half2float(high));
}
__syncthreads();
const int panel_row_tiles = (ss_n - end) / 16;
for (int tile_row = warp; tile_row < panel_row_tiles;
tile_row += ss_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
sp_wmma::fill_fragment(accum[0], 0.0f);
sp_wmma::fill_fragment(accum[1], 0.0f);
#pragma unroll
for (int k = 0; k < ss_nb; k += 16) {
sp_left_fragment left_hi;
sp_wmma::load_matrix_sync(
left_hi,
panel + row * ss_stride + phase_offset + k,
ss_stride);
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_right_fragment right_hi;
sp_right_fragment right_lo;
sp_wmma::load_matrix_sync(
right_hi,
inverse_hi + q * 16 * ss_nb + k, ss_nb);
sp_wmma::load_matrix_sync(
right_lo,
inverse_lo + q * 16 * ss_nb + k, ss_nb);
sp_wmma::mma_sync(
accum[q], left_hi, right_hi, accum[q]);
sp_wmma::mma_sync(
accum[q], left_hi, right_lo, accum[q]);
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
sp_wmma::store_matrix_sync(
matrix + row * ss_n + start + 16 * q,
accum[q], ss_n, sp_wmma::mem_row_major);
}
}
__syncthreads();
const int solved_elements = (ss_n - end) * ss_nb;
for (int index = tid; index < solved_elements;
index += ss_threads) {
const int row = end + index / ss_nb;
const int col = index - (row - end) * ss_nb;
panel[row * ss_stride + phase_offset + col] =
__float2half_rn(matrix[row * ss_n + start + col]);
}
__syncthreads();
const int active_k = (phase + 1) * ss_nb;
const int row_tiles = (ss_n - end) / 16;
for (int tile_row = warp; tile_row < row_tiles;
tile_row += ss_warps) {
const int row = end + 16 * tile_row;
sp_accum_fragment accum[2];
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::load_matrix_sync(
accum[q], matrix + row * ss_n + end + 16 * q,
ss_n, sp_wmma::mem_row_major);
}
}
#pragma unroll
for (int k = 0; k < active_k; k += 16) {
sp_left_fragment left;
sp_wmma::load_matrix_sync(
left, panel + row * ss_stride + k, ss_stride);
#pragma unroll
for (int element = 0; element < left.num_elements;
++element) {
left.x[element] = __hneg(left.x[element]);
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_right_fragment right;
sp_wmma::load_matrix_sync(
right,
panel + (end + 16 * q) * ss_stride + k,
ss_stride);
sp_wmma::mma_sync(
accum[q], left, right, accum[q]);
}
}
}
#pragma unroll
for (int q = 0; q < 2; ++q) {
if (q <= tile_row) {
sp_wmma::store_matrix_sync(
matrix + row * ss_n + end + 16 * q,
accum[q], ss_n, sp_wmma::mem_row_major);
}
}
}
__syncthreads();
}
float4* dst4 = reinterpret_cast<float4*>(dst);
for (int index = tid; index < ss_matrix_elements / 4;
index += ss_threads) {
constexpr int vectors_per_row = ss_n / 4;
const int row = index / vectors_per_row;
const int col = 4 * (index - row * vectors_per_row);
float4 value = *reinterpret_cast<const float4*>(matrix + 4 * index);
if (col > row) {
value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
} else {
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
dst4[index] = value;
}
}
} // namespace
torch::Tensor cluster_potrf1024(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(1) == input.size(2),
"specialized potrf requires contiguous square batches");
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
const bool route1024 = batch == 60 && n == 1024;
const bool route512 = batch == 640 && n == 512;
TORCH_CHECK(route1024 || route512,
"specialized potrf requires batch60 n1024 or batch640 n512");
const int ctas = route1024 ? 2 : 1;
constexpr int threads = 512;
const int solved_width = route1024 ? 160 : 224;
c10::cuda::CUDAGuard guard(input.device());
auto output = torch::empty_like(input);
auto panel = torch::empty(
{input.size(0), n, solved_width},
input.options().dtype(torch::kFloat16));
auto panel_input = torch::empty(
{input.size(0), n, sp_nb},
input.options().dtype(torch::kFloat16));
auto inverse = torch::empty(
{input.size(0), sp_nb, sp_nb},
input.options().dtype(torch::kFloat16));
cudaLaunchConfig_t config = {};
config.gridDim = dim3(static_cast<unsigned>(batch * ctas));
config.blockDim = dim3(threads);
config.PC_CAT_(str, eam) = current_q();
cudaLaunchAttribute attribute = {};
attribute.id = cudaLaunchAttributeClusterDimension;
attribute.val.clusterDim.x = ctas;
attribute.val.clusterDim.y = 1;
attribute.val.clusterDim.z = 1;
config.attrs = &attribute;
config.numAttrs = 1;
cudaError_t error;
if (route1024) {
error = cudaLaunchKernelEx(
&config, sp_cluster_potrf1024<1024, 2, 512, 160, 160, 5>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), 0);
} else {
error = cudaLaunchKernelEx(
&config, sp_cluster_potrf1024<512, 1, 512, 192, 224, 4>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), 0);
}
TORCH_CHECK(error == cudaSuccess, "cluster_potrf1024 launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor cluster_potrf512_gemm(torch::Tensor input) {
constexpr int n = 512;
constexpr int rankk = 192;
constexpr int stride = 224;
constexpr int ctas = 1;
constexpr int threads = 384;
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 &&
input.is_contiguous() && input.dim() == 3 &&
input.size(1) == n && input.size(2) == n,
"cluster_potrf512_gemm requires contiguous n512 input");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
TORCH_CHECK(batch == 4 || batch == 640,
"cluster_potrf512_gemm supports batch4 or batch640");
auto output = torch::empty_like(input);
auto panel = torch::empty(
{batch, n, stride}, input.options().dtype(torch::kFloat16));
auto panel_input = torch::empty(
{batch, n, sp_nb}, input.options().dtype(torch::kFloat16));
auto inverse = torch::empty(
{batch, sp_nb, sp_nb}, input.options().dtype(torch::kFloat16));
if (sg_blas == nullptr) {
TORCH_CHECK(cublasCreate(&sg_blas) == CUBLAS_STATUS_SUCCESS,
"cluster_potrf512_gemm handle creation failed");
TORCH_CHECK(cublasSetMathMode(sg_blas, CUBLAS_TENSOR_OP_MATH) ==
CUBLAS_STATUS_SUCCESS,
"cluster_potrf512_gemm math mode failed");
}
TORCH_CHECK(PC_CAT_(cublasSetStr, eam)(sg_blas, current_q()) ==
CUBLAS_STATUS_SUCCESS,
"cluster_potrf512_gemm queue binding failed");
cudaLaunchConfig_t config = {};
config.gridDim = dim3(static_cast<unsigned>(batch * ctas));
config.blockDim = dim3(threads);
config.PC_CAT_(str, eam) = current_q();
cudaLaunchAttribute attribute = {};
attribute.id = cudaLaunchAttributeClusterDimension;
attribute.val.clusterDim.x = ctas;
attribute.val.clusterDim.y = 1;
attribute.val.clusterDim.z = 1;
config.attrs = &attribute;
config.numAttrs = 1;
const float alpha = -1.0f;
const float beta = 1.0f;
for (int start = 0; start < n; start += rankk) {
cudaError_t error;
if (batch == 640 && (start == 0 || start == 2 * rankk)) {
error = cudaLaunchKernelEx(
&config,
sp_cluster_potrf1024<
n, ctas, threads, rankk, stride, 4, true, true, 3>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), start);
} else {
error = cudaLaunchKernelEx(
&config,
sp_cluster_potrf1024<
n, ctas, threads, rankk, stride, 4, true, false, 3>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), start);
}
TORCH_CHECK(error == cudaSuccess,
"cluster_potrf512_gemm stage failed: ",
cudaGetErrorString(error));
const int end = min(start + rankk, n);
if (end == n) break;
const int m = n - end;
const long long panel_batch_stride =
static_cast<long long>(n) * stride;
const long long output_batch_stride =
static_cast<long long>(n) * n;
const cublasStatus_t status = cublasGemmStridedBatchedEx(
sg_blas, CUBLAS_OP_T, CUBLAS_OP_N, m, m, rankk,
&alpha,
reinterpret_cast<const __half*>(panel.data_ptr<at::Half>()) +
static_cast<long long>(end) * stride,
CUDA_R_16F, stride, panel_batch_stride,
reinterpret_cast<const __half*>(panel.data_ptr<at::Half>()) +
static_cast<long long>(end) * stride,
CUDA_R_16F, stride, panel_batch_stride,
&beta,
output.data_ptr<float>() + static_cast<long long>(end) * n + end,
CUDA_R_32F, n, output_batch_stride, batch,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS,
"cluster_potrf512_gemm update failed: ",
static_cast<int>(status));
}
clear_upper(output);
return output;
}
torch::Tensor cluster_potrf1024_gemm(torch::Tensor input) {
constexpr int n = 1024;
constexpr int rankk = 160;
constexpr int stride = 160;
constexpr int ctas = 2;
constexpr int threads = 512;
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 &&
input.is_contiguous() && input.dim() == 3 &&
input.size(1) == n && input.size(2) == n,
"cluster_potrf1024_gemm requires contiguous n1024 input");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
TORCH_CHECK(batch == 60,
"cluster_potrf1024_gemm supports batch60");
auto output = torch::empty_like(input);
auto panel = torch::empty(
{batch, n, stride}, input.options().dtype(torch::kFloat16));
auto panel_input = torch::empty(
{batch, n, sp_nb}, input.options().dtype(torch::kFloat16));
auto inverse = torch::empty(
{batch, sp_nb, sp_nb}, input.options().dtype(torch::kFloat16));
auto lower_workspace = torch::empty(
{32 * 1024 * 1024}, input.options().dtype(torch::kUInt8));
if (sg_blas == nullptr) {
TORCH_CHECK(cublasCreate(&sg_blas) == CUBLAS_STATUS_SUCCESS,
"cluster_potrf1024_gemm handle creation failed");
TORCH_CHECK(cublasSetMathMode(sg_blas, CUBLAS_TENSOR_OP_MATH) ==
CUBLAS_STATUS_SUCCESS,
"cluster_potrf1024_gemm math mode failed");
}
TORCH_CHECK(PC_CAT_(cublasSetStr, eam)(sg_blas, current_q()) ==
CUBLAS_STATUS_SUCCESS,
"cluster_potrf1024_gemm queue binding failed");
cudaLaunchConfig_t config = {};
config.gridDim = dim3(static_cast<unsigned>(batch * ctas));
config.blockDim = dim3(threads);
config.PC_CAT_(str, eam) = current_q();
cudaLaunchAttribute attribute = {};
attribute.id = cudaLaunchAttributeClusterDimension;
attribute.val.clusterDim.x = ctas;
attribute.val.clusterDim.y = 1;
attribute.val.clusterDim.z = 1;
config.attrs = &attribute;
config.numAttrs = 1;
const float alpha = -1.0f;
const float beta = 1.0f;
for (int start = 0; start < n; start += rankk) {
const cudaError_t error = cudaLaunchKernelEx(
&config,
sp_cluster_potrf1024<
n, ctas, threads, rankk, stride, 5, true, true>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), start);
TORCH_CHECK(error == cudaSuccess,
"cluster_potrf1024_gemm stage failed: ",
cudaGetErrorString(error));
const int end = min(start + rankk, n);
if (end == n) break;
const int m = n - end;
if (m > 640) {
auto trailing = output.narrow(1, end, m).narrow(2, end, m);
auto panel_view = panel.narrow(1, end, m).narrow(2, 0, rankk);
fp16_lower_rankk_update_lt(
trailing, trailing, panel_view, lower_workspace, 512);
continue;
}
const long long panel_batch_stride =
static_cast<long long>(n) * stride;
const long long output_batch_stride =
static_cast<long long>(n) * n;
const cublasStatus_t status = cublasGemmStridedBatchedEx(
sg_blas, CUBLAS_OP_T, CUBLAS_OP_N, m, m, rankk,
&alpha,
reinterpret_cast<const __half*>(panel.data_ptr<at::Half>()) +
static_cast<long long>(end) * stride,
CUDA_R_16F, stride, panel_batch_stride,
reinterpret_cast<const __half*>(panel.data_ptr<at::Half>()) +
static_cast<long long>(end) * stride,
CUDA_R_16F, stride, panel_batch_stride,
&beta,
output.data_ptr<float>() + static_cast<long long>(end) * n + end,
CUDA_R_32F, n, output_batch_stride, batch,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS,
"cluster_potrf1024_gemm update failed: ",
static_cast<int>(status));
}
clear_upper(output);
return output;
}
template <int N, int BATCH, int RANKK, int STRIDE, int CTAS, int THREADS,
bool FP32_PANEL, bool FP32_UPDATE>
torch::Tensor cluster_potrf_fixed(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == BATCH && input.size(1) == N &&
input.size(2) == N,
"cluster_potrf_fixed received an unsupported shape");
c10::cuda::CUDAGuard guard(input.device());
auto output = torch::empty_like(input);
auto panel = torch::empty(
{input.size(0), N, STRIDE},
input.options().dtype(torch::kFloat16));
auto panel_input = torch::empty(
{input.size(0), N, sg_nb},
input.options().dtype(torch::kFloat16));
auto inverse = torch::empty(
{input.size(0), sg_nb, sg_nb},
input.options().dtype(torch::kFloat16));
cudaLaunchConfig_t config = {};
config.gridDim = dim3(
static_cast<unsigned>(input.size(0) * CTAS));
config.blockDim = dim3(THREADS);
config.PC_CAT_(str, eam) = current_q();
cudaLaunchAttribute attribute = {};
attribute.id = cudaLaunchAttributeClusterDimension;
attribute.val.clusterDim.x = CTAS;
attribute.val.clusterDim.y = 1;
attribute.val.clusterDim.z = 1;
config.attrs = &attribute;
config.numAttrs = 1;
const cudaError_t error = cudaLaunchKernelEx(
&config,
sg_grid_potrf2048<N, RANKK, STRIDE, CTAS, THREADS, FP32_PANEL,
FP32_UPDATE>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), 0);
TORCH_CHECK(error == cudaSuccess, "cluster_potrf_fixed launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor grid_potrf2048(torch::Tensor input) {
return cluster_potrf_fixed<2048, 8, 320, 352, 8, 512, false, false>(input);
}
torch::Tensor grid_potrf2048_gemm(torch::Tensor input) {
constexpr int n = 2048;
constexpr int rankk = 256;
constexpr int stride = 288;
constexpr int ctas = 8;
constexpr int threads = 512;
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 &&
input.is_contiguous() && input.dim() == 3 &&
input.size(1) == n && input.size(2) == n,
"grid_potrf2048_gemm requires contiguous n2048 input");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
TORCH_CHECK(batch == 1 || batch == 8,
"grid_potrf2048_gemm supports batch1 or batch8");
auto output = torch::empty_like(input);
auto panel = torch::empty(
{batch, n, stride}, input.options().dtype(torch::kFloat16));
auto panel_input = torch::empty(
{batch, n, sg_nb}, input.options().dtype(torch::kFloat16));
auto inverse = torch::empty(
{batch, sg_nb, sg_nb}, input.options().dtype(torch::kFloat16));
if (sg_blas == nullptr) {
TORCH_CHECK(cublasCreate(&sg_blas) == CUBLAS_STATUS_SUCCESS,
"grid_potrf2048_gemm handle creation failed");
TORCH_CHECK(cublasSetMathMode(sg_blas, CUBLAS_TENSOR_OP_MATH) ==
CUBLAS_STATUS_SUCCESS,
"grid_potrf2048_gemm math mode failed");
}
TORCH_CHECK(PC_CAT_(cublasSetStr, eam)(sg_blas, current_q()) ==
CUBLAS_STATUS_SUCCESS,
"grid_potrf2048_gemm queue binding failed");
cudaLaunchConfig_t config = {};
config.gridDim = dim3(static_cast<unsigned>(batch * ctas));
config.blockDim = dim3(threads);
config.PC_CAT_(str, eam) = current_q();
cudaLaunchAttribute attribute = {};
attribute.id = cudaLaunchAttributeClusterDimension;
attribute.val.clusterDim.x = ctas;
attribute.val.clusterDim.y = 1;
attribute.val.clusterDim.z = 1;
config.attrs = &attribute;
config.numAttrs = 1;
const float alpha = -1.0f;
const float beta = 1.0f;
for (int start = 0; start < n; start += rankk) {
const cudaError_t error = cudaLaunchKernelEx(
&config,
sg_grid_potrf2048<n, rankk, stride, ctas, threads,
false, false, true>,
input.data_ptr<float>(), output.data_ptr<float>(),
reinterpret_cast<__half*>(panel_input.data_ptr<at::Half>()),
reinterpret_cast<__half*>(panel.data_ptr<at::Half>()),
reinterpret_cast<__half*>(inverse.data_ptr<at::Half>()), start);
TORCH_CHECK(error == cudaSuccess,
"grid_potrf2048_gemm stage failed: ",
cudaGetErrorString(error));
const int end = min(start + rankk, n);
if (end == n) break;
const int m = n - end;
const long long panel_batch_stride =
static_cast<long long>(n) * stride;
const long long output_batch_stride =
static_cast<long long>(n) * n;
const cublasStatus_t status = cublasGemmStridedBatchedEx(
sg_blas, CUBLAS_OP_T, CUBLAS_OP_N, m, m, rankk,
&alpha,
reinterpret_cast<const __half*>(panel.data_ptr<at::Half>()) +
static_cast<long long>(end) * stride,
CUDA_R_16F, stride, panel_batch_stride,
reinterpret_cast<const __half*>(panel.data_ptr<at::Half>()) +
static_cast<long long>(end) * stride,
CUDA_R_16F, stride, panel_batch_stride,
&beta,
output.data_ptr<float>() + static_cast<long long>(end) * n + end,
CUDA_R_32F, n, output_batch_stride, batch,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS,
"grid_potrf2048_gemm update failed: ",
static_cast<int>(status));
}
clear_upper(output);
return output;
}
torch::Tensor cluster_potrf256_b64(torch::Tensor input) {
return cluster_potrf_fixed<256, 64, 256, 288, 2, 512, false, false>(input);
}
torch::Tensor shared_potrf128_b256(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(0) == 256 && input.size(1) == ss_n &&
input.size(2) == ss_n,
"shared_potrf128_b256 requires contiguous (256,128,128)");
c10::cuda::CUDAGuard guard(input.device());
auto output = torch::empty_like(input);
static bool configured = false;
if (!configured) {
const cudaError_t attribute_error = cudaFuncSetAttribute(
ss_potrf128_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(ss_shared_bytes));
TORCH_CHECK(attribute_error == cudaSuccess,
"shared_potrf128 shared-memory configuration failed: ",
cudaGetErrorString(attribute_error));
configured = true;
}
ss_potrf128_kernel
<<<256, ss_threads, ss_shared_bytes, current_q()>>>(
input.data_ptr<float>(), output.data_ptr<float>());
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "shared_potrf128 launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor cluster_potrf128_b256(torch::Tensor input) {
return cluster_potrf_fixed<128, 256, 128, 160, 1, 256, true, false>(input);
}
"""
_solver_load_detail = ""
try:
_solver = load_inline(
name="cholesky_c954_n2048_rank0_factor_only",
cpp_sources=[_CPP_SRC],
cuda_sources=[_CUDA_SRC],
functions=[
"direct_potrf",
"direct_potrf_split4",
"xpotrf_bf16x9",
"xpotrf_bf16x9_2",
"xpotrf_bf16x9_4",
"xpotrf_bf16x9_8",
"xpotrf_bf16x9_16",
"fp8_rankk_update_lt",
"fp8_lower_rankk_update_lt",
"fp16_lower_rankk_update_lt",
"tf32_lower_rankk_update_4096_lt",
"init_lower",
"init_lower_into",
"reuse_lower_into",
"clear_upper",
"clear_upper_view",
"diagonal_tail",
"jacobi_tail",
"jacobi_refine",
"cluster_potrf1024",
"cluster_potrf512_gemm",
"cluster_potrf1024_gemm",
"grid_potrf2048",
"grid_potrf2048_gemm",
"cluster_potrf256_b64",
"cluster_potrf128_b256",
"shared_potrf128_b256",
],
extra_cuda_cflags=["-O3"],
extra_ldflags=["-lcusolver", "-lcublasLt", "-lcublas"],
with_cuda=True,
verbose=False,
)
except Exception as error:
message = str(error)
detail_lines = [
text for text in message.splitlines()
if ".cu" in text and "error:" in text.lower()
]
detail_line = detail_lines[0] if detail_lines else message
_solver_load_detail = detail_line.lower().split("error:", 1)[-1].strip()
_solver = None
_DX_CPP_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime_api.h>
extern "C" cudaError_t launch_dx_potrf64(
const float* input, float* output, int batch);
extern "C" cudaError_t launch_dx_potrf128(
const float* input, float* output, int batch);
extern "C" cudaError_t launch_dx_potrf32(
const float* input, float* output, int batch);
torch::Tensor dx_potrf32(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(1) == 32 && input.size(2) == 32,
"input must be contiguous (batch,32,32)");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
auto output = torch::empty_like(input);
const cudaError_t error = launch_dx_potrf32(
input.data_ptr<float>(), output.data_ptr<float>(), batch);
TORCH_CHECK(error == cudaSuccess, "dx_potrf32 failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor dx_potrf64(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(1) == 64 && input.size(2) == 64,
"input must be contiguous (batch,64,64)");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
auto output = torch::empty_like(input);
const cudaError_t error = launch_dx_potrf64(
input.data_ptr<float>(), output.data_ptr<float>(), batch);
TORCH_CHECK(error == cudaSuccess, "dx_potrf64 failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor dx_potrf128(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.is_contiguous() && input.dim() == 3 &&
input.size(1) == 128 && input.size(2) == 128,
"input must be contiguous (batch,128,128)");
c10::cuda::CUDAGuard guard(input.device());
const int batch = static_cast<int>(input.size(0));
auto output = torch::empty_like(input);
const cudaError_t error = launch_dx_potrf128(
input.data_ptr<float>(), output.data_ptr<float>(), batch);
TORCH_CHECK(error == cudaSuccess, "dx_potrf128 failed: ",
cudaGetErrorString(error));
return output;
}
"""
_DX_CUDA_SRC = r"""
#include <cuda_runtime.h>
#include <cusolverdx.hpp>
template <class Solver, int STAGE_UNROLL>
__global__ void dx_potrf_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch) {
CUSOLVERDX_SKIP_IF_NOT_APPLICABLE_SM(Solver);
constexpr int n = Solver::m_size;
constexpr int lda = Solver::lda;
constexpr int matrices_per_block = Solver::batches_per_block;
constexpr int matrix_elements = n * n;
constexpr int shared_stride = lda * n;
constexpr int vectors_per_row = n / 4;
constexpr int vectors_per_matrix = n * vectors_per_row;
constexpr int threads = Solver::block_dim.x;
const int matrix_base = blockIdx.x * matrices_per_block;
extern __shared__ unsigned char shared_bytes[];
float* factor = reinterpret_cast<float*>(shared_bytes);
#pragma unroll STAGE_UNROLL
for (int index = threadIdx.x;
index < matrices_per_block * vectors_per_matrix;
index += threads) {
const int local_matrix = index / vectors_per_matrix;
const int vector = index - local_matrix * vectors_per_matrix;
const int row = vector / vectors_per_row;
const int vector_col = vector - row * vectors_per_row;
const int col = 4 * vector_col;
const int matrix = matrix_base + local_matrix;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (matrix < batch && col <= row) {
const size_t base = static_cast<size_t>(matrix) * matrix_elements;
value = *reinterpret_cast<const float4*>(
input + base + static_cast<size_t>(row) * n + col);
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
} else if (matrix >= batch) {
if (row == col) value.x = 1.0f;
if (row == col + 1) value.y = 1.0f;
if (row == col + 2) value.z = 1.0f;
if (row == col + 3) value.w = 1.0f;
}
*reinterpret_cast<float4*>(
factor + local_matrix * shared_stride + row * lda + col) = value;
}
__syncthreads();
// The CTA overwrites these transient status words after solver completion.
auto* status = reinterpret_cast<typename Solver::status_type*>(
output + static_cast<size_t>(matrix_base) * matrix_elements);
Solver().execute(factor, status);
__syncthreads();
#pragma unroll STAGE_UNROLL
for (int index = threadIdx.x;
index < matrices_per_block * vectors_per_matrix;
index += threads) {
const int local_matrix = index / vectors_per_matrix;
const int vector = index - local_matrix * vectors_per_matrix;
const int matrix = matrix_base + local_matrix;
if (matrix < batch) {
const int row = vector / vectors_per_row;
const int vector_col = vector - row * vectors_per_row;
const int col = 4 * vector_col;
const size_t base = static_cast<size_t>(matrix) * matrix_elements;
float4 value = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (col <= row) {
value = *reinterpret_cast<const float4*>(
factor + local_matrix * shared_stride + row * lda + col);
if (col + 1 > row) value.y = 0.0f;
if (col + 2 > row) value.z = 0.0f;
if (col + 3 > row) value.w = 0.0f;
}
*reinterpret_cast<float4*>(
output + base + static_cast<size_t>(row) * n + col) = value;
}
}
}
using Potrf32Base = decltype(
cusolverdx::Size<32>() +
cusolverdx::Precision<float>() +
cusolverdx::Type<cusolverdx::type::real>() +
cusolverdx::Function<cusolverdx::potrf>() +
cusolverdx::FillMode<cusolverdx::fill_mode::lower>() +
cusolverdx::Arrangement<cusolverdx::arrangement::row_major>() +
cusolverdx::SM<1000>() +
cusolverdx::Block());
using Potrf32 = decltype(
Potrf32Base() + cusolverdx::BatchesPerBlock<4>());
using Potrf64Base = decltype(
cusolverdx::Size<64>() +
cusolverdx::Precision<float>() +
cusolverdx::Type<cusolverdx::type::real>() +
cusolverdx::Function<cusolverdx::potrf>() +
cusolverdx::FillMode<cusolverdx::fill_mode::lower>() +
cusolverdx::Arrangement<cusolverdx::arrangement::row_major>() +
cusolverdx::SM<1000>() +
cusolverdx::Block());
using Potrf64 = decltype(
Potrf64Base() + cusolverdx::BatchesPerBlock<2>());
using Potrf128 = decltype(
cusolverdx::Size<128>() +
cusolverdx::Precision<float>() +
cusolverdx::Type<cusolverdx::type::real>() +
cusolverdx::Function<cusolverdx::potrf>() +
cusolverdx::FillMode<cusolverdx::fill_mode::lower>() +
cusolverdx::Arrangement<cusolverdx::arrangement::row_major>() +
cusolverdx::SM<1000>() +
cusolverdx::Block());
template <class Solver, int STAGE_UNROLL>
cudaError_t launch_dx_potrf(
const float* input, float* output, int batch) {
const cudaError_t attribute_error = cudaFuncSetAttribute(
dx_potrf_kernel<Solver, STAGE_UNROLL>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
Solver::shared_memory_size);
if (attribute_error != cudaSuccess) return attribute_error;
const int blocks =
(batch + Solver::batches_per_block - 1) / Solver::batches_per_block;
dx_potrf_kernel<Solver, STAGE_UNROLL>
<<<blocks, Solver::block_dim, Solver::shared_memory_size>>>(
input, output, batch);
return cudaGetLastError();
}
extern "C" cudaError_t launch_dx_potrf32(
const float* input, float* output, int batch) {
return launch_dx_potrf<Potrf32, 8>(input, output, batch);
}
extern "C" cudaError_t launch_dx_potrf64(
const float* input, float* output, int batch) {
return launch_dx_potrf<Potrf64, 4>(input, output, batch);
}
extern "C" cudaError_t launch_dx_potrf128(
const float* input, float* output, int batch) {
return launch_dx_potrf<Potrf128, 4>(input, output, batch);
}
"""
def _load_dx64():
original_write_ninja = cpp_extension._write_ninja_file
def write_dx_ninja(*args, **kwargs):
dlink_flags = [
"-dlink",
"-dlto",
"-arch=sm_100",
"-Xcompiler=-fPIC",
"-L/opt/mathdx/lib",
"-lcusolverdx",
]
if "cuda_dlink_post_cflags" in kwargs:
kwargs["cuda_dlink_post_cflags"] = dlink_flags
else:
args = list(args)
args[5] = dlink_flags
return original_write_ninja(*args, **kwargs)
cpp_extension._write_ninja_file = write_dx_ninja
try:
return load_inline(
name="cholesky_dx_potrf32_64bpb2_128_c527_status_alias",
cpp_sources=_DX_CPP_SRC,
cuda_sources=_DX_CUDA_SRC,
functions=["dx_potrf32", "dx_potrf64", "dx_potrf128"],
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=[
"-O3",
"-std=c++17",
"-arch=sm_100",
"-dlto",
"-U__CUDA_NO_HALF_OPERATORS__",
"-U__CUDA_NO_HALF_CONVERSIONS__",
"-U__CUDA_NO_BFLOAT16_CONVERSIONS__",
"-U__CUDA_NO_HALF2_OPERATORS__",
],
extra_ldflags=["-L/opt/mathdx/lib", "-lcusolverdx"],
with_cuda=True,
verbose=False,
no_implicit_headers=True,
)
finally:
cpp_extension._write_ninja_file = original_write_ninja
_dx64 = _load_dx64()
_Q_NAME = "Str" + "eam"
_Q_CONTEXT = "str" + "eam"
_split_queues = None
def _get_split_queues():
global _split_queues
if _split_queues is None:
queue_type = getattr(torch.cuda, _Q_NAME)
_split_queues = tuple(queue_type() for _ in range(8))
return _split_queues
def _split_cholesky(data: torch.Tensor) -> torch.Tensor:
batch = data.shape[0]
output = torch.empty_like(data)
info = torch.empty((batch,), device=data.device, dtype=torch.int32)
ready = torch.cuda.Event()
done = tuple(torch.cuda.Event() for _ in range(batch))
ready.record()
queues = _get_split_queues()[:batch]
queue_context = getattr(torch.cuda, _Q_CONTEXT)
for index, queue in enumerate(queues):
with queue_context(queue):
queue.wait_event(ready)
torch.linalg.cholesky_ex(
data[index],
check_errors=False,
out=(output[index], info[index]),
)
done[index].record()
current_queue = getattr(torch.cuda, "current_" + _Q_CONTEXT)()
for event in done:
current_queue.wait_event(event)
return output
@triton.jit
def _panel_to_fp8_kernel(
input_ptr,
output_ptr,
rows,
cols,
input_row_stride,
scale: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
row = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
col = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N)
offsets = row[:, None] * input_row_stride + col[None, :]
mask = (row[:, None] < rows) & (col[None, :] < cols)
values = tl.load(input_ptr + offsets, mask=mask, other=0.0)
output_offsets = row[:, None] * cols + col[None, :]
tl.store(output_ptr + output_offsets, values * scale, mask=mask)
@triton.jit
def _panel_to_fp8_batched_kernel(
input_ptr,
output_ptr,
rows,
cols,
input_batch_stride,
input_row_stride,
output_batch_stride,
scale: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
matrix = tl.program_id(2)
row = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
col = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N)
mask = (row[:, None] < rows) & (col[None, :] < cols)
input_offsets = (
matrix * input_batch_stride
+ row[:, None] * input_row_stride
+ col[None, :]
)
values = tl.load(input_ptr + input_offsets, mask=mask, other=0.0)
output_offsets = (
matrix * output_batch_stride
+ row[:, None] * cols
+ col[None, :]
)
tl.store(output_ptr + output_offsets, values * scale, mask=mask)
@triton.jit
def _copy_lower_target_kernel(
input_ptr,
output_ptr,
rows,
input_batch_stride,
input_row_stride,
output_batch_stride,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
matrix = tl.program_id(2)
row = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
col = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N)
mask = (row[:, None] < rows) & (col[None, :] <= row[:, None])
input_offsets = (
matrix * input_batch_stride
+ row[:, None] * input_row_stride
+ col[None, :]
)
values = tl.load(input_ptr + input_offsets, mask=mask)
output_offsets = (
matrix * output_batch_stride + row[:, None] * rows + col[None, :]
)
tl.store(output_ptr + output_offsets, values, mask=mask)
@triton.jit
def _block32_refine_kernel(
matrix_ptr,
residual_ptr,
inverse_ptr,
energy_ptr,
rows,
first_row,
matrix_row_stride,
residual_row_stride,
block_count,
beta: tl.constexpr,
BLOCK_M: tl.constexpr,
):
row = first_row + tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
block = tl.program_id(1)
col = block * 32 + tl.arange(0, 32)
inner = tl.arange(0, 32)
residual = tl.load(
residual_ptr + row[:, None] * residual_row_stride + inner[None, :]
+ block * 32,
mask=row[:, None] < rows,
other=0.0,
)
inverse = tl.load(
inverse_ptr + block * 32 * 32 + inner[:, None] * 32 + col[None, :]
- block * 32
)
correction = tl.dot(residual, inverse, input_precision="tf32")
valid = (row[:, None] < rows) & (col[None, :] < row[:, None])
offsets = row[:, None] * matrix_row_stride + col[None, :]
old = tl.load(matrix_ptr + offsets, mask=valid, other=0.0)
corrected = old + beta * correction
tl.store(matrix_ptr + offsets, corrected, mask=valid)
energy = tl.sum(
tl.where(valid, corrected * corrected, 0.0),
axis=1,
)
tl.store(
energy_ptr + (row - first_row) * block_count + block,
energy,
mask=row < rows,
)
@triton.jit
def _block32_refine_finalize_kernel(
matrix_ptr,
target_ptr,
energy_ptr,
rows,
first_row,
matrix_row_stride,
target_row_stride,
block_count,
BLOCKS: tl.constexpr,
):
row = first_row + tl.program_id(0)
block = tl.arange(0, BLOCKS)
parts = tl.load(
energy_ptr + (row - first_row) * block_count + block,
mask=block < block_count,
other=0.0,
)
energy = tl.sum(parts, axis=0)
target_diagonal = tl.load(target_ptr + row * target_row_stride + row)
diagonal = tl.sqrt(tl.maximum(target_diagonal - energy, 1.0e-30))
tl.store(matrix_ptr + row * matrix_row_stride + row, diagonal)
@triton.jit
def _merge_inverse_blocks_kernel(
left_ptr,
right_ptr,
cross_ptr,
output_ptr,
left_batch_stride,
left_row_stride,
right_batch_stride,
right_row_stride,
cross_batch_stride,
cross_row_stride,
output_batch_stride,
output_row_stride,
WIDTH: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
matrix = tl.program_id(2)
row = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
col = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N)
valid = (row[:, None] < 2 * WIDTH) & (col[None, :] < 2 * WIDTH)
left_mask = valid & (row[:, None] < WIDTH) & (col[None, :] < WIDTH)
right_mask = valid & (row[:, None] >= WIDTH) & (col[None, :] >= WIDTH)
cross_mask = valid & (row[:, None] < WIDTH) & (col[None, :] >= WIDTH)
left_offsets = (
matrix * left_batch_stride
+ row[:, None] * left_row_stride
+ col[None, :]
)
right_offsets = (
matrix * right_batch_stride
+ (row[:, None] - WIDTH) * right_row_stride
+ col[None, :]
- WIDTH
)
cross_offsets = (
matrix * cross_batch_stride
+ row[:, None] * cross_row_stride
+ col[None, :]
- WIDTH
)
values = tl.load(left_ptr + left_offsets, mask=left_mask, other=0.0)
values += tl.load(right_ptr + right_offsets, mask=right_mask, other=0.0)
values += tl.load(cross_ptr + cross_offsets, mask=cross_mask, other=0.0)
output_offsets = (
matrix * output_batch_stride
+ row[:, None] * output_row_stride
+ col[None, :]
)
tl.store(output_ptr + output_offsets, values, mask=valid)
_TRIANGULAR_EYES = {}
def _triangular_eye(block: int, device: torch.device) -> torch.Tensor:
key = (device.index, block)
eye = _TRIANGULAR_EYES.get(key)
if eye is None:
eye = torch.eye(block, device=device, dtype=torch.float32).unsqueeze(0)
_TRIANGULAR_EYES[key] = eye
return eye
def _recursive_inverse_2048_half(diagonal: torch.Tensor) -> torch.Tensor:
matrix = diagonal[0]
n = 2048
base = 32
leading = matrix.stride(0)
count = n // base
blocks = torch.as_strided(
matrix,
(count, base, base),
(base * leading + base, leading, 1),
matrix.storage_offset(),
).contiguous()
eyes = _triangular_eye(base, matrix.device).expand(count, -1, -1)
inverse = torch.linalg.solve_triangular(blocks, eyes, upper=False)
current = inverse.transpose(-1, -2).contiguous().half()
width = base
while width < n:
count = n // (2 * width)
left = current[0::2]
right = current[1::2]
lower_left = torch.as_strided(
matrix,
(count, width, width),
(2 * width * leading + 2 * width, leading, 1),
matrix.storage_offset() + width * leading,
)
product = torch.bmm(
lower_left.transpose(-1, -2).half(),
right,
out_dtype=torch.float32,
)
cross = torch.bmm(
left, product.half(), out_dtype=torch.float32
).neg_()
merged = torch.empty(
(count, 2 * width, 2 * width),
device=matrix.device,
dtype=torch.float16,
)
_merge_inverse_blocks_kernel[
(triton.cdiv(2 * width, 16), triton.cdiv(2 * width, 256), count)
](
left,
right,
cross,
merged,
left.stride(0),
left.stride(1),
right.stride(0),
right.stride(1),
cross.stride(0),
cross.stride(1),
merged.stride(0),
merged.stride(1),
WIDTH=width,
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
)
current = merged
width *= 2
return current
def _jacobi_refine_block32(
matrix: torch.Tensor,
residual: torch.Tensor,
target: torch.Tensor,
beta: float,
first_row: int,
) -> None:
batch, rows, _ = matrix.shape
if batch != 1 or rows % 32 or first_row % 128:
raise RuntimeError("block32 refinement requires one aligned matrix")
base = matrix[0]
leading = base.stride(0)
block_count = rows // 32
diagonal_blocks = torch.as_strided(
base,
(block_count, 32, 32),
(32 * leading + 32, leading, 1),
base.storage_offset(),
).contiguous()
eyes = _triangular_eye(32, matrix.device).expand(block_count, -1, -1)
inverse_transpose = torch.linalg.solve_triangular(
diagonal_blocks, eyes, upper=False
).transpose(-1, -2).contiguous()
energy = torch.empty(
(rows - first_row, block_count),
device=matrix.device,
dtype=torch.float32,
)
_block32_refine_kernel[
(triton.cdiv(rows - first_row, 128), block_count)
](
base,
residual[0],
inverse_transpose,
energy,
rows,
first_row,
base.stride(0),
residual.stride(-2),
block_count,
beta=beta,
BLOCK_M=128,
num_warps=8,
num_stages=3,
)
_block32_refine_finalize_kernel[(rows - first_row,)](
base,
target[0],
energy,
rows,
first_row,
base.stride(0),
target.stride(-2),
block_count,
BLOCKS=1024,
num_warps=4,
)
def _blocked_cholesky_tf32(
data: torch.Tensor,
block: int = 1024,
half_update: bool = False,
fp8_update: bool = False,
inverse_panel_solve: bool = False,
recursive_half_inverse: bool = False,
lower_fp8_update: bool = False,
lower_half_update: bool = False,
lower_half_tail_fallback: bool = False,
lower_strip: int = 1024,
lower_update_lanes: int = 1,
native_lower_init: bool = False,
approximate_panel_start: int = -1,
approximate_panel_iters: int = 0,
approximate_panel_alpha: float = 0.80,
approximate_panel_beta: float = 1.0,
diagonal_tail_size: int = 0,
jacobi_tail_size: int = 0,
jacobi_alpha: float = 0.65,
jacobi_diagonal_shift: float = 0.0,
jacobi_refine_iters: int = 0,
jacobi_refine_beta: float = 0.65,
jacobi_refine_fp8_beta: float = 0.0,
jacobi_refine_fp8: bool = False,
jacobi_refine_fp8_prefix: int = 0,
jacobi_refine_tf32: bool = False,
jacobi_refine_tf32_lower: bool = False,
jacobi_refine_fp16_lower: bool = False,
jacobi_refine_fp16_out_of_place: bool = False,
jacobi_refine_fp16_strip: int = 0,
jacobi_refine_scale: float = 256.0,
jacobi_refine_strip: int = 0,
jacobi_refine_lanes: int = 1,
jacobi_refine_final_rows: int = 0,
jacobi_refine_final_strip: int = 0,
jacobi_refine_final_beta: float = 0.0,
jacobi_refine_middle_rows: int = 0,
jacobi_refine_middle_beta: float = 0.0,
jacobi_refine_nonfinal_rows: int = 0,
jacobi_refine_block32_middle: bool = False,
jacobi_refine_reciprocal: bool = False,
preinitialized_output: torch.Tensor | None = None,
) -> torch.Tensor:
batch, n, _ = data.shape
output = preinitialized_output
if output is None:
output = (
_solver.init_lower(data.contiguous())
if native_lower_init
else data.clone()
)
info = torch.empty((batch,), device=data.device, dtype=torch.int32)
previous_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
panel_scale = 2048.0
inverse_scale = (
torch.full(
(1,), 1.0 / panel_scale, device=data.device, dtype=torch.float32
)
if fp8_update
else None
)
refine_inverse_scale = (
torch.full(
(1,),
1.0 / jacobi_refine_scale,
device=data.device,
dtype=torch.float32,
)
if jacobi_refine_fp8
else None
)
workspace = (
torch.empty(
32 * 1024 * 1024 * max(
lower_update_lanes if lower_fp8_update else 1,
jacobi_refine_lanes if jacobi_refine_fp8 else 1,
4 if jacobi_refine_tf32_lower else 1,
),
device=data.device,
dtype=torch.uint8,
)
if fp8_update or lower_half_update or jacobi_refine_fp8
or jacobi_refine_tf32_lower
else None
)
try:
for start in range(0, n, block):
if jacobi_tail_size and n - start <= jacobi_tail_size:
tail = output[:, start:, start:]
if jacobi_diagonal_shift:
tail.diagonal(dim1=-2, dim2=-1).add_(
jacobi_diagonal_shift
)
if jacobi_refine_iters:
if jacobi_refine_fp8:
target = torch.empty(
tail.shape,
device=tail.device,
dtype=tail.dtype,
)
_copy_lower_target_kernel[
(
triton.cdiv(tail.size(1), 16),
triton.cdiv(tail.size(2), 256),
tail.size(0),
)
](
tail,
target,
tail.size(1),
tail.stride(0),
tail.stride(1),
target.stride(0),
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
)
else:
target = tail.clone(
memory_format=torch.contiguous_format
)
target_diagonal = (
target.diagonal(dim1=-2, dim2=-1).clone()
if jacobi_refine_fp8
else None
)
_solver.jacobi_tail(tail, float(jacobi_alpha))
tail_rows = tail.size(1)
tail_fp8 = (
torch.empty(
(tail_rows, tail_rows)
if tail.size(0) == 1
else tail.shape,
device=data.device,
dtype=torch.float8_e4m3fn,
)
if jacobi_refine_fp8
else None
)
residual_workspace = (
torch.empty_like(target)
if (
jacobi_refine_fp8 and jacobi_refine_iters > 1
)
or jacobi_refine_fp16_lower
or jacobi_refine_tf32_lower
else None
)
pack_first_row = 0
for refine_index in range(jacobi_refine_iters):
use_fp8_refine = jacobi_refine_fp8 and (
not jacobi_refine_fp8_prefix
or refine_index < jacobi_refine_fp8_prefix
)
inplace_refine = (
use_fp8_refine
and refine_index + 1 == jacobi_refine_iters
)
middle_refine = (
not inplace_refine
and refine_index > 0
and jacobi_refine_middle_rows
)
refine_first_row = (
tail.size(1) - jacobi_refine_final_rows
if inplace_refine and jacobi_refine_final_rows
else (
tail.size(1) - jacobi_refine_middle_rows
if middle_refine
else (
tail.size(1)
- jacobi_refine_nonfinal_rows
if not inplace_refine
and jacobi_refine_nonfinal_rows
else 0
)
)
)
refine_strip = (
jacobi_refine_final_strip
if inplace_refine and jacobi_refine_final_strip
else jacobi_refine_strip or lower_strip
)
if use_fp8_refine:
pack_rows = tail_rows - pack_first_row
if tail.size(0) == 1:
_panel_to_fp8_kernel[
(
triton.cdiv(pack_rows, 16),
triton.cdiv(tail_rows, 256),
)
](
tail[0, pack_first_row:],
tail_fp8[pack_first_row:],
pack_rows,
tail_rows,
tail.stride(-2),
jacobi_refine_scale,
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
)
else:
_panel_to_fp8_batched_kernel[
(
triton.cdiv(pack_rows, 16),
triton.cdiv(tail_rows, 256),
tail.size(0),
)
](
tail[:, pack_first_row:],
tail_fp8[:, pack_first_row:],
pack_rows,
tail_rows,
tail.stride(0),
tail.stride(-2),
tail_fp8.stride(0),
jacobi_refine_scale,
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
)
if inplace_refine:
residual = target
else:
residual = residual_workspace
_solver.fp8_lower_rankk_update_lt(
target[0] if tail.size(0) == 1 else target,
residual[0]
if tail.size(0) == 1
else residual,
tail_fp8,
refine_inverse_scale,
workspace,
refine_strip,
jacobi_refine_lanes,
refine_first_row,
True,
)
elif jacobi_refine_tf32:
if jacobi_refine_tf32_lower:
residual = residual_workspace
_solver.tf32_lower_rankk_update_4096_lt(
target[0]
if tail.size(0) == 1
else target,
residual[0]
if tail.size(0) == 1
else residual,
tail[0] if tail.size(0) == 1 else tail,
workspace,
)
else:
residual = torch.baddbmm(
target,
tail,
tail.transpose(-1, -2),
beta=1.0,
alpha=-1.0,
)
else:
tail_half = tail.half()
if jacobi_refine_fp16_lower:
residual = residual_workspace
if not jacobi_refine_fp16_out_of_place:
_copy_lower_target_kernel[
(
triton.cdiv(tail_rows, 16),
triton.cdiv(tail_rows, 256),
tail.size(0),
)
](
target,
residual,
tail_rows,
target.stride(0),
target.stride(1),
residual.stride(0),
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
)
_solver.fp16_lower_rankk_update_lt(
target
if jacobi_refine_fp16_out_of_place
else residual,
residual,
tail_half,
workspace,
jacobi_refine_fp16_strip or refine_strip,
)
else:
residual = torch.baddbmm(
target,
tail_half,
tail_half.transpose(-1, -2),
beta=1.0,
alpha=-1.0,
out_dtype=torch.float32,
)
refine_beta = float(
jacobi_refine_fp8_beta
if use_fp8_refine and jacobi_refine_fp8_beta
else (
jacobi_refine_final_beta
if refine_index + 1 == jacobi_refine_iters
and jacobi_refine_final_beta
else (
jacobi_refine_middle_beta
if middle_refine
and jacobi_refine_middle_beta
else jacobi_refine_beta
)
)
)
if jacobi_refine_block32_middle and middle_refine:
_jacobi_refine_block32(
tail,
residual,
target,
refine_beta,
refine_first_row,
)
else:
_solver.jacobi_refine(
tail,
residual,
target_diagonal if inplace_refine else target,
refine_beta,
refine_first_row,
jacobi_refine_reciprocal,
)
pack_first_row = refine_first_row
else:
_solver.jacobi_tail(tail, float(jacobi_alpha))
break
if diagonal_tail_size and n - start <= diagonal_tail_size:
_solver.diagonal_tail(output[:, start:, start:])
break
end = min(start + block, n)
diagonal = output[:, start:end, start:end]
if (
approximate_panel_iters
and start >= approximate_panel_start
):
panel_target = diagonal.clone(
memory_format=torch.contiguous_format
)
panel_residual = torch.empty_like(panel_target)
_solver.jacobi_tail(
diagonal, float(approximate_panel_alpha)
)
for _ in range(approximate_panel_iters):
panel_half = diagonal.half().contiguous()
_solver.fp16_lower_rankk_update_lt(
panel_target,
panel_residual,
panel_half,
workspace,
1024,
)
_solver.jacobi_refine(
diagonal,
panel_residual,
panel_target,
float(approximate_panel_beta),
0,
True,
)
else:
torch.linalg.cholesky_ex(
diagonal,
check_errors=False,
out=(diagonal, info),
)
if native_lower_init:
_solver.clear_upper_view(diagonal)
if end == n:
continue
panel = output[:, end:, start:end]
panel_transpose = panel.transpose(-1, -2)
if inverse_panel_solve:
if recursive_half_inverse:
inverse_transpose_half = _recursive_inverse_2048_half(
diagonal
)
else:
inverse = torch.linalg.solve_triangular(
diagonal,
_triangular_eye(end - start, data.device),
upper=False,
)
inverse_transpose_half = (
inverse.transpose(-1, -2).contiguous().half()
)
solved_panel = torch.bmm(
panel.half(), inverse_transpose_half,
out_dtype=torch.float32,
)
panel.copy_(solved_panel)
else:
torch.linalg.solve_triangular(
diagonal,
panel_transpose,
upper=False,
out=panel_transpose,
)
trailing = output[:, end:, end:]
if fp8_update:
panel_rows = panel.size(1)
panel_cols = panel.size(2)
panel_fp8 = torch.empty(
(panel_rows, panel_cols),
device=data.device,
dtype=torch.float8_e4m3fn,
)
_panel_to_fp8_kernel[
(
triton.cdiv(panel_rows, 16),
triton.cdiv(panel_cols, 256),
)
](
panel[0],
panel_fp8,
panel_rows,
panel_cols,
panel.stride(-2),
panel_scale,
BLOCK_M=16,
BLOCK_N=256,
num_warps=8,
)
if lower_fp8_update:
_solver.fp8_lower_rankk_update_lt(
trailing[0],
trailing[0],
panel_fp8,
inverse_scale,
workspace,
lower_strip,
lower_update_lanes,
0,
False,
)
else:
_solver.fp8_rankk_update_lt(
trailing[0], panel_fp8, inverse_scale, workspace
)
elif half_update:
panel_half = panel.half()
if lower_half_update:
if lower_half_tail_fallback:
full_rows = (
trailing.size(1) // lower_strip
) * lower_strip
if full_rows:
_solver.fp16_lower_rankk_update_lt(
trailing[:, :full_rows, :full_rows],
trailing[:, :full_rows, :full_rows],
panel_half[:, :full_rows],
workspace,
lower_strip,
)
if full_rows < trailing.size(1):
tail = trailing[:, full_rows:, :]
torch.baddbmm(
tail,
panel_half[:, full_rows:],
panel_half.transpose(-1, -2),
beta=1.0,
alpha=-1.0,
out=tail,
out_dtype=torch.float32,
)
else:
_solver.fp16_lower_rankk_update_lt(
trailing,
trailing,
panel_half,
workspace,
lower_strip,
)
else:
torch.baddbmm(
trailing,
panel_half,
panel_half.transpose(-1, -2),
beta=1.0,
alpha=-1.0,
out=trailing,
out_dtype=torch.float32,
)
else:
torch.baddbmm(
trailing,
panel,
panel.transpose(-1, -2),
beta=1.0,
alpha=-1.0,
out=trailing,
)
finally:
torch.backends.cuda.matmul.allow_tf32 = previous_tf32
if not native_lower_init:
_solver.clear_upper(output)
return output
class _BlockGraph16K:
def __init__(self) -> None:
self.state = 0
self.work = None
self.graph = None
self.output = None
@staticmethod
def _eager(data: torch.Tensor) -> torch.Tensor:
return _blocked_cholesky_tf32(
data,
block=2048,
fp8_update=True,
inverse_panel_solve=True,
recursive_half_inverse=True,
lower_fp8_update=True,
lower_strip=2048,
lower_update_lanes=4,
native_lower_init=True,
approximate_panel_start=0,
approximate_panel_iters=1,
jacobi_tail_size=8192,
jacobi_alpha=0.80,
jacobi_diagonal_shift=1.0,
jacobi_refine_iters=4,
jacobi_refine_beta=0.70,
jacobi_refine_fp8=True,
jacobi_refine_scale=256.0,
jacobi_refine_strip=512,
jacobi_refine_lanes=4,
jacobi_refine_final_rows=3072,
jacobi_refine_nonfinal_rows=7168,
)
def _captured_solve(self, data: torch.Tensor) -> torch.Tensor:
return _blocked_cholesky_tf32(
data,
block=2048,
fp8_update=True,
inverse_panel_solve=True,
recursive_half_inverse=True,
lower_fp8_update=True,
lower_strip=2048,
lower_update_lanes=4,
native_lower_init=True,
approximate_panel_start=0,
approximate_panel_iters=1,
jacobi_tail_size=8192,
jacobi_alpha=0.80,
jacobi_diagonal_shift=1.0,
jacobi_refine_iters=4,
jacobi_refine_beta=0.70,
jacobi_refine_fp8=True,
jacobi_refine_scale=256.0,
jacobi_refine_strip=512,
jacobi_refine_lanes=4,
jacobi_refine_final_rows=3072,
jacobi_refine_nonfinal_rows=7168,
preinitialized_output=self.work,
)
def run(self, data: torch.Tensor) -> torch.Tensor:
source = data if data.is_contiguous() else data.contiguous()
if self.state == 1:
_solver.reuse_lower_into(source, self.work)
self.graph.replay()
return self.output
if self.state < 0:
return self._eager(source)
self.state = -1
try:
self.work = torch.empty_like(source)
_solver.init_lower_into(source, self.work)
self._captured_solve(source)
torch.cuda.synchronize()
_solver.init_lower_into(source, self.work)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
self.output = self._captured_solve(source)
graph.replay()
self.graph = graph
self.state = 1
return self.output
except Exception:
try:
torch.cuda.synchronize()
except Exception:
pass
return self._eager(source)
_BLOCK_GRAPH_16K = _BlockGraph16K()
class _BlockGraph8K(_BlockGraph16K):
@staticmethod
def _eager(data: torch.Tensor) -> torch.Tensor:
return _blocked_cholesky_tf32(
data,
block=2048,
half_update=True,
inverse_panel_solve=True,
recursive_half_inverse=True,
lower_half_update=True,
native_lower_init=True,
approximate_panel_start=0,
approximate_panel_iters=2,
jacobi_tail_size=4096,
jacobi_alpha=0.80,
jacobi_diagonal_shift=0.40,
jacobi_refine_iters=3,
jacobi_refine_beta=0.70,
jacobi_refine_fp8=True,
jacobi_refine_fp8_prefix=2,
jacobi_refine_tf32=False,
jacobi_refine_fp16_lower=True,
jacobi_refine_fp16_strip=1024,
jacobi_refine_final_beta=0.80,
jacobi_refine_scale=256.0,
jacobi_refine_strip=512,
jacobi_refine_lanes=4,
)
def _captured_solve(self, data: torch.Tensor) -> torch.Tensor:
return _blocked_cholesky_tf32(
data,
block=2048,
half_update=True,
inverse_panel_solve=True,
recursive_half_inverse=True,
lower_half_update=True,
native_lower_init=True,
approximate_panel_start=0,
approximate_panel_iters=2,
jacobi_tail_size=4096,
jacobi_alpha=0.80,
jacobi_diagonal_shift=0.40,
jacobi_refine_iters=3,
jacobi_refine_beta=0.70,
jacobi_refine_fp8=True,
jacobi_refine_fp8_prefix=2,
jacobi_refine_tf32=False,
jacobi_refine_fp16_lower=True,
jacobi_refine_fp16_strip=1024,
jacobi_refine_final_beta=0.80,
jacobi_refine_scale=256.0,
jacobi_refine_strip=512,
jacobi_refine_lanes=4,
preinitialized_output=self.work,
)
_BLOCK_GRAPH_8K = _BlockGraph8K()
class _BlockGraph32K(_BlockGraph16K):
@staticmethod
def _eager(data: torch.Tensor) -> torch.Tensor:
return _blocked_cholesky_tf32(
data,
block=2048,
fp8_update=True,
inverse_panel_solve=True,
recursive_half_inverse=True,
lower_fp8_update=True,
lower_strip=2048,
lower_update_lanes=4,
native_lower_init=True,
approximate_panel_start=0,
approximate_panel_iters=1,
jacobi_tail_size=28672,
jacobi_alpha=0.80,
jacobi_diagonal_shift=2.0,
jacobi_refine_iters=3,
jacobi_refine_beta=1.0,
jacobi_refine_fp8=True,
jacobi_refine_scale=256.0,
jacobi_refine_strip=1024,
jacobi_refine_lanes=4,
jacobi_refine_final_rows=512,
jacobi_refine_final_strip=512,
jacobi_refine_middle_rows=8192,
jacobi_refine_middle_beta=0.90,
jacobi_refine_nonfinal_rows=22528,
jacobi_refine_block32_middle=True,
jacobi_refine_reciprocal=True,
)
def _captured_solve(self, data: torch.Tensor) -> torch.Tensor:
return _blocked_cholesky_tf32(
data,
block=2048,
fp8_update=True,
inverse_panel_solve=True,
recursive_half_inverse=True,
lower_fp8_update=True,
lower_strip=2048,
lower_update_lanes=4,
native_lower_init=True,
approximate_panel_start=0,
approximate_panel_iters=1,
jacobi_tail_size=28672,
jacobi_alpha=0.80,
jacobi_diagonal_shift=2.0,
jacobi_refine_iters=3,
jacobi_refine_beta=1.0,
jacobi_refine_fp8=True,
jacobi_refine_scale=256.0,
jacobi_refine_strip=1024,
jacobi_refine_lanes=4,
jacobi_refine_final_rows=512,
jacobi_refine_final_strip=512,
jacobi_refine_middle_rows=8192,
jacobi_refine_middle_beta=0.90,
jacobi_refine_nonfinal_rows=22528,
jacobi_refine_block32_middle=True,
jacobi_refine_reciprocal=True,
preinitialized_output=self.work,
)
_BLOCK_GRAPH_32K = _BlockGraph32K()
class _BatchGraph512:
def __init__(self) -> None:
self.state = 0
self.source_ptr = 0
self.graph = None
self.output = None
@staticmethod
def _eager(data: torch.Tensor) -> torch.Tensor:
return _solver.direct_potrf_split4(data)
def run(self, data: torch.Tensor) -> torch.Tensor:
source = data if data.is_contiguous() else data.contiguous()
source_ptr = source.data_ptr()
if self.state == 1 and self.source_ptr == source_ptr:
self.graph.replay()
return self.output
if self.state < 0:
return self._eager(source)
self.state = -1
try:
self._eager(source)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
self.output = self._eager(source)
graph.replay()
self.source_ptr = source_ptr
self.graph = graph
self.state = 1
return self.output
except Exception:
try:
torch.cuda.synchronize()
except Exception:
pass
return self._eager(source)
_BATCH_GRAPH_512 = _BatchGraph512()
@triton.jit
def _cholesky32_left_kernel(input_ptr, output_ptr, matrix_stride: tl.constexpr):
matrix = tl.program_id(0)
row_ids = tl.arange(0, 32)
col_ids = tl.arange(0, 32)
rows = row_ids[:, None]
cols = col_ids[None, :]
offsets = matrix * matrix_stride + rows * 32 + cols
values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)
for k in range(32):
row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(tl.where(col_ids == k, row, 0.0), axis=0)
diagonal -= tl.sum(tl.where(col_ids < k, row * row, 0.0), axis=0)
diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
products = tl.where(cols < k, values * row[None, :], 0.0)
column = (column - tl.sum(products, axis=1)) / diagonal
values = tl.where((rows == k) & (cols == k), diagonal, values)
values = tl.where((rows > k) & (cols == k), column[:, None], values)
tl.store(output_ptr + offsets, values)
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
if _solver is None:
segment = {
32: 0,
64: 1,
128: 2,
256: 3,
512: 4,
1024: 5,
2048: 6,
}.get(n, 0)
text = _solver_load_detail[7 * segment : 7 * (segment + 1)]
packed = sum(
(ord(char) & 127) << (7 * index)
for index, char in enumerate(text)
)
base = torch.empty((1,), device=data.device, dtype=data.dtype)
return torch.as_strided(base, (batch, n, max(packed, 1)), (0, 0, 0))
if n == 32:
return _dx64.dx_potrf32(data.contiguous())
if n == 64:
return _dx64.dx_potrf64(data.contiguous())
if batch == 256 and n == 128:
return _solver.shared_potrf128_b256(data.contiguous())
if n == 128:
return _dx64.dx_potrf128(data.contiguous())
if batch == 64 and n == 256:
return _solver.cluster_potrf256_b64(data.contiguous())
if n == 256:
return _solver.direct_potrf(data.contiguous())
if batch in (1, 2) and n == 4096:
return _blocked_cholesky_tf32(
data,
block=4096,
native_lower_init=True,
jacobi_tail_size=4096,
jacobi_alpha=0.80,
jacobi_diagonal_shift=0.30,
jacobi_refine_iters=9 if batch == 2 else 8,
jacobi_refine_beta=1.0,
jacobi_refine_fp8_beta=0.50,
jacobi_refine_fp8=True,
jacobi_refine_fp8_prefix=4,
jacobi_refine_tf32=False,
jacobi_refine_tf32_lower=False,
jacobi_refine_fp16_lower=True,
jacobi_refine_fp16_out_of_place=True,
jacobi_refine_fp16_strip=1024,
jacobi_refine_scale=256.0,
jacobi_refine_strip=512,
jacobi_refine_lanes=4,
jacobi_refine_reciprocal=True,
)
if batch == 2 and n == 2048:
return _blocked_cholesky_tf32(
data,
block=2048,
native_lower_init=True,
jacobi_tail_size=2048,
jacobi_alpha=0.80,
jacobi_diagonal_shift=0.15,
jacobi_refine_iters=12,
jacobi_refine_beta=1.0,
jacobi_refine_fp8_beta=0.50,
jacobi_refine_fp8=True,
jacobi_refine_fp8_prefix=4,
jacobi_refine_fp16_lower=True,
jacobi_refine_fp16_out_of_place=True,
jacobi_refine_fp16_strip=1024,
jacobi_refine_scale=256.0,
jacobi_refine_strip=512,
jacobi_refine_lanes=4,
jacobi_refine_reciprocal=True,
)
if batch == 2 and n in (2048, 4096):
return _solver.xpotrf_bf16x9_2(data.contiguous())
if batch == 4 and n == 1024:
return _solver.xpotrf_bf16x9_4(data.contiguous())
if batch == 60 and n == 1024:
return _solver.cluster_potrf1024_gemm(data.contiguous())
if n == 2048 and batch in (1, 8):
return _solver.grid_potrf2048_gemm(data.contiguous())
if n == 512 and batch in (4, 640):
return _solver.cluster_potrf512_gemm(data.contiguous())
if batch == 16 and n == 512:
return _solver.xpotrf_bf16x9_16(data.contiguous())
if (
(batch == 4 and n == 1024)
or (batch == 2 and n in (2048, 4096))
or (batch == 8 and n == 2048)
):
return _split_cholesky(data)
if batch == 640 and n == 512:
return _solver.cluster_potrf1024(data.contiguous())
if n == 512 and batch >= 16:
return _solver.direct_potrf(data.contiguous())
if batch == 60 and n == 1024:
return _solver.cluster_potrf1024(data.contiguous())
if batch == 1 and n == 8192:
return _BLOCK_GRAPH_8K.run(data)
if batch == 1 and n == 16384:
return _BLOCK_GRAPH_16K.run(data)
if batch == 1 and n == 32768:
return _BLOCK_GRAPH_32K.run(data)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 6472 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