submission 898138
Praneeth Veligeti · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 833 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-898138?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:b555a464c5181cd52fca9e588461061a6265b32b4c33f7e64d52bfe40dcf4d22
license declaredunknown
license concludedunknown
authorsPraneeth Veligeti
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(shared-memory
__shared__ float pivot[128];Kernel source
submission.py833 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True
def _blocked_tf32(matrix: torch.Tensor) -> torch.Tensor:
n = matrix.shape[-1]
if n <= 256:
return torch.linalg.cholesky_ex(matrix, check_errors=False).L
half = n // 2
a11 = matrix[..., :half, :half]
a21 = matrix[..., half:, :half]
a22 = matrix[..., half:, half:]
l11 = _blocked_tf32(a11)
l21 = torch.linalg.solve_triangular(l11.mT, a21, upper=True, left=False)
schur = torch.baddbmm(a22, l21, l21.mT, beta=1.0, alpha=-1.0)
l22 = _blocked_tf32(schur)
output = torch.zeros_like(matrix)
output[..., :half, :half] = l11
output[..., half:, :half] = l21
output[..., half:, half:] = l22
return output
def _blocked_large(matrix: torch.Tensor, base: int) -> torch.Tensor:
"""Recursive tensor-core Schur updates with cuSOLVER base panels."""
if matrix.shape[-1] <= base:
return torch.linalg.cholesky_ex(matrix, check_errors=False).L
half = matrix.shape[-1] // 2
a11 = matrix[..., :half, :half]
a21 = matrix[..., half:, :half]
a22 = matrix[..., half:, half:]
l11 = _blocked_large(a11, base)
l21 = torch.linalg.solve_triangular(l11.mT, a21, upper=True, left=False)
schur = torch.baddbmm(a22, l21, l21.mT, beta=1.0, alpha=-1.0)
l22 = _blocked_large(schur, base)
output = torch.empty_like(matrix)
_assemble_large[(triton.cdiv(output.numel(), 512),)](
l11, l21, l22, output,
l11.stride(0), l11.stride(1), l11.stride(2),
l21.stride(0), l21.stride(1), l21.stride(2),
l22.stride(0), l22.stride(1), l22.stride(2),
N=matrix.shape[-1], H=half, TOTAL=output.numel(), BLOCK=512)
return output
@triton.jit
def _assemble_large(l11, l21, l22, output,
l11_s0, l11_s1, l11_s2,
l21_s0, l21_s1, l21_s2,
l22_s0, l22_s1, l22_s2,
N: tl.constexpr, H: tl.constexpr,
TOTAL: tl.constexpr, BLOCK: tl.constexpr):
index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
valid = index < TOTAL
matrix_elements = N * N
matrix = index // matrix_elements
local = index - matrix * matrix_elements
row = local // N
column = local - row * N
top_left = (row < H) & (column < H) & valid
bottom_left = (row >= H) & (column < H) & valid
bottom_right = (row >= H) & (column >= H) & valid
value = tl.zeros((BLOCK,), tl.float32)
value += tl.load(
l11 + matrix * l11_s0 + row * l11_s1 + column * l11_s2,
mask=top_left, other=0.0)
value += tl.load(
l21 + matrix * l21_s0 + (row - H) * l21_s1 + column * l21_s2,
mask=bottom_left, other=0.0)
value += tl.load(
l22 + matrix * l22_s0 + (row - H) * l22_s1
+ (column - H) * l22_s2,
mask=bottom_right, other=0.0)
tl.store(output + index, value, mask=valid)
@triton.jit
def _assemble_1024(l11, l21, l22, output,
l11_s0, l11_s1, l11_s2,
l21_s0, l21_s1, l21_s2,
l22_s0, l22_s1, l22_s2,
TOTAL: tl.constexpr, BLOCK: tl.constexpr):
index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
valid = index < TOTAL
matrix = index // (1024 * 1024)
local = index - matrix * (1024 * 1024)
row = local // 1024
column = local - row * 1024
top_left = (row < 512) & (column < 512) & valid
bottom_left = (row >= 512) & (column < 512) & valid
bottom_right = (row >= 512) & (column >= 512) & valid
value = tl.zeros((BLOCK,), tl.float32)
value += tl.load(
l11 + matrix * l11_s0 + row * l11_s1 + column * l11_s2,
mask=top_left, other=0.0)
value += tl.load(
l21 + matrix * l21_s0 + (row - 512) * l21_s1 + column * l21_s2,
mask=bottom_left, other=0.0)
value += tl.load(
l22 + matrix * l22_s0 + (row - 512) * l22_s1
+ (column - 512) * l22_s2,
mask=bottom_right, other=0.0)
tl.store(output + index, value, mask=valid)
def _blocked_1024(data: torch.Tensor) -> torch.Tensor:
a11 = data[..., :512, :512]
a21 = data[..., 512:, :512]
a22 = data[..., 512:, 512:]
l11 = torch.linalg.cholesky_ex(a11, check_errors=False).L
l21 = torch.linalg.solve_triangular(
l11.mT, a21, upper=True, left=False
)
schur = a22 - l21 @ l21.mT
l22 = torch.linalg.cholesky_ex(schur, check_errors=False).L
output = torch.empty_like(data)
_assemble_1024[(triton.cdiv(output.numel(), 256),)](
l11, l21, l22, output,
l11.stride(0), l11.stride(1), l11.stride(2),
l21.stride(0), l21.stride(1), l21.stride(2),
l22.stride(0), l22.stride(1), l22.stride(2),
TOTAL=output.numel(), BLOCK=256)
return output
CPP_SRC = r"""
torch::Tensor emulated_potrf(torch::Tensor input);
torch::Tensor blocked_potrf(torch::Tensor input, int64_t block);
torch::Tensor register_row_potrf64(torch::Tensor input);
torch::Tensor register_row_potrf128(torch::Tensor input);
torch::Tensor blocked256_register(torch::Tensor input);
torch::Tensor blocked512_register(torch::Tensor input);
torch::Tensor blocked1024_register(torch::Tensor input);
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <unordered_map>
#include <vector>
namespace {
cusolverDnHandle_t solver_handle() {
static cusolverDnHandle_t handle = [] {
cusolverDnHandle_t value;
TORCH_CHECK(cusolverDnCreate(&value) == CUSOLVER_STATUS_SUCCESS,
"cusolverDnCreate failed");
TORCH_CHECK(
cusolverDnSetMathMode(value, CUSOLVER_FP32_EMULATED_BF16X9_MATH)
== CUSOLVER_STATUS_SUCCESS,
"cusolverDnSetMathMode failed");
return value;
}();
return handle;
}
cusolverDnParams_t solver_params() {
static cusolverDnParams_t params = [] {
cusolverDnParams_t value;
TORCH_CHECK(cusolverDnCreateParams(&value) == CUSOLVER_STATUS_SUCCESS,
"cusolverDnCreateParams failed");
return value;
}();
return params;
}
cusolverDnHandle_t blocked_solver_handle() {
static cusolverDnHandle_t handle = [] {
cusolverDnHandle_t value;
TORCH_CHECK(cusolverDnCreate(&value) == CUSOLVER_STATUS_SUCCESS,
"blocked cusolverDnCreate failed");
return value;
}();
return handle;
}
cublasHandle_t blocked_blas_handle() {
static cublasHandle_t handle = [] {
cublasHandle_t value;
TORCH_CHECK(cublasCreate(&value) == CUBLAS_STATUS_SUCCESS,
"cublasCreate failed");
TORCH_CHECK(cublasSetMathMode(value, CUBLAS_TF32_TENSOR_OP_MATH)
== CUBLAS_STATUS_SUCCESS,
"cublasSetMathMode failed");
return value;
}();
return handle;
}
cublasHandle_t precise_blas_handle() {
static cublasHandle_t handle = [] {
cublasHandle_t value;
TORCH_CHECK(cublasCreate(&value) == CUBLAS_STATUS_SUCCESS,
"precise cublasCreate failed");
TORCH_CHECK(cublasSetMathMode(value, CUBLAS_DEFAULT_MATH)
== CUBLAS_STATUS_SUCCESS,
"precise cublasSetMathMode failed");
return value;
}();
return handle;
}
__global__ void make_block_pointers(
float* matrices, float** diagonal, float** panel,
int batch, int64_t matrix_stride, int64_t diagonal_offset,
int64_t panel_offset) {
int matrix = blockIdx.x * blockDim.x + threadIdx.x;
if (matrix < batch) {
float* base = matrices + int64_t(matrix) * matrix_stride;
diagonal[matrix] = base + diagonal_offset;
panel[matrix] = base + panel_offset;
}
}
__global__ void clear_batched_row_major_upper(
float* matrices, int64_t total, int n, int64_t matrix_stride) {
int64_t index = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
if (index < total) {
int64_t local = index % matrix_stride;
int row = local / n;
int column = local - int64_t(row) * n;
if (column > row) matrices[index] = 0.0f;
}
}
__global__ void clear_batched_row_major_upper_2d(
float* matrices, int n, int64_t matrix_stride) {
int column = blockIdx.x * blockDim.x + threadIdx.x;
int row = blockIdx.y * blockDim.y + threadIdx.y;
int matrix = blockIdx.z;
if (row < n && column < n && column > row) {
matrices[int64_t(matrix) * matrix_stride + int64_t(row) * n + column]
= 0.0f;
}
}
__global__ void clear_upper(float* matrix, int64_t n) {
int64_t index = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t count = n * n;
if (index < count) {
int64_t row = index / n;
int64_t col = index - row * n;
if (col > row) matrix[index] = 0.0f;
}
}
template <int K>
struct RegisterRowStep64 {
__device__ __forceinline__ static void run(
float (&row0)[64], float (&row1)[64], int lane) {
constexpr unsigned mask = 0xffffffffu;
constexpr int owner = K & 31;
float diagonal = 0.0f;
if (lane == owner) {
diagonal = K < 32 ? row0[K] : row1[K];
#pragma unroll
for (int j = 0; j < K; ++j) {
float value = K < 32 ? row0[j] : row1[j];
diagonal -= value * value;
}
}
diagonal = __shfl_sync(mask, diagonal, owner);
float inverse = rsqrtf(diagonal);
float sum0 = row0[K];
float sum1 = row1[K];
#pragma unroll
for (int j = 0; j < K; ++j) {
float owned = K < 32 ? row0[j] : row1[j];
float pivot = __shfl_sync(mask, owned, owner);
sum0 -= row0[j] * pivot;
sum1 -= row1[j] * pivot;
}
if (lane == K) row0[K] = sqrtf(diagonal);
else if (lane > K) row0[K] = sum0 * inverse;
if (lane + 32 == K) row1[K] = sqrtf(diagonal);
else if (lane + 32 > K) row1[K] = sum1 * inverse;
RegisterRowStep64<K + 1>::run(row0, row1, lane);
}
};
template <>
struct RegisterRowStep64<64> {
__device__ __forceinline__ static void run(
float (&)[64], float (&)[64], int) {}
};
__global__ __launch_bounds__(32) void register_row_kernel64(
const float* __restrict__ input, float* __restrict__ output) {
int matrix = blockIdx.x;
int lane = threadIdx.x;
constexpr int64_t stride = 64 * 64;
const float* source = input + int64_t(matrix) * stride;
float* destination = output + int64_t(matrix) * stride;
float row0[64];
float row1[64];
#pragma unroll
for (int column = 0; column < 64; ++column) {
row0[column] = column <= lane ? source[lane * 64 + column] : 0.0f;
row1[column] = column <= lane + 32
? source[(lane + 32) * 64 + column] : 0.0f;
}
RegisterRowStep64<0>::run(row0, row1, lane);
#pragma unroll
for (int column = 0; column < 64; ++column) {
destination[lane * 64 + column] =
column <= lane ? row0[column] : 0.0f;
destination[(lane + 32) * 64 + column] =
column <= lane + 32 ? row1[column] : 0.0f;
}
}
template <int K>
struct RegisterRowStep128 {
__device__ __forceinline__ static void run(
float (&row)[128], float* pivot, int matrix_row) {
if (matrix_row == K) {
float diagonal = row[K];
#pragma unroll
for (int j = 0; j < K; ++j) {
diagonal -= row[j] * row[j];
pivot[j] = row[j];
}
row[K] = sqrtf(diagonal);
pivot[K] = row[K];
}
__syncthreads();
if (matrix_row > K) {
float value = row[K];
#pragma unroll
for (int j = 0; j < K; ++j) {
value -= row[j] * pivot[j];
}
row[K] = value / pivot[K];
}
__syncthreads();
RegisterRowStep128<K + 1>::run(row, pivot, matrix_row);
}
};
template <>
struct RegisterRowStep128<128> {
__device__ __forceinline__ static void run(
float (&)[128], float*, int) {}
};
__global__ __launch_bounds__(128, 1) void register_row_kernel128(
const float* __restrict__ input, float* __restrict__ output) {
int matrix = blockIdx.x;
int matrix_row = threadIdx.x;
constexpr int64_t stride = 128 * 128;
const float* source = input + int64_t(matrix) * stride;
float* destination = output + int64_t(matrix) * stride;
__shared__ float pivot[128];
float row[128];
#pragma unroll
for (int column = 0; column < 128; ++column) {
row[column] = column <= matrix_row
? source[matrix_row * 128 + column] : 0.0f;
}
RegisterRowStep128<0>::run(row, pivot, matrix_row);
#pragma unroll
for (int column = 0; column < 128; ++column) {
destination[matrix_row * 128 + column] =
column <= matrix_row ? row[column] : 0.0f;
}
}
__global__ __launch_bounds__(128, 1) void factor128_inplace(
float* matrices, int leading, int64_t stride, int64_t offset) {
int matrix_row = threadIdx.x;
float* tile = matrices + int64_t(blockIdx.x) * stride + offset;
__shared__ float pivot[128];
float row[128];
#pragma unroll
for (int column = 0; column < 128; ++column) {
row[column] = column <= matrix_row
? tile[matrix_row * leading + column] : 0.0f;
}
RegisterRowStep128<0>::run(row, pivot, matrix_row);
#pragma unroll
for (int column = 0; column < 128; ++column) {
if (column <= matrix_row) {
tile[matrix_row * leading + column] = row[column];
}
}
}
} // namespace
torch::Tensor emulated_potrf(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be FP32");
TORCH_CHECK(input.dim() == 3 && input.size(0) == 1, "batch must be one");
auto output = input.clone();
const int64_t n = output.size(1);
auto handle = solver_handle();
auto params = solver_params();
size_t device_bytes = 0;
size_t host_bytes = 0;
auto status = cusolverDnXpotrf_bufferSize(
handle, params, CUBLAS_FILL_MODE_UPPER, n,
CUDA_R_32F, output.data_ptr<float>(), n, CUDA_R_32F,
&device_bytes, &host_bytes);
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, "bufferSize failed: ", int(status));
static std::unordered_map<int64_t, torch::Tensor> workspaces;
static std::unordered_map<int64_t, torch::Tensor> infos;
static std::unordered_map<int64_t, std::vector<unsigned char>> host_workspaces;
auto byte_options = torch::TensorOptions().device(input.device()).dtype(torch::kUInt8);
if (!workspaces.count(n) || workspaces.at(n).numel() < static_cast<int64_t>(device_bytes)) {
workspaces[n] = torch::empty({static_cast<int64_t>(device_bytes)}, byte_options);
infos[n] = torch::empty(
{1}, torch::TensorOptions().device(input.device()).dtype(torch::kInt32));
host_workspaces[n] = std::vector<unsigned char>(host_bytes);
}
auto& workspace = workspaces.at(n);
auto& info = infos.at(n);
auto& host_workspace = host_workspaces.at(n);
status = cusolverDnXpotrf(
handle, params, CUBLAS_FILL_MODE_UPPER, n,
CUDA_R_32F, output.data_ptr<float>(), n, CUDA_R_32F,
workspace.data_ptr(), device_bytes,
host_workspace.data(), host_bytes,
info.data_ptr<int>());
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, "Xpotrf failed: ", int(status));
int threads = 256;
int64_t count = n * n;
clear_upper<<<(count + threads - 1) / threads, threads>>>(output.data_ptr<float>(), n);
return output;
}
torch::Tensor blocked_potrf(torch::Tensor input, int64_t block_value) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.dim() == 3 && input.size(1) == input.size(2),
"input must be batched square matrices");
const int batch = input.size(0);
const int n = input.size(1);
const int block = static_cast<int>(block_value);
TORCH_CHECK(n % block == 0, "block must divide n");
auto output = input.clone();
auto pointers = torch::empty(
{2, batch}, torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
auto info = torch::empty(
{batch}, torch::TensorOptions().device(input.device()).dtype(torch::kInt32));
auto diagonal = reinterpret_cast<float**>(pointers[0].data_ptr<int64_t>());
auto panel = reinterpret_cast<float**>(pointers[1].data_ptr<int64_t>());
const int64_t stride = int64_t(n) * n;
const float one = 1.0f;
const float minus_one = -1.0f;
for (int k = 0; k < n; k += block) {
const int trailing = n - k - block;
const int64_t diagonal_offset = int64_t(k) * n + k;
const int64_t panel_offset = trailing > 0
? int64_t(k + block) * n + k : diagonal_offset;
make_block_pointers<<<(batch + 127) / 128, 128>>>(
output.data_ptr<float>(), diagonal, panel, batch, stride,
diagonal_offset, panel_offset);
auto solver_status = cusolverDnSpotrfBatched(
blocked_solver_handle(), CUBLAS_FILL_MODE_UPPER, block,
diagonal, n, info.data_ptr<int>(), batch);
TORCH_CHECK(solver_status == CUSOLVER_STATUS_SUCCESS,
"batched panel potrf failed: ", int(solver_status));
if (trailing > 0) {
auto trsm_status = cublasStrsmBatched(
blocked_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
block, trailing, &one,
const_cast<const float**>(diagonal), n, panel, n, batch);
TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
"batched TRSM failed: ", int(trsm_status));
float* panel_base = output.data_ptr<float>() + panel_offset;
float* trailing_base = output.data_ptr<float>()
+ int64_t(k + block) * n + (k + block);
auto gemm_status = cublasSgemmStridedBatched(
blocked_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
trailing, trailing, block, &minus_one,
panel_base, n, stride, panel_base, n, stride,
&one, trailing_base, n, stride, batch);
TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
"batched update GEMM failed: ", int(gemm_status));
}
}
const int64_t total = int64_t(batch) * stride;
dim3 clear_threads(32, 8);
dim3 clear_grid((n + clear_threads.x - 1) / clear_threads.x,
(n + clear_threads.y - 1) / clear_threads.y,
batch);
clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
output.data_ptr<float>(), n, stride);
auto error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess,
"blocked launch failed: ", cudaGetErrorString(error));
return output;
}
torch::Tensor register_row_potrf64(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.dim() == 3 && input.size(1) == 64 && input.size(2) == 64,
"register-row kernel requires 64x64 matrices");
auto output = torch::empty_like(input);
register_row_kernel64<<<input.size(0), 32>>>(
input.data_ptr<float>(), output.data_ptr<float>());
auto error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "register-row launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor register_row_potrf128(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.dim() == 3 && input.size(1) == 128
&& input.size(2) == 128,
"register-row kernel requires 128x128 matrices");
auto output = torch::empty_like(input);
register_row_kernel128<<<input.size(0), 128>>>(
input.data_ptr<float>(), output.data_ptr<float>());
auto error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "register-row launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor blocked256_register(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.dim() == 3 && input.size(1) == 256
&& input.size(2) == 256,
"blocked256 requires 256x256 matrices");
const int batch = input.size(0);
constexpr int n = 256;
constexpr int64_t stride = int64_t(n) * n;
auto output = input.clone();
auto pointers = torch::empty(
{2, batch},
torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
auto diagonal = reinterpret_cast<float**>(
pointers[0].data_ptr<int64_t>());
auto panel = reinterpret_cast<float**>(
pointers[1].data_ptr<int64_t>());
factor128_inplace<<<batch, 128>>>(
output.data_ptr<float>(), n, stride, 0);
make_block_pointers<<<(batch + 127) / 128, 128>>>(
output.data_ptr<float>(), diagonal, panel, batch, stride,
0, 128 * n);
const float one = 1.0f;
const float minus_one = -1.0f;
auto trsm_status = cublasStrsmBatched(
precise_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
128, 128, &one,
const_cast<const float**>(diagonal), n,
panel, n, batch);
TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
"precise batched TRSM failed: ", int(trsm_status));
float* panel_base = output.data_ptr<float>() + 128 * n;
float* trailing_base = panel_base + 128;
auto gemm_status = cublasSgemmStridedBatched(
precise_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
128, 128, 128, &minus_one,
panel_base, n, stride,
panel_base, n, stride,
&one, trailing_base, n, stride, batch);
TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
"precise batched GEMM failed: ", int(gemm_status));
factor128_inplace<<<batch, 128>>>(
output.data_ptr<float>(), n, stride, 128 * n + 128);
dim3 clear_threads(32, 8);
dim3 clear_grid(8, 32, batch);
clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
output.data_ptr<float>(), n, stride);
auto error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "blocked256 launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor blocked512_register(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.dim() == 3 && input.size(1) == 512
&& input.size(2) == 512,
"blocked512 requires 512x512 matrices");
const int batch = input.size(0);
constexpr int n = 512;
constexpr int block = 128;
constexpr int64_t stride = int64_t(n) * n;
auto output = input.clone();
auto pointers = torch::empty(
{2, batch},
torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
auto diagonal = reinterpret_cast<float**>(
pointers[0].data_ptr<int64_t>());
auto panel = reinterpret_cast<float**>(
pointers[1].data_ptr<int64_t>());
const float one = 1.0f;
const float minus_one = -1.0f;
for (int k = 0; k < n; k += block) {
factor128_inplace<<<batch, 128>>>(
output.data_ptr<float>(), n, stride, int64_t(k) * n + k);
const int trailing = n - k - block;
if (trailing == 0) continue;
const int64_t diagonal_offset = int64_t(k) * n + k;
const int64_t panel_offset = int64_t(k + block) * n + k;
make_block_pointers<<<(batch + 127) / 128, 128>>>(
output.data_ptr<float>(), diagonal, panel, batch, stride,
diagonal_offset, panel_offset);
auto trsm_status = cublasStrsmBatched(
precise_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
block, trailing, &one,
const_cast<const float**>(diagonal), n,
panel, n, batch);
TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
"blocked512 TRSM failed: ", int(trsm_status));
float* panel_base = output.data_ptr<float>() + panel_offset;
float* trailing_base = output.data_ptr<float>()
+ int64_t(k + block) * n + (k + block);
auto gemm_status = cublasSgemmStridedBatched(
precise_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
trailing, trailing, block, &minus_one,
panel_base, n, stride,
panel_base, n, stride,
&one, trailing_base, n, stride, batch);
TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
"blocked512 GEMM failed: ", int(gemm_status));
}
dim3 clear_threads(32, 8);
dim3 clear_grid(16, 64, batch);
clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
output.data_ptr<float>(), n, stride);
auto error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "blocked512 launch failed: ",
cudaGetErrorString(error));
return output;
}
torch::Tensor blocked1024_register(torch::Tensor input) {
TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32,
"input must be CUDA FP32");
TORCH_CHECK(input.dim() == 3 && input.size(1) == 1024
&& input.size(2) == 1024,
"blocked1024 requires 1024x1024 matrices");
const int batch = input.size(0);
constexpr int n = 1024;
constexpr int block = 128;
constexpr int64_t stride = int64_t(n) * n;
auto output = input.clone();
auto pointers = torch::empty(
{2, batch},
torch::TensorOptions().device(input.device()).dtype(torch::kInt64));
auto diagonal = reinterpret_cast<float**>(
pointers[0].data_ptr<int64_t>());
auto panel = reinterpret_cast<float**>(
pointers[1].data_ptr<int64_t>());
const float one = 1.0f;
const float minus_one = -1.0f;
for (int k = 0; k < n; k += block) {
factor128_inplace<<<batch, 128>>>(
output.data_ptr<float>(), n, stride, int64_t(k) * n + k);
const int trailing = n - k - block;
if (trailing == 0) continue;
const int64_t diagonal_offset = int64_t(k) * n + k;
const int64_t panel_offset = int64_t(k + block) * n + k;
make_block_pointers<<<(batch + 127) / 128, 128>>>(
output.data_ptr<float>(), diagonal, panel, batch, stride,
diagonal_offset, panel_offset);
auto trsm_status = cublasStrsmBatched(
precise_blas_handle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT,
block, trailing, &one,
const_cast<const float**>(diagonal), n,
panel, n, batch);
TORCH_CHECK(trsm_status == CUBLAS_STATUS_SUCCESS,
"blocked1024 TRSM failed: ", int(trsm_status));
float* panel_base = output.data_ptr<float>() + panel_offset;
float* trailing_base = output.data_ptr<float>()
+ int64_t(k + block) * n + (k + block);
auto gemm_status = cublasSgemmStridedBatched(
precise_blas_handle(), CUBLAS_OP_T, CUBLAS_OP_N,
trailing, trailing, block, &minus_one,
panel_base, n, stride,
panel_base, n, stride,
&one, trailing_base, n, stride, batch);
TORCH_CHECK(gemm_status == CUBLAS_STATUS_SUCCESS,
"blocked1024 GEMM failed: ", int(gemm_status));
}
dim3 clear_threads(32, 8);
dim3 clear_grid(32, 128, batch);
clear_batched_row_major_upper_2d<<<clear_grid, clear_threads>>>(
output.data_ptr<float>(), n, stride);
auto error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, "blocked1024 launch failed: ",
cudaGetErrorString(error));
return output;
}
"""
native = load_inline(
name="chol_ranked_hybrid_v3",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=[
"emulated_potrf",
"blocked_potrf",
"register_row_potrf64",
"register_row_potrf128",
"blocked256_register",
"blocked512_register",
"blocked1024_register",
],
extra_ldflags=["-lcusolver"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
@triton.autotune(
configs=[triton.Config({}, num_warps=w) for w in (1, 2, 4, 8)],
key=["N"],
)
@triton.jit
def _chol_small_kernel(A_ptr, L_ptr, N: tl.constexpr, BS: tl.constexpr):
pid = tl.program_id(0)
r = tl.arange(0, BS)
row = r[:, None]
col = r[None, :]
bounds = (row < N) & (col < N)
base = pid * N * N
a = tl.load(A_ptr + base + row * N + col, mask=bounds, other=0.0)
for k in tl.static_range(N): # compile-time unroll: N is a constexpr
dk = tl.sum(tl.where((row == k) & (col == k), a, 0.0))
inv = 1.0 / tl.sqrt(dk)
v = tl.sum(tl.where(col == k, a, 0.0), axis=1) * inv
v = tl.where(r >= k, v, 0.0)
a = a - v[:, None] * v[None, :] # update; also zeroes row/col k
a = tl.where(col == k, v[:, None], a) # write finished L column in place
out = tl.where(row >= col, a, 0.0)
tl.store(L_ptr + base + row * N + col, out, mask=bounds)
def _triton_cholesky(data):
x = data.contiguous()
out = torch.empty_like(x)
n = x.shape[-1]
_chol_small_kernel[(x.shape[0],)](x, out, N=n, BS=triton.next_power_of_2(n))
return out
def custom_kernel(data: input_t) -> output_t:
batch = data.shape[0]
n = data.shape[-1]
if n == 32:
return _triton_cholesky(data)
if n == 64:
return native.register_row_potrf64(data)
if n == 128:
return native.register_row_potrf128(data)
if n == 256:
return native.blocked256_register(data)
if batch == 16 and n == 512:
return native.blocked512_register(data)
if batch == 4 and n == 1024:
return native.blocked1024_register(data)
if batch == 640 and n == 512:
return native.blocked_potrf(data, 128)
if batch == 60 and n == 1024:
return native.blocked_potrf(data, 256)
if batch == 1 and n == 8192:
return _blocked_large(data, 4096)
if batch == 1 and n == 16384:
return _blocked_large(data, 4096)
if batch == 1 and n == 32768:
return _blocked_large(data, 4096)
if batch == 8 and n == 2048:
return native.blocked_potrf(data, 128)
if 2 <= batch <= 4 and n >= 2048:
results = []
for i in range(batch):
L = torch.linalg.cholesky_ex(data[i], check_errors=False).L
results.append(L)
return torch.stack(results)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 833 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