submission 889544
debadree25 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1792 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-889544?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:07edb4089ed2db81a7dda8397899d34130d1c7d97b7a7612c2a99a4f4c82a5e0
license declaredunknown
license concludedunknown
authorsdebadree25
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
cholesky64_smem_persist(const float* __restrict__ A,vector-width = float4
float4 v = *reinterpret_cast<const float4*>(a + lane * 32 + t * 4);Kernel source
submission.py1792 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
"""Hybrid: n32/64/128 + n256 dense-panel + medium high-batch + large3 TF32.
- n=32/64/128: champ CUDA
- medium high-batch blocked batched (512@≥64, 1024@≥16, 2048@≥8)
- large TF32 selective (n≥16384 or n≥4096&batch≥2)
- else torch
"""
from __future__ import annotations
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
SMALL_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <algorithm>
#include <stdexcept>
#include <string>
#define CHECK_CUDA(expr) do { \
cudaError_t _e = (expr); \
if (_e != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(_e)); \
} \
} while (0)
// -------------------- n=32: one warp / matrix, rows in registers --------------------
template <int MPB>
__global__ void __launch_bounds__(32 * MPB, 4)
cholesky32_warp_persist(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int lane = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 1024u;
float* out = L + (size_t)b * 1024u;
float row[32];
#pragma unroll
for (int t = 0; t < 8; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + lane * 32 + t * 4);
row[t * 4 + 0] = v.x;
row[t * 4 + 1] = v.y;
row[t * 4 + 2] = v.z;
row[t * 4 + 3] = v.w;
}
#pragma unroll
for (int k = 0; k < 32; ++k) {
float sum = 0.f;
#pragma unroll
for (int j = 0; j < 32; ++j) {
if (j < k) {
float rkj = __shfl_sync(0xffffffffu, row[j], k);
sum = fmaf(row[j], rkj, sum);
}
}
float aik = row[k] - sum;
if (lane == k) {
row[k] = sqrtf(fmaxf(aik, 0.f));
}
float diag = __shfl_sync(0xffffffffu, row[k], k);
float inv = 1.f / diag;
if (lane > k) {
row[k] = aik * inv;
} else if (lane < k) {
row[k] = 0.f;
}
}
#pragma unroll
for (int t = 0; t < 8; ++t) {
float4 v;
int j0 = t * 4;
v.x = (lane >= (j0 + 0)) ? row[j0 + 0] : 0.f;
v.y = (lane >= (j0 + 1)) ? row[j0 + 1] : 0.f;
v.z = (lane >= (j0 + 2)) ? row[j0 + 2] : 0.f;
v.w = (lane >= (j0 + 3)) ? row[j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + lane * 32 + j0) = v;
}
}
}
torch::Tensor cholesky_n32(torch::Tensor A) {
TORCH_CHECK(A.is_cuda(), "A must be CUDA");
TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
TORCH_CHECK(A.dim() == 3 && A.size(1) == 32 && A.size(2) == 32, "expected (B,32,32)");
auto Ac = A.contiguous();
auto L = torch::empty_like(Ac);
const int batch = static_cast<int>(Ac.size(0));
if (batch == 0) return L;
constexpr int mpb = 8;
dim3 block(32, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 1024 ? need : 1024;
if (grid < 1) grid = 1;
cholesky32_warp_persist<mpb><<<grid, block>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
CHECK_CUDA(cudaGetLastError());
return L;
}
// -------------------- n=64: left-looking, SMEM, full unroll, multi-matrix --------------------
// Each of 64 threads owns one matrix row in shared memory (bank-padded stride 65).
// Outer k and inner j fully unrolled so j<k is compile-time (n32-style quality).
// Dynamic SMEM: MPB * 64 * 65 floats (static 48KB cap is too small for MPB>=3).
template <int MPB>
__global__ void __launch_bounds__(64 * MPB, 2)
cholesky64_smem_persist(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x; // 0..63 row owner
const int slot = threadIdx.y; // matrix slot in CTA
const int stride = gridDim.x * MPB;
extern __shared__ float smem[];
float* tile = smem + slot * (64 * 65);
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 4096u;
float* out = L + (size_t)b * 4096u;
// Cooperative float4 load of full 64x64 into padded rows.
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + tid * 64 + t * 4);
int j0 = t * 4;
tile[tid * 65 + j0 + 0] = v.x;
tile[tid * 65 + j0 + 1] = v.y;
tile[tid * 65 + j0 + 2] = v.z;
tile[tid * 65 + j0 + 3] = v.w;
}
__syncthreads();
#pragma unroll
for (int k = 0; k < 64; ++k) {
float sum = 0.f;
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < k) {
sum = fmaf(tile[tid * 65 + j], tile[k * 65 + j], sum);
}
}
float aik = tile[tid * 65 + k] - sum;
if (tid == k) {
tile[k * 65 + k] = sqrtf(fmaxf(aik, 0.f));
}
__syncthreads();
float inv = 1.f / tile[k * 65 + k];
if (tid > k) {
tile[tid * 65 + k] = aik * inv;
} else if (tid < k) {
tile[tid * 65 + k] = 0.f;
}
__syncthreads();
}
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v;
int j0 = t * 4;
v.x = (tid >= (j0 + 0)) ? tile[tid * 65 + j0 + 0] : 0.f;
v.y = (tid >= (j0 + 1)) ? tile[tid * 65 + j0 + 1] : 0.f;
v.z = (tid >= (j0 + 2)) ? tile[tid * 65 + j0 + 2] : 0.f;
v.w = (tid >= (j0 + 3)) ? tile[tid * 65 + j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + tid * 64 + j0) = v;
}
if (b + stride < batch) {
__syncthreads();
}
}
}
// Register-resident row variant: lower SMEM traffic (only diag broadcast).
// Higher register pressure; occupancy kept low via launch_bounds.
template <int MPB>
__global__ void __launch_bounds__(64 * MPB, 1)
cholesky64_reg_persist(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
__shared__ float sk[MPB][64]; // row-k lower + diag publish
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 4096u;
float* out = L + (size_t)b * 4096u;
float row[64];
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + tid * 64 + t * 4);
row[t * 4 + 0] = v.x;
row[t * 4 + 1] = v.y;
row[t * 4 + 2] = v.z;
row[t * 4 + 3] = v.w;
}
#pragma unroll
for (int k = 0; k < 64; ++k) {
// Publish L[k, 0:k] (and later diag at sk[k]).
if (tid == k) {
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < k) sk[slot][j] = row[j];
}
}
__syncthreads();
float sum = 0.f;
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < k) {
sum = fmaf(row[j], sk[slot][j], sum);
}
}
float aik = row[k] - sum;
if (tid == k) {
float d = sqrtf(fmaxf(aik, 0.f));
row[k] = d;
sk[slot][k] = d;
}
__syncthreads();
float inv = 1.f / sk[slot][k];
if (tid > k) {
row[k] = aik * inv;
} else if (tid < k) {
row[k] = 0.f;
}
__syncthreads();
}
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v;
int j0 = t * 4;
v.x = (tid >= (j0 + 0)) ? row[j0 + 0] : 0.f;
v.y = (tid >= (j0 + 1)) ? row[j0 + 1] : 0.f;
v.z = (tid >= (j0 + 2)) ? row[j0 + 2] : 0.f;
v.w = (tid >= (j0 + 3)) ? row[j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + tid * 64 + j0) = v;
}
}
}
// Right-looking register row + SMEM column broadcast, full unroll.
template <int MPB>
__global__ void __launch_bounds__(64 * MPB, 1)
cholesky64_right_persist(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
__shared__ float col[MPB][64];
__shared__ float diag_s[MPB];
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 4096u;
float* out = L + (size_t)b * 4096u;
float row[64];
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + tid * 64 + t * 4);
row[t * 4 + 0] = v.x;
row[t * 4 + 1] = v.y;
row[t * 4 + 2] = v.z;
row[t * 4 + 3] = v.w;
}
#pragma unroll
for (int k = 0; k < 64; ++k) {
if (tid == k) {
float d = sqrtf(fmaxf(row[k], 0.f));
row[k] = d;
diag_s[slot] = d;
}
__syncthreads();
float inv = 1.f / diag_s[slot];
float lik;
if (tid > k) {
lik = row[k] * inv;
row[k] = lik;
} else if (tid == k) {
lik = diag_s[slot];
} else {
lik = 0.f;
row[k] = 0.f;
}
col[slot][tid] = lik;
__syncthreads();
float lik_i = col[slot][tid];
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j > k && tid >= j) {
row[j] = fmaf(-lik_i, col[slot][j], row[j]);
}
}
__syncthreads();
}
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v;
int j0 = t * 4;
v.x = (tid >= (j0 + 0)) ? row[j0 + 0] : 0.f;
v.y = (tid >= (j0 + 1)) ? row[j0 + 1] : 0.f;
v.z = (tid >= (j0 + 2)) ? row[j0 + 2] : 0.f;
v.w = (tid >= (j0 + 3)) ? row[j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + tid * 64 + j0) = v;
}
}
}
// Variant selector: 0=smem left, 1=reg left, 2=right
#ifndef CHOLESKY64_VARIANT
#define CHOLESKY64_VARIANT 0
#endif
torch::Tensor cholesky_n64(torch::Tensor A) {
TORCH_CHECK(A.is_cuda(), "A must be CUDA");
TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
TORCH_CHECK(A.dim() == 3 && A.size(1) == 64 && A.size(2) == 64, "expected (B,64,64)");
auto Ac = A.contiguous();
auto L = torch::empty_like(Ac);
const int batch = static_cast<int>(Ac.size(0));
if (batch == 0) return L;
constexpr int mpb = 2; // Modal sweep: mpb2 ~37us > mpb4 ~43 > mpb8 ~50 @1024x64
dim3 block(64, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
#if CHOLESKY64_VARIANT == 1
cholesky64_reg_persist<mpb><<<grid, block>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
#elif CHOLESKY64_VARIANT == 2
cholesky64_right_persist<mpb><<<grid, block>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
#else
{
// 4 * 64 * 65 * 4 = 66560 bytes dynamic SMEM (over static 48KB cap).
const int shmem = mpb * 64 * 65 * (int)sizeof(float);
CHECK_CUDA(cudaFuncSetAttribute(
cholesky64_smem_persist<mpb>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shmem));
cholesky64_smem_persist<mpb><<<grid, block, shmem>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#endif
CHECK_CUDA(cudaGetLastError());
return L;
}
"""
SMALL_CPP_SRC = r"""
torch::Tensor cholesky_n32(torch::Tensor A);
torch::Tensor cholesky_n64(torch::Tensor A);
"""
_mod_small = load_inline(
name="chol_s3264_fast_v1", # same name as champ for compile cache hit
cpp_sources=[SMALL_CPP_SRC],
cuda_sources=[SMALL_CUDA_SRC],
functions=["cholesky_n32", "cholesky_n64"],
verbose=False,
extra_cuda_cflags=["-O3", "--use_fast_math", "-DCHOLESKY64_VARIANT=1"],
)
N128_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <algorithm>
#include <stdexcept>
#include <string>
#define CHECK_CUDA(expr) do { \
cudaError_t _e = (expr); \
if (_e != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(_e)); \
} \
} while (0)
// -------------------- n=128 research kernels (VARIANT 0..4) --------------------
// Failed baselines (do not regress to these alone):
// full row[128]/thr ~665us; naive blocked BS32 ~258us; hier 2x64 ~310us.
// Target: Modal beat ~151us torch; Popcorn <=130us @256x128.
//
// VARIANT 0: phase-split halfreg left (known best ~113us Modal @ MPB=2)
// VARIANT 1: blocked BS multi-matrix SMEM SYRK
// VARIANT 2: interleaved even/odd halfreg (load-balanced dots)
// VARIANT 3: right-looking halfreg
// VARIANT 4: 256-thread blocked BS=16, 2 thr/row SYRK
#ifndef CHOLESKY128_BS
#define CHOLESKY128_BS 16
#endif
#ifndef CHOLESKY128_VARIANT
#define CHOLESKY128_VARIANT 0
#endif
#ifndef CHOLESKY128_MPB
#define CHOLESKY128_MPB 2
#endif
// ---- V0: phase-split halfreg left-looking (optimized) ----
// 2 threads/row hold col[64]. Phase0 k in [0,64): only half0 owns columns;
// no shfl (half1 partial always 0). Phase1 k in [64,128): both halves + xor.
template <int MPB>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_halfreg_phase(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x; // 0..255
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
const int half = tid & 1;
const int row = tid >> 1;
const int col0 = half << 6;
__shared__ float sk[MPB][128];
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 16384u;
float* out = L + (size_t)b * 16384u;
float col[64];
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + row * 128 + col0 + t * 4);
col[t * 4 + 0] = v.x;
col[t * 4 + 1] = v.y;
col[t * 4 + 2] = v.z;
col[t * 4 + 3] = v.w;
}
// Phase 0: k in [0,64). half1 never contributes; skip shfl_xor.
#pragma unroll
for (int k = 0; k < 64; ++k) {
if (row == k && half == 0) {
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < k) sk[slot][j] = col[j];
}
}
__syncthreads();
float sum = 0.f;
if (half == 0) {
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < k) sum = fmaf(col[j], sk[slot][j], sum);
}
float aik = col[k] - sum;
if (row == k) {
float d = sqrtf(fmaxf(aik, 0.f));
col[k] = d;
sk[slot][k] = d;
} else {
// stash aik in a register; need inv after sync — use sk write only for diag
// For non-diag rows keep aik in col temporarily via reinterpret: store in unused
// upper half slot area is free for phase0 — use sk[slot][64+row] is wrong size.
// Keep classic path: write diag, sync, then scale using recomputed aik...
// Recompute is free enough with j<k already done; stash aik in col[k] temp:
if (row > k) col[k] = aik; // temp: pre-scale residual
}
}
__syncthreads();
if (half == 0) {
float inv = 1.f / sk[slot][k];
if (row > k) col[k] = col[k] * inv;
else if (row < k) col[k] = 0.f;
}
__syncthreads();
}
// Phase 1: k in [64,128)
#pragma unroll
for (int kk = 0; kk < 64; ++kk) {
const int k = 64 + kk;
if (row == k) {
if (half == 0) {
#pragma unroll
for (int j = 0; j < 64; ++j) sk[slot][j] = col[j];
} else {
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < kk) sk[slot][64 + j] = col[j];
}
}
}
__syncthreads();
float sum = 0.f;
if (half == 0) {
#pragma unroll
for (int j = 0; j < 64; ++j) {
sum = fmaf(col[j], sk[slot][j], sum);
}
} else {
#pragma unroll
for (int j = 0; j < 64; ++j) {
if (j < kk) sum = fmaf(col[j], sk[slot][64 + j], sum);
}
}
sum += __shfl_xor_sync(0xffffffffu, sum, 1);
if (half == 1) {
float aik = col[kk] - sum;
if (row == k) {
float d = sqrtf(fmaxf(aik, 0.f));
col[kk] = d;
sk[slot][k] = d;
} else if (row > k) {
col[kk] = aik; // temp residual
}
}
__syncthreads();
if (half == 1) {
float inv = 1.f / sk[slot][k];
if (row > k) col[kk] = col[kk] * inv;
else if (row < k) col[kk] = 0.f;
}
__syncthreads();
}
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v;
int j0 = col0 + t * 4;
v.x = (row >= (j0 + 0)) ? col[t * 4 + 0] : 0.f;
v.y = (row >= (j0 + 1)) ? col[t * 4 + 1] : 0.f;
v.z = (row >= (j0 + 2)) ? col[t * 4 + 2] : 0.f;
v.w = (row >= (j0 + 3)) ? col[t * 4 + 3] : 0.f;
*reinterpret_cast<float4*>(out + row * 128 + j0) = v;
}
}
}
// ---- V2: interleaved even/odd column halfreg ----
// col[j] holds global column (2*j + half). Both threads contribute ~k/2 terms.
template <int MPB>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_halfreg_interleave(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
const int half = tid & 1;
const int row = tid >> 1;
__shared__ float sk[MPB][128];
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 16384u;
float* out = L + (size_t)b * 16384u;
float col[64];
#pragma unroll
for (int j = 0; j < 64; ++j) {
col[j] = a[row * 128 + (j * 2 + half)];
}
#pragma unroll
for (int k = 0; k < 128; ++k) {
// Publish L[k, 0:k] from row-k's two halves
if (row == k) {
#pragma unroll
for (int j = 0; j < 64; ++j) {
int c = j * 2 + half;
if (c < k) sk[slot][c] = col[j];
}
}
__syncthreads();
float sum = 0.f;
#pragma unroll
for (int j = 0; j < 64; ++j) {
int c = j * 2 + half;
if (c < k) sum = fmaf(col[j], sk[slot][c], sum);
}
sum += __shfl_xor_sync(0xffffffffu, sum, 1);
// Owner of column k updates A[row,k]
if (half == (k & 1)) {
int jk = k >> 1;
float aik = col[jk] - sum;
if (row == k) {
float d = sqrtf(fmaxf(aik, 0.f));
col[jk] = d;
sk[slot][k] = d;
} else if (row > k) {
col[jk] = aik;
}
}
__syncthreads();
if (half == (k & 1)) {
int jk = k >> 1;
float inv = 1.f / sk[slot][k];
if (row > k) col[jk] = col[jk] * inv;
else if (row < k) col[jk] = 0.f;
}
__syncthreads();
}
// Store (scatter even/odd)
#pragma unroll
for (int j = 0; j < 64; ++j) {
int c = j * 2 + half;
out[row * 128 + c] = (row >= c) ? col[j] : 0.f;
}
}
}
// ---- V3: right-looking halfreg ----
// Publish L[:,k] to SMEM, rank-1 update trailing columns in both halves.
template <int MPB>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_halfreg_right(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
const int half = tid & 1;
const int row = tid >> 1;
const int col0 = half << 6;
__shared__ float colk[MPB][128];
__shared__ float diag_s[MPB];
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 16384u;
float* out = L + (size_t)b * 16384u;
float col[64];
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + row * 128 + col0 + t * 4);
col[t * 4 + 0] = v.x;
col[t * 4 + 1] = v.y;
col[t * 4 + 2] = v.z;
col[t * 4 + 3] = v.w;
}
#pragma unroll
for (int k = 0; k < 128; ++k) {
// Diag from owner of column k
const int kh = k >> 6; // which half owns k
const int kk = k & 63; // local index
if (half == kh && row == k) {
float d = sqrtf(fmaxf(col[kk], 0.f));
col[kk] = d;
diag_s[slot] = d;
}
__syncthreads();
float inv = 1.f / diag_s[slot];
float lik = 0.f;
if (half == kh) {
if (row > k) {
lik = col[kk] * inv;
col[kk] = lik;
} else if (row == k) {
lik = diag_s[slot];
} else {
col[kk] = 0.f;
}
}
// Broadcast full column k via SMEM (only half owner writes active rows)
if (half == kh) {
colk[slot][row] = (row >= k) ? ((row == k) ? diag_s[slot] : lik) : 0.f;
}
__syncthreads();
float lik_i = colk[slot][row];
// Rank-1: for all j > k owned by this half, update if row >= j
#pragma unroll
for (int jloc = 0; jloc < 64; ++jloc) {
int j = col0 + jloc;
if (j > k && row >= j) {
col[jloc] = fmaf(-lik_i, colk[slot][j], col[jloc]);
}
}
__syncthreads();
}
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v;
int j0 = col0 + t * 4;
v.x = (row >= (j0 + 0)) ? col[t * 4 + 0] : 0.f;
v.y = (row >= (j0 + 1)) ? col[t * 4 + 1] : 0.f;
v.z = (row >= (j0 + 2)) ? col[t * 4 + 2] : 0.f;
v.w = (row >= (j0 + 3)) ? col[t * 4 + 3] : 0.f;
*reinterpret_cast<float4*>(out + row * 128 + j0) = v;
}
}
}
// ---- V1: blocked BS multi-matrix, 128 thr/matrix, SYRK panel-in-regs ----
// Right-looking blocked: potf2 panel (intra-panel left-looking), rank-BS SYRK.
// Modal best: BS=32 MPB=1 ~75.6us @256x128 (torch ~151). Prior naive BS32 ~258.
// NOTE: do NOT float4-cast tile[row*129+col] — 129-stride breaks 16B alignment.
template <int MPB, int BS>
__global__ void __launch_bounds__(128 * MPB, 2)
cholesky128_blocked_mpb(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
extern __shared__ float smem[];
constexpr int TILE = 128 * 129;
float* tile = smem + slot * TILE;
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 16384u;
float* out = L + (size_t)b * 16384u;
#pragma unroll
for (int t = 0; t < 32; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + tid * 128 + t * 4);
int j0 = t * 4;
tile[tid * 129 + j0 + 0] = v.x;
tile[tid * 129 + j0 + 1] = v.y;
tile[tid * 129 + j0 + 2] = v.z;
tile[tid * 129 + j0 + 3] = v.w;
}
__syncthreads();
#pragma unroll
for (int kb = 0; kb < 128; kb += BS) {
#pragma unroll
for (int kk = 0; kk < BS; ++kk) {
const int k = kb + kk;
float sum = 0.f;
#pragma unroll
for (int j = 0; j < BS; ++j) {
if (j < kk) {
sum = fmaf(tile[tid * 129 + (kb + j)],
tile[k * 129 + (kb + j)], sum);
}
}
float aik = tile[tid * 129 + k] - sum;
if (tid == k) {
tile[k * 129 + k] = sqrtf(fmaxf(aik, 0.f));
}
__syncthreads();
float inv = 1.f / tile[k * 129 + k];
if (tid > k) tile[tid * 129 + k] = aik * inv;
else if (tid < k) tile[tid * 129 + k] = 0.f;
__syncthreads();
}
if (kb + BS < 128) {
float panel[BS];
#pragma unroll
for (int t = 0; t < BS; ++t) panel[t] = tile[tid * 129 + kb + t];
__syncthreads();
#pragma unroll 1
for (int j = kb + BS; j < 128; ++j) {
if (tid >= j) {
float s = 0.f;
#pragma unroll
for (int t = 0; t < BS; ++t) {
s = fmaf(panel[t], tile[j * 129 + kb + t], s);
}
tile[tid * 129 + j] -= s;
}
}
__syncthreads();
}
}
#pragma unroll
for (int t = 0; t < 32; ++t) {
float4 v;
int j0 = t * 4;
v.x = (tid >= (j0 + 0)) ? tile[tid * 129 + j0 + 0] : 0.f;
v.y = (tid >= (j0 + 1)) ? tile[tid * 129 + j0 + 1] : 0.f;
v.z = (tid >= (j0 + 2)) ? tile[tid * 129 + j0 + 2] : 0.f;
v.w = (tid >= (j0 + 3)) ? tile[tid * 129 + j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + tid * 128 + j0) = v;
}
if (b + stride < batch) __syncthreads();
}
}
// ---- V4: 256-thread CTA, BS=16, 2 threads/row for vectorized SYRK ----
// thr = row*2+lane; lane 0/1 split panel BS/2 for SYRK partial dots + xor reduce.
// Panel factor still uses lane0 as row owner (lane1 assists SYRK only).
template <int MPB, int BS>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_blocked_vecsyrk(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x; // 0..255
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
const int lane = tid & 1;
const int row = tid >> 1; // 0..127
extern __shared__ float smem[];
constexpr int TILE = 128 * 129;
float* tile = smem + slot * TILE;
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 16384u;
float* out = L + (size_t)b * 16384u;
// Cooperative load: 256 thr * 64 elems? each thr loads 64 floats of its row half
{
const int col0 = lane << 6;
#pragma unroll
for (int t = 0; t < 16; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + row * 128 + col0 + t * 4);
int j0 = col0 + t * 4;
tile[row * 129 + j0 + 0] = v.x;
tile[row * 129 + j0 + 1] = v.y;
tile[row * 129 + j0 + 2] = v.z;
tile[row * 129 + j0 + 3] = v.w;
}
}
__syncthreads();
#pragma unroll 1
for (int kb = 0; kb < 128; kb += BS) {
// Panel potf2/trsm: only lane0 rows (one thr per row)
#pragma unroll
for (int kk = 0; kk < BS; ++kk) {
const int k = kb + kk;
float sum = 0.f;
if (lane == 0) {
#pragma unroll
for (int j = 0; j < BS; ++j) {
if (j < kk) {
sum = fmaf(tile[row * 129 + (kb + j)],
tile[k * 129 + (kb + j)], sum);
}
}
float aik = tile[row * 129 + k] - sum;
if (row == k) {
tile[k * 129 + k] = sqrtf(fmaxf(aik, 0.f));
} else if (row > k) {
// stash
tile[row * 129 + k] = aik;
}
}
__syncthreads();
if (lane == 0) {
float inv = 1.f / tile[k * 129 + k];
if (row > k) tile[row * 129 + k] = tile[row * 129 + k] * inv;
else if (row < k) tile[row * 129 + k] = 0.f;
}
__syncthreads();
}
if (kb + BS < 128) {
// Each thr holds BS/2 panel elems for its row
float panel_half[BS / 2];
const int t0 = lane * (BS / 2);
#pragma unroll
for (int t = 0; t < BS / 2; ++t) {
panel_half[t] = tile[row * 129 + kb + t0 + t];
}
__syncthreads();
#pragma unroll 1
for (int j = kb + BS; j < 128; ++j) {
float s = 0.f;
if (row >= j) {
#pragma unroll
for (int t = 0; t < BS / 2; ++t) {
s = fmaf(panel_half[t], tile[j * 129 + kb + t0 + t], s);
}
}
s += __shfl_xor_sync(0xffffffffu, s, 1);
if (lane == 0 && row >= j) {
tile[row * 129 + j] -= s;
}
}
__syncthreads();
}
}
if (lane == 0) {
#pragma unroll
for (int t = 0; t < 32; ++t) {
float4 v;
int j0 = t * 4;
v.x = (row >= (j0 + 0)) ? tile[row * 129 + j0 + 0] : 0.f;
v.y = (row >= (j0 + 1)) ? tile[row * 129 + j0 + 1] : 0.f;
v.z = (row >= (j0 + 2)) ? tile[row * 129 + j0 + 2] : 0.f;
v.w = (row >= (j0 + 3)) ? tile[row * 129 + j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + row * 128 + j0) = v;
}
}
if (b + stride < batch) __syncthreads();
}
}
// ---- V5: pure right-looking rank-1, 128 thr, full SMEM tile ----
template <int MPB>
__global__ void __launch_bounds__(128 * MPB, 2)
cholesky128_right_rank1(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x;
const int slot = threadIdx.y;
const int stride = gridDim.x * MPB;
extern __shared__ float smem[];
constexpr int TILE = 128 * 129;
float* tile = smem + slot * (TILE + 128);
float* colk = tile + TILE;
for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
const float* a = A + (size_t)b * 16384u;
float* out = L + (size_t)b * 16384u;
#pragma unroll
for (int t = 0; t < 32; ++t) {
float4 v = *reinterpret_cast<const float4*>(a + tid * 128 + t * 4);
int j0 = t * 4;
tile[tid * 129 + j0 + 0] = v.x;
tile[tid * 129 + j0 + 1] = v.y;
tile[tid * 129 + j0 + 2] = v.z;
tile[tid * 129 + j0 + 3] = v.w;
}
__syncthreads();
#pragma unroll
for (int k = 0; k < 128; ++k) {
if (tid == k) {
tile[k * 129 + k] = sqrtf(fmaxf(tile[k * 129 + k], 0.f));
}
__syncthreads();
float inv = 1.f / tile[k * 129 + k];
float lik;
if (tid > k) {
lik = tile[tid * 129 + k] * inv;
tile[tid * 129 + k] = lik;
} else if (tid == k) {
lik = tile[k * 129 + k];
} else {
lik = 0.f;
tile[tid * 129 + k] = 0.f;
}
colk[tid] = lik;
__syncthreads();
float li = colk[tid];
#pragma unroll 1
for (int j = k + 1; j < 128; ++j) {
if (tid >= j) {
tile[tid * 129 + j] = fmaf(-li, colk[j], tile[tid * 129 + j]);
}
}
__syncthreads();
}
#pragma unroll
for (int t = 0; t < 32; ++t) {
float4 v;
int j0 = t * 4;
v.x = (tid >= (j0 + 0)) ? tile[tid * 129 + j0 + 0] : 0.f;
v.y = (tid >= (j0 + 1)) ? tile[tid * 129 + j0 + 1] : 0.f;
v.z = (tid >= (j0 + 2)) ? tile[tid * 129 + j0 + 2] : 0.f;
v.w = (tid >= (j0 + 3)) ? tile[tid * 129 + j0 + 3] : 0.f;
*reinterpret_cast<float4*>(out + tid * 128 + j0) = v;
}
if (b + stride < batch) __syncthreads();
}
}
torch::Tensor cholesky_n128(torch::Tensor A) {
TORCH_CHECK(A.is_cuda(), "A must be CUDA");
TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
TORCH_CHECK(A.dim() == 3 && A.size(1) == 128 && A.size(2) == 128, "expected (B,128,128)");
auto Ac = A.contiguous();
auto L = torch::empty_like(Ac);
const int batch = static_cast<int>(Ac.size(0));
if (batch == 0) return L;
constexpr int mpb = CHOLESKY128_MPB;
constexpr int bs = CHOLESKY128_BS;
#if CHOLESKY128_VARIANT == 1
{
dim3 block(128, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
const int shmem = mpb * (128 * 129) * (int)sizeof(float);
CHECK_CUDA(cudaFuncSetAttribute(
cholesky128_blocked_mpb<mpb, bs>,
cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
cholesky128_blocked_mpb<mpb, bs><<<grid, block, shmem>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#elif CHOLESKY128_VARIANT == 2
{
dim3 block(256, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
cholesky128_halfreg_interleave<mpb><<<grid, block>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#elif CHOLESKY128_VARIANT == 3
{
dim3 block(256, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
cholesky128_halfreg_right<mpb><<<grid, block>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#elif CHOLESKY128_VARIANT == 4
{
dim3 block(256, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
const int shmem = mpb * (128 * 129) * (int)sizeof(float);
CHECK_CUDA(cudaFuncSetAttribute(
cholesky128_blocked_vecsyrk<mpb, bs>,
cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
cholesky128_blocked_vecsyrk<mpb, bs><<<grid, block, shmem>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#elif CHOLESKY128_VARIANT == 5
{
dim3 block(128, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
const int shmem = mpb * (128 * 129 + 128) * (int)sizeof(float);
CHECK_CUDA(cudaFuncSetAttribute(
cholesky128_right_rank1<mpb>,
cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
cholesky128_right_rank1<mpb><<<grid, block, shmem>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#else
{
dim3 block(256, mpb);
int need = (batch + mpb - 1) / mpb;
int grid = need < 2048 ? need : 2048;
if (grid < 1) grid = 1;
cholesky128_halfreg_phase<mpb><<<grid, block>>>(
Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
}
#endif
CHECK_CUDA(cudaGetLastError());
return L;
}
// ---- n=256 dense-panel SYRK (beam_mid V1) ----
#ifndef CHOLESKY256_BS
#define CHOLESKY256_BS 32
#endif
__device__ __forceinline__ int lt256(int i, int j) {
return (i * (i + 1)) / 2 + j;
}
// ---- V1 STRUCTURAL: dense-panel SYRK + multi-j ILP (JU=4) ----
// After each panel potf, publish L[:, kb:kb+BS] into dense dens[row*(BS+1)+t]
// so SYRK is two dense vectors (far fewer bank conflicts than packed gather).
// Multi-j unrolling amortizes dens row loads for L[j,*].
template <int BS>
__global__ void __launch_bounds__(256, 1)
cholesky256_dense_syrk(const float* __restrict__ A,
float* __restrict__ L,
int batch) {
const int tid = threadIdx.x; // row
extern __shared__ float smem[];
constexpr int PACK = 256 * 257 / 2;
constexpr int DENS_STRIDE = BS + 1; // bank pad
float* tile = smem;
float* dens = smem + PACK; // 256 * DENS_STRIDE
for (int b = blockIdx.x; b < batch; b += gridDim.x) {
const float* a = A + (size_t)b * 65536u;
float* out = L + (size_t)b * 65536u;
#pragma unroll 1
for (int t = 0; t < 64; ++t) {
int j0 = t * 4;
float4 v = *reinterpret_cast<const float4*>(a + tid * 256 + j0);
if (j0 + 0 <= tid) tile[lt256(tid, j0 + 0)] = v.x;
if (j0 + 1 <= tid) tile[lt256(tid, j0 + 1)] = v.y;
if (j0 + 2 <= tid) tile[lt256(tid, j0 + 2)] = v.z;
if (j0 + 3 <= tid) tile[lt256(tid, j0 + 3)] = v.w;
}
__syncthreads();
#pragma unroll 1
for (int kb = 0; kb < 256; kb += BS) {
// Panel potf2 (left-looking within panel)
#pragma unroll
for (int kk = 0; kk < BS; ++kk) {
const int k = kb + kk;
float sum = 0.f;
float aik = 0.f;
if (tid >= k) {
#pragma unroll
for (int j = 0; j < BS; ++j) {
if (j < kk) {
sum = fmaf(tile[lt256(tid, kb + j)],
tile[lt256(k, kb + j)], sum);
}
}
aik = tile[lt256(tid, k)] - sum;
if (tid == k) {
tile[lt256(k, k)] = sqrtf(fmaxf(aik, 0.f));
}
}
__syncthreads();
if (tid > k) {
tile[lt256(tid, k)] = aik * (1.f / tile[lt256(k, k)]);
} else if (tid < k) {
// keep upper-of-panel zero in packed (already unused)
}
__syncthreads();
}
// Publish dense panel
#pragma unroll
for (int t = 0; t < BS; ++t) {
dens[tid * DENS_STRIDE + t] =
(tid >= (kb + t)) ? tile[lt256(tid, kb + t)] : 0.f;
}
__syncthreads();
if (kb + BS < 256) {
float my_panel[BS];
#pragma unroll
for (int t = 0; t < BS; ++t) my_panel[t] = dens[tid * DENS_STRIDE + t];
// Multi-j SYRK: update columns j in steps of 4
const int j0 = kb + BS;
const int nfull = ((256 - j0) / 4) * 4;
#pragma unroll 1
for (int jj = 0; jj < nfull; jj += 4) {
const int j = j0 + jj;
float s0 = 0.f, s1 = 0.f, s2 = 0.f, s3 = 0.f;
float p0[BS], p1[BS], p2[BS], p3[BS];
#pragma unroll
for (int t = 0; t < BS; ++t) {
p0[t] = dens[j * DENS_STRIDE + t];
p1[t] = dens[(j + 1) * DENS_STRIDE + t];
p2[t] = dens[(j + 2) * DENS_STRIDE + t];
p3[t] = dens[(j + 3) * DENS_STRIDE + t];
}
#pragma unroll
for (int t = 0; t < BS; ++t) {
float mp = my_panel[t];
s0 = fmaf(mp, p0[t], s0);
s1 = fmaf(mp, p1[t], s1);
s2 = fmaf(mp, p2[t], s2);
s3 = fmaf(mp, p3[t], s3);
}
if (tid >= j) tile[lt256(tid, j)] -= s0;
if (tid >= j + 1) tile[lt256(tid, j + 1)] -= s1;
if (tid >= j + 2) tile[lt256(tid, j + 2)] -= s2;
if (tid >= j + 3) tile[lt256(tid, j + 3)] -= s3;
}
#pragma unroll 1
for (int j = j0 + nfull; j < 256; ++j) {
float s = 0.f;
#pragma unroll
for (int t = 0; t < BS; ++t) {
s = fmaf(my_panel[t], dens[j * DENS_STRIDE + t], s);
}
if (tid >= j) tile[lt256(tid, j)] -= s;
}
__syncthreads();
}
}
#pragma unroll 1
for (int t = 0; t < 64; ++t) {
int j0 = t * 4;
float4 v;
v.x = (tid >= (j0 + 0)) ? tile[lt256(tid, j0 + 0)] : 0.f;
v.y = (tid >= (j0 + 1)) ? tile[lt256(tid, j0 + 1)] : 0.f;
v.z = (tid >= (j0 + 2)) ? tile[lt256(tid, j0 + 2)] : 0.f;
v.w = (tid >= (j0 + 3)) ? tile[lt256(tid, j0 + 3)] : 0.f;
*reinterpret_cast<float4*>(out + tid * 256 + j0) = v;
}
}
}
static void launch_n256(const float* A, float* L, int batch) {
constexpr int bs = CHOLESKY256_BS;
int grid = batch < 2048 ? batch : 2048;
if (grid < 1) grid = 1;
const int shmem = (256 * 257 / 2 + 256 * (bs + 1)) * (int)sizeof(float);
CHECK_CUDA(cudaFuncSetAttribute(
cholesky256_dense_syrk<bs>,
cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
cholesky256_dense_syrk<bs><<<grid, 256, shmem>>>(A, L, batch);
CHECK_CUDA(cudaGetLastError());
}
torch::Tensor cholesky_n256(torch::Tensor A) {
TORCH_CHECK(A.is_cuda(), "A must be CUDA");
TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
TORCH_CHECK(A.dim() == 3 && A.size(1) == 256 && A.size(2) == 256,
"expected (B,256,256)");
auto Ac = A.contiguous();
auto L = torch::empty_like(Ac);
const int batch = static_cast<int>(Ac.size(0));
if (batch == 0) return L;
launch_n256(Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
return L;
}
"""
N128_CPP_SRC = r"""
torch::Tensor cholesky_n128(torch::Tensor A);
torch::Tensor cholesky_n256(torch::Tensor A);
"""
_mod_n128 = load_inline(
name="chol_n128_n256_dense_v1",
cpp_sources=[N128_CPP_SRC],
cuda_sources=[N128_CUDA_SRC],
functions=["cholesky_n128", "cholesky_n256"],
verbose=False,
extra_cuda_cflags=[
"-O3",
"-DCHOLESKY128_VARIANT=1",
"-DCHOLESKY128_MPB=1",
"-DCHOLESKY128_BS=32",
"-DCHOLESKY256_BS=32",
],
)
# ---------------------------------------------------------------------------
# Medium high-batch: blocked right-looking with batched potf / TRSM / GEMM
# Row-major SPD as col-major FILL_UPPER → L in lower triangle (champ convention).
# ---------------------------------------------------------------------------
MEDIUM2_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cuda_runtime.h>
#include <algorithm>
#include <mutex>
#include <stdexcept>
#include <string>
// large3-style TF32 trailing knobs
#ifndef LARGE3_GEMMEX_M_MIN
#define LARGE3_GEMMEX_M_MIN 4096
#endif
#ifndef LARGE3_REFINE_TAIL
#define LARGE3_REFINE_TAIL 0
#endif
#ifndef LARGE3_TF32_MATH
#define LARGE3_TF32_MATH 1
#endif
#define CHECK_CUDA(expr) do { \
cudaError_t _err = (expr); \
if (_err != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(_err)); \
} \
} while (0)
#define CHECK_CUBLAS(expr) do { \
cublasStatus_t _st = (expr); \
if (_st != CUBLAS_STATUS_SUCCESS) { \
throw std::runtime_error(std::string("cuBLAS error ") + std::to_string((int)_st)); \
} \
} while (0)
#define CHECK_CUSOLVER(expr) do { \
cusolverStatus_t _st = (expr); \
if (_st != CUSOLVER_STATUS_SUCCESS) { \
throw std::runtime_error(std::string("cuSOLVER error ") + std::to_string((int)_st)); \
} \
} while (0)
static cublasHandle_t g_blas = nullptr;
static cusolverDnHandle_t g_solver = nullptr;
static std::once_flag g_once;
// Persistent device pointer arrays + potf workspace.
static float** g_d_A = nullptr;
static float** g_d_B = nullptr;
static int g_ptr_cap = 0;
static float* g_work = nullptr;
static int g_work_cap = 0;
static int* g_info = nullptr;
static int g_info_cap = 0;
static void init_handles() {
std::call_once(g_once, []() {
CHECK_CUBLAS(cublasCreate(&g_blas));
CHECK_CUBLAS(cublasSetMathMode(g_blas, CUBLAS_TF32_TENSOR_OP_MATH));
CHECK_CUSOLVER(cusolverDnCreate(&g_solver));
});
}
static void ensure_ptrs(int batch) {
if (batch > g_ptr_cap) {
if (g_d_A) CHECK_CUDA(cudaFree(g_d_A));
if (g_d_B) CHECK_CUDA(cudaFree(g_d_B));
CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_d_A),
(size_t)batch * sizeof(float*)));
CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_d_B),
(size_t)batch * sizeof(float*)));
g_ptr_cap = batch;
}
}
static float* ensure_work(int lwork) {
if (lwork > g_work_cap) {
if (g_work) CHECK_CUDA(cudaFree(g_work));
CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_work),
(size_t)std::max(lwork, 1) * sizeof(float)));
g_work_cap = std::max(lwork, 1);
}
return g_work;
}
static int* ensure_info(int ninfo) {
if (ninfo > g_info_cap) {
if (g_info) CHECK_CUDA(cudaFree(g_info));
CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_info),
(size_t)std::max(ninfo, 1) * sizeof(int)));
g_info_cap = std::max(ninfo, 1);
}
return g_info;
}
__global__ void zero_upper_kernel(float* __restrict__ L, int n, int batch) {
const int b = blockIdx.z;
const int i = blockIdx.y * blockDim.y + threadIdx.y;
const int j = blockIdx.x * blockDim.x + threadIdx.x;
if (b >= batch || i >= n || j >= n) return;
if (j > i) {
L[(size_t)b * (size_t)n * (size_t)n + (size_t)i * (size_t)n + (size_t)j] = 0.f;
}
}
// ptrs[b] = base + b * mat_stride + offset
__global__ void fill_offset_ptrs_kernel(
float** __restrict__ ptrs,
float* base,
long long mat_stride,
long long offset,
int batch) {
int b = blockIdx.x * blockDim.x + threadIdx.x;
if (b < batch) {
ptrs[b] = base + (long long)b * mat_stride + offset;
}
}
static void fill_ptrs(float** d_ptrs, float* base, long long mat_stride,
long long offset, int batch) {
int threads = 256;
int blocks = (batch + threads - 1) / threads;
if (blocks < 1) blocks = 1;
fill_offset_ptrs_kernel<<<blocks, threads>>>(
d_ptrs, base, mat_stride, offset, batch);
CHECK_CUDA(cudaGetLastError());
}
// mode: 0 = SpotrfBatched on each diagonal tile; 1 = host loop potf reused workspace
torch::Tensor cholesky_blocked_batched(torch::Tensor A, int64_t nb_in, int64_t mode_in) {
TORCH_CHECK(A.is_cuda() && A.scalar_type() == torch::kFloat32 && A.dim() == 3);
TORCH_CHECK(A.size(1) == A.size(2));
init_handles();
auto L = A.contiguous().clone();
const int batch = static_cast<int>(L.size(0));
const int n = static_cast<int>(L.size(1));
if (batch == 0) return L;
int nb = static_cast<int>(nb_in);
if (nb < 32) nb = 32;
if (nb > n) nb = n;
const int mode = static_cast<int>(mode_in);
float* Lptr = L.data_ptr<float>();
const long long mat_stride = (long long)n * (long long)n;
ensure_ptrs(batch);
int* d_info = ensure_info(std::max(batch, 1));
// Clear info once; potf writes per call (batched fills all).
CHECK_CUDA(cudaMemset(d_info, 0, (size_t)std::max(batch, 1) * sizeof(int)));
int lwork = 0;
CHECK_CUSOLVER(cusolverDnSpotrf_bufferSize(
g_solver, CUBLAS_FILL_MODE_UPPER, nb, Lptr, n, &lwork));
float* work = ensure_work(lwork);
const float one = 1.f;
const float minus_one = -1.f;
for (int k = 0; k < n; k += nb) {
const int kb = std::min(nb, n - k);
const long long off_kk = (long long)k + (long long)k * (long long)n;
// 1) Factor diagonal tiles Akk (kb x kb, lda = n)
if (mode == 0) {
fill_ptrs(g_d_A, Lptr, mat_stride, off_kk, batch);
CHECK_CUSOLVER(cusolverDnSpotrfBatched(
g_solver,
CUBLAS_FILL_MODE_UPPER,
kb,
g_d_A,
n,
d_info,
batch));
} else {
// Tight host loop, single workspace reused across batch items.
for (int b = 0; b < batch; ++b) {
float* Akk = Lptr + (long long)b * mat_stride + off_kk;
CHECK_CUSOLVER(cusolverDnSpotrf(
g_solver, CUBLAS_FILL_MODE_UPPER, kb, Akk, n,
work, lwork, d_info));
}
}
const int m = n - (k + kb);
if (m <= 0) continue;
const long long off_12 = (long long)k + (long long)(k + kb) * (long long)n;
const long long off_22 = (long long)(k + kb) + (long long)(k + kb) * (long long)n;
// 2) Panel TRSM: A12 := inv(Akk)^T * A12 (batched)
// SIDE_LEFT, UPPER, OP_T: Akk^T X = A12, X is kb x m col-major
fill_ptrs(g_d_A, Lptr, mat_stride, off_kk, batch);
fill_ptrs(g_d_B, Lptr, mat_stride, off_12, batch);
CHECK_CUBLAS(cublasStrsmBatched(
g_blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
kb,
m,
&one,
(const float**)g_d_A,
n,
g_d_B,
n,
batch));
// 3) Trailing GEMM: A22 -= A12^T * A12 (strided batched, TF32)
float* A12 = Lptr + off_12;
float* A22 = Lptr + off_22;
CHECK_CUBLAS(cublasSgemmStridedBatched(
g_blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
m,
m,
kb,
&minus_one,
A12,
n,
mat_stride,
A12,
n,
mat_stride,
&one,
A22,
n,
mat_stride,
batch));
}
dim3 block(16, 16);
dim3 grid((n + 15) / 16, (n + 15) / 16, batch);
zero_upper_kernel<<<grid, block>>>(Lptr, n, batch);
CHECK_CUDA(cudaGetLastError());
// No end-sync: leave work async like torch path.
return L;
}
// Trailing A22 -= A12^T @ A12.
// Large m: GemmEx FAST_TF32. Small m / refine tail: Sgemm under handle math mode
// (TF32_TENSOR_OP when LARGE3_TF32_MATH=1 - same residual as champ/large2).
static inline void trailing_update(
cublasHandle_t blas,
int m,
int kb,
int gemmex_m_min,
int force_fp32,
const float* A12,
float* A22,
int ld) {
const float one = 1.f;
const float minus_one = -1.f;
if (!force_fp32 && m >= gemmex_m_min) {
CHECK_CUBLAS(cublasGemmEx(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
m,
m,
kb,
&minus_one,
A12,
CUDA_R_32F,
ld,
A12,
CUDA_R_32F,
ld,
&one,
A22,
CUDA_R_32F,
ld,
CUBLAS_COMPUTE_32F_FAST_TF32,
CUBLAS_GEMM_DEFAULT_TENSOR_OP));
} else {
CHECK_CUBLAS(cublasSgemm(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
m,
m,
kb,
&minus_one,
A12,
ld,
A12,
ld,
&one,
A22,
ld));
}
}
static void potf_blocked_one(
float* A,
int n,
int nb,
int gemmex_m_min,
int refine_tail,
float* work,
int lwork,
int* d_info) {
const float one = 1.f;
cublasHandle_t blas = g_blas;
cusolverDnHandle_t solver = g_solver;
for (int k = 0; k < n; k += nb) {
const int kb = std::min(nb, n - k);
float* Akk = A + (size_t)k + (size_t)k * (size_t)n;
// Panel always FP32 cuSOLVER.
CHECK_CUSOLVER(cusolverDnSpotrf(
solver, CUBLAS_FILL_MODE_UPPER, kb, Akk, n, work, lwork, d_info));
const int m = n - (k + kb);
if (m <= 0) {
continue;
}
float* A12 = A + (size_t)k + (size_t)(k + kb) * (size_t)n;
float* A22 = A + (size_t)(k + kb) + (size_t)(k + kb) * (size_t)n;
CHECK_CUBLAS(cublasStrsm(
blas,
CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER,
CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT,
kb,
m,
&one,
Akk,
n,
A12,
n));
// refine_tail: prefer Sgemm path for small remaining m (residual knob;
// under TF32 math mode residual stays ~champ - true FP32 needs DEFAULT math).
const int force_fp32 = (refine_tail > 0 && m <= refine_tail) ? 1 : 0;
trailing_update(blas, m, kb, gemmex_m_min, force_fp32, A12, A22, n);
}
}
torch::Tensor cholesky_tf32_blocked(
torch::Tensor A,
int64_t nb_in,
int64_t gemmex_m_min_in,
int64_t refine_tail_in) {
TORCH_CHECK(A.is_cuda() && A.is_contiguous());
TORCH_CHECK(A.scalar_type() == torch::kFloat32);
TORCH_CHECK(A.dim() == 3 && A.size(1) == A.size(2));
init_handles();
const int batch = static_cast<int>(A.size(0));
const int n = static_cast<int>(A.size(1));
int nb = static_cast<int>(nb_in);
if (nb < 32) nb = 32;
if (nb > n) nb = n;
int gemmex_m_min = static_cast<int>(gemmex_m_min_in);
if (gemmex_m_min < 0) gemmex_m_min = LARGE3_GEMMEX_M_MIN;
int refine_tail = static_cast<int>(refine_tail_in);
if (refine_tail < 0) refine_tail = LARGE3_REFINE_TAIL;
auto L = A.clone();
float* Lptr = L.data_ptr<float>();
const long long mat_stride = (long long)n * (long long)n;
int lwork = 0;
CHECK_CUSOLVER(cusolverDnSpotrf_bufferSize(
g_solver, CUBLAS_FILL_MODE_UPPER, nb, Lptr, n, &lwork));
auto work = torch::empty({std::max(lwork, 1)}, A.options());
auto info = torch::zeros(
{1}, torch::TensorOptions().dtype(torch::kInt32).device(A.device()));
float* wptr = work.data_ptr<float>();
int* iptr = info.data_ptr<int>();
if (batch == 1) {
potf_blocked_one(Lptr, n, nb, gemmex_m_min, refine_tail, wptr, lwork, iptr);
} else {
for (int b = 0; b < batch; ++b) {
potf_blocked_one(
Lptr + (long long)b * mat_stride,
n,
nb,
gemmex_m_min,
refine_tail,
wptr,
lwork,
iptr);
}
}
dim3 block(16, 16);
dim3 grid((n + 15) / 16, (n + 15) / 16, batch);
zero_upper_kernel<<<grid, block>>>(Lptr, n, batch);
CHECK_CUDA(cudaGetLastError());
// No host device-sync (large2 speed path); torch/popcorn sync externally.
return L;
}
"""
MEDIUM2_CPP_SRC = r"""
torch::Tensor cholesky_blocked_batched(torch::Tensor A, int64_t nb_in, int64_t mode_in);
torch::Tensor cholesky_tf32_blocked(
torch::Tensor A,
int64_t nb_in,
int64_t gemmex_m_min_in,
int64_t refine_tail_in);
"""
# Single extension (medium batched + large TF32) — 3 load_inline total like champ.
_mod_lib = load_inline(
name="chol_hybrid_m2l3_v3",
cpp_sources=[MEDIUM2_CPP_SRC],
cuda_sources=[MEDIUM2_CUDA_SRC],
functions=["cholesky_blocked_batched", "cholesky_tf32_blocked"],
# large3 gemmex + medium2 batched
verbose=False,
extra_cuda_cflags=[
"-O3",
"-DLARGE3_GEMMEX_M_MIN=4096",
"-DLARGE3_REFINE_TAIL=0",
"-DLARGE3_TF32_MATH=1",
],
extra_ldflags=["-lcublas", "-lcusolver"],
)
def _choose_nb_large(n: int) -> int:
# large3: 2048@32k was best public 88.7ms; 512@16k; 256@4k
if n >= 32768:
return 2048
if n >= 16384:
return 512
if n >= 4096:
return 256
return 128
def _choose_gemmex_m_min(n: int) -> int:
if n >= 4096:
return 4096 if n >= 16384 else 2048
return 10**9
def _choose_nb_medium(n: int) -> int:
# medium3 Modal: nb=256@512 edges 128 (~1% @640×512); same for 1024/2048
if n >= 512:
return 256
return 64
def _use_tf32_blocked(batch: int, n: int) -> bool:
if n >= 16384:
return True
if n >= 4096 and batch >= 2:
return True
return False
def _use_medium_blocked(batch: int, n: int) -> bool:
"""High-batch only. 2×2048 stays torch (Popcorn torch faster than medium)."""
if n == 512 and batch >= 64:
return True
if n == 1024 and batch >= 16:
return True
if n == 2048 and batch >= 8:
return True
return False
_MEDIUM_POTF_MODE = 0
def custom_kernel(data: input_t) -> output_t:
A = data if data.is_contiguous() else data.contiguous()
batch, n, n2 = A.shape
if n != n2:
raise ValueError("expected square matrices")
if n == 32:
return _mod_small.cholesky_n32(A)
if n == 64:
return _mod_small.cholesky_n64(A)
if n == 128:
return _mod_n128.cholesky_n128(A)
if n == 256:
return _mod_n128.cholesky_n256(A)
b_i, n_i = int(batch), int(n)
if _use_medium_blocked(b_i, n_i):
return _mod_lib.cholesky_blocked_batched(
A, _choose_nb_medium(n_i), _MEDIUM_POTF_MODE
)
if _use_tf32_blocked(b_i, n_i):
return _mod_lib.cholesky_tf32_blocked(
A, _choose_nb_large(n_i), _choose_gemmex_m_min(n_i), 0
)
return torch.linalg.cholesky_ex(A, check_errors=False).L
scrolls · 1792 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