submission 923062
Shaman Shetty · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 301 lines, June 9 Researcher Reciprocity License v1.0.
submission2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-923062?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:52057f386ef74dfb040a0b303df07449051356b8e8c8223ee7e711a6c7e91505
license declaredunknown
license concludedunknown
authorsShaman Shetty
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float smem[];Kernel source
submission2.py301 lines
import torch
try:
from task import input_t, output_t
except Exception:
input_t = torch.Tensor
output_t = torch.Tensor
try:
torch.backends.cuda.matmul.allow_tf32 = False
torch.set_float32_matmul_precision("highest")
except Exception:
pass
_EXT = None
_BASE = 128
_TRSM_OUT = None
_PLAN = {}
_CPP_DECL = "void bchol(torch::Tensor dst, torch::Tensor src, bool zero_upper);"
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>
__global__ void bchol_kernel(
float* __restrict__ dst, const float* __restrict__ src,
const long batch, const int n,
const long d_srow, const long d_sbatch,
const long s_srow, const long s_sbatch,
const int ld, const int zero_upper)
{
extern __shared__ float smem[];
float* a = smem;
float* dia = smem + (size_t)n * ld;
const long b = blockIdx.x;
if (b >= batch) return;
const float* S = src + b * s_sbatch;
float* D = dst + b * d_sbatch;
const int t = threadIdx.x;
const int T = blockDim.x;
for (int idx = t; idx < n * n; idx += T) {
const int r = idx / n;
const int c = idx - r * n;
a[r * ld + c] = S[(long)r * s_srow + c];
}
__syncthreads();
const int i = t;
for (int j = 0; j < n; ++j) {
const float ajj = a[j * ld + j];
const float dj = sqrtf(ajj);
const float inv = 1.0f / dj;
if (i == j) dia[j] = dj;
if (i > j && i < n) a[i * ld + j] *= inv;
__syncthreads();
if (i > j && i < n) {
const float lij = a[i * ld + j];
float* ai = a + i * ld;
#pragma unroll 4
for (int k = j + 1; k <= i; ++k)
ai[k] -= lij * a[k * ld + j];
}
__syncthreads();
}
for (int idx = t; idx < n * n; idx += T) {
const int r = idx / n;
const int c = idx - r * n;
if (c < r) D[(long)r * d_srow + c] = a[r * ld + c];
else if (c == r) D[(long)r * d_srow + c] = dia[r];
else if (zero_upper) D[(long)r * d_srow + c] = 0.0f;
}
}
void bchol(torch::Tensor dst, torch::Tensor src, bool zero_upper)
{
TORCH_CHECK(dst.is_cuda() && src.is_cuda(), "bchol: CUDA tensors required");
TORCH_CHECK(dst.scalar_type() == at::kFloat && src.scalar_type() == at::kFloat,
"bchol: float32 required");
TORCH_CHECK(dst.dim() == 3 && src.dim() == 3, "bchol: (B,n,n) required");
const long batch = dst.size(0);
const long n_l = dst.size(1);
TORCH_CHECK(src.size(0) == batch && src.size(1) == n_l &&
src.size(2) == n_l && dst.size(2) == n_l, "bchol: shape mismatch");
TORCH_CHECK(n_l >= 1 && n_l <= 128, "bchol: 1 <= n <= 128 required");
TORCH_CHECK(dst.stride(2) == 1 && src.stride(2) == 1,
"bchol: innermost stride must be 1");
if (batch == 0) return;
const int n = (int)n_l;
int ld = n + 1;
if ((ld & 1) == 0) ld += 1;
const int T = ((n + 31) / 32) * 32;
const size_t smem = ((size_t)n * ld + n) * sizeof(float);
if (smem > 48 * 1024) {
static size_t smem_set = 0;
if (smem > smem_set) {
cudaError_t err = cudaFuncSetAttribute(
bchol_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
TORCH_CHECK(err == cudaSuccess, "bchol: shared memory opt-in failed");
smem_set = smem;
}
}
const at::cuda::OptionalCUDAGuard guard(at::device_of(dst));
bchol_kernel<<<dim3((unsigned int)batch), dim3(T), smem>>>(
dst.data_ptr<float>(), src.data_ptr<float>(),
batch, n,
dst.stride(1), dst.stride(0),
src.stride(1), src.stride(0),
ld, zero_upper ? 1 : 0);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
"""
def _base_chol(V, zero_upper):
if _EXT is not None:
_EXT.bchol(V, V, zero_upper)
else:
V.copy_(torch.linalg.cholesky_ex(V, check_errors=False).L)
def _panel_nb(n):
if n <= 2048:
return _BASE
if n <= 4096:
return 256
return 512
def _probe_trsm_out(device):
try:
g = torch.Generator(device="cpu").manual_seed(0)
L = torch.randn(2, 5, 5, generator=g).tril() + 5.0 * torch.eye(5)
Bf = torch.randn(2, 7, 9, generator=g)
L, Bf = L.to(device), Bf.to(device)
B = Bf[:, :, 2:7]
U = L.transpose(-1, -2)
ref = torch.linalg.solve_triangular(U, B, upper=True, left=False)
torch.linalg.solve_triangular(U, B, upper=True, left=False, out=B)
return bool(torch.allclose(B, ref, rtol=1e-4, atol=1e-5))
except Exception:
return False
def _trsm_right(L11, B):
global _TRSM_OUT
U = L11.transpose(-1, -2)
if _TRSM_OUT is None:
_TRSM_OUT = _probe_trsm_out(B.device)
if _TRSM_OUT:
torch.linalg.solve_triangular(U, B, upper=True, left=False, out=B)
else:
B.copy_(torch.linalg.solve_triangular(U, B, upper=True, left=False))
def _update_trailing(A22, P):
m = A22.size(-1)
Pt = P.transpose(-1, -2)
if m <= 256:
A22.baddbmm_(P, Pt, alpha=-1.0)
return
if m <= 1024:
cb = 256
elif m <= 4096:
cb = 512
elif m <= 16384:
cb = 1024
else:
cb = 2048
for c0 in range(0, m, cb):
c1 = min(c0 + cb, m)
A22[:, c0:, c0:c1].baddbmm_(P[:, c0:, :], Pt[:, :, c0:c1], alpha=-1.0)
def _chol_ip(A):
n = A.size(-1)
if n <= _BASE:
_base_chol(A, False)
return
nb0 = _panel_nb(n)
j = 0
while j < n:
nb = min(nb0, n - j)
k1 = j + nb
D = A[:, j:k1, j:k1]
if nb <= _BASE:
_base_chol(D, False)
else:
_chol_ip(D)
if k1 < n:
P = A[:, k1:, j:k1]
_trsm_right(D, P)
_update_trailing(A[:, k1:, k1:], P)
j = k1
def _run_custom(data):
n = data.size(-1)
if n <= _BASE and _EXT is not None and data.stride(-1) == 1:
L = torch.empty(data.shape, dtype=data.dtype, device=data.device)
_EXT.bchol(L, data, True)
return L
A = data.contiguous()
L = A.clone() if A is data else A
_chol_ip(L)
return L.tril_()
def _run_torch(data):
return torch.linalg.cholesky_ex(data, check_errors=False).L
def _quick_ok(A, L):
try:
if not bool(torch.isfinite(L).all()):
return False
n = A.size(-1)
s = min(n, 128)
idx = torch.randperm(n, device=A.device)[:s]
rec = torch.matmul(L[:, idx, :], L.transpose(-1, -2))
num = (rec - A[:, idx, :]).norm()
den = A[:, idx, :].norm().clamp_min(1e-30)
return bool((num / den) < 2e-3)
except Exception:
return False
def _select(A):
try:
Lc = _run_custom(A)
ok = _quick_ok(A, Lc)
del Lc
if ok:
return _run_custom
except Exception:
pass
return _run_torch
def custom_kernel(data: input_t) -> output_t:
A = data
if not torch.is_tensor(A):
return _run_torch(A)
if A.dim() == 2:
return custom_kernel(A.unsqueeze(0)).squeeze(0)
if A.dim() != 3 or (not A.is_cuda) or A.dtype != torch.float32:
return _run_torch(A)
key = (A.size(0), A.size(1), A.device.index)
fn = _PLAN.get(key)
if fn is None:
fn = _select(A)
_PLAN[key] = fn
return fn(A)
def _init():
global _EXT, _BASE
if not torch.cuda.is_available():
return
try:
from torch.utils.cpp_extension import load_inline
ext = load_inline(
name="bchol_ext_v1",
cpp_sources=_CPP_DECL,
cuda_sources=_CUDA_SRC,
functions=["bchol"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
except Exception:
return
dev = torch.device("cuda", torch.cuda.current_device())
for base in (128, 96):
try:
g = torch.Generator(device=dev).manual_seed(0)
R = torch.randn(3, base, base, device=dev, generator=g)
A = R @ R.transpose(-1, -2) + base * torch.eye(base, device=dev)
L = torch.empty_like(A)
ext.bchol(L, A, True)
Ld = L.double()
resid = (Ld @ Ld.transpose(-1, -2) - A.double()).norm() / A.double().norm()
ok = (bool(torch.isfinite(L).all())
and float(resid) < 1e-4
and float(L.triu(1).abs().max()) == 0.0
and bool((L.diagonal(dim1=-2, dim2=-1) > 0).all()))
if ok:
_EXT, _BASE = ext, base
return
except Exception:
continue
_init()
scrolls · 301 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