submission 885534
thinkhard101 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 590 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-885534?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:37ae2ae1d2c64b0d9af2f26c43d044398300c5aa630e873eb79c53113705acb0
license declaredunknown
license concludedunknown
authorsthinkhard101
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float shared[];vector-width = float4
const float4* a4 = reinterpret_cast<const float4*>(a);Kernel source
submission.py590 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
torch.set_float32_matmul_precision("high")
CPP_SRC = r"""
torch::Tensor cholesky_small_cuda(torch::Tensor input);
torch::Tensor zero_upper_cuda(torch::Tensor output);
torch::Tensor store_panel_cuda(
torch::Tensor input,
torch::Tensor output,
torch::Tensor half_output,
torch::Tensor fp8_output,
int64_t row_offset,
int64_t col_offset,
double scale);
torch::Tensor store_panel_fp8_cuda(
torch::Tensor input,
torch::Tensor output,
torch::Tensor fp8_output,
int64_t row_offset,
int64_t col_offset,
double scale);
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
template <int N>
__device__ __forceinline__ int lower_offset(int row, int col) {
return row * (N + 1) + col;
}
template <int N>
__global__ void cholesky_small_kernel(
const float* __restrict__ input,
float* __restrict__ output,
int batch) {
extern __shared__ float shared[];
constexpr int matrices_per_block = N == 32 ? 4 : (N == 64 ? 2 : 1);
const int local_matrix = threadIdx.x / N;
const int matrix =
blockIdx.x * matrices_per_block + local_matrix;
const int tid = threadIdx.x - local_matrix * N;
if (matrix >= batch) {
return;
}
float* lower = shared + local_matrix * N * (N + 1);
const float* a = input + static_cast<long long>(matrix) * N * N;
float* l = output + static_cast<long long>(matrix) * N * N;
if constexpr (N == 32) {
const float4* a4 = reinterpret_cast<const float4*>(a);
for (int vector = tid; vector < N * N / 4; vector += N) {
const int row = vector / (N / 4);
const int col = (vector - row * (N / 4)) * 4;
const float4 values = a4[vector];
if (row >= col) {
lower[lower_offset<N>(row, col)] = values.x;
}
if (row >= col + 1) {
lower[lower_offset<N>(row, col + 1)] = values.y;
}
if (row >= col + 2) {
lower[lower_offset<N>(row, col + 2)] = values.z;
}
if (row >= col + 3) {
lower[lower_offset<N>(row, col + 3)] = values.w;
}
}
} else {
for (int index = tid; index < N * N; index += N) {
const int row = index / N;
const int col = index - row * N;
if (row >= col) {
lower[lower_offset<N>(row, col)] = a[index];
} else {
l[index] = 0.0f;
}
}
}
if (N == 32) {
__syncwarp();
} else {
__syncthreads();
}
for (int k = 0; k < N; ++k) {
const int k_base = lower_offset<N>(k, 0);
float diagonal_value = 0.0f;
if constexpr (N <= 64) {
if (tid < 32) {
float sum = 0.0f;
for (int j = tid; j < k; j += 32) {
const float value = lower[k_base + j];
sum += value * value;
}
for (int offset = 16; offset > 0; offset >>= 1) {
sum += __shfl_down_sync(0xffffffff, sum, offset);
}
if (tid == 0) {
const float diagonal = lower[k_base + k] - sum;
lower[k_base + k] = sqrtf(fmaxf(diagonal, 0.0f));
}
}
if constexpr (N == 32) {
__syncwarp();
} else {
__syncthreads();
}
} else {
const int lane = tid & 31;
float sum = 0.0f;
for (int j = lane; j < k; j += 32) {
const float value = lower[k_base + j];
sum += value * value;
}
for (int offset = 16; offset > 0; offset >>= 1) {
sum += __shfl_down_sync(0xffffffff, sum, offset);
}
if (lane == 0) {
const float diagonal = lower[k_base + k] - sum;
diagonal_value = sqrtf(fmaxf(diagonal, 0.0f));
}
diagonal_value =
__shfl_sync(0xffffffff, diagonal_value, 0);
if (tid < 32) {
lower[k_base + k] = diagonal_value;
}
}
const float diagonal =
N == 128 ? diagonal_value : lower[k_base + k];
for (int row = k + 1 + tid; row < N; row += N) {
const int row_base = lower_offset<N>(row, 0);
float value = lower[row_base + k];
if constexpr (N == 32) {
#pragma unroll 4
for (int j = 0; j < k; ++j) {
value -= lower[row_base + j] * lower[k_base + j];
}
} else {
for (int j = 0; j < k; ++j) {
value -= lower[row_base + j] * lower[k_base + j];
}
}
lower[row_base + k] = value / diagonal;
}
if (N == 32) {
__syncwarp();
} else {
__syncthreads();
}
}
if constexpr (N == 32) {
float4* l4 = reinterpret_cast<float4*>(l);
for (int vector = tid; vector < N * N / 4; vector += N) {
const int row = vector / (N / 4);
const int col = (vector - row * (N / 4)) * 4;
float4 values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (row >= col) {
values.x = lower[lower_offset<N>(row, col)];
}
if (row >= col + 1) {
values.y = lower[lower_offset<N>(row, col + 1)];
}
if (row >= col + 2) {
values.z = lower[lower_offset<N>(row, col + 2)];
}
if (row >= col + 3) {
values.w = lower[lower_offset<N>(row, col + 3)];
}
l4[vector] = values;
}
} else {
for (int index = tid; index < N * N; index += N) {
const int row = index / N;
const int col = index - row * N;
if (row >= col) {
l[index] = lower[lower_offset<N>(row, col)];
}
}
}
}
__global__ void zero_upper_kernel(float* output, int n) {
const int row = blockIdx.x;
const int matrix = blockIdx.y;
float* matrix_output =
output + static_cast<long long>(matrix) * n * n;
for (int col = row + 1 + threadIdx.x; col < n; col += blockDim.x) {
matrix_output[static_cast<long long>(row) * n + col] = 0.0f;
}
}
template <bool WRITE_HALF>
__global__ void store_panel_kernel(
const float* input,
float* output,
__half* half_output,
unsigned char* fp8_output,
long long rows,
long long cols,
long long input_stride_0,
long long input_stride_1,
int n,
int row_offset,
int col_offset,
float scale) {
const long long total = rows * cols;
const long long thread =
static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
const long long stride = static_cast<long long>(gridDim.x) * blockDim.x;
for (long long index = thread; index < total; index += stride) {
const long long row = index / cols;
const long long col = index - row * cols;
const float value = input[row * input_stride_0 + col * input_stride_1];
const long long destination =
static_cast<long long>(row_offset + row) * n + col_offset + col;
output[destination] = value;
if constexpr (WRITE_HALF) {
half_output[destination] = __float2half_rn(value);
}
fp8_output[destination] = __nv_cvt_float_to_fp8(
value * scale, __NV_SATFINITE, __NV_E4M3);
}
}
torch::Tensor store_panel_cuda(
torch::Tensor input,
torch::Tensor output,
torch::Tensor half_output,
torch::Tensor fp8_output,
int64_t row_offset,
int64_t col_offset,
double scale) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(input.dim() == 2, "input must have two dimensions");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
TORCH_CHECK(half_output.is_contiguous(), "half output must be contiguous");
TORCH_CHECK(fp8_output.is_contiguous(), "fp8 output must be contiguous");
const long long total = input.numel();
int blocks = static_cast<int>((total + 255) / 256);
if (blocks > 65535) {
blocks = 65535;
}
const int n = output.size(0);
store_panel_kernel<true><<<blocks, 256>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
reinterpret_cast<__half*>(half_output.data_ptr<at::Half>()),
reinterpret_cast<unsigned char*>(fp8_output.data_ptr()),
input.size(0),
input.size(1),
input.stride(0),
input.stride(1),
n,
static_cast<int>(row_offset),
static_cast<int>(col_offset),
static_cast<float>(scale));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
return output;
}
torch::Tensor store_panel_fp8_cuda(
torch::Tensor input,
torch::Tensor output,
torch::Tensor fp8_output,
int64_t row_offset,
int64_t col_offset,
double scale) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(input.dim() == 2, "input must have two dimensions");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
TORCH_CHECK(fp8_output.is_contiguous(), "fp8 output must be contiguous");
const long long total = input.numel();
int blocks = static_cast<int>((total + 255) / 256);
if (blocks > 65535) {
blocks = 65535;
}
const int n = output.size(0);
store_panel_kernel<false><<<blocks, 256>>>(
input.data_ptr<float>(),
output.data_ptr<float>(),
nullptr,
reinterpret_cast<unsigned char*>(fp8_output.data_ptr()),
input.size(0),
input.size(1),
input.stride(0),
input.stride(1),
n,
static_cast<int>(row_offset),
static_cast<int>(col_offset),
static_cast<float>(scale));
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
return output;
}
torch::Tensor zero_upper_cuda(torch::Tensor output) {
TORCH_CHECK(output.is_cuda(), "output must be CUDA");
TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be float32");
TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
TORCH_CHECK(output.dim() == 3, "output must have three dimensions");
const int batch = output.size(0);
const int n = output.size(1);
zero_upper_kernel<<<dim3(n, batch), 512>>>(output.data_ptr<float>(), n);
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
return output;
}
torch::Tensor cholesky_small_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
TORCH_CHECK(input.dim() == 3, "input must have three dimensions");
const int batch = input.size(0);
const int n = input.size(1);
auto output = torch::empty_like(input);
if (n == 32) {
constexpr int shared_bytes = 4 * 32 * 33 * sizeof(float);
cholesky_small_kernel<32><<<(batch + 3) / 4, 128, shared_bytes>>>(
input.data_ptr<float>(), output.data_ptr<float>(), batch);
} else if (n == 64) {
constexpr int shared_bytes = 2 * 64 * 65 * sizeof(float);
cholesky_small_kernel<64><<<(batch + 1) / 2, 128, shared_bytes>>>(
input.data_ptr<float>(), output.data_ptr<float>(), batch);
} else if (n == 128) {
constexpr int shared_bytes = 128 * 129 * sizeof(float);
static const bool configured = []() {
const cudaError_t status = cudaFuncSetAttribute(
cholesky_small_kernel<128>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
128 * 129 * sizeof(float));
return status == cudaSuccess;
}();
TORCH_CHECK(configured, "failed to configure shared memory");
cholesky_small_kernel<128><<<batch, 128, shared_bytes>>>(
input.data_ptr<float>(), output.data_ptr<float>(), batch);
} else {
TORCH_CHECK(false, "cholesky_small_cuda only supports n=32, n=64, or n=128");
}
const cudaError_t error = cudaGetLastError();
TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
return output;
}
"""
_native = load_inline(
name="cholesky_small_b200_v4",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=[
"cholesky_small_cuda",
"zero_upper_cuda",
"store_panel_cuda",
"store_panel_fp8_cuda",
],
extra_cuda_cflags=["-O3", "--use_fast_math"],
verbose=False,
)
_large_info = None
_medium_info = None
_fp8_scale = None
def _blocked_large_cholesky(data: torch.Tensor) -> torch.Tensor:
global _fp8_scale
matrix = data[0]
n = matrix.shape[0]
use_fp8 = n >= 16384
use_half = n >= 32768
block = 2048 if n >= 32768 else 4096
output = torch.empty_like(data)
lower = output[0]
lower_half = (
torch.empty_like(matrix, dtype=torch.float16) if use_half else None
)
lower_fp8 = (
torch.empty_like(matrix, dtype=torch.float8_e4m3fn) if use_fp8 else None
)
if use_fp8 and (_fp8_scale is None or _fp8_scale.device != data.device):
_fp8_scale = torch.full((), 1.0 / 16.0, device=data.device)
fp8_scale = _fp8_scale if use_fp8 else None
for start in range(0, n, block):
end = start + block
diagonal = matrix[start:end, start:end].clone()
if start:
if lower_fp8 is None:
previous = lower[start:end, :start]
diagonal.addmm_(previous, previous.T, beta=1.0, alpha=-1.0)
else:
if n == 16384 or start == block:
previous = lower_fp8[start:end, :start]
diagonal.add_(
torch._scaled_mm(
previous,
previous.T,
scale_a=fp8_scale,
scale_b=fp8_scale,
out_dtype=torch.float32,
),
alpha=-1.0,
)
else:
previous = lower_half[start:end, :start]
diagonal.add_(
torch.mm(previous, previous.T, out_dtype=torch.float32),
alpha=-1.0,
)
factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
if lower_fp8 is None:
lower[start:end, start:end] = factor
elif lower_half is None:
_native.store_panel_fp8_cuda(
factor,
lower,
lower_fp8,
start,
start,
16.0,
)
else:
_native.store_panel_cuda(
factor,
lower,
lower_half,
lower_fp8,
start,
start,
16.0,
)
if end < n:
panel = matrix[end:, start:end].clone()
if start:
if lower_fp8 is None:
panel.addmm_(
lower[end:, :start],
lower[start:end, :start].T,
beta=1.0,
alpha=-1.0,
)
else:
panel.add_(
torch._scaled_mm(
lower_fp8[end:, :start],
lower_fp8[start:end, :start].T,
scale_a=fp8_scale,
scale_b=fp8_scale,
out_dtype=torch.float32,
),
alpha=-1.0,
)
solved = torch.linalg.solve_triangular(
factor.T,
panel,
upper=True,
left=False,
)
if lower_fp8 is None:
lower[end:, start:end] = solved
elif lower_half is None:
_native.store_panel_fp8_cuda(
solved,
lower,
lower_fp8,
end,
start,
16.0,
)
else:
_native.store_panel_cuda(
solved,
lower,
lower_half,
lower_fp8,
end,
start,
16.0,
)
return _native.zero_upper_cuda(output)
def _blocked_batched(data: torch.Tensor, block: int) -> torch.Tensor:
n = data.shape[1]
output = torch.zeros_like(data)
for start in range(0, n, block):
end = min(start + block, n)
diagonal = data[:, start:end, start:end].clone()
if start:
previous = output[:, start:end, :start]
diagonal.baddbmm_(
previous,
previous.transpose(1, 2),
beta=1.0,
alpha=-1.0,
)
factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
output[:, start:end, start:end] = factor
if end < n:
panel = data[:, end:, start:end].clone()
if start:
panel.baddbmm_(
output[:, end:, :start],
output[:, start:end, :start].transpose(1, 2),
beta=1.0,
alpha=-1.0,
)
solved = torch.linalg.solve_triangular(
factor.transpose(1, 2),
panel,
upper=True,
left=False,
out=panel,
)
output[:, end:, start:end] = solved
return output
def custom_kernel(data: input_t) -> output_t:
global _large_info, _medium_info
batch, n, _ = data.shape
if n == 32:
return _native.cholesky_small_cuda(data)
if n == 64 and batch % 2 == 0:
return _native.cholesky_small_cuda(data)
if n == 128:
return _native.cholesky_small_cuda(data)
if n == 1024 and batch == 60:
return _blocked_batched(data, 128)
if n == 1024 and batch == 4:
output = torch.empty_like(data)
if _medium_info is None or _medium_info.device != data.device:
_medium_info = torch.empty((4,), dtype=torch.int32, device=data.device)
for index in range(batch):
torch.linalg.cholesky_ex(
data[index],
check_errors=False,
out=(output[index], _medium_info[index]),
)
return output
if n >= 8192 and batch == 1:
return _blocked_large_cholesky(data)
if n >= 2048 and batch == 2:
output = torch.empty_like(data)
if _large_info is None or _large_info.device != data.device:
_large_info = torch.empty((2,), dtype=torch.int32, device=data.device)
for index in range(batch):
torch.linalg.cholesky_ex(
data[index],
check_errors=False,
out=(output[index], _large_info[index]),
)
return output
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 590 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