submission 926538
zdy23 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 148 lines, June 9 Researcher Reciprocity License v1.0.
v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-926538?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:0a3afac042cae056869ee2bcdc96a96d8ad74310084a14aa6eb20d6fb80498a2
license declaredunknown
license concludedunknown
authorszdy23
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float S[32][33];Kernel source
v2.py148 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
_cpp_src = r"""
torch::Tensor cholesky_n32(torch::Tensor A);
torch::Tensor cholesky_n64(torch::Tensor A);
"""
_cuda_src = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
__global__ void cholesky_n32_kernel(const float* __restrict__ A,
float* __restrict__ L) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const size_t off = (size_t)b * 1024u;
__shared__ float S[32][33];
const float* Ab = A + off;
float* Lb = L + off;
#pragma unroll
for (int j = 0; j < 32; ++j) {
S[tid][j] = Ab[tid * 32 + j];
}
__syncthreads();
#pragma unroll
for (int k = 0; k < 32; ++k) {
if (tid == k) {
S[k][k] = sqrtf(S[k][k]);
}
__syncthreads();
const float diag = S[k][k];
if (tid > k) {
S[tid][k] /= diag;
}
__syncthreads();
if (tid > k) {
const float lik = S[tid][k];
#pragma unroll
for (int j = k + 1; j <= tid; ++j) {
S[tid][j] -= lik * S[j][k];
}
}
__syncthreads();
}
#pragma unroll
for (int j = 0; j < 32; ++j) {
Lb[tid * 32 + j] = (j <= tid) ? S[tid][j] : 0.f;
}
}
__global__ void cholesky_n64_kernel(const float* __restrict__ A,
float* __restrict__ L) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const size_t off = (size_t)b * 4096u;
__shared__ float S[64][65];
const float* Ab = A + off;
float* Lb = L + off;
#pragma unroll 4
for (int j = 0; j < 64; ++j) {
S[tid][j] = Ab[tid * 64 + j];
}
__syncthreads();
for (int k = 0; k < 64; ++k) {
if (tid == k) {
S[k][k] = sqrtf(S[k][k]);
}
__syncthreads();
const float diag = S[k][k];
if (tid > k) {
S[tid][k] /= diag;
}
__syncthreads();
if (tid > k) {
const float lik = S[tid][k];
for (int j = k + 1; j <= tid; ++j) {
S[tid][j] -= lik * S[j][k];
}
}
__syncthreads();
}
#pragma unroll 4
for (int j = 0; j < 64; ++j) {
Lb[tid * 64 + j] = (j <= tid) ? S[tid][j] : 0.f;
}
}
torch::Tensor cholesky_n32(torch::Tensor A) {
TORCH_CHECK(A.is_cuda() && A.scalar_type() == torch::kFloat32);
TORCH_CHECK(A.dim() == 3 && A.size(1) == 32 && A.size(2) == 32);
auto Ac = A.contiguous();
auto L = torch::empty_like(Ac);
const int B = (int)Ac.size(0);
cholesky_n32_kernel<<<B, 32>>>(Ac.data_ptr<float>(), L.data_ptr<float>());
return L;
}
torch::Tensor cholesky_n64(torch::Tensor A) {
TORCH_CHECK(A.is_cuda() && A.scalar_type() == torch::kFloat32);
TORCH_CHECK(A.dim() == 3 && A.size(1) == 64 && A.size(2) == 64);
auto Ac = A.contiguous();
auto L = torch::empty_like(Ac);
const int B = (int)Ac.size(0);
cholesky_n64_kernel<<<B, 64>>>(Ac.data_ptr<float>(), L.data_ptr<float>());
return L;
}
"""
_mod = load_inline(
name="cholesky_v2_ultimate",
cpp_sources=_cpp_src,
cuda_sources=_cuda_src,
functions=["cholesky_n32", "cholesky_n64"],
extra_cuda_cflags=["-O3", "--use_fast_math"],
with_cuda=True,
verbose=False,
)
def custom_kernel(data: input_t) -> output_t:
A = data
if not A.is_contiguous() or A.dtype != torch.float32:
A = A.contiguous().float()
n = A.shape[-1]
if n == 32:
return _mod.cholesky_n32(A)
if n == 64:
return _mod.cholesky_n64(A)
return torch.linalg.cholesky(A)
scrolls · 148 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