submission 926178
.grayveyard · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1460 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-926178?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:955f53abf22ee8c6c69a1a7cea613d20f149d4ba6f323a6f0532dab4187ba452
license declaredunknown
license concludedunknown
authors.grayveyard
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float smem[4][32 * 33];vector-width = float4
const float4* A4 = reinterpret_cast<const float4*>(Am);Kernel source
submission.py1460 lines
import os
import torch
from task import input_t, output_t
_CPP_SRC = r"""
torch::Tensor chol(torch::Tensor A);
"""
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cublas_v2.h>
#include <cstdlib>
#include <mutex>
#include <unordered_map>
#define CUBLAS_OK(x) TORCH_CHECK((x) == CUBLAS_STATUS_SUCCESS, "cuBLAS error at ", __FILE__, ":", __LINE__)
namespace {
constexpr float kTiny = 1.175494351e-38f;
// ---------------------------------------------------------------------------
// Kernel: batched Cholesky for n <= 32, one warp per matrix, 4 matrices/block.
// Deferred-scaling right-looking: tile holds unscaled Schur values t_ij = L_ij * d_j.
// ---------------------------------------------------------------------------
__global__ void chol32_kernel(const float* __restrict__ A, float* __restrict__ L,
int batch, int n) {
__shared__ float smem[4][32 * 33];
__shared__ float sinvd[4][32];
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const long long m = (long long)blockIdx.x * 4 + warp;
if (m >= batch) return;
const float* Am = A + m * (long long)n * n;
float* Lm = L + m * (long long)n * n;
float* t = smem[warp];
float* invd = sinvd[warp];
const int nn = n * n;
if ((n & 3) == 0) {
const float4* A4 = reinterpret_cast<const float4*>(Am);
const int nn4 = nn >> 2;
for (int v = lane; v < nn4; v += 32) {
float4 x = A4[v];
int e = v << 2;
int r = e / n, c = e - r * n;
float* dst = &t[r * 33 + c];
dst[0] = x.x; dst[1] = x.y; dst[2] = x.z; dst[3] = x.w;
}
} else {
for (int v = lane; v < nn; v += 32) {
int r = v / n, c = v - r * n;
t[r * 33 + c] = Am[v];
}
}
__syncwarp();
for (int j = 0; j < n; ++j) {
const float s = fmaxf(t[j * 33 + j], kTiny);
const float inv2 = 1.0f / s;
if (lane == j) invd[j] = rsqrtf(s);
const int k = lane;
if (k > j && k < n) {
const float ckj = t[k * 33 + j] * inv2;
for (int i = k; i < n; ++i) {
t[i * 33 + k] -= t[i * 33 + j] * ckj;
}
}
__syncwarp();
}
if ((n & 3) == 0) {
float4* L4 = reinterpret_cast<float4*>(Lm);
const int nn4 = nn >> 2;
for (int v = lane; v < nn4; v += 32) {
int e = v << 2;
int r = e / n, c = e - r * n;
float4 y;
y.x = (r >= c ) ? t[r * 33 + c ] * invd[c ] : 0.f;
y.y = (r >= c + 1) ? t[r * 33 + c + 1] * invd[c + 1] : 0.f;
y.z = (r >= c + 2) ? t[r * 33 + c + 2] * invd[c + 2] : 0.f;
y.w = (r >= c + 3) ? t[r * 33 + c + 3] * invd[c + 3] : 0.f;
L4[v] = y;
}
} else {
for (int v = lane; v < nn; v += 32) {
int r = v / n, c = v - r * n;
Lm[v] = (r >= c) ? t[r * 33 + c] * invd[c] : 0.f;
}
}
}
// ---------------------------------------------------------------------------
// Device: in-smem Cholesky of an n x n tile (n <= NMAX) held in t (row stride
// LDT = NMAX + 1). Left-looking 16-column panels: ILP-friendly block updates,
// warp-0 shuffle factor of the 16x16 diagonal, independent row solves below.
// Produces TRUE L in the lower triangle of t; invd[j] = 1 / L[j][j].
// ---------------------------------------------------------------------------
template <int NMAX, int THREADS>
__device__ __forceinline__ void factor_tile_smem(float* t, float* invd, int n,
int tid, int lane, int wid) {
constexpr int LDT = NMAX + 1;
for (int p = 0; p < n; p += 16) {
const int pe = min(p + 16, n);
// (a) update panel cols [p,pe) rows [p,n) against columns [0,p)
if (p > 0) {
const int rows = n - p, cols = pe - p;
for (int v = tid; v < rows * cols; v += THREADS) {
const int i = p + v / cols;
const int c = p + v % cols;
if (c > i) continue;
const float* ri = &t[i * LDT];
const float* rc = &t[c * LDT];
float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
int k = 0;
#pragma unroll 1
for (; k + 3 < p; k += 4) {
a0 += ri[k] * rc[k];
a1 += ri[k + 1] * rc[k + 1];
a2 += ri[k + 2] * rc[k + 2];
a3 += ri[k + 3] * rc[k + 3];
}
float acc = (a0 + a1) + (a2 + a3);
for (; k < p; ++k) acc += ri[k] * rc[k];
t[i * LDT + c] -= acc;
}
}
__syncthreads();
// (b) warp 0 factors the diagonal block; lane l owns row p + l
if (wid == 0) {
const int i = p + lane;
for (int j = p; j < pe; ++j) {
const float dj = fmaxf(t[j * LDT + j], kTiny);
const float lj = sqrtf(dj);
const float rj = 1.0f / lj;
float x = 0.f;
if (i > j && i < pe) x = t[i * LDT + j];
if (lane == j - p) {
t[j * LDT + j] = lj;
invd[j] = rj;
}
const float li = x * rj; // L[i][j] for this lane's row
if (i > j && i < pe) t[i * LDT + j] = li;
__syncwarp();
#pragma unroll 1
for (int c = j + 1; c < pe; ++c) {
const float lc = __shfl_sync(0xffffffffu, li, c - p);
if (i >= c && i < pe) t[i * LDT + c] -= li * lc;
}
__syncwarp();
}
}
__syncthreads();
// (c) rows below the block: independent row solves against the 16x16 L
for (int i = pe + tid; i < n; i += THREADS) {
#pragma unroll 1
for (int j = p; j < pe; ++j) {
float acc = t[i * LDT + j];
#pragma unroll 1
for (int k = p; k < j; ++k) {
acc -= t[i * LDT + k] * t[j * LDT + k];
}
t[i * LDT + j] = acc * invd[j];
}
}
__syncthreads();
}
}
// ---------------------------------------------------------------------------
// Kernel: batched Cholesky of an n x n tile (n <= NMAX), one block per matrix.
// Generalized src/dst with leading dims so it also factors diagonal blocks of a
// larger matrix in-place. Deferred scaling, one __syncthreads per column.
// ---------------------------------------------------------------------------
template <int NMAX, int THREADS>
__global__ void chol_tile_kernel(const float* __restrict__ src, long long src_stride, int src_ld,
float* __restrict__ dst, long long dst_stride, int dst_ld,
int n, int zero_upper) {
constexpr int LDT = NMAX + 1;
extern __shared__ float sh[];
float* t = sh; // NMAX * LDT
float* invd = sh + NMAX * LDT; // NMAX
const int tid = threadIdx.x;
const int lane = tid & 31;
const int wid = tid >> 5;
constexpr int NWARPS = THREADS / 32;
const float* Sm = src + (long long)blockIdx.x * src_stride;
float* Dm = dst + (long long)blockIdx.x * dst_stride;
// load
if ((n & 3) == 0 && (src_ld & 3) == 0) {
const int rowv = n >> 2;
const int tot = n * rowv;
for (int v = tid; v < tot; v += THREADS) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
const float4 x = *reinterpret_cast<const float4*>(Sm + (long long)r * src_ld + c4);
float* d = &t[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
} else {
for (int v = tid; v < n * n; v += THREADS) {
int r = v / n, c = v - r * n;
t[r * LDT + c] = Sm[(long long)r * src_ld + c];
}
}
__syncthreads();
factor_tile_smem<NMAX, THREADS>(t, invd, n, tid, lane, wid);
// store: t already holds true L in its lower triangle
if ((n & 3) == 0 && (dst_ld & 3) == 0) {
const int rowv = n >> 2;
const int tot = n * rowv;
for (int v = tid; v < tot; v += THREADS) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
float4 y;
y.x = (r >= c4 ) ? t[r * LDT + c4 ] : (zero_upper ? 0.f : Dm[(long long)r * dst_ld + c4 ]);
y.y = (r >= c4 + 1) ? t[r * LDT + c4 + 1] : (zero_upper ? 0.f : Dm[(long long)r * dst_ld + c4 + 1]);
y.z = (r >= c4 + 2) ? t[r * LDT + c4 + 2] : (zero_upper ? 0.f : Dm[(long long)r * dst_ld + c4 + 2]);
y.w = (r >= c4 + 3) ? t[r * LDT + c4 + 3] : (zero_upper ? 0.f : Dm[(long long)r * dst_ld + c4 + 3]);
*reinterpret_cast<float4*>(Dm + (long long)r * dst_ld + c4) = y;
}
} else {
for (int v = tid; v < n * n; v += THREADS) {
int r = v / n, c = v - r * n;
if (r >= c) Dm[(long long)r * dst_ld + c] = t[r * LDT + c];
else if (zero_upper) Dm[(long long)r * dst_ld + c] = 0.f;
}
}
}
// ---------------------------------------------------------------------------
// Kernel: fused batched Cholesky for n == 2*TILE, one block per matrix.
// All three quadrant tiles live in dynamic smem simultaneously:
// phase A: factor A11; phase B: TRSM L21 (while warps load A22);
// phase C: SYRK t22 -= L21 L21^T; phase D: factor t22. One launch total.
// ---------------------------------------------------------------------------
template <int TILE, int THREADS>
__global__ void chol2x2_kernel(const float* __restrict__ A, float* __restrict__ L) {
constexpr int LDT = TILE + 1;
constexpr int N = 2 * TILE;
extern __shared__ float sh[];
float* t11 = sh;
float* t21 = sh + TILE * LDT;
float* t22 = sh + 2 * TILE * LDT;
float* invd = sh + 3 * TILE * LDT; // TILE entries
const int tid = threadIdx.x;
const int lane = tid & 31;
const int wid = tid >> 5;
constexpr int NWARPS = THREADS / 32;
const float* Am = A + (long long)blockIdx.x * N * N;
float* Lm = L + (long long)blockIdx.x * N * N;
// ---- load A11 and A21 (row-vectorized float4) ----
{
constexpr int rowv = TILE / 4;
for (int v = tid; v < TILE * rowv; v += THREADS) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
float4 x = *reinterpret_cast<const float4*>(Am + (long long)r * N + c4);
float* d = &t11[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
x = *reinterpret_cast<const float4*>(Am + (long long)(r + TILE) * N + c4);
d = &t21[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
}
__syncthreads();
// ---- phase A: factor t11 (produces true L11, invd = 1/diag) ----
factor_tile_smem<TILE, THREADS>(t11, invd, TILE, tid, lane, wid);
// ---- store L11 quadrant + zero A12 quadrant; load A22 with spare warps ----
{
constexpr int rowv = TILE / 4;
for (int v = tid; v < TILE * rowv; v += THREADS) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
float4 y;
y.x = (r >= c4 ) ? t11[r * LDT + c4 ] : 0.f;
y.y = (r >= c4 + 1) ? t11[r * LDT + c4 + 1] : 0.f;
y.z = (r >= c4 + 2) ? t11[r * LDT + c4 + 2] : 0.f;
y.w = (r >= c4 + 3) ? t11[r * LDT + c4 + 3] : 0.f;
*reinterpret_cast<float4*>(Lm + (long long)r * N + c4) = y;
float* tb = &t11[r * LDT + c4];
tb[0] = y.x; tb[1] = y.y; tb[2] = y.z; tb[3] = y.w;
const float4 z = {0.f, 0.f, 0.f, 0.f};
*reinterpret_cast<float4*>(Lm + (long long)r * N + TILE + c4) = z;
// A22 tile load
float4 x = *reinterpret_cast<const float4*>(Am + (long long)(r + TILE) * N + TILE + c4);
float* d = &t22[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
}
__syncthreads();
// ---- phase B: TRSM X * L11^T = A21, thread r owns row r ----
if (tid < TILE) {
const int r = tid;
#pragma unroll 1
for (int jb = 0; jb < TILE; jb += 16) {
#pragma unroll 1
for (int j = jb; j < jb + 16; ++j) {
float acc = t21[r * LDT + j];
#pragma unroll 1
for (int k = jb; k < j; ++k) {
acc -= t21[r * LDT + k] * t11[j * LDT + k];
}
t21[r * LDT + j] = acc * invd[j];
}
int c = jb + 16;
#pragma unroll 1
for (; c + 3 < TILE; c += 4) {
float s0 = 0.f, s1 = 0.f, s2 = 0.f, s3 = 0.f;
#pragma unroll
for (int k = jb; k < jb + 16; ++k) {
const float xv = t21[r * LDT + k];
s0 += xv * t11[(c ) * LDT + k];
s1 += xv * t11[(c + 1) * LDT + k];
s2 += xv * t11[(c + 2) * LDT + k];
s3 += xv * t11[(c + 3) * LDT + k];
}
t21[r * LDT + c ] -= s0;
t21[r * LDT + c + 1] -= s1;
t21[r * LDT + c + 2] -= s2;
t21[r * LDT + c + 3] -= s3;
}
#pragma unroll 1
for (; c < TILE; ++c) {
float s = 0.f;
#pragma unroll
for (int k = jb; k < jb + 16; ++k) {
s += t21[r * LDT + k] * t11[c * LDT + k];
}
t21[r * LDT + c] -= s;
}
}
}
__syncthreads();
// ---- store L21 quadrant ----
{
constexpr int rowv = TILE / 4;
for (int v = tid; v < TILE * rowv; v += THREADS) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
const float* s2 = &t21[r * LDT + c4];
float4 y{s2[0], s2[1], s2[2], s2[3]};
*reinterpret_cast<float4*>(Lm + (long long)(r + TILE) * N + c4) = y;
}
}
// ---- phase C: t22 -= L21 * L21^T (lower triangle only) ----
// enumerate (i, jp): i in [0,TILE), jp in [0, i/2], j = 2*jp; item count
// before row i is f(i) = floor((i+1)^2 / 4).
{
for (int p = tid; p < TILE * (TILE + 2) / 4; p += THREADS) {
int i = (int)(2.f * sqrtf((float)p)) - 1;
if (i < 0) i = 0;
while ((i + 2) * (i + 2) / 4 <= p) ++i;
while ((i + 1) * (i + 1) / 4 > p) --i;
const int jp = p - (i + 1) * (i + 1) / 4;
const int j = jp * 2;
float acc0 = 0.f, acc1 = 0.f;
const float* ri = &t21[i * LDT];
const float* rj0 = &t21[j * LDT];
const float* rj1 = &t21[(j + 1) * LDT];
#pragma unroll 8
for (int k = 0; k < TILE; ++k) {
const float a = ri[k];
acc0 += a * rj0[k];
acc1 += a * rj1[k];
}
t22[i * LDT + j] -= acc0;
if (j + 1 <= i) t22[i * LDT + j + 1] -= acc1;
}
}
__syncthreads();
// ---- phase D: factor t22 (true L) ----
factor_tile_smem<TILE, THREADS>(t22, invd, TILE, tid, lane, wid);
// ---- store L22 quadrant ----
{
constexpr int rowv = TILE / 4;
for (int v = tid; v < TILE * rowv; v += THREADS) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
float4 y;
y.x = (r >= c4 ) ? t22[r * LDT + c4 ] : 0.f;
y.y = (r >= c4 + 1) ? t22[r * LDT + c4 + 1] : 0.f;
y.z = (r >= c4 + 2) ? t22[r * LDT + c4 + 2] : 0.f;
y.w = (r >= c4 + 3) ? t22[r * LDT + c4 + 3] : 0.f;
*reinterpret_cast<float4*>(Lm + (long long)(r + TILE) * N + TILE + c4) = y;
}
}
}
// ---------------------------------------------------------------------------
// Kernel: persistent whole-matrix Cholesky, ONE launch for 512 <= n <= 4096
// (n multiple of 128). Grid = (G, B): G cooperating blocks per matrix using a
// device-side sense barrier; right-looking nb=128 with register-tiled fp32
// SYRK (8x8 per thread). Assumes out holds a full copy of A's lower triangle
// is NOT needed: blocks copy tril(A) in-kernel first.
// ---------------------------------------------------------------------------
__device__ __forceinline__ void gbarrier(int* cnt, int* gen, int G, int tid) {
__syncthreads();
if (tid == 0) {
__threadfence();
const int g = *(volatile int*)gen;
if (atomicAdd(cnt, 1) == G - 1) {
*cnt = 0;
__threadfence();
atomicAdd(gen, 1);
} else {
while (*(volatile int*)gen == g) { __nanosleep(64); }
}
__threadfence();
}
__syncthreads();
}
__global__ void chol_persist_kernel(const float* __restrict__ A, float* __restrict__ out,
int* __restrict__ bars, int n, int G) {
constexpr int LDT = 129;
extern __shared__ float sh[];
float* s1 = sh; // 128*129 tile (L11 / Li)
float* s2 = sh + 128 * LDT; // 128*129 tile (rows / Lj)
float* invd = sh + 2 * 128 * LDT; // 128
const int tid = threadIdx.x; // 256 threads
const int lane = tid & 31;
const int wid = tid >> 5;
const int g = blockIdx.x;
const long long m = blockIdx.y;
const float* Am = A + m * (long long)n * n;
float* Om = out + m * (long long)n * n;
int* cnt = bars + m * 2;
int* gen = cnt + 1;
const int nt = n >> 7; // 128-tiles per side
// ---- copy tril(A) -> out, distributed over blocks (row-tile granularity) ----
{
const int rowv = n >> 2;
for (int rt = g; rt < nt; rt += G) {
const int r0 = rt << 7;
const long long base = (long long)r0 * n;
// rows r0..r0+127, columns 0..r0+127 (full tiles up to diagonal)
const int quads = ((r0 + 128) >> 2);
for (int v = tid; v < 128 * quads; v += 256) {
const int r = r0 + v / quads;
const int c4 = (v - (v / quads) * quads) << 2;
if (c4 > r) continue;
float4 x = *reinterpret_cast<const float4*>(Am + (long long)r * n + c4);
if (c4 + 3 > r) {
if (c4 + 1 > r) x.y = 0.f;
if (c4 + 2 > r) x.z = 0.f;
x.w = 0.f;
}
*reinterpret_cast<float4*>(Om + (long long)r * n + c4) = x;
}
(void)base; (void)rowv;
}
}
if (G > 1) gbarrier(cnt, gen, G, tid); else __syncthreads();
for (int k = 0; k < nt; ++k) {
const long long q = (long long)k << 7;
// ---- factor diagonal tile (block g == k % G does it) ----
if (g == (k % G)) {
for (int v = tid; v < 128 * 32; v += 256) {
const int r = v >> 5, c4 = (v & 31) << 2;
const float4 x = *reinterpret_cast<const float4*>(Om + (q + r) * n + q + c4);
float* d = &s1[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
__syncthreads();
factor_tile_smem<128, 256>(s1, invd, 128, tid, lane, wid);
for (int v = tid; v < 128 * 32; v += 256) {
const int r = v >> 5, c4 = (v & 31) << 2;
float4 y;
y.x = (r >= c4 ) ? s1[r * LDT + c4 ] : 0.f;
y.y = (r >= c4 + 1) ? s1[r * LDT + c4 + 1] : 0.f;
y.z = (r >= c4 + 2) ? s1[r * LDT + c4 + 2] : 0.f;
y.w = (r >= c4 + 3) ? s1[r * LDT + c4 + 3] : 0.f;
*reinterpret_cast<float4*>(Om + (q + r) * n + q + c4) = y;
}
}
if (G > 1) gbarrier(cnt, gen, G, tid); else __syncthreads();
if (k + 1 == nt) break;
// ---- TRSM: row-tiles below diagonal, distributed ----
// stage L11 into s1 (every block; the factoring block already has it)
for (int v = tid; v < 128 * 32; v += 256) {
const int r = v >> 5, c4 = (v & 31) << 2;
const float4 x = *reinterpret_cast<const float4*>(Om + (q + r) * n + q + c4);
float* d = &s1[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
if (tid < 128) invd[tid] = 1.0f / s1[tid * LDT + tid];
__syncthreads();
for (int rt = k + 1 + g; rt < nt; rt += G) {
const long long r0 = (long long)rt << 7;
// load 128 rows x 128 cols into s2
for (int v = tid; v < 128 * 32; v += 256) {
const int r = v >> 5, c4 = (v & 31) << 2;
const float4 x = *reinterpret_cast<const float4*>(Om + (r0 + r) * n + q + c4);
float* d = &s2[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
__syncthreads();
// two threads... one row per thread pair is complex; rows 0..127 over 256 thr:
if (tid < 128) {
const int r = tid;
#pragma unroll 1
for (int jb = 0; jb < 128; jb += 16) {
#pragma unroll 1
for (int j = jb; j < jb + 16; ++j) {
float acc = s2[r * LDT + j];
#pragma unroll 1
for (int kk = jb; kk < j; ++kk) {
acc -= s2[r * LDT + kk] * s1[j * LDT + kk];
}
s2[r * LDT + j] = acc * invd[j];
}
int c = jb + 16;
#pragma unroll 1
for (; c + 3 < 128; c += 4) {
float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
#pragma unroll
for (int kk = jb; kk < jb + 16; ++kk) {
const float xv = s2[r * LDT + kk];
a0 += xv * s1[(c ) * LDT + kk];
a1 += xv * s1[(c + 1) * LDT + kk];
a2 += xv * s1[(c + 2) * LDT + kk];
a3 += xv * s1[(c + 3) * LDT + kk];
}
s2[r * LDT + c ] -= a0;
s2[r * LDT + c + 1] -= a1;
s2[r * LDT + c + 2] -= a2;
s2[r * LDT + c + 3] -= a3;
}
}
}
__syncthreads();
for (int v = tid; v < 128 * 32; v += 256) {
const int r = v >> 5, c4 = (v & 31) << 2;
const float* sp = &s2[r * LDT + c4];
float4 y{sp[0], sp[1], sp[2], sp[3]};
*reinterpret_cast<float4*>(Om + (r0 + r) * n + q + c4) = y;
}
__syncthreads();
}
if (G > 1) gbarrier(cnt, gen, G, tid); else __syncthreads();
// ---- SYRK: trailing tile-pairs (it >= jt), distributed ----
const int T = nt - k - 1;
const int npair = T * (T + 1) / 2;
for (int pj = g; pj < npair; pj += G) {
// enumerate (it, jt): f(it) = it*(it+1)/2 <= pj
int it = (int)((-1.0f + sqrtf(1.0f + 8.0f * (float)pj)) * 0.5f);
while ((it + 1) * (it + 2) / 2 <= pj) ++it;
while (it * (it + 1) / 2 > pj) --it;
const int jt = pj - it * (it + 1) / 2;
const long long ri = (long long)(k + 1 + it) << 7;
const long long rj = (long long)(k + 1 + jt) << 7;
// stage Li -> s1, Lj -> s2 (panel columns at q)
for (int v = tid; v < 128 * 32; v += 256) {
const int r = v >> 5, c4 = (v & 31) << 2;
float4 x = *reinterpret_cast<const float4*>(Om + (ri + r) * n + q + c4);
float* d = &s1[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
x = *reinterpret_cast<const float4*>(Om + (rj + r) * n + q + c4);
d = &s2[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
__syncthreads();
// 8x8 per thread: 16x16 thread grid covers 128x128
const int ty = tid >> 4, tx = tid & 15;
const int rb = ty << 3, cb = tx << 3;
float acc[8][8];
#pragma unroll
for (int u = 0; u < 8; ++u)
#pragma unroll
for (int v2 = 0; v2 < 8; ++v2) acc[u][v2] = 0.f;
#pragma unroll 1
for (int kk = 0; kk < 128; ++kk) {
float a[8], b[8];
#pragma unroll
for (int u = 0; u < 8; ++u) a[u] = s1[(rb + u) * LDT + kk];
#pragma unroll
for (int v2 = 0; v2 < 8; ++v2) b[v2] = s2[(cb + v2) * LDT + kk];
#pragma unroll
for (int u = 0; u < 8; ++u)
#pragma unroll
for (int v2 = 0; v2 < 8; ++v2) acc[u][v2] += a[u] * b[v2];
}
// apply: C[ri..][rj..] -= acc (only lower part when it == jt)
const bool diag_tile = (it == jt);
#pragma unroll
for (int u = 0; u < 8; ++u) {
const long long grow = ri + rb + u;
#pragma unroll
for (int v2 = 0; v2 < 8; ++v2) {
const long long gcol = rj + cb + v2;
if (!diag_tile || gcol <= grow) {
Om[grow * n + gcol] -= acc[u][v2];
}
}
}
__syncthreads();
}
if (G > 1) gbarrier(cnt, gen, G, tid); else __syncthreads();
}
// zero float4 quads in tiles strictly above the diagonal tile row
// (within-tile uppers were already zeroed by the factor store)
{
const int rowv = n >> 2;
const long long total4 = (long long)n * rowv;
for (long long v = (long long)g * 256 + tid; v < total4; v += (long long)G * 256) {
const int r = (int)(v / rowv);
const int c4 = (int)(v - (long long)r * rowv) << 2;
if ((c4 >> 7) > (r >> 7)) {
const float4 z = {0.f, 0.f, 0.f, 0.f};
*reinterpret_cast<float4*>(Om + (long long)r * n + c4) = z;
}
}
}
}
// ---------------------------------------------------------------------------
// Kernel: batched lower-triangular inverse of an n x n block (n <= 128).
// One block per matrix, thread j computes column j by forward substitution.
// T column j is stored transposed into row j of the smem tile (upper part).
// Output written row-major (full square, upper zeroed).
// ---------------------------------------------------------------------------
__global__ void trtri_kernel(const float* __restrict__ src, long long src_stride, int src_ld,
float* __restrict__ T, long long t_stride, int n) {
constexpr int LDT = 129;
extern __shared__ float sh[];
float* t = sh; // 128 * 129
float* dinv = sh + 128 * LDT; // 128
const int tid = threadIdx.x;
const float* Sm = src + (long long)blockIdx.x * src_stride;
float* Tm = T + (long long)blockIdx.x * t_stride;
if ((n & 3) == 0 && (src_ld & 3) == 0) {
const int rowv = n >> 2;
const int tot = n * rowv;
for (int v = tid; v < tot; v += 128) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
const float4 x = *reinterpret_cast<const float4*>(Sm + (long long)r * src_ld + c4);
float* d = &t[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
} else {
for (int v = tid; v < n * n; v += 128) {
int r = v / n, c = v - r * n;
t[r * LDT + c] = Sm[(long long)r * src_ld + c];
}
}
__syncthreads();
if (tid < n) dinv[tid] = 1.0f / t[tid * LDT + tid];
__syncthreads();
const int j = tid;
if (j < n) {
// T[j][j] = dinv[j]; store T[i][j] (i>j) at t[j*LDT + i]
for (int i = j + 1; i < n; ++i) {
float s = t[i * LDT + j] * dinv[j];
for (int k = j + 1; k < i; ++k) {
s += t[i * LDT + k] * t[j * LDT + k];
}
t[j * LDT + i] = -s * dinv[i];
}
}
__syncthreads();
// write out: column j
if (j < n) {
for (int i = 0; i < n; ++i) {
float v = 0.f;
if (i == j) v = dinv[j];
else if (i > j) v = t[j * LDT + i];
Tm[(long long)i * 128 + j] = v;
}
}
}
// ---------------------------------------------------------------------------
// Kernel: exact in-place triangular solve X * L11^T = B for the fp32 zone.
// Rows of the panel are independent; each thread owns one row, solving in
// 16-column blocks (register-blocked substitution), L11 read via __ldg.
// ---------------------------------------------------------------------------
template <bool LSMEM>
__global__ void trsm_rows_kernel(float* __restrict__ X, long long x_stride, int x_ld,
const float* __restrict__ L11, long long l_stride, int l_ld,
int m, int dj) {
constexpr int LDT = 129;
extern __shared__ float sh[];
float* t = sh; // 128 * 129 row tile
float* dinv = sh + 128 * LDT; // 128
float* ls = LSMEM ? (sh + 128 * LDT + 128) : nullptr; // optional L11 staging
const int tid = threadIdx.x; // 128 threads
const int b = blockIdx.y;
const int row0 = blockIdx.x * 128;
const int rows = min(128, m - row0);
if (rows <= 0) return;
float* Xm = X + (long long)b * x_stride + (long long)row0 * x_ld;
const float* Lg = L11 + (long long)b * l_stride;
if (LSMEM) {
if ((dj & 3) == 0 && (l_ld & 3) == 0) {
const int rowv = dj >> 2;
for (int v = tid; v < dj * rowv; v += 128) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
const float4 x = *reinterpret_cast<const float4*>(Lg + (long long)r * l_ld + c4);
float* d = &ls[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
} else {
for (int v = tid; v < dj * dj; v += 128) {
int r = v / dj, c = v - r * dj;
ls[r * LDT + c] = Lg[(long long)r * l_ld + c];
}
}
}
const float* Lm = LSMEM ? ls : Lg;
const long long l_row = LSMEM ? (long long)LDT : (long long)l_ld;
if (tid < dj) dinv[tid] = 1.0f / __ldg(&Lg[(long long)tid * l_ld + tid]);
if ((dj & 3) == 0 && (x_ld & 3) == 0) {
const int rowv = dj >> 2;
const int tot = rows * rowv;
for (int v = tid; v < tot; v += 128) {
int r = v / rowv, c4 = (v - r * rowv) << 2;
const float4 x = *reinterpret_cast<const float4*>(Xm + (long long)r * x_ld + c4);
float* d = &t[r * LDT + c4];
d[0] = x.x; d[1] = x.y; d[2] = x.z; d[3] = x.w;
}
} else {
for (int v = tid; v < rows * dj; v += 128) {
int r = v / dj, c = v - r * dj;
t[r * LDT + c] = Xm[(long long)r * x_ld + c];
}
}
__syncthreads();
const int r = tid;
if (r < rows) {
// blocked substitution: short 16-wide solve chains + parallel updates
#pragma unroll 1
for (int jb = 0; jb < dj; jb += 16) {
const int jt = min(jb + 16, dj);
#pragma unroll 1
for (int j = jb; j < jt; ++j) {
float acc = t[r * LDT + j];
#pragma unroll 1
for (int k = jb; k < j; ++k) {
acc -= t[r * LDT + k] * Lm[(long long)j * l_row + k];
}
t[r * LDT + j] = acc * dinv[j];
}
int c = jt;
#pragma unroll 1
for (; c + 3 < dj; c += 4) {
float s0 = 0.f, s1 = 0.f, s2 = 0.f, s3 = 0.f;
#pragma unroll
for (int k = jb; k < jb + 16; ++k) {
const float xv = t[r * LDT + k];
s0 += xv * Lm[(long long)(c ) * l_row + k];
s1 += xv * Lm[(long long)(c + 1) * l_row + k];
s2 += xv * Lm[(long long)(c + 2) * l_row + k];
s3 += xv * Lm[(long long)(c + 3) * l_row + k];
}
t[r * LDT + c ] -= s0;
t[r * LDT + c + 1] -= s1;
t[r * LDT + c + 2] -= s2;
t[r * LDT + c + 3] -= s3;
}
#pragma unroll 1
for (; c < dj; ++c) {
float s = 0.f;
#pragma unroll
for (int k = jb; k < jb + 16; ++k) {
s += t[r * LDT + k] * Lm[(long long)c * l_row + k];
}
t[r * LDT + c] -= s;
}
}
}
__syncthreads();
if ((dj & 3) == 0 && (x_ld & 3) == 0) {
const int rowv = dj >> 2;
const int tot = rows * rowv;
for (int v = tid; v < tot; v += 128) {
int r2 = v / rowv, c4 = (v - r2 * rowv) << 2;
float4 y;
const float* s2 = &t[r2 * LDT + c4];
y.x = s2[0]; y.y = s2[1]; y.z = s2[2]; y.w = s2[3];
*reinterpret_cast<float4*>(Xm + (long long)r2 * x_ld + c4) = y;
}
} else {
for (int v = tid; v < rows * dj; v += 128) {
int r2 = v / dj, c = v - r2 * dj;
Xm[(long long)r2 * x_ld + c] = t[r2 * LDT + c];
}
}
}
// ---------------------------------------------------------------------------
// Kernel: copy lower triangle of A into out, zero strict upper. Grid-stride.
// Avoids reading the strictly-upper source elements.
// ---------------------------------------------------------------------------
__global__ void copy_tril_kernel(const float* __restrict__ A, float* __restrict__ out,
long long total4, int n, int rowv) {
const long long stride = (long long)gridDim.x * blockDim.x;
for (long long v = (long long)blockIdx.x * blockDim.x + threadIdx.x; v < total4; v += stride) {
const long long rowid = v / rowv;
const int c4 = (int)(v - rowid * rowv) << 2;
const int r = (int)(rowid % n);
const long long base = rowid * n + c4; // == (b*n + r)*n + c4
float4 y;
if (c4 + 3 <= r) {
y = *reinterpret_cast<const float4*>(A + base);
} else if (c4 > r) {
continue; // upper quad: left for the final zeroing pass
} else {
const float* a = A + base;
y.x = (c4 <= r) ? a[0] : 0.f;
y.y = (c4 + 1 <= r) ? a[1] : 0.f;
y.z = (c4 + 2 <= r) ? a[2] : 0.f;
y.w = (c4 + 3 <= r) ? a[3] : 0.f;
}
*reinterpret_cast<float4*>(out + base) = y;
}
}
__global__ void copy_tril_scalar_kernel(const float* __restrict__ A, float* __restrict__ out,
long long total, int n) {
const long long stride = (long long)gridDim.x * blockDim.x;
for (long long v = (long long)blockIdx.x * blockDim.x + threadIdx.x; v < total; v += stride) {
const long long rc = v % ((long long)n * n);
const int r = (int)(rc / n), c = (int)(rc % n);
if (c <= r) out[v] = A[v];
}
}
// ---------------------------------------------------------------------------
// Kernel: write zeros to every strictly-upper element (final cleanup pass).
// Vectorized over float4 quads; quads fully in the lower-inclusive triangle
// are skipped, boundary quads store zeros per-lane.
// ---------------------------------------------------------------------------
__global__ void zero_upper_kernel(float* __restrict__ out, long long total4, int n, int rowv) {
const long long stride = (long long)gridDim.x * blockDim.x;
for (long long v = (long long)blockIdx.x * blockDim.x + threadIdx.x; v < total4; v += stride) {
const long long rowid = v / rowv;
const int c4 = (int)(v - rowid * rowv) << 2;
const int r = (int)(rowid % n);
if (c4 + 3 <= r) continue;
const long long base = rowid * n + c4;
if (c4 > r) {
const float4 z = {0.f, 0.f, 0.f, 0.f};
*reinterpret_cast<float4*>(out + base) = z;
} else {
float* o = out + base;
if (c4 + 1 > r) o[1] = 0.f;
if (c4 + 2 > r) o[2] = 0.f;
o[3] = 0.f; // c4+3 > r guaranteed here
}
}
}
__global__ void zero_upper_scalar_kernel(float* __restrict__ out, long long total, int n) {
const long long stride = (long long)gridDim.x * blockDim.x;
for (long long v = (long long)blockIdx.x * blockDim.x + threadIdx.x; v < total; v += stride) {
const long long rc = v % ((long long)n * n);
const int r = (int)(rc / n), c = (int)(rc % n);
if (c > r) out[v] = 0.f;
}
}
// ---------------------------------------------------------------------------
// Kernel: copy a strided block (m x dj, ld = src_ld) into compact W (ld = 128).
// ---------------------------------------------------------------------------
__global__ void copy_block_kernel(const float* __restrict__ src, long long src_stride, int src_ld,
float* __restrict__ W, long long w_stride,
int m, int dj) {
const int b = blockIdx.y;
const float* Sm = src + (long long)b * src_stride;
float* Wm = W + (long long)b * w_stride;
const int rowv = dj >> 2; // dj multiple of 4
const long long tot = (long long)m * rowv;
const long long stride = (long long)gridDim.x * blockDim.x;
for (long long v = (long long)blockIdx.x * blockDim.x + threadIdx.x; v < tot; v += stride) {
const int r = (int)(v / rowv);
const int c4 = (int)(v - (long long)r * rowv) << 2;
*reinterpret_cast<float4*>(Wm + (long long)r * 128 + c4) =
*reinterpret_cast<const float4*>(Sm + (long long)r * src_ld + c4);
}
}
// ---------------------------------------------------------------------------
// Kernel: cast a strided fp32 panel (m x d, ld) to compact fp16 (ld = p_ld).
// ---------------------------------------------------------------------------
__global__ void cast_half_kernel(const float* __restrict__ src, long long src_stride, int src_ld,
__half* __restrict__ dst, long long dst_stride, int dst_ld,
int m, int d) {
const int b = blockIdx.y;
const float* Sm = src + (long long)b * src_stride;
__half* Dm = dst + (long long)b * dst_stride;
const int rowv = d >> 3; // d multiple of 8
const long long tot = (long long)m * rowv;
const long long stride = (long long)gridDim.x * blockDim.x;
for (long long v = (long long)blockIdx.x * blockDim.x + threadIdx.x; v < tot; v += stride) {
const int r = (int)(v / rowv);
const int c8 = (int)(v - (long long)r * rowv) << 3;
const float4 a = *reinterpret_cast<const float4*>(Sm + (long long)r * src_ld + c8);
const float4 c = *reinterpret_cast<const float4*>(Sm + (long long)r * src_ld + c8 + 4);
union { float4 f; __half2 h[4]; } u;
u.h[0] = __floats2half2_rn(a.x, a.y);
u.h[1] = __floats2half2_rn(a.z, a.w);
u.h[2] = __floats2half2_rn(c.x, c.y);
u.h[3] = __floats2half2_rn(c.z, c.w);
*reinterpret_cast<float4*>(Dm + (long long)r * dst_ld + c8) = u.f;
}
}
// ===========================================================================
// Host side
// ===========================================================================
struct Handles {
cublasHandle_t exact = nullptr;
at::Tensor ws_exact;
};
Handles& handles() {
static Handles h;
static std::once_flag flag;
std::call_once(flag, [] {
CUBLAS_OK(cublasCreate(&h.exact));
auto opts = at::TensorOptions().dtype(at::kByte).device(at::kCUDA);
h.ws_exact = at::empty({(long)(64 * 1024 * 1024)}, opts);
CUBLAS_OK(cublasSetWorkspace(h.exact, h.ws_exact.data_ptr(), 64 * 1024 * 1024));
});
return h;
}
struct Workspace {
at::Tensor W; // (B, maxm, 128) fp32 TRSM staging
at::Tensor T; // (B, 128, 128) fp32 inverted diag blocks
at::Tensor P16; // (B, maxm, nb) fp16 panel (fp16 zone only)
};
Workspace& get_workspace(long B, long n, long nb, bool fp16) {
static std::unordered_map<long long, Workspace> cache;
static std::mutex mu;
std::lock_guard<std::mutex> lock(mu);
long long key = (B << 24) ^ n ^ ((long long)fp16 << 62);
auto it = cache.find(key);
if (it != cache.end()) return it->second;
Workspace ws;
auto opts = at::TensorOptions().dtype(at::kFloat).device(at::kCUDA);
long maxm = std::max<long>(n - 128, 1);
ws.W = at::empty({B, maxm, 128}, opts);
ws.T = at::empty({B, 128, 128}, opts);
if (fp16) {
ws.P16 = at::empty({B, std::max<long>(n - nb, 1), nb}, opts.dtype(at::kHalf));
}
return cache.emplace(key, std::move(ws)).first->second;
}
enum class Zone { FP32, TF32, FP16 };
// row-major gemm helper: C(MxN) = alpha * opA(A) * opB(B) + beta * C
// op flags refer to the row-major views. Strided-batched.
void gemm_rm(cublasHandle_t hd, cublasOperation_t opA, cublasOperation_t opB,
int M, int N, int K, float alpha,
const void* A, cudaDataType_t Atype, int lda, long long strideA,
const void* B, cudaDataType_t Btype, int ldb, long long strideB,
float beta, void* C, int ldc, long long strideC, int batch,
cublasComputeType_t compute) {
const size_t esA = (Atype == CUDA_R_16F) ? 2 : 4;
if (batch <= 4) {
for (int b = 0; b < batch; ++b) {
CUBLAS_OK(cublasGemmEx(
hd, opB, opA, N, M, K, &alpha,
(const char*)B + (size_t)b * strideB * esA, Btype, ldb,
(const char*)A + (size_t)b * strideA * esA, Atype, lda,
&beta, (char*)C + (size_t)b * strideC * 4, CUDA_R_32F, ldc,
compute, CUBLAS_GEMM_DEFAULT));
}
return;
}
CUBLAS_OK(cublasGemmStridedBatchedEx(
hd, opB, opA, N, M, K, &alpha,
B, Btype, ldb, strideB,
A, Atype, lda, strideA,
&beta, C, CUDA_R_32F, ldc, strideC,
batch, compute, CUBLAS_GEMM_DEFAULT));
}
int chol_arch() {
static int arch = [] {
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, 0);
return prop.major * 10 + prop.minor;
}();
return arch;
}
bool g_fused256_ok = false;
bool g_trsm_lsmem_ok = false;
bool g_persist_ok = false;
int g_sm_count = 0;
void setup_smem_attrs() {
static std::once_flag flag;
std::call_once(flag, [] {
int max_optin = 0;
cudaDeviceGetAttribute(&max_optin, cudaDevAttrMaxSharedMemoryPerBlockOptin, 0);
constexpr int fused_bytes = 3 * 128 * 129 * 4 + 128 * 4;
if (max_optin >= fused_bytes) {
g_fused256_ok = cudaFuncSetAttribute((const void*)chol2x2_kernel<128, 256>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
fused_bytes) == cudaSuccess;
}
constexpr int bytes = 128 * 129 * 4 + 128 * 4;
cudaFuncSetAttribute((const void*)chol_tile_kernel<128, 256>,
cudaFuncAttributeMaxDynamicSharedMemorySize, bytes);
cudaFuncSetAttribute((const void*)chol_tile_kernel<128, 512>,
cudaFuncAttributeMaxDynamicSharedMemorySize, bytes);
cudaFuncSetAttribute((const void*)trtri_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, bytes);
cudaFuncSetAttribute((const void*)trsm_rows_kernel<false>,
cudaFuncAttributeMaxDynamicSharedMemorySize, bytes);
constexpr int trsm_big = 128 * 129 * 4 * 2 + 128 * 4;
if (max_optin >= trsm_big) {
g_trsm_lsmem_ok = cudaFuncSetAttribute((const void*)trsm_rows_kernel<true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
trsm_big) == cudaSuccess;
g_persist_ok = cudaFuncSetAttribute((const void*)chol_persist_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
trsm_big) == cudaSuccess;
}
cudaDeviceGetAttribute(&g_sm_count, cudaDevAttrMultiProcessorCount, 0);
cudaFuncSetAttribute((const void*)chol_tile_kernel<64, 128>,
cudaFuncAttributeMaxDynamicSharedMemorySize, 64 * 65 * 4 + 64 * 4);
});
}
void factor_tile128(const float* src, long long sstride, int sld,
float* dst, long long dstride, int dld,
int n, int zero_upper, int batch) {
constexpr int bytes = 128 * 129 * 4 + 128 * 4;
// 512 threads/block: more warps to hide the sequential-column latency of a
// single matrix's diagonal-block factorization (helps single & batched).
chol_tile_kernel<128, 512><<<batch, 512, bytes>>>(
src, sstride, sld, dst, dstride, dld, n, zero_upper);
}
int pick_nb(long n) {
if (n <= 2048) return 256;
if (n <= 4096) return 512;
if (n <= 16384) return 1024;
return 1024;
}
void blocked_chol(const at::Tensor& A, at::Tensor& out, long B, long n) {
setup_smem_attrs();
Handles& h = handles();
// diagnostic staging: 1=tile factor only, 2=+trsm, 3=+gemms, 4=full (default)
const char* dz = getenv("CHOL_DIAG");
const int diag = dz ? atoi(dz) : 4;
Zone zone;
if (n >= 8192) zone = Zone::FP16;
else if (n >= 2048 || (n >= 512 && B >= 8)) zone = Zone::TF32;
else zone = Zone::FP32;
const int nb = (zone == Zone::FP32) ? 128 : pick_nb(n);
Workspace& ws = get_workspace(B, n, nb, zone == Zone::FP16);
const float* Ap = A.data_ptr<float>();
float* O = out.data_ptr<float>();
const long long mstride = (long long)n * n;
// out = tril(A)
{
if ((n & 3) == 0) {
const int rowv = (int)(n >> 2);
long long total4 = (long long)B * n * rowv;
int grid = (int)std::min<long long>((total4 + 255) / 256, 32768);
copy_tril_kernel<<<grid, 256, 0>>>(Ap, O, total4, (int)n, rowv);
} else {
long long total = (long long)B * n * n;
int grid = (int)std::min<long long>((total + 255) / 256, 32768);
copy_tril_scalar_kernel<<<grid, 256, 0>>>(Ap, O, total, (int)n);
}
}
const cublasComputeType_t cc =
(zone == Zone::FP32) ? CUBLAS_COMPUTE_32F : CUBLAS_COMPUTE_32F_FAST_TF32;
for (long k = 0; k < n; k += nb) {
const int d = (int)std::min<long>(nb, n - k);
// ---- panel factorization: columns [k, k+d), all rows to n ----
for (int j = 0; j < d; j += 128) {
const long q = k + j;
const int dj = std::min(128, d - j);
factor_tile128(O + q * n + q, mstride, (int)n,
O + q * n + q, mstride, (int)n, dj, /*zero_upper=*/0, (int)B);
const long long mbelow = n - q - dj;
if (mbelow > 0 && diag >= 2) {
{
dim3 grid((unsigned)((mbelow + 127) / 128), (unsigned)B);
if (g_trsm_lsmem_ok && diag >= 4 && getenv("CHOL_NO_LSMEM") == nullptr) {
constexpr int bytes = 128 * 129 * 4 * 2 + 128 * 4;
trsm_rows_kernel<true><<<grid, 128, bytes>>>(
O + (q + dj) * n + q, mstride, (int)n,
O + q * n + q, mstride, (int)n, (int)mbelow, dj);
} else {
constexpr int bytes = 128 * 129 * 4 + 128 * 4;
trsm_rows_kernel<false><<<grid, 128, bytes>>>(
O + (q + dj) * n + q, mstride, (int)n,
O + q * n + q, mstride, (int)n, (int)mbelow, dj);
}
}
// inner update of remaining panel columns
const int wcols = (diag >= 3) ? (int)(k + d - q - dj) : 0;
if (wcols > 0) {
gemm_rm(h.exact, CUBLAS_OP_N, CUBLAS_OP_T,
(int)mbelow, wcols, dj, -1.f,
O + (q + dj) * n + q, CUDA_R_32F, (int)n, mstride,
O + (q + dj) * n + q, CUDA_R_32F, (int)n, mstride,
1.f, O + (q + dj) * n + (q + dj), (int)n, mstride, (int)B, cc);
}
}
}
// ---- trailing update ----
const long long mout = (diag >= 3) ? (n - k - d) : 0;
if (mout <= 0) continue;
const float* P = O + (k + d) * n + k;
float* C = O + (k + d) * n + (k + d);
if (zone == Zone::FP16) {
__half* P16 = reinterpret_cast<__half*>(ws.P16.data_ptr());
const long long pstride = ws.P16.size(1) * (long long)nb;
{
const int rowv = d >> 3;
long long tot = mout * rowv;
dim3 grid((unsigned)std::min<long long>((tot + 255) / 256, 32768), (unsigned)B);
cast_half_kernel<<<grid, 256, 0>>>(
P, mstride, (int)n, P16, pstride, nb, (int)mout, d);
}
const int bc = 4096;
for (long c0 = 0; c0 < mout; c0 += bc) {
const int bcc = (int)std::min<long>(bc, mout - c0);
gemm_rm(h.exact, CUBLAS_OP_N, CUBLAS_OP_T,
(int)(mout - c0), bcc, d, -1.f,
P16 + c0 * (long long)nb, CUDA_R_16F, nb, pstride,
P16 + c0 * (long long)nb, CUDA_R_16F, nb, pstride,
1.f, C + c0 * n + c0, (int)n, mstride, (int)B,
CUBLAS_COMPUTE_32F);
}
} else if (mout > 2048) {
// triangular trailing update as block-column gemms (halves the flops)
const long bc = 2048;
for (long c0 = 0; c0 < mout; c0 += bc) {
const int bcc = (int)std::min<long>(bc, mout - c0);
gemm_rm(h.exact, CUBLAS_OP_N, CUBLAS_OP_T,
(int)(mout - c0), bcc, d, -1.f,
P + c0 * n, CUDA_R_32F, (int)n, mstride,
P + c0 * n, CUDA_R_32F, (int)n, mstride,
1.f, C + c0 * n + c0, (int)n, mstride, (int)B, cc);
}
} else {
gemm_rm(h.exact, CUBLAS_OP_N, CUBLAS_OP_T,
(int)mout, (int)mout, d, -1.f,
P, CUDA_R_32F, (int)n, mstride,
P, CUDA_R_32F, (int)n, mstride,
1.f, C, (int)n, mstride, (int)B, cc);
}
}
// final pass: restore the strict upper triangle to exact zeros (trailing
// gemm updates write full squares and clobber the initial zeros)
if ((n & 3) == 0) {
const int rowv = (int)(n >> 2);
long long total4 = (long long)B * n * rowv;
int grid = (int)std::min<long long>((total4 + 255) / 256, 32768);
zero_upper_kernel<<<grid, 256>>>(O, total4, (int)n, rowv);
} else {
long long total = (long long)B * n * n;
int grid = (int)std::min<long long>((total + 255) / 256, 32768);
zero_upper_scalar_kernel<<<grid, 256>>>(O, total, (int)n);
}
}
} // namespace
torch::Tensor chol(torch::Tensor A) {
if (!A.is_cuda() || A.dtype() != torch::kFloat32 || A.dim() != 3 ||
(A.size(1) > 128 && (A.size(1) & 7) != 0)) {
return std::get<0>(at::linalg_cholesky_ex(A, /*upper=*/false, /*check_errors=*/false));
}
const at::cuda::OptionalCUDAGuard guard(A.device());
if (!A.is_contiguous()) A = A.contiguous();
const long B = A.size(0), n = A.size(1);
auto out = at::empty_like(A);
if (n <= 32) {
chol32_kernel<<<(unsigned)((B + 3) / 4), 128, 0>>>(
A.data_ptr<float>(), out.data_ptr<float>(), (int)B, (int)n);
} else if (n <= 64) {
setup_smem_attrs();
constexpr int bytes = 64 * 65 * 4 + 64 * 4;
chol_tile_kernel<64, 128><<<(unsigned)B, 128, bytes>>>(
A.data_ptr<float>(), (long long)n * n, (int)n,
out.data_ptr<float>(), (long long)n * n, (int)n, (int)n, 1);
} else if (n <= 128) {
setup_smem_attrs();
constexpr int bytes = 128 * 129 * 4 + 128 * 4;
chol_tile_kernel<128, 256><<<(unsigned)B, 256, bytes>>>(
A.data_ptr<float>(), (long long)n * n, (int)n,
out.data_ptr<float>(), (long long)n * n, (int)n, (int)n, 1);
} else if (n == 256) {
setup_smem_attrs();
if (g_fused256_ok && getenv("CHOL_NO_FUSED256") == nullptr) {
constexpr int fused_bytes = 3 * 128 * 129 * 4 + 128 * 4;
chol2x2_kernel<128, 256><<<(unsigned)B, 256, fused_bytes>>>(
A.data_ptr<float>(), out.data_ptr<float>());
} else {
blocked_chol(A, out, B, n);
}
} else if (n >= 512 && n <= 4096 && (n & 127) == 0) {
setup_smem_attrs();
if (g_persist_ok && getenv("CHOL_PERSIST") != nullptr) {
static std::unordered_map<long, at::Tensor> bar_cache;
static std::mutex bar_mu;
int* bars;
{
std::lock_guard<std::mutex> lock(bar_mu);
auto it = bar_cache.find(B);
if (it == bar_cache.end()) {
it = bar_cache.emplace(B, at::zeros({B * 2},
at::TensorOptions().dtype(at::kInt).device(at::kCUDA))).first;
}
bars = it->second.data_ptr<int>();
}
const int nt = (int)(n >> 7);
long Gl = g_sm_count / B;
if (Gl < 1) Gl = 1;
if (Gl > 2L * nt) Gl = 2L * nt;
const int G = (int)Gl;
constexpr int bytes = 2 * 128 * 129 * 4 + 128 * 4;
dim3 grid((unsigned)G, (unsigned)B);
chol_persist_kernel<<<grid, 256, bytes>>>(
A.data_ptr<float>(), out.data_ptr<float>(), bars, (int)n, G);
} else {
blocked_chol(A, out, B, n);
}
} else {
blocked_chol(A, out, B, n);
}
return out;
}
"""
def _build_ext():
import tempfile
from torch.utils.cpp_extension import load_inline
# Let torch pick the arch for the visible device; force a writable build dir.
os.environ.pop("TORCH_CUDA_ARCH_LIST", None)
os.environ.setdefault(
"TORCH_EXTENSIONS_DIR", tempfile.mkdtemp(prefix="chol_ext_")
)
os.environ.setdefault("MAX_JOBS", "8")
return load_inline(
name="chol_fast_ext",
cpp_sources=[_CPP_SRC],
cuda_sources=[_CUDA_SRC],
functions=["chol"],
extra_cuda_cflags=["-O3"],
extra_ldflags=["-lcublas"],
verbose=True,
)
def _self_test(ext) -> None:
"""Factor a few SPD matrices through every code path and verify against torch.
Runs once at import (outside any timed region). Raises on any mismatch so a
broken build/binary can never produce wrong ranked results.
"""
gen = torch.Generator(device="cuda").manual_seed(0)
cases = (
(4, 16, {}), (3, 32, {}), (2, 64, {}), (2, 128, {}),
(2, 384, {"CHOL_DIAG": "1", "CHOL_NO_FUSED256": "1"}), # tile factor only
(2, 384, {"CHOL_DIAG": "2", "CHOL_NO_FUSED256": "1"}), # + plain trsm
(2, 384, {"CHOL_DIAG": "3", "CHOL_NO_FUSED256": "1"}), # + gemms
(2, 384, {}), # full (smem trsm)
(2, 256, {"CHOL_NO_FUSED256": "1"}), # blocked at 256
(2, 256, {}), # fused chol2x2
(2, 512, {}), (1, 1024, {}),
)
for batch, n, env in cases:
for k, v in env.items():
os.environ[k] = v
try:
g = torch.randn((batch, n, n), device="cuda", generator=gen)
a = g @ g.transpose(-1, -2) / n
a.diagonal(dim1=-2, dim2=-1).add_(0.01)
out = ext.chol(a)
torch.cuda.synchronize()
finally:
for k in env:
os.environ.pop(k, None)
print(f"SELFTEST ok batch={batch} n={n} env={env}", flush=True)
if "CHOL_DIAG" in env:
continue # staged diagnostic phases produce intentionally partial results
if not torch.isfinite(out).all():
raise RuntimeError(f"self-test nonfinite at n={n}")
if torch.triu(out, diagonal=1).abs().max().item() != 0.0:
raise RuntimeError(f"self-test upper nonzero at n={n}")
recon = out @ out.transpose(-1, -2)
rel = (recon - a).abs().max().item() / a.abs().max().item()
if rel > 1e-2:
raise RuntimeError(f"self-test reconstruction off at n={n}: rel={rel}")
def _perf_probe(ext) -> None:
"""Print a per-component timing breakdown for hot shapes (import-time only)."""
import time
for b, n in ((1, 4096), (1, 8192), (16, 512), (640, 512)):
g = torch.randn((b, n, n), device="cuda")
a = g @ g.transpose(-1, -2) / n
a.diagonal(dim1=-2, dim2=-1).add_(0.01)
del g
for diag in ("1", "2", "3", ""):
if diag:
os.environ["CHOL_DIAG"] = diag
else:
os.environ.pop("CHOL_DIAG", None)
for _ in range(2):
ext.chol(a)
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(5):
ext.chol(a)
torch.cuda.synchronize()
dt = (time.perf_counter() - t0) / 5
print(f"PROBE b={b} n={n} diag={diag or 'full'}: {dt * 1e6:.0f} us", flush=True)
del a
os.environ.pop("CHOL_DIAG", None)
try:
_ext = _build_ext()
if torch.cuda.is_available():
_self_test(_ext)
if os.environ.get("CHOL_PROBE", "0") == "1":
try:
_perf_probe(_ext)
except Exception:
import traceback
traceback.print_exc()
except Exception:
import traceback
traceback.print_exc()
_ext = None
# ---------------------------------------------------------------------------
# Auto-tuner: for each (batch, n) shape, pick whichever backend is fastest —
# our custom CUDA extension or cuSOLVER (torch). The choice is made ONCE per
# shape on the first (untimed warm-up) call and cached, so the timed calls do
# only a dict lookup. This can never be slower than cuSOLVER, and keeps every
# shape where our kernels win.
# ---------------------------------------------------------------------------
def _torch_chol(data: input_t) -> output_t:
return torch.linalg.cholesky_ex(data, check_errors=False).L
def _valid_enough(data: torch.Tensor, out: torch.Tensor) -> bool:
"""Cheap sanity gate mirroring the competition checker (loosely)."""
if out is None or out.shape != data.shape or not torch.isfinite(out).all():
return False
n = data.shape[-1]
eps = torch.finfo(torch.float32).eps
diag = torch.diagonal(out, dim1=-2, dim2=-1)
if (diag <= 0).any():
return False
scale = torch.linalg.matrix_norm(data, ord=1, dim=(-2, -1)).clamp_min(1e-30)
recon = out @ out.transpose(-1, -2)
resid = torch.linalg.matrix_norm(recon - data, ord=1, dim=(-2, -1))
return bool((resid <= 20.0 * n * eps * scale).all())
def _time_backend(fn, data: torch.Tensor, iters: int = 6) -> float:
"""Milliseconds per call, or +inf if the backend errors/misbehaves."""
try:
for _ in range(3):
fn(data)
torch.cuda.synchronize()
except Exception:
return float("inf")
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(iters):
fn(data)
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) / iters
_impl_cache: dict = {}
def _select_backend(data: torch.Tensor):
candidates = [_torch_chol]
if _ext is not None and data.is_cuda and data.dtype == torch.float32 and data.dim() == 3:
probe = data.clone()
try:
out = _ext.chol(probe)
torch.cuda.synchronize()
if _valid_enough(data, out):
candidates.append(_ext.chol)
except Exception:
pass
best, best_t = _torch_chol, float("inf")
for fn in candidates:
t = _time_backend(fn, data.clone())
if t < best_t:
best_t, best = t, fn
return best
def custom_kernel(data: input_t) -> output_t:
if not (isinstance(data, torch.Tensor) and data.is_cuda and data.dim() == 3):
return _torch_chol(data)
key = (int(data.shape[0]), int(data.shape[1]), str(data.dtype))
fn = _impl_cache.get(key)
if fn is None:
fn = _select_backend(data)
_impl_cache[key] = fn
return fn(data)
scrolls · 1460 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