submission 911904
patwrall · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2290 lines, June 9 Researcher Reciprocity License v1.0.
submission_1b26b37.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-911904?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:c6e76cea8aa4c6a06b3b8429046395bb2956c1a7f4f3abe7eefb496e776424a2
license declaredunknown
license concludedunknown
authorspatwrall
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
__device__ __forceinline__ void cp_async16(void* dst, const void* src) {mbarrier
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));mma
"mma.sync.aligned.m16n8k8.row.col.f32.tf32.tf32.f32 "persistent-kernel
persistent_cholesky_kernel(float* __restrict__ A, int n, long long bstride,shared-memory
extern __shared__ float smem[];tcgen05
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"vector-width = float4
for (int q = 0; q < 32; q += 4) *(float4*)&r[q] = *(const float4*)&src[q];Kernel source
submission_1b26b37.py2290 lines
import sys
from pathlib import Path
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
P = lambda *a: print(*a, file=sys.stderr)
CPP_SRC = r"""
"""
# Pure device code, no torch dependency — extracted verbatim by the local
# sm_86 test harness. Keep torch/ATen includes out of this string.
CUDA_KERNELS = r"""
// Factor one B x B SPD tile (dst <- L, upper zeroed); one CTA, unblocked
// Cholesky staged through smem (B*(B+1) floats, padded). src may equal dst.
template <int B, int NT>
__device__ void potrf_tile(const float* __restrict__ src, float* __restrict__ dst,
int lda, float* smem) {
const int tid = threadIdx.x;
auto s = [&](int i, int j) -> float& { return smem[i * (B + 1) + j]; };
for (int idx = tid; idx < B * B; idx += NT)
s(idx / B, idx % B) = src[(idx / B) * lda + (idx % B)];
__syncthreads();
for (int k = 0; k < B; ++k) {
if (tid == 0) s(k, k) = sqrtf(s(k, k));
__syncthreads();
const float dinv = 1.0f / s(k, k);
for (int i = k + 1 + tid; i < B; i += NT) s(i, k) *= dinv;
__syncthreads();
for (int i = k + 1 + tid; i < B; i += NT) {
const float lik = s(i, k);
for (int j = k + 1; j <= i; ++j) s(i, j) -= lik * s(j, k);
}
__syncthreads();
}
for (int idx = tid; idx < B * B; idx += NT) {
const int i = idx / B, j = idx % B;
dst[i * lda + j] = (j <= i) ? s(i, j) : 0.0f;
}
__syncthreads();
}
// rho regime, n == B: one CTA per matrix, whole matrix is a single tile,
// output written out of place (no clone pass).
template <int B, int NT>
__global__ void potrf_batched_kernel(const float* __restrict__ A, float* __restrict__ L,
long long batch_stride) {
extern __shared__ float smem[];
potrf_tile<B, NT>(A + blockIdx.x * batch_stride, L + blockIdx.x * batch_stride, B, smem);
}
template <int B, int NT>
void launch_potrf_batched(const float* A, float* L, long long batch) {
constexpr int SMEM = B * (B + 1) * sizeof(float);
cudaFuncSetAttribute(potrf_batched_kernel<B, NT>,
cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
potrf_batched_kernel<B, NT>
<<<batch, NT, SMEM>>>(A, L, (long long)B * B);
}
// n == 32 special case: one warp per matrix, each lane owns a row in
// registers, active column published through smem; no block syncs.
__global__ void potrf32_warp_kernel(const float* __restrict__ A, float* __restrict__ L,
long long batch) {
__shared__ float col[4][32];
const int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
const long long m = (long long)blockIdx.x * 4 + warp;
if (m >= batch) return;
const float* src = A + m * 32 * 32 + lane * 32;
float* dst = L + m * 32 * 32 + lane * 32;
float r[32];
#pragma unroll
for (int q = 0; q < 32; q += 4) *(float4*)&r[q] = *(const float4*)&src[q];
#pragma unroll
for (int k = 0; k < 32; ++k) {
const float dk = sqrtf(__shfl_sync(0xffffffffu, r[k], k));
if (lane == k) r[k] = dk;
else if (lane > k) r[k] /= dk;
col[warp][lane] = r[k];
__syncwarp();
#pragma unroll
for (int j = 0; j < 32; ++j)
if (j > k && lane >= j) r[j] -= r[k] * col[warp][j];
__syncwarp();
}
#pragma unroll
for (int j = 0; j < 32; ++j)
if (j > lane) r[j] = 0.0f;
#pragma unroll
for (int q = 0; q < 32; q += 4) *(float4*)&dst[q] = *(float4*)&r[q];
}
// n == 64: one 2-warp CTA per matrix, each thread owns a row in registers.
// The unscaled column is published pre-scaling (1/d^2 folded into the
// update FMA) so each k step needs a single block sync; colbuf is double
// buffered to avoid a second one.
__global__ void __launch_bounds__(64)
potrf64_rows_kernel(const float* __restrict__ A, float* __restrict__ L, long long batch) {
__shared__ float colbuf[2][64];
const int i = threadIdx.x;
const float* src = A + blockIdx.x * 64 * 64 + i * 64;
float* dst = L + blockIdx.x * 64 * 64 + i * 64;
float r[64];
#pragma unroll
for (int q = 0; q < 64; q += 4) *(float4*)&r[q] = *(const float4*)&src[q];
#pragma unroll
for (int k = 0; k < 64; ++k) {
colbuf[k & 1][i] = r[k];
__syncthreads();
const float dkk = colbuf[k & 1][k];
if (i > k) {
const float t = r[k] * (1.0f / dkk);
#pragma unroll
for (int j = 0; j < 64; ++j)
if (j > k && j <= i) r[j] -= t * colbuf[k & 1][j];
r[k] *= rsqrtf(dkk);
} else if (i == k) {
r[k] = sqrtf(dkk);
}
}
#pragma unroll
for (int j = 0; j < 64; ++j)
if (j > i) r[j] = 0.0f;
#pragma unroll
for (int q = 0; q < 64; q += 4) *(float4*)&dst[q] = *(float4*)&r[q];
}
// n == 128: one 128-thread CTA per matrix, rows resident in smem (rows in
// registers would spill); same single-sync-per-step structure.
__global__ void __launch_bounds__(128)
potrf128_rows_kernel(const float* __restrict__ A, float* __restrict__ L, long long batch) {
extern __shared__ float s[]; // 128 x 129
const int i = threadIdx.x;
const float* src = A + blockIdx.x * 128 * 128;
float* dst = L + blockIdx.x * 128 * 128;
for (int idx = i; idx < 128 * 128; idx += 128)
s[(idx / 128) * 129 + idx % 128] = src[idx];
__syncthreads();
// elimination keeps unscaled values (diag holds d^2); scaling by
// rsqrt(d^2) happens per column in the store pass, so nothing mutates
// a column while other threads still read it
float* row = s + i * 129;
for (int k = 0; k < 128; ++k) {
if (i > k) {
const float t = row[k] * (1.0f / s[k * 129 + k]);
for (int j = k + 1; j <= i; ++j) row[j] -= t * s[j * 129 + k];
}
__syncthreads();
}
for (int idx = i; idx < 128 * 128; idx += 128) {
const int rr = idx / 128, cc = idx % 128;
float v = 0.0f;
if (cc < rr) v = s[rr * 129 + cc] * rsqrtf(s[cc * 129 + cc]);
else if (cc == rr) v = sqrtf(s[cc * 129 + cc]);
dst[idx] = v;
}
}
// n == 128, register scheme: one 256-thread CTA per matrix, threads i and
// i+128 own the two 64-column halves of row i in registers. Unscaled column
// k is published through double-buffered smem (all rows, so colbuf[i] also
// serves as each row's own a_ik); one block sync per step like the n=64
// scheme, rsqrt scaling folded in after publication.
// Rank-4 variant of the sweep: publish 4 RAW columns per barrier, factor
// the 4x4 pivot block in registers (uniform work, redundant per thread),
// then back-substitute the four combined coefficients so each j costs four
// FMAs against raw columns. Quarter the barriers of rank-1; deeper raw-
// column cancellation than rank-2, so accuracy must be re-validated on any
// new shape it is applied to.
// Column ownership is INTERLEAVED by 4-wide group: the h-half of a row owns
// groups h, h+2, ... (register slot r[4q+u] holds column 8q+4h+u), so both
// halves keep equal live work at every step — column-half ownership left
// half-0 idle for the whole second half of the sweep (30% barrier stall).
template <int ROWS>
__device__ __forceinline__ void panel_sweep128_r4(float (&r)[64], float (*colbuf)[4][ROWS],
int i, int h, bool active) {
#pragma unroll
for (int k = 0; k < 128; k += 4) {
const int lq = (k >> 3) * 4; // r[] base of pivot group
const bool own = active && (h == ((k >> 2) & 1));
const int p = (k >> 2) & 1;
if (own) {
#pragma unroll
for (int b = 0; b < 4; ++b) colbuf[p][b][i] = r[lq + b];
}
__syncthreads();
const float M00 = colbuf[p][0][k], M10 = colbuf[p][0][k + 1],
M20 = colbuf[p][0][k + 2], M30 = colbuf[p][0][k + 3],
M11 = colbuf[p][1][k + 1], M21 = colbuf[p][1][k + 2],
M31 = colbuf[p][1][k + 3], M22 = colbuf[p][2][k + 2],
M32 = colbuf[p][2][k + 3], M33 = colbuf[p][3][k + 3];
const float rd0 = 1.0f / M00;
const float s10 = M10 * rd0, s20 = M20 * rd0, s30 = M30 * rd0;
const float d1 = M11 - s10 * M10;
const float rd1 = 1.0f / d1;
const float m21 = M21 - s20 * M10, m31 = M31 - s30 * M10;
const float s21 = m21 * rd1, s31 = m31 * rd1;
const float d2 = M22 - s20 * M20 - s21 * m21;
const float rd2 = 1.0f / d2;
const float m32 = M32 - s30 * M20 - s31 * m21;
const float s32 = m32 * rd2;
const float d3 = M33 - s30 * M30 - s31 * m31 - s32 * m32;
const float e0 = colbuf[p][0][i], e1 = colbuf[p][1][i],
e2 = colbuf[p][2][i], e3 = colbuf[p][3][i];
const float t0 = e0 * rd0;
const float e1c = e1 - s10 * e0;
const float t1 = e1c * rd1;
const float e2c = e2 - s20 * e0 - s21 * e1c;
const float t2 = e2c * rd2;
const float e3c = e3 - s30 * e0 - s31 * e1c - s32 * e2c;
const float t3 = e3c * (1.0f / d3);
if (active && i > k + 3) {
const float a3 = t3;
const float a2 = t2 - a3 * s32;
const float a1 = t1 - a2 * s21 - a3 * s31;
const float a0 = t0 - a1 * s10 - a2 * s20 - a3 * s30;
if (h == 0) {
#pragma unroll
for (int q = 0; q < 16; ++q) {
if (q * 8 + 3 > k + 3) {
const float4 v0 = *(const float4*)&colbuf[p][0][q * 8];
const float4 v1 = *(const float4*)&colbuf[p][1][q * 8];
const float4 v2 = *(const float4*)&colbuf[p][2][q * 8];
const float4 v3 = *(const float4*)&colbuf[p][3][q * 8];
const float w0[4] = {v0.x, v0.y, v0.z, v0.w};
const float w1[4] = {v1.x, v1.y, v1.z, v1.w};
const float w2[4] = {v2.x, v2.y, v2.z, v2.w};
const float w3[4] = {v3.x, v3.y, v3.z, v3.w};
#pragma unroll
for (int u = 0; u < 4; ++u) {
const int j = q * 8 + u;
if (j > k + 3 && j <= i)
r[q * 4 + u] -= a0 * w0[u] + a1 * w1[u] + a2 * w2[u] + a3 * w3[u];
}
}
}
} else {
#pragma unroll
for (int q = 0; q < 16; ++q) {
if (q * 8 + 4 + 3 > k + 3) {
const float4 v0 = *(const float4*)&colbuf[p][0][q * 8 + 4];
const float4 v1 = *(const float4*)&colbuf[p][1][q * 8 + 4];
const float4 v2 = *(const float4*)&colbuf[p][2][q * 8 + 4];
const float4 v3 = *(const float4*)&colbuf[p][3][q * 8 + 4];
const float w0[4] = {v0.x, v0.y, v0.z, v0.w};
const float w1[4] = {v1.x, v1.y, v1.z, v1.w};
const float w2[4] = {v2.x, v2.y, v2.z, v2.w};
const float w3[4] = {v3.x, v3.y, v3.z, v3.w};
#pragma unroll
for (int u = 0; u < 4; ++u) {
const int j = q * 8 + 4 + u;
if (j > k + 3 && j <= i)
r[q * 4 + u] -= a0 * w0[u] + a1 * w1[u] + a2 * w2[u] + a3 * w3[u];
}
}
}
}
}
if (own) {
if (i == k) r[lq] = sqrtf(M00);
else if (i > k) r[lq] = e0 * rsqrtf(M00);
if (i == k + 1) r[lq + 1] = sqrtf(d1);
else if (i > k + 1) r[lq + 1] = e1c * rsqrtf(d1);
if (i == k + 2) r[lq + 2] = sqrtf(d2);
else if (i > k + 2) r[lq + 2] = e2c * rsqrtf(d2);
if (i == k + 3) r[lq + 3] = sqrtf(d3);
else if (i > k + 3) r[lq + 3] = e3c * rsqrtf(d3);
}
}
}
// interleaved-group column index of register slot r[4q+u] for half h
__device__ __forceinline__ int ilv_col(int q, int h, int u) { return q * 8 + h * 4 + u; }
__global__ void __launch_bounds__(256, 2)
potrf128_regs_kernel(const float* __restrict__ A, float* __restrict__ L, long long batch) {
__shared__ float colbuf[2][4][128]; // [parity][column][row]
const int i = threadIdx.x & 127;
const int h = threadIdx.x >> 7;
const float* src = A + blockIdx.x * 128 * 128 + i * 128;
float* dst = L + blockIdx.x * 128 * 128 + i * 128;
float r[64];
#pragma unroll
for (int q = 0; q < 16; ++q) *(float4*)&r[q * 4] = *(const float4*)&src[ilv_col(q, h, 0)];
panel_sweep128_r4<128>(r, colbuf, i, h, true);
#pragma unroll
for (int q = 0; q < 16; ++q) {
#pragma unroll
for (int u = 0; u < 4; ++u)
if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
*(float4*)&dst[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
}
}
// n == 256: one 512-thread CTA per matrix. Phase A sweeps the whole 256x128
// left panel (L11 + L21) held in registers (two 64-col halves per row);
// L21 is then staged to smem ([129] pad -> conflict-free column reads) for
// an in-CTA SYRK into registers, and phase D sweeps the trailing 128x128
// with threads 0..255 while the rest only participate in barriers.
__global__ void __launch_bounds__(512, 1)
fused256_kernel(const float* __restrict__ A, float* __restrict__ Lg, long long bstride) {
extern __shared__ float smem[]; // sL[128][129] | colbuf 2*4*256
float (*sL)[129] = (float (*)[129])smem;
float* colb = smem + 128 * 129;
const float* src = A + blockIdx.x * bstride;
float* dst = Lg + blockIdx.x * bstride;
float r[64];
const int i = threadIdx.x & 255;
const int h = (threadIdx.x >> 8) & 1;
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&r[q * 4] = *(const float4*)&src[i * 256 + ilv_col(q, h, 0)];
panel_sweep128_r4<256>(r, (float (*)[4][256])colb, i, h, true);
if (i < 128) {
#pragma unroll
for (int q = 0; q < 16; ++q)
#pragma unroll
for (int u = 0; u < 4; ++u)
if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
} else {
#pragma unroll
for (int q = 0; q < 16; ++q)
#pragma unroll
for (int u = 0; u < 4; ++u) sL[i - 128][ilv_col(q, h, u)] = r[q * 4 + u];
}
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&dst[i * 256 + ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
// zero the upper-right 128x128 block
for (int idx = threadIdx.x; idx < 128 * 32; idx += 512) {
const int rr = idx >> 5, cc = (idx & 31) * 4;
*(float4*)&dst[rr * 256 + 128 + cc] = make_float4(0.f, 0.f, 0.f, 0.f);
}
__syncthreads();
// SYRK: 4 rows (lane-stride 1 -> conflict-free column reads of sL) x
// 8 cols (constant per warp -> broadcast) per thread
const int rb = threadIdx.x & 31, cb = (threadIdx.x >> 5) * 8;
float acc[4][8] = {};
for (int k = 0; k < 128; ++k) {
float a4[4], b8[8];
#pragma unroll
for (int u = 0; u < 4; ++u) a4[u] = sL[rb + 32 * u][k];
#pragma unroll
for (int v = 0; v < 8; ++v) b8[v] = sL[cb + v][k];
#pragma unroll
for (int u = 0; u < 4; ++u)
#pragma unroll
for (int v = 0; v < 8; ++v) acc[u][v] += a4[u] * b8[v];
}
float c22[4][8];
#pragma unroll
for (int u = 0; u < 4; ++u)
#pragma unroll
for (int v = 0; v < 8; v += 4) {
const float4 a4 = *(const float4*)&src[(128 + rb + 32 * u) * 256 + 128 + cb + v];
c22[u][v] = a4.x - acc[u][v]; c22[u][v + 1] = a4.y - acc[u][v + 1];
c22[u][v + 2] = a4.z - acc[u][v + 2]; c22[u][v + 3] = a4.w - acc[u][v + 3];
}
__syncthreads(); // all SYRK reads of sL done before overwrite
#pragma unroll
for (int u = 0; u < 4; ++u)
#pragma unroll
for (int v = 0; v < 8; ++v) sL[rb + 32 * u][cb + v] = c22[u][v];
__syncthreads();
// phase D: trailing 128x128 out of smem
const bool act = threadIdx.x < 256;
const int i2 = threadIdx.x & 127;
const int h2 = (threadIdx.x >> 7) & 1;
if (act) {
#pragma unroll
for (int q = 0; q < 16; ++q)
#pragma unroll
for (int u = 0; u < 4; ++u) r[q * 4 + u] = sL[i2][ilv_col(q, h2, u)];
}
__syncthreads();
panel_sweep128_r4<128>(r, (float (*)[4][128])colb, i2, h2, act);
if (act) {
#pragma unroll
for (int q = 0; q < 16; ++q) {
#pragma unroll
for (int u = 0; u < 4; ++u)
if (ilv_col(q, h2, u) > i2) r[q * 4 + u] = 0.0f;
*(float4*)&dst[(128 + i2) * 256 + 128 + ilv_col(q, h2, 0)] = *(float4*)&r[q * 4];
}
}
}
void launch_fused256(const float* A, float* L, long long batch) {
constexpr int SMEM = (128 * 129 + 2 * 4 * 256) * sizeof(float);
cudaFuncSetAttribute(fused256_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
fused256_kernel<<<batch, 512, SMEM>>>(A, L, (long long)256 * 256);
}
// Blocked-128 panel step for n in {512, 1024, ...}: one CTA per consumer
// 128-row block (blockIdx.x = c), each CTA holds pivot block P plus its
// consumer block in registers and runs the fused potrf+trsm sweep; the
// pivot factorization is recomputed redundantly per CTA (cheap) so no
// cross-CTA sync is needed. CTA c==0 stores the pivot rows.
__global__ void __launch_bounds__(512, 1)
panel512_kernel(float* __restrict__ Lg, int n, long long bstride, int P) {
__shared__ float colbuf[2][4][256];
float* M = Lg + blockIdx.y * bstride;
const int i = threadIdx.x & 255;
const int h = (threadIdx.x >> 8) & 1;
const int nblk = n / 128;
const int cons = nblk - P - 1;
const int grow = (i < 128) ? 128 * P + i
: 128 * (P + 1 + (int)blockIdx.x) + (i - 128);
float* rowp = M + (long long)grow * n + 128 * P;
float r[64];
const bool have = (i < 128) || ((int)blockIdx.x < cons);
if (have) {
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&r[q * 4] = *(const float4*)&rowp[ilv_col(q, h, 0)];
}
panel_sweep128_r4<256>(r, colbuf, i, h, have);
// Pivot rows are stored ONLY when this launch has a single CTA (last
// column): with consumers in flight, sibling CTAs still read the raw
// pivot block from global, so the store moves to the update launch.
if (i < 128) {
if (blockIdx.x == 0 && cons == 0) {
#pragma unroll
for (int q = 0; q < 16; ++q) {
#pragma unroll
for (int u = 0; u < 4; ++u)
if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
*(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
}
}
} else if (have) {
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
}
}
// Left-looking panel step: identical sweep to panel512_kernel, but the
// factored pivot block goes to the side buffer `piv` (one 128x128 slot per
// (matrix, column)) instead of L — in left-looking form nothing reads the
// pivot block after its own column, so sibling CTAs can keep reading the
// raw pivot from L race-free; copy_pivots_kernel patches L at the end.
__global__ void __launch_bounds__(512, 1)
panel512_left_kernel(float* __restrict__ Lg, float* __restrict__ piv,
int n, long long bstride, int Q) {
__shared__ float colbuf[2][4][256];
float* M = Lg + blockIdx.y * bstride;
const int i = threadIdx.x & 255;
const int h = (threadIdx.x >> 8) & 1;
const int nblk = n / 128;
const int cons = nblk - Q - 1;
const int grow = (i < 128) ? 128 * Q + i
: 128 * (Q + 1 + (int)blockIdx.x) + (i - 128);
float* rowp = M + (long long)grow * n + 128 * Q;
float r[64];
const bool have = (i < 128) || ((int)blockIdx.x < cons);
if (have) {
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&r[q * 4] = *(const float4*)&rowp[ilv_col(q, h, 0)];
}
panel_sweep128_r4<256>(r, colbuf, i, h, have);
if (i < 128) {
if (blockIdx.x == 0) {
#pragma unroll
for (int q = 0; q < 16; ++q)
#pragma unroll
for (int u = 0; u < 4; ++u)
if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
float* pdst = (cons == 0)
? rowp
: piv + ((long long)blockIdx.y * nblk + Q) * 128 * 128 + i * 128;
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&pdst[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
}
} else if (have) {
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
}
}
// Patch the stashed pivot blocks (already tril-zeroed) into L; also
// overwrites the above-diagonal garbage baddbmm left inside pivot blocks.
__global__ void copy_pivots_kernel(float* __restrict__ Lg, const float* __restrict__ piv,
int n, long long bstride) {
const int nblk = n / 128, Q = blockIdx.x;
const float* src = piv + ((long long)blockIdx.y * nblk + Q) * 128 * 128;
float* dst = Lg + blockIdx.y * bstride + (long long)(128 * Q) * n + 128 * Q;
for (int idx = threadIdx.x; idx < 128 * 32; idx += 256) {
const int rr = idx >> 5, cc = (idx & 31) * 4;
*(float4*)&dst[(long long)rr * n + cc] = *(const float4*)&src[rr * 128 + cc];
}
}
__device__ __forceinline__ void pair_decode(int p, int J, int& I, int& K);
// Trailing update for the blocked-128 path: C(I,K) -= L(I,P) * L(K,P)^T for
// 128-tile pair p (I >= K > P); operands staged transposed ([k][row]) in
// 64-wide k chunks; 256 threads, 8x8 outputs each. Diagonal tiles skip
// above-diagonal stores to preserve the tril zeros.
__global__ void __launch_bounds__(256, 2)
update128_regs_kernel(float* __restrict__ Lg, int n, long long bstride, int P) {
extern __shared__ float smem[];
float (*sIT)[132] = (float (*)[132])smem;
float (*sKT)[132] = (float (*)[132])(smem + 64 * 132);
float* M = Lg + blockIdx.y * bstride;
const int cons = n / 128 - P - 1;
if ((int)blockIdx.x == cons * (cons + 1) / 2) {
// extra CTA: factor + store the pivot block (raw in global — panel
// CTAs never wrote it, and no update CTA touches block column P)
__shared__ float cb[2][4][128];
const int i = threadIdx.x & 127;
const int h = (threadIdx.x >> 7) & 1;
float* rowp = M + (long long)(128 * P + i) * n + 128 * P;
float r[64];
#pragma unroll
for (int q = 0; q < 16; ++q)
*(float4*)&r[q * 4] = *(const float4*)&rowp[ilv_col(q, h, 0)];
panel_sweep128_r4<128>(r, cb, i, h, true);
#pragma unroll
for (int q = 0; q < 16; ++q) {
#pragma unroll
for (int u = 0; u < 4; ++u)
if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
*(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
}
return;
}
int I, K;
pair_decode(blockIdx.x, P, I, K);
const float* LI = M + (long long)I * 128 * n + 128 * P;
const float* LK = M + (long long)K * 128 * n + 128 * P;
float* C = M + (long long)I * 128 * n + 128 * K;
const int tx = threadIdx.x & 15, ty = threadIdx.x >> 4;
float acc[8][8] = {};
for (int kc = 0; kc < 128; kc += 64) {
__syncthreads();
for (int idx = threadIdx.x; idx < 64 * 128; idx += 256) {
const int rr = idx >> 6, kk = idx & 63;
sIT[kk][rr] = LI[rr * n + kc + kk];
sKT[kk][rr] = LK[rr * n + kc + kk];
}
__syncthreads();
#pragma unroll 8
for (int k = 0; k < 64; ++k) {
float a8[8], b8[8];
#pragma unroll
for (int q = 0; q < 8; q += 4) {
*(float4*)&a8[q] = *(const float4*)&sIT[k][ty * 8 + q];
*(float4*)&b8[q] = *(const float4*)&sKT[k][tx * 8 + q];
}
#pragma unroll
for (int u = 0; u < 8; ++u)
#pragma unroll
for (int v = 0; v < 8; ++v) acc[u][v] += a8[u] * b8[v];
}
}
#pragma unroll
for (int u = 0; u < 8; ++u) {
const int rr = ty * 8 + u;
#pragma unroll
for (int v = 0; v < 8; v += 4) {
const int cc = tx * 8 + v;
if (I == K && cc + 3 > rr) {
#pragma unroll
for (int w = 0; w < 4; ++w)
if (cc + w <= rr) C[rr * n + cc + w] -= acc[u][v + w];
} else {
float4* cp = (float4*)&C[rr * n + cc];
float4 c4 = *cp;
c4.x -= acc[u][v]; c4.y -= acc[u][v + 1];
c4.z -= acc[u][v + 2]; c4.w -= acc[u][v + 3];
*cp = c4;
}
}
}
}
// float4 tril copy (out-of-place): the at::tril clone runs ~4x below
// memory speed and showed up at 9% of the 8x2048 profile
__global__ void tril_f4_kernel(const float* __restrict__ A, float* __restrict__ L,
int n, long long bstride) {
const float* src = A + blockIdx.y * bstride;
float* dst = L + blockIdx.y * bstride;
const int nq = n >> 2;
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < n * nq;
idx += gridDim.x * blockDim.x) {
const int i = idx / nq, c = (idx % nq) * 4;
float4 v = *(const float4*)&src[(long long)i * n + c];
if (c + 3 > i) {
if (c + 0 > i) v.x = 0.0f;
if (c + 1 > i) v.y = 0.0f;
if (c + 2 > i) v.z = 0.0f;
v.w = 0.0f;
}
*(float4*)&dst[(long long)i * n + c] = v;
}
}
// Right-looking blocked-128 factorization from the register panel sweep;
// per column: fused potrf+trsm panel CTAs, then the 128-tile update.
void cholesky_b128regs(float* L, int n, long long batch) {
constexpr int USM = 2 * 64 * 132 * sizeof(float);
cudaFuncSetAttribute(update128_regs_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, USM);
const int nblk = n / 128;
const long long bstride = (long long)n * n;
for (int P = 0; P < nblk; ++P) {
const int cons = nblk - P - 1;
panel512_kernel<<<dim3(cons > 0 ? cons : 1, batch), 512>>>(L, n, bstride, P);
if (cons > 0) // +1 CTA factors and stores the pivot block
update128_regs_kernel
<<<dim3(cons * (cons + 1) / 2 + 1, batch), 256, USM>>>(L, n, bstride, P);
}
}
// Factor the diagonal tile of block-column J; one CTA per matrix.
template <int B, int NT>
__global__ void potrf_diag_kernel(float* __restrict__ A, int n, long long bstride, int J) {
extern __shared__ float smem[];
float* tile = A + blockIdx.x * bstride + (long long)J * B * (n + 1);
potrf_tile<B, NT>(tile, tile, n, smem);
}
// Solve L_IJ * L_JJ^T = A_IJ for one tile below the diagonal; smem holds
// the diagonal tile and the tile being solved, one thread per row.
template <int B, int NT>
__device__ void trsm_tile_dev(float* __restrict__ M, int n, int J, int tile_idx, float* smem) {
float (*sL)[B + 1] = (float (*)[B + 1])smem;
float (*sX)[B + 1] = (float (*)[B + 1])(smem + B * (B + 1));
const float* diag = M + (long long)J * B * (n + 1);
float* tile = M + (long long)(J + 1 + tile_idx) * B * n + (long long)J * B;
for (int idx = threadIdx.x; idx < B * B; idx += NT) {
sL[idx / B][idx % B] = diag[(idx / B) * n + idx % B];
sX[idx / B][idx % B] = tile[(idx / B) * n + idx % B];
}
__syncthreads();
const int r = threadIdx.x;
if (r < B) {
for (int k = 0; k < B; ++k) {
float acc = sX[r][k];
for (int j = 0; j < k; ++j) acc -= sX[r][j] * sL[k][j];
sX[r][k] = acc / sL[k][k];
}
for (int c = 0; c < B; ++c) tile[r * n + c] = sX[r][c];
}
__syncthreads();
}
template <int B>
__global__ void trsm_tile_kernel(float* __restrict__ A, int n, long long bstride, int J) {
__shared__ float smem[2 * B * (B + 1)];
trsm_tile_dev<B, B>(A + blockIdx.y * bstride, n, J, blockIdx.x, smem);
}
// Trailing update C_IK -= L_IJ * L_KJ^T for tile pair p (I >= K, I == K
// covers SYRK); 16x16 threads, 4x4 register tile per thread. Operands
// staged transposed ([k][row]) so each k step is two float4 loads.
template <int B>
__device__ void update_tile_dev(float* __restrict__ M, int n, int J, int p, float* smem) {
float (*sI)[B + 4] = (float (*)[B + 4])smem;
float (*sK)[B + 4] = (float (*)[B + 4])(smem + B * (B + 4));
int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
while ((Ip + 1) * (Ip + 2) / 2 <= p) ++Ip;
while (Ip * (Ip + 1) / 2 > p) --Ip;
const int I = J + 1 + Ip, K = J + 1 + (p - Ip * (Ip + 1) / 2);
const float* LI = M + (long long)I * B * n + (long long)J * B;
const float* LK = M + (long long)K * B * n + (long long)J * B;
float* C = M + (long long)I * B * n + (long long)K * B;
for (int idx = threadIdx.x; idx < B * B; idx += 256) {
const int r = idx / B, k = idx % B;
sI[k][r] = LI[r * n + k];
sK[k][r] = LK[r * n + k];
}
__syncthreads();
const int tx = threadIdx.x % 16, ty = threadIdx.x / 16;
float acc[4][4] = {};
#pragma unroll 8
for (int k = 0; k < B; ++k) {
const float4 a = *(const float4*)&sI[k][ty * 4];
const float4 b = *(const float4*)&sK[k][tx * 4];
const float ar[4] = {a.x, a.y, a.z, a.w};
const float br[4] = {b.x, b.y, b.z, b.w};
#pragma unroll
for (int i = 0; i < 4; ++i)
#pragma unroll
for (int j = 0; j < 4; ++j) acc[i][j] += ar[i] * br[j];
}
#pragma unroll
for (int i = 0; i < 4; ++i) {
float4* cp = (float4*)&C[(ty * 4 + i) * n + tx * 4];
float4 c = *cp;
c.x -= acc[i][0]; c.y -= acc[i][1]; c.z -= acc[i][2]; c.w -= acc[i][3];
*cp = c;
}
__syncthreads();
}
template <int B>
__global__ void update_tile_kernel(float* __restrict__ A, int n, long long bstride, int J) {
__shared__ float smem[2 * B * (B + 4)];
update_tile_dev<B>(A + blockIdx.y * bstride, n, J, blockIdx.x, smem);
}
// Whole factorization fused into one kernel, one CTA per matrix: no
// per-column launches, matrix stays hot in L2, only block-level syncs.
// Composes the existing tile device functions sequentially per CTA;
// batch supplies the parallelism. Also writes the tril copy itself.
template <int B, int NT>
__global__ void fused_small_kernel(const float* __restrict__ A, float* __restrict__ Lg,
int n, long long bstride) {
extern __shared__ float scratch[];
const float* src = A + blockIdx.x * bstride;
float* L = Lg + blockIdx.x * bstride;
for (int idx = threadIdx.x; idx < n * n; idx += NT) {
const int i = idx / n, j = idx % n;
L[idx] = (j <= i) ? src[idx] : 0.0f;
}
__syncthreads();
const int N = n / B;
for (int J = 0; J < N; ++J) {
float* diag = L + (long long)J * B * (n + 1);
potrf_tile<B, NT>(diag, diag, n, scratch);
const int T = N - J - 1;
for (int t = 0; t < T; ++t) trsm_tile_dev<B, NT>(L, n, J, t, scratch);
const int pairs = T * (T + 1) / 2;
for (int p = 0; p < pairs; ++p) update_tile_dev<B>(L, n, J, p, scratch);
}
}
template <int B, int NT>
void launch_fused_small(const float* A, float* L, int n, long long batch) {
constexpr int SMEM = 2 * B * (B + 4) * sizeof(float);
cudaFuncSetAttribute(fused_small_kernel<B, NT>,
cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
fused_small_kernel<B, NT><<<batch, NT, SMEM>>>(A, L, n, (long long)n * n);
}
// In-CTA n=64 diagonal potrf: threads 0..63 own rows in registers, one
// sync per step (all CTA threads participate in syncs). Result goes to
// global and to sDiag (64x68, float4-aligned rows) for the following trsm.
__device__ __forceinline__ void potrf64_regs_dev(float* __restrict__ tile, int lda,
float* __restrict__ sDiag, float* __restrict__ colbuf,
float (&r)[64]) {
const int i = threadIdx.x;
if (i < 64) {
#pragma unroll
for (int q = 0; q < 64; q += 4) *(float4*)&r[q] = *(const float4*)&tile[i * lda + q];
}
#pragma unroll
for (int k = 0; k < 64; ++k) {
if (i < 64) colbuf[(k & 1) * 64 + i] = r[k];
__syncthreads();
if (i < 64) {
const float dkk = colbuf[(k & 1) * 64 + k];
if (i > k) {
const float t = r[k] * (1.0f / dkk);
#pragma unroll
for (int j = 0; j < 64; ++j)
if (j > k && j <= i) r[j] -= t * colbuf[(k & 1) * 64 + j];
r[k] *= rsqrtf(dkk);
} else if (i == k) {
r[k] = sqrtf(dkk);
}
}
}
if (i < 64) {
#pragma unroll
for (int j = 0; j < 64; ++j)
if (j > i) r[j] = 0.0f;
#pragma unroll
for (int q = 0; q < 64; q += 4) {
*(float4*)&tile[i * lda + q] = *(float4*)&r[q];
*(float4*)&sDiag[i * 68 + q] = *(float4*)&r[q];
}
}
__syncthreads();
}
// In-CTA trsm for up to four 64-tiles at once: 256 threads, one row per
// thread in registers, diagonal tile read from smem; no internal syncs.
__device__ __forceinline__ void trsm4_regs_dev(float* __restrict__ Lg, int n, int J, int t0, int cnt,
const float* __restrict__ sDiag, float (&x)[64]) {
const int tl = threadIdx.x >> 6, r = threadIdx.x & 63;
if (tl >= cnt) return;
float* rowp = Lg + (long long)(J + 1 + t0 + tl) * 64 * n + (long long)J * 64 + (long long)r * n;
#pragma unroll
for (int q = 0; q < 64; q += 4) *(float4*)&x[q] = *(const float4*)&rowp[q];
#pragma unroll
for (int k = 0; k < 64; ++k) {
float acc = x[k];
#pragma unroll
for (int j = 0; j < 64; ++j)
if (j < k) acc -= x[j] * sDiag[k * 68 + j];
x[k] = acc * (1.0f / sDiag[k * 68 + k]);
}
#pragma unroll
for (int q = 0; q < 64; q += 4) *(float4*)&rowp[q] = *(float4*)&x[q];
}
// pair decode shared by the update helpers
__device__ __forceinline__ void pair_decode(int p, int J, int& I, int& K) {
int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
while ((Ip + 1) * (Ip + 2) <= 2 * p) ++Ip;
while (Ip * (Ip + 1) > 2 * p) --Ip;
I = J + 1 + Ip;
K = J + 1 + (p - Ip * (Ip + 1) / 2);
}
// prefetch pair p's two 64x64 panels into the shared register array
__device__ __forceinline__ void update_prefetch_dev(const float* __restrict__ M, int n,
int J, int p, float (&rr)[64]) {
int I, K;
pair_decode(p, J, I, K);
const float* LI = M + (long long)I * 64 * n + (long long)J * 64;
const float* LK = M + (long long)K * 64 * n + (long long)J * 64;
#pragma unroll
for (int j = 0; j < 16; ++j) {
const int idx = threadIdx.x + j * 256;
const int r = idx / 64, k = idx % 64;
rr[j] = LI[r * n + k];
rr[j + 16] = LK[r * n + k];
}
}
// spill the prefetched registers into the transposed staging layout
__device__ __forceinline__ void update_stage_dev(const float (&rr)[64], float* smem) {
float (*sI)[68] = (float (*)[68])smem;
float (*sK)[68] = (float (*)[68])(smem + 64 * 68);
#pragma unroll
for (int j = 0; j < 16; ++j) {
const int idx = threadIdx.x + j * 256;
const int r = idx / 64, k = idx % 64;
sI[k][r] = rr[j];
sK[k][r] = rr[j + 16];
}
}
// compute pair p's update from already-staged smem operands
__device__ __forceinline__ void update_compute_dev(float* __restrict__ M, int n, int J,
int p, float* smem) {
float (*sI)[68] = (float (*)[68])smem;
float (*sK)[68] = (float (*)[68])(smem + 64 * 68);
int I, K;
pair_decode(p, J, I, K);
float* C = M + (long long)I * 64 * n + (long long)K * 64;
const int tx = threadIdx.x % 16, ty = threadIdx.x / 16;
float acc[4][4] = {};
#pragma unroll 8
for (int k = 0; k < 64; ++k) {
const float4 a = *(const float4*)&sI[k][ty * 4];
const float4 b = *(const float4*)&sK[k][tx * 4];
const float ar[4] = {a.x, a.y, a.z, a.w};
const float br[4] = {b.x, b.y, b.z, b.w};
#pragma unroll
for (int i = 0; i < 4; ++i)
#pragma unroll
for (int j = 0; j < 4; ++j) acc[i][j] += ar[i] * br[j];
}
#pragma unroll
for (int i = 0; i < 4; ++i) {
float4* cp = (float4*)&C[(ty * 4 + i) * n + tx * 4];
float4 c = *cp;
c.x -= acc[i][0]; c.y -= acc[i][1]; c.z -= acc[i][2]; c.w -= acc[i][3];
*cp = c;
}
}
// Fused-fast: one CTA per matrix, composing the register-resident potrf
// and 4-wide register trsm with the register-tiled update; only the
// per-step syncs those primitives need. For n in {256, 512, 1024}.
template <int NT, int MINB>
__global__ void __launch_bounds__(256, MINB)
fused_fast_kernel(const float* __restrict__ A, float* __restrict__ Lg,
int n, long long bstride) {
extern __shared__ float scratch[]; // update scratch | sDiag | colbuf
float* sDiag = scratch + 2 * 64 * 68;
float* colbuf = sDiag + 64 * 68;
const float* src = A + blockIdx.x * bstride;
float* L = Lg + blockIdx.x * bstride;
for (int idx = threadIdx.x; idx < n * n; idx += NT) {
const int i = idx / n, j = idx % n;
L[idx] = (j <= i) ? src[idx] : 0.0f;
}
__syncthreads();
const int N = n / 64;
float rr[64];
for (int J = 0; J < N; ++J) {
potrf64_regs_dev(L + (long long)J * 64 * (n + 1), n, sDiag, colbuf, rr);
const int T = N - J - 1;
for (int g = 0; g < T; g += 4)
trsm4_regs_dev(L, n, J, g, (T - g < 4) ? (T - g) : 4, sDiag, rr);
__syncthreads();
const int pairs = T * (T + 1) / 2;
for (int p = 0; p < pairs; ++p) update_tile_dev<64>(L, n, J, p, scratch);
}
}
template <int NT, int MINB>
void launch_fused_fast(const float* A, float* L, int n, long long batch) {
constexpr int SMEM = (2 * 64 * 68 + 64 * 68 + 128) * sizeof(float);
cudaFuncSetAttribute(fused_fast_kernel<NT, MINB>,
cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
fused_fast_kernel<NT, MINB><<<batch, NT, SMEM>>>(A, L, n, (long long)n * n);
}
// Grid-wide software barrier for the persistent kernel: generation counter
// plus arrive count. Valid because the grid never exceeds resident capacity.
__device__ void grid_barrier(int* count, volatile int* gen, int nctas) {
__syncthreads();
if (threadIdx.x == 0) {
const int g = *gen;
if (atomicAdd(count, 1) == nctas - 1) {
*count = 0;
__threadfence();
*gen = g + 1;
} else {
while (*gen == g) __nanosleep(64);
}
__threadfence();
}
__syncthreads();
}
// Whole factorization in one launch: persistent CTAs sweep the potrf, trsm
// and update phases of each block-column, separated by grid barriers.
// Kills the per-column launch spine that dominates low-batch mid sizes.
template <int B>
__global__ void __launch_bounds__(256)
persistent_cholesky_kernel(float* __restrict__ A, int n, long long bstride,
int batch, int nctas, int* bar) {
__shared__ float smem[2 * B * (B + 4)];
const int N = n / B;
for (int J = 0; J < N; ++J) {
for (int it = blockIdx.x; it < batch; it += nctas) {
float* tile = A + it * bstride + (long long)J * B * (n + 1);
potrf_tile<B, 256>(tile, tile, n, smem);
}
grid_barrier(bar, bar + 1, nctas);
const int T = N - J - 1;
if (T == 0) break;
for (int it = blockIdx.x; it < T * batch; it += nctas)
trsm_tile_dev<B, 256>(A + (it / T) * bstride, n, J, it % T, smem);
grid_barrier(bar, bar + 1, nctas);
const int pairs = T * (T + 1) / 2;
for (int it = blockIdx.x; it < pairs * batch; it += nctas)
update_tile_dev<B>(A + (it / pairs) * bstride, n, J, it % pairs, smem);
grid_barrier(bar, bar + 1, nctas);
}
}
// Launch config: as many CTAs as can be simultaneously resident. The
// barrier buffer is allocated once; count self-resets and gen only needs
// monotonicity, so no zeroing between calls.
template <int B>
void cholesky_persistent(float* A, int n, long long batch) {
static int nctas = 0;
static int* bar = nullptr;
if (nctas == 0) {
int dev, sms, per_sm;
cudaGetDevice(&dev);
cudaDeviceGetAttribute(&sms, cudaDevAttrMultiProcessorCount, dev);
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&per_sm, persistent_cholesky_kernel<B>, 256, 0);
nctas = sms * per_sm;
cudaMalloc(&bar, 2 * sizeof(int));
cudaMemset(bar, 0, 2 * sizeof(int));
}
persistent_cholesky_kernel<B><<<nctas, 256>>>(A, n, (long long)n * n, batch, nctas, bar);
}
// Round-to-nearest fp32 -> tf32 conversion.
__device__ __forceinline__ float to_tf32(float x) {
unsigned u;
asm("cvt.rna.tf32.f32 %0, %1;" : "=r"(u) : "f"(x));
return __uint_as_float(u);
}
// One m16n8k8 TF32 tensor-core MMA, fp32 accumulate.
__device__ __forceinline__ void mma_tf32(float d[4], const unsigned a[4], const unsigned b[2]) {
asm volatile(
"mma.sync.aligned.m16n8k8.row.col.f32.tf32.tf32.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
: "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]));
}
// 16-byte async global->shared copy plus commit/wait wrappers.
__device__ __forceinline__ void cp_async16(void* dst, const void* src) {
const unsigned d = (unsigned)__cvta_generic_to_shared(dst);
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;\n" :: "r"(d), "l"(src));
}
__device__ __forceinline__ void cp_commit() { asm volatile("cp.async.commit_group;\n"); }
template <int N>
__device__ __forceinline__ void cp_wait() { asm volatile("cp.async.wait_group %0;\n" :: "n"(N)); }
// TF32 trailing update over 128x128 supertiles (2x2 groups of 64-tiles).
// 8 warps per CTA in a 2x4 grid, each warp computes a 64x32 slice with
// m16n8k8 MMAs. 16-wide k-slabs are cp.async double-buffered so the next
// slab loads in while the current one feeds the tensor cores. Edge
// subtiles (beyond T or above the diagonal) are masked.
template <int B>
__global__ void __launch_bounds__(256, 1)
update_supertile_tf32_kernel(float* __restrict__ A, int n, long long bstride, int J) {
constexpr int KW = 16, PAD = 20;
__shared__ float sI[2][2 * B][PAD], sK[2][2 * B][PAD];
const int T = n / B - J - 1;
const int p = blockIdx.x;
int Sp = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
while ((Sp + 1) * (Sp + 2) / 2 <= p) ++Sp;
while (Sp * (Sp + 1) / 2 > p) --Sp;
const int si = Sp, sk = p - Sp * (Sp + 1) / 2;
float* M = A + blockIdx.y * bstride;
const long long panel_col = (long long)J * B;
const long long row_i = (long long)(J + 1 + 2 * si) * B;
const long long row_k = (long long)(J + 1 + 2 * sk) * B;
auto stage = [&](int buf, int ks) {
for (int idx = threadIdx.x; idx < 2 * B * (KW / 4); idx += 256) {
const int r = idx / (KW / 4), q = (idx % (KW / 4)) * 4;
if (2 * si + r / B < T)
cp_async16(&sI[buf][r][q], M + (row_i + r) * n + panel_col + ks + q);
else
*(float4*)&sI[buf][r][q] = make_float4(0.f, 0.f, 0.f, 0.f);
if (2 * sk + r / B < T)
cp_async16(&sK[buf][r][q], M + (row_k + r) * n + panel_col + ks + q);
else
*(float4*)&sK[buf][r][q] = make_float4(0.f, 0.f, 0.f, 0.f);
}
};
const int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
const int wm = warp >> 2, wn = warp & 3;
const int g = lane >> 2, t = lane & 3;
float acc[4][4][4] = {};
stage(0, 0);
cp_commit();
constexpr int NSLAB = B / KW;
for (int s = 0; s < NSLAB; ++s) {
if (s + 1 < NSLAB) {
stage((s + 1) & 1, (s + 1) * KW);
cp_commit();
cp_wait<1>();
} else {
cp_wait<0>();
}
__syncthreads();
const int buf = s & 1;
#pragma unroll
for (int koff = 0; koff < KW; koff += 8) {
unsigned a[4][4], b[4][2];
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
const int r0 = wm * 64 + mt * 16 + g;
a[mt][0] = __float_as_uint(to_tf32(sI[buf][r0][koff + t]));
a[mt][1] = __float_as_uint(to_tf32(sI[buf][r0 + 8][koff + t]));
a[mt][2] = __float_as_uint(to_tf32(sI[buf][r0][koff + t + 4]));
a[mt][3] = __float_as_uint(to_tf32(sI[buf][r0 + 8][koff + t + 4]));
}
#pragma unroll
for (int nt = 0; nt < 4; ++nt) {
const int c0 = wn * 32 + nt * 8 + g;
b[nt][0] = __float_as_uint(to_tf32(sK[buf][c0][koff + t]));
b[nt][1] = __float_as_uint(to_tf32(sK[buf][c0][koff + t + 4]));
}
#pragma unroll
for (int mt = 0; mt < 4; ++mt)
#pragma unroll
for (int nt = 0; nt < 4; ++nt) mma_tf32(acc[mt][nt], a[mt], b[nt]);
}
__syncthreads();
}
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
#pragma unroll
for (int nt = 0; nt < 4; ++nt) {
#pragma unroll
for (int e = 0; e < 4; ++e) {
const int r = wm * 64 + mt * 16 + g + (e >= 2 ? 8 : 0);
const int c = wn * 32 + nt * 8 + t * 2 + (e & 1);
const int rt = 2 * si + r / B, ct = 2 * sk + c / B;
if (rt < T && ct < T && ct <= rt)
M[(row_i + r) * n + row_k + c] -= acc[mt][nt][e];
}
}
}
}
#if __CUDA_ARCH__ >= 1000
// 64-bit shared-memory matrix descriptor (no swizzle, K-major,
// core-matrix-tiled storage).
__device__ __forceinline__ unsigned long long smem_desc(unsigned addr, unsigned lbo, unsigned sbo) {
return (unsigned long long)((addr & 0x3FFFFu) >> 4)
| ((unsigned long long)((lbo >> 4) & 0x3FFFu) << 16)
| ((unsigned long long)((sbo >> 4) & 0x3FFFu) << 32)
| (1ULL << 46);
}
#endif
// Split-TF32 trailing update on 5th-gen tensor cores: same 128x128
// supertile scheme as update_supertile_tf32_kernel, but the 128x128x64
// product runs as K=8 tcgen05 MMA triples (hi*hi + hi*lo + lo*hi, ~fp32
// accuracy) accumulating in tensor memory. Operands are staged in
// core-matrix-tiled smem (8x4 tf32 blocks of 128 contiguous bytes).
template <int B>
__global__ void __launch_bounds__(128)
update_supertile_tc5_kernel(float* __restrict__ A, int n, long long bstride, int J) {
#if __CUDA_ARCH__ >= 1000
constexpr int M = 2 * B;
extern __shared__ __align__(1024) float dsm[];
float* sA = dsm;
float* sB = dsm + M * 16;
__shared__ __align__(8) unsigned long long mbar;
__shared__ unsigned taddr_slot;
const int T = n / B - J - 1;
const int p = blockIdx.x;
int Sp = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
while ((Sp + 1) * (Sp + 2) / 2 <= p) ++Sp;
while (Sp * (Sp + 1) / 2 > p) --Sp;
const int si = Sp, sk = p - Sp * (Sp + 1) / 2;
float* Mm = A + blockIdx.y * bstride;
const long long panel_col = (long long)J * B;
const long long row_i = (long long)(J + 1 + 2 * si) * B;
const long long row_k = (long long)(J + 1 + 2 * sk) * B;
const int tid = threadIdx.x;
const unsigned mbar_a = (unsigned)__cvta_generic_to_shared(&mbar);
if (tid == 0) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
asm volatile("fence.mbarrier_init.release.cluster;");
}
if (tid < 32) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(M) : "memory");
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}
// K is processed in 16-wide slabs so smem stays at 32KB and several
// CTAs stay resident per SM (TMEM caps residency at 4); hi/lo split
// panels give ~fp32 accuracy via 3-term MMA (hi*hi + hi*lo + lo*hi)
constexpr int KS = 16, NSLAB = 64 / KS, PT = M * KS / 128;
float* sAlo = dsm + 2 * M * KS;
float* sBlo = dsm + 3 * M * KS;
__syncthreads();
const unsigned taddr = taddr_slot;
const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
| ((unsigned)(M >> 3) << 17) | ((unsigned)(M >> 4) << 24);
// threads prefetch the next slab into registers while the current
// slab's MMAs run, then wait for the MMAs before rewriting smem
float ra[PT], rb[PT];
auto load_regs = [&](int ks) {
#pragma unroll
for (int j = 0; j < PT; ++j) {
const int idx = tid + j * 128;
const int m = idx / KS, k = idx % KS;
ra[j] = (2 * si + m / B < T) ? Mm[(row_i + m) * n + panel_col + ks + k] : 0.0f;
rb[j] = (2 * sk + m / B < T) ? Mm[(row_k + m) * n + panel_col + ks + k] : 0.0f;
}
};
auto wait_mbar = [&](unsigned phase) {
unsigned done = 0;
while (!done)
asm volatile(
"{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
"selp.u32 %0, 1, 0, p;\n}\n"
: "=r"(done) : "r"(mbar_a), "r"(phase));
};
load_regs(0);
for (int s = 0; s < NSLAB; ++s) {
if (s > 0) {
wait_mbar((unsigned)((s - 1) & 1));
asm volatile("tcgen05.fence::before_thread_sync;");
__syncthreads();
}
#pragma unroll
for (int j = 0; j < PT; ++j) {
const int idx = tid + j * 128;
const int m = idx / KS, k = idx % KS;
const int off = (k / 4) * (M * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
const float ah = to_tf32(ra[j]), bh = to_tf32(rb[j]);
sA[off] = ah;
sB[off] = bh;
sAlo[off] = to_tf32(ra[j] - ah);
sBlo[off] = to_tf32(rb[j] - bh);
}
asm volatile("fence.proxy.async.shared::cta;"); // generic->async visibility
__syncthreads();
if (tid == 0) {
asm volatile("tcgen05.fence::after_thread_sync;");
const unsigned aa = (unsigned)__cvta_generic_to_shared(sA);
const unsigned bb = (unsigned)__cvta_generic_to_shared(sB);
const unsigned al = (unsigned)__cvta_generic_to_shared(sAlo);
const unsigned bl = (unsigned)__cvta_generic_to_shared(sBlo);
#pragma unroll
for (int c = 0; c < KS / 8; ++c) {
const unsigned koff = (unsigned)(c * 2 * M * 16);
const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
#pragma unroll
for (int t = 0; t < 3; ++t) {
const unsigned long long da = smem_desc(ops[t][0] + koff, M * 16, 128);
const unsigned long long db = smem_desc(ops[t][1] + koff, M * 16, 128);
asm volatile(
"{\n.reg .pred p;\nsetp.ne.b32 p, %4, 0;\n"
"tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, p;\n}\n"
:: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
"r"((s != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
}
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mbar_a) : "memory");
}
if (s + 1 < NSLAB) load_regs((s + 1) * KS);
}
wait_mbar((unsigned)((NSLAB - 1) & 1));
asm volatile("tcgen05.fence::before_thread_sync;");
__syncthreads();
// each warp reads its 32-row slice of the accumulator, four x8 loads
// per wait, and subtracts from C with edge masking
const int warp = tid >> 5, lane = tid & 31;
const int r = warp * 32 + lane;
const int rt = 2 * si + r / B;
const bool interior = (2 * si + 1 < T) && (2 * sk + 1 < T);
for (int c0 = 0; c0 < M; c0 += 32) {
float v[32];
#pragma unroll
for (int q = 0; q < 4; ++q)
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
"=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
: "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
asm volatile("tcgen05.wait::ld.sync.aligned;");
float* crow = Mm + (row_i + r) * n + row_k + c0;
if (interior && si != sk) {
#pragma unroll
for (int q = 0; q < 8; ++q) {
float4* cp = (float4*)(crow + q * 4);
float4 c4 = *cp;
c4.x -= v[q * 4]; c4.y -= v[q * 4 + 1];
c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
*cp = c4;
}
} else if (interior) {
// diagonal supertile: subtile above the diagonal is masked out
const int climit = (r < B) ? B : M;
#pragma unroll
for (int q = 0; q < 8; ++q) {
if (c0 + q * 4 < climit) {
float4* cp = (float4*)(crow + q * 4);
float4 c4 = *cp;
c4.x -= v[q * 4]; c4.y -= v[q * 4 + 1];
c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
*cp = c4;
}
}
} else if (rt < T) {
#pragma unroll
for (int e = 0; e < 32; ++e) {
const int c = c0 + e;
const int ct = 2 * sk + c / B;
if (ct < T && ct <= rt)
Mm[(row_i + r) * n + row_k + c] -= v[e];
}
}
}
__syncthreads();
if (tid < 32)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(M) : "memory");
#endif
}
// b=128 split-TF32 trailing update: one CTA per 128x128 tile pair (I >= K)
// over a K=128 panel — twice the FLOP per pair of the b=64 supertile path
// and, since n % 128 == 0, no edge masking anywhere. Same slab/MMA/ld
// machinery as update_supertile_tc5_kernel.
__global__ void __launch_bounds__(128)
update_tile_tc5_128_kernel(float* __restrict__ A, int n, long long bstride, int J) {
#if __CUDA_ARCH__ >= 1000
constexpr int M = 128, KS = 16, NSLAB = 128 / KS, PT = M * KS / 128;
extern __shared__ __align__(1024) float dsm[];
float* sA = dsm;
float* sB = dsm + M * KS;
float* sAlo = dsm + 2 * M * KS;
float* sBlo = dsm + 3 * M * KS;
__shared__ __align__(8) unsigned long long mbar;
__shared__ unsigned taddr_slot;
const int p = blockIdx.x;
int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
while ((Ip + 1) * (Ip + 2) / 2 <= p) ++Ip;
while (Ip * (Ip + 1) / 2 > p) --Ip;
const int I = J + 1 + Ip, K = J + 1 + (p - Ip * (Ip + 1) / 2);
float* Mm = A + blockIdx.y * bstride;
const long long panel_col = (long long)J * M;
const long long row_i = (long long)I * M;
const long long row_k = (long long)K * M;
const int tid = threadIdx.x;
const unsigned mbar_a = (unsigned)__cvta_generic_to_shared(&mbar);
if (tid == 0) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
asm volatile("fence.mbarrier_init.release.cluster;");
}
if (tid < 32) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(M) : "memory");
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}
__syncthreads();
const unsigned taddr = taddr_slot;
const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
| ((unsigned)(M >> 3) << 17) | ((unsigned)(M >> 4) << 24);
float ra[PT], rb[PT];
auto load_regs = [&](int ks) {
#pragma unroll
for (int j = 0; j < PT; ++j) {
const int idx = tid + j * 128;
const int m = idx / KS, k = idx % KS;
ra[j] = Mm[(row_i + m) * n + panel_col + ks + k];
rb[j] = Mm[(row_k + m) * n + panel_col + ks + k];
}
};
auto wait_mbar = [&](unsigned phase) {
unsigned done = 0;
while (!done)
asm volatile(
"{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
"selp.u32 %0, 1, 0, p;\n}\n"
: "=r"(done) : "r"(mbar_a), "r"(phase));
};
load_regs(0);
for (int s = 0; s < NSLAB; ++s) {
if (s > 0) {
wait_mbar((unsigned)((s - 1) & 1));
asm volatile("tcgen05.fence::before_thread_sync;");
__syncthreads();
}
#pragma unroll
for (int j = 0; j < PT; ++j) {
const int idx = tid + j * 128;
const int m = idx / KS, k = idx % KS;
const int off = (k / 4) * (M * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
const float ah = to_tf32(ra[j]), bh = to_tf32(rb[j]);
sA[off] = ah;
sB[off] = bh;
sAlo[off] = to_tf32(ra[j] - ah);
sBlo[off] = to_tf32(rb[j] - bh);
}
asm volatile("fence.proxy.async.shared::cta;");
__syncthreads();
if (tid == 0) {
asm volatile("tcgen05.fence::after_thread_sync;");
const unsigned aa = (unsigned)__cvta_generic_to_shared(sA);
const unsigned bb = (unsigned)__cvta_generic_to_shared(sB);
const unsigned al = (unsigned)__cvta_generic_to_shared(sAlo);
const unsigned bl = (unsigned)__cvta_generic_to_shared(sBlo);
#pragma unroll
for (int c = 0; c < KS / 8; ++c) {
const unsigned koff = (unsigned)(c * 2 * M * 16);
const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
#pragma unroll
for (int t = 0; t < 3; ++t) {
const unsigned long long da = smem_desc(ops[t][0] + koff, M * 16, 128);
const unsigned long long db = smem_desc(ops[t][1] + koff, M * 16, 128);
asm volatile(
"{\n.reg .pred q;\nsetp.ne.b32 q, %4, 0;\n"
"tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, q;\n}\n"
:: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
"r"((s != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
}
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mbar_a) : "memory");
}
if (s + 1 < NSLAB) load_regs((s + 1) * KS);
}
wait_mbar((unsigned)((NSLAB - 1) & 1));
asm volatile("tcgen05.fence::before_thread_sync;");
__syncthreads();
const int warp = tid >> 5, lane = tid & 31;
const int r = warp * 32 + lane;
for (int c0 = 0; c0 < M; c0 += 32) {
float v[32];
#pragma unroll
for (int q = 0; q < 4; ++q)
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
"=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
: "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
asm volatile("tcgen05.wait::ld.sync.aligned;");
float* crow = Mm + (row_i + r) * n + row_k + c0;
#pragma unroll
for (int q = 0; q < 8; ++q) {
float4* cp = (float4*)(crow + q * 4);
float4 c4 = *cp;
c4.x -= v[q * 4]; c4.y -= v[q * 4 + 1];
c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
*cp = c4;
}
}
__syncthreads();
if (tid < 32)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(M) : "memory");
#endif
}
// 128-wide trsm needs 132KB smem, so dynamic (B200 only).
__global__ void trsm_tile_128_kernel(float* __restrict__ A, int n, long long bstride, int J) {
extern __shared__ float smem[];
trsm_tile_dev<128, 128>(A + blockIdx.y * bstride, n, J, blockIdx.x, smem);
}
// b=128 blocked driver: fp32 potrf/trsm panels (B200 has the smem for the
// 128-wide trsm), tc5 trailing update. Requires n % 128 == 0.
void cholesky_blocked_128(float* A, int n, long long batch) {
constexpr int B = 128;
constexpr int PSMEM = B * (B + 1) * sizeof(float);
constexpr int TSMEM = 2 * B * (B + 1) * sizeof(float);
constexpr int USMEM = 56 * 1024;
const int N = n / B;
const long long bstride = (long long)n * n;
cudaFuncSetAttribute(potrf_diag_kernel<B, 256>,
cudaFuncAttributeMaxDynamicSharedMemorySize, PSMEM);
cudaFuncSetAttribute(trsm_tile_128_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, TSMEM);
cudaFuncSetAttribute(update_tile_tc5_128_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, USMEM);
for (int J = 0; J < N; ++J) {
potrf_diag_kernel<B, 256><<<batch, 256, PSMEM>>>(A, n, bstride, J);
const int T = N - J - 1;
if (T > 0) {
trsm_tile_128_kernel<<<dim3(T, batch), B, TSMEM>>>(A, n, bstride, J);
update_tile_tc5_128_kernel
<<<dim3(T * (T + 1) / 2, batch), 128, USMEM>>>(A, n, bstride, J);
}
}
}
// Right-looking blocked Cholesky, one launch trio per block-column.
// Expects the strict upper triangle of every matrix to already be zero.
template <int B>
void cholesky_blocked(float* A, int n, long long batch, int mode) {
constexpr int SMEM = B * (B + 1) * sizeof(float);
const int N = n / B;
const long long bstride = (long long)n * n;
for (int J = 0; J < N; ++J) {
potrf_diag_kernel<B, 128><<<batch, 128, SMEM>>>(A, n, bstride, J);
const int T = N - J - 1;
if (T > 0) {
trsm_tile_kernel<B><<<dim3(T, batch), B>>>(A, n, bstride, J);
const int S = (T + 1) / 2;
if (mode == 3) {
// padded above the 32KB actually used so at most 4 CTAs
// are resident per SM (4 x 128 TMEM columns = capacity)
constexpr int SMEM5 = 56 * 1024;
cudaFuncSetAttribute(update_supertile_tc5_kernel<B>,
cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM5);
update_supertile_tc5_kernel<B>
<<<dim3(S * (S + 1) / 2, batch), 128, SMEM5>>>(A, n, bstride, J);
} else if (mode == 1) {
update_supertile_tf32_kernel<B>
<<<dim3(S * (S + 1) / 2, batch), 256>>>(A, n, bstride, J);
} else {
update_tile_kernel<B>
<<<dim3(T * (T + 1) / 2, batch), 256>>>(A, n, bstride, J);
}
}
}
}
// Dynamic dataflow scheduler: persistent CTAs claim nodes from a global
// ticket in topological order (per column: potrf, trsms, update pairs;
// round-robin over matrices) and poll per-tile dependency state. No phase
// barriers, so column J+1 panel work overlaps column J's update tail.
// ucount[I][K] counts columns applied to tile (I,K); readiness is:
// potrf(J): ucount[J][J]==J; trsm(I,J): potrf done && ucount[I][J]==J;
// update(I,K,J): both trsms done && ucount[I][K]==J.
__global__ void __launch_bounds__(256)
sched_cholesky_kernel(float* __restrict__ A, int n, long long bstride, int batch,
int* __restrict__ ticket, int* __restrict__ potrf_done,
int* __restrict__ trsm_done, int* __restrict__ ucount,
long long total_per_m) {
extern __shared__ float scratch[];
__shared__ int cur;
const int N = n / 64;
const int tid = threadIdx.x;
for (;;) {
if (tid == 0) cur = atomicAdd(ticket, 1);
__syncthreads();
const long long t = cur;
__syncthreads();
if (t >= total_per_m * batch) return;
const int m = (int)(t % batch);
long long off = t / batch;
int J = 0, T, P;
for (;; ++J) {
T = N - J - 1;
P = T * (T + 1) / 2;
if (off < 1 + T + P) break;
off -= 1 + T + P;
}
float* M = A + (long long)m * bstride;
volatile int* pd = (volatile int*)(potrf_done + (long long)m * N);
volatile int* td = (volatile int*)(trsm_done + (long long)m * N * N);
volatile int* uc = (volatile int*)(ucount + (long long)m * N * N);
if (off == 0) {
if (tid == 0) while (uc[J * N + J] < J) __nanosleep(64);
__syncthreads();
float* diag = M + (long long)J * 64 * (n + 1);
potrf_tile<64, 256>(diag, diag, n, scratch);
if (tid == 0) { __threadfence(); pd[J] = 1; }
} else if (off <= T) {
const int I = J + 1 + (int)(off - 1);
if (tid == 0) while (!pd[J] || uc[I * N + J] < J) __nanosleep(64);
__syncthreads();
trsm_tile_dev<64, 256>(M, n, J, I - J - 1, scratch);
if (tid == 0) { __threadfence(); td[J * N + I] = 1; }
} else {
const int p = (int)(off - 1 - T);
int I, K;
pair_decode(p, J, I, K);
if (tid == 0)
while (!td[J * N + I] || !td[J * N + K] || uc[I * N + K] < J) __nanosleep(64);
__syncthreads();
update_tile_dev<64>(M, n, J, p, scratch);
if (tid == 0) { __threadfence(); atomicAdd((int*)&uc[I * N + K], 1); }
}
__syncthreads();
}
}
// Scheduler state buffers are cached and grown as needed; one memset per
// call resets ticket + flags.
void cholesky_sched(float* A, int n, long long batch) {
const int N = n / 64;
long long total = 0;
for (int J = 0; J < N; ++J) {
const long long T = N - J - 1;
total += 1 + T + T * (T + 1) / 2;
}
static int* buf = nullptr;
static long long cap = 0;
const long long need = 1 + batch * (long long)N + 2 * batch * (long long)N * N;
if (need > cap) {
if (buf) cudaFree(buf);
cudaMalloc(&buf, need * sizeof(int));
cap = need;
}
cudaMemset(buf, 0, need * sizeof(int));
int* ticket = buf;
int* pd = buf + 1;
int* td = pd + batch * N;
int* uc = td + batch * (long long)N * N;
const long long work = total * batch;
const int grid = (int)((work < 444) ? work : 444);
constexpr int SMEM = 2 * 64 * 68 * sizeof(float);
sched_cholesky_kernel<<<grid, 256, SMEM>>>(A, n, (long long)n * n, (int)batch,
ticket, pd, td, uc, total);
}
// tc5 split-tf32 supertile update as a scheduler node: the TMEM
// accumulator and mbarrier are CTA-persistent (allocated once by the
// scheduler); `commits` tracks mbarrier phase parity across nodes.
__device__ void sched_tc5_update_dev(float* __restrict__ M, int n, int J, int sp, int T,
float* dsm, unsigned taddr, unsigned mbar_a,
int& commits) {
#if __CUDA_ARCH__ >= 1000
constexpr int B = 64, MW = 128, KS = 16;
int Sp = (int)((sqrt(8.0 * sp + 1.0) - 1.0) / 2.0);
while ((Sp + 1) * (Sp + 2) <= 2 * sp) ++Sp;
while (Sp * (Sp + 1) > 2 * sp) --Sp;
const int si = Sp, sk = sp - Sp * (Sp + 1) / 2;
const long long panel_col = (long long)J * B;
const long long row_i = (long long)(J + 1 + 2 * si) * B;
const long long row_k = (long long)(J + 1 + 2 * sk) * B;
const int tid = threadIdx.x;
float* sA = dsm;
float* sB = dsm + MW * KS;
float* sAlo = dsm + 2 * MW * KS;
float* sBlo = dsm + 3 * MW * KS;
const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
| ((unsigned)(MW >> 3) << 17) | ((unsigned)(MW >> 4) << 24);
for (int ks = 0; ks < B; ks += KS) {
for (int idx = tid; idx < MW * KS; idx += 128) {
const int m = idx / KS, k = idx % KS;
const int off = (k / 4) * (MW * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
const float av = (2 * si + m / B < T) ? M[(row_i + m) * n + panel_col + ks + k] : 0.0f;
const float bv = (2 * sk + m / B < T) ? M[(row_k + m) * n + panel_col + ks + k] : 0.0f;
const float ah = to_tf32(av), bh = to_tf32(bv);
sA[off] = ah;
sB[off] = bh;
sAlo[off] = to_tf32(av - ah);
sBlo[off] = to_tf32(bv - bh);
}
asm volatile("fence.proxy.async.shared::cta;");
__syncthreads();
if (tid == 0) {
asm volatile("tcgen05.fence::after_thread_sync;");
const unsigned aa = (unsigned)__cvta_generic_to_shared(sA);
const unsigned bb = (unsigned)__cvta_generic_to_shared(sB);
const unsigned al = (unsigned)__cvta_generic_to_shared(sAlo);
const unsigned bl = (unsigned)__cvta_generic_to_shared(sBlo);
#pragma unroll
for (int c = 0; c < KS / 8; ++c) {
const unsigned koff = (unsigned)(c * 2 * MW * 16);
const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
#pragma unroll
for (int t = 0; t < 3; ++t) {
const unsigned long long da = smem_desc(ops[t][0] + koff, MW * 16, 128);
const unsigned long long db = smem_desc(ops[t][1] + koff, MW * 16, 128);
asm volatile(
"{\n.reg .pred q;\nsetp.ne.b32 q, %4, 0;\n"
"tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, q;\n}\n"
:: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
"r"((ks != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
}
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mbar_a) : "memory");
}
++commits;
{
unsigned done = 0;
const unsigned phase = (unsigned)((commits - 1) & 1);
while (!done)
asm volatile(
"{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
"selp.u32 %0, 1, 0, p;\n}\n"
: "=r"(done) : "r"(mbar_a), "r"(phase));
}
asm volatile("tcgen05.fence::before_thread_sync;");
__syncthreads();
}
const int warp = tid >> 5, lane = tid & 31;
const int r = warp * 32 + lane;
const int rt = 2 * si + r / B;
const bool interior = (2 * si + 1 < T) && (2 * sk + 1 < T);
for (int c0 = 0; c0 < MW; c0 += 32) {
float v[32];
#pragma unroll
for (int q = 0; q < 4; ++q)
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
"=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
: "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
asm volatile("tcgen05.wait::ld.sync.aligned;");
float* crow = M + (row_i + r) * n + row_k + c0;
if (interior && si != sk) {
#pragma unroll
for (int q = 0; q < 8; ++q) {
float4* cp = (float4*)(crow + q * 4);
float4 c4 = *cp;
c4.x -= v[q * 4]; c4.y -= v[q * 4 + 1];
c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
*cp = c4;
}
} else if (interior) {
const int climit = (r < B) ? B : MW;
#pragma unroll
for (int q = 0; q < 8; ++q) {
if (c0 + q * 4 < climit) {
float4* cp = (float4*)(crow + q * 4);
float4 c4 = *cp;
c4.x -= v[q * 4]; c4.y -= v[q * 4 + 1];
c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
*cp = c4;
}
}
} else if (rt < T) {
#pragma unroll
for (int e = 0; e < 32; ++e) {
const int c = c0 + e;
const int ct = 2 * sk + c / B;
if (ct < T && ct <= rt)
M[(row_i + r) * n + row_k + c] -= v[e];
}
}
}
__syncthreads();
#endif
}
// Dataflow scheduler with tc5 supertile update nodes: same ticket/poll
// design as sched_cholesky_kernel, but updates are 128x128 supertiles on
// the tensor cores, and each supertile completion bumps all its pair
// counters. TMEM and the mbarrier live for the whole CTA.
__global__ void __launch_bounds__(128)
sched5_cholesky_kernel(float* __restrict__ A, int n, long long bstride, int batch,
int* __restrict__ ticket, int* __restrict__ potrf_done,
int* __restrict__ trsm_done, int* __restrict__ ucount,
long long total_per_m) {
#if __CUDA_ARCH__ >= 1000
extern __shared__ __align__(1024) float dsm[];
__shared__ __align__(8) unsigned long long mbar;
__shared__ unsigned taddr_slot;
__shared__ int cur;
const int N = n / 64;
const int tid = threadIdx.x;
const unsigned mbar_a = (unsigned)__cvta_generic_to_shared(&mbar);
if (tid == 0) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
asm volatile("fence.mbarrier_init.release.cluster;");
}
if (tid < 32) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(128) : "memory");
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}
__syncthreads();
const unsigned taddr = taddr_slot;
int commits = 0;
for (;;) {
if (tid == 0) cur = atomicAdd(ticket, 1);
__syncthreads();
const long long t = cur;
__syncthreads();
if (t >= total_per_m * batch) break;
const int m = (int)(t % batch);
long long off = t / batch;
int J = 0, T, S;
for (;; ++J) {
T = N - J - 1;
S = (T + 1) / 2;
const long long cnt = 1 + T + (long long)S * (S + 1) / 2;
if (off < cnt) break;
off -= cnt;
}
float* M = A + (long long)m * bstride;
volatile int* pd = (volatile int*)(potrf_done + (long long)m * N);
volatile int* td = (volatile int*)(trsm_done + (long long)m * N * N);
volatile int* uc = (volatile int*)(ucount + (long long)m * N * N);
if (off == 0) {
if (tid == 0) while (uc[J * N + J] < J) __nanosleep(64);
__syncthreads();
float* diag = M + (long long)J * 64 * (n + 1);
potrf_tile<64, 128>(diag, diag, n, dsm);
if (tid == 0) { __threadfence(); pd[J] = 1; }
} else if (off <= T) {
const int I = J + 1 + (int)(off - 1);
if (tid == 0) while (!pd[J] || uc[I * N + J] < J) __nanosleep(64);
__syncthreads();
trsm_tile_dev<64, 128>(M, n, J, I - J - 1, dsm);
if (tid == 0) { __threadfence(); td[J * N + I] = 1; }
} else {
const int sp = (int)(off - 1 - T);
int Sp = (int)((sqrt(8.0 * sp + 1.0) - 1.0) / 2.0);
while ((Sp + 1) * (Sp + 2) <= 2 * sp) ++Sp;
while (Sp * (Sp + 1) > 2 * sp) --Sp;
const int si = Sp, sk = sp - Sp * (Sp + 1) / 2;
if (tid == 0) {
for (int d1 = 0; d1 < 2; ++d1)
for (int d2 = 0; d2 < 2; ++d2) {
const int It = 2 * si + d1, Kt = 2 * sk + d2;
if (It < T && Kt < T && Kt <= It) {
const int I = J + 1 + It, K = J + 1 + Kt;
while (!td[J * N + I] || !td[J * N + K] || uc[I * N + K] < J)
__nanosleep(64);
}
}
}
__syncthreads();
sched_tc5_update_dev(M, n, J, sp, T, dsm, taddr, mbar_a, commits);
if (tid == 0) {
__threadfence();
for (int d1 = 0; d1 < 2; ++d1)
for (int d2 = 0; d2 < 2; ++d2) {
const int It = 2 * si + d1, Kt = 2 * sk + d2;
if (It < T && Kt < T && Kt <= It)
atomicAdd((int*)&uc[(J + 1 + It) * N + (J + 1 + Kt)], 1);
}
}
}
__syncthreads();
}
__syncthreads();
if (tid < 32)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(128) : "memory");
#endif
}
// 56KB smem pad keeps at most 4 CTAs per SM resident, matching the 512
// TMEM columns so no CTA ever blocks in tcgen05.alloc.
void cholesky_sched5(float* A, int n, long long batch) {
const int N = n / 64;
long long total = 0;
for (int J = 0; J < N; ++J) {
const long long T = N - J - 1, S = (T + 1) / 2;
total += 1 + T + S * (S + 1) / 2;
}
static int* buf = nullptr;
static long long cap = 0;
const long long need = 1 + batch * (long long)N + 2 * batch * (long long)N * N;
if (need > cap) {
if (buf) cudaFree(buf);
cudaMalloc(&buf, need * sizeof(int));
cap = need;
}
cudaMemset(buf, 0, need * sizeof(int));
int* ticket = buf;
int* pd = buf + 1;
int* td = pd + batch * N;
int* uc = td + batch * (long long)N * N;
const long long work = total * batch;
const int grid = (int)((work < 592) ? work : 592);
constexpr int SMEM = 56 * 1024;
cudaFuncSetAttribute(sched5_cholesky_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
sched5_cholesky_kernel<<<grid, 128, SMEM>>>(A, n, (long long)n * n, (int)batch,
ticket, pd, td, uc, total);
}
// Split one 128-wide block-column panel into tf32 hi/lo pairs stored
// slab-major in core-matrix-tiled order so pipeline stages can be filled
// with plain contiguous async copies. One CTA per (trailing rowblock,
// 16-wide slab): scratch offset = ((rb * 8 + s) * 2 + {hi,lo}) * 2048.
__global__ void presplit_panel_kernel(const float* __restrict__ L, int n, long long bstride,
int J, float* __restrict__ scratch, long long sstride) {
const int rb = blockIdx.x, sl = blockIdx.y;
const float* M = L + blockIdx.z * bstride;
float* out = scratch + blockIdx.z * sstride + ((long long)(rb * 8 + sl) * 2) * 2048;
const long long row0 = (long long)(J + 1) * 128 + (long long)rb * 128;
const long long col0 = (long long)J * 128 + sl * 16;
for (int idx = threadIdx.x; idx < 128 * 16; idx += 128) {
const int m = idx / 16, k = idx % 16;
const int off = (k / 4) * (128 * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
const float v = M[(row0 + m) * n + col0 + k];
const float h = to_tf32(v);
out[off] = h;
out[2048 + off] = to_tf32(v - h);
}
}
// Warp-specialized-style pipelined trailing update for b=128 tiles:
// 3-stage cp.async pipeline over pre-split slabs; the tensor pipe never
// drains (mbarrier wait is only for buffer reuse two slabs back). One CTA
// per 128x128 tile pair, K = 128 in eight 16-wide slabs, split-tf32.
__global__ void __launch_bounds__(128)
update128_pipe_kernel(float* __restrict__ A, int n, long long bstride, int J,
const float* __restrict__ scratch, long long sstride) {
#if __CUDA_ARCH__ >= 1000
constexpr int MW = 128, KS = 16, NSLAB = 8, STG = 3;
constexpr int CH = MW * KS; // floats per hi or lo chunk (2048)
extern __shared__ __align__(1024) float dsm[]; // STG stages x 4 chunks
__shared__ __align__(8) unsigned long long mbar[STG];
__shared__ unsigned taddr_slot;
const int p = blockIdx.x;
int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
while ((Ip + 1) * (Ip + 2) <= 2 * p) ++Ip;
while (Ip * (Ip + 1) > 2 * p) --Ip;
const int Kp = p - Ip * (Ip + 1) / 2;
float* M = A + blockIdx.y * bstride;
const float* scr = scratch + blockIdx.y * sstride;
const long long row_i = (long long)(J + 1 + Ip) * MW;
const long long row_k = (long long)(J + 1 + Kp) * MW;
const int tid = threadIdx.x;
unsigned mb[STG];
if (tid == 0) {
#pragma unroll
for (int i = 0; i < STG; ++i)
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;"
:: "r"((unsigned)__cvta_generic_to_shared(&mbar[i])));
asm volatile("fence.mbarrier_init.release.cluster;");
}
#pragma unroll
for (int i = 0; i < STG; ++i) mb[i] = (unsigned)__cvta_generic_to_shared(&mbar[i]);
if (tid < 32) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(MW) : "memory");
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}
__syncthreads();
const unsigned taddr = taddr_slot;
const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
| ((unsigned)(MW >> 3) << 17) | ((unsigned)(MW >> 4) << 24);
// stage loader: copy 4 pre-tiled chunks (Ahi, Alo, Khi, Klo) for slab sl
auto load_stage = [&](int buf, int sl) {
float* dst = dsm + (long long)buf * 4 * CH;
const float* si = scr + ((long long)(Ip * 8 + sl) * 2) * CH;
const float* sk = scr + ((long long)(Kp * 8 + sl) * 2) * CH;
for (int q = tid; q < CH / 4; q += 128) {
cp_async16(dst + q * 4, si + q * 4);
cp_async16(dst + CH + q * 4, si + CH + q * 4);
cp_async16(dst + 2 * CH + q * 4, sk + q * 4);
cp_async16(dst + 3 * CH + q * 4, sk + CH + q * 4);
}
cp_commit();
};
auto wait_mbar = [&](unsigned a, unsigned phase) {
unsigned done = 0;
while (!done)
asm volatile(
"{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
"selp.u32 %0, 1, 0, p;\n}\n"
: "=r"(done) : "r"(a), "r"(phase));
};
int uses[STG] = {0, 0, 0};
load_stage(0, 0);
load_stage(1, 1);
for (int sl = 0; sl < NSLAB; ++sl) {
const int buf = sl % STG;
// loads for slab sl were issued 2 stages ago (or in the prologue);
// allow one newer group to stay in flight — except at the tail,
// where the newest group IS this slab's and must complete
if (sl + 2 < NSLAB) cp_wait<1>(); else cp_wait<0>();
asm volatile("fence.proxy.async.shared::cta;");
__syncthreads();
if (tid == 0) {
asm volatile("tcgen05.fence::after_thread_sync;");
const unsigned aa = (unsigned)__cvta_generic_to_shared(dsm + (long long)buf * 4 * CH);
const unsigned al = aa + CH * 4;
const unsigned bb = aa + 2 * CH * 4;
const unsigned bl = aa + 3 * CH * 4;
#pragma unroll
for (int c = 0; c < KS / 8; ++c) {
const unsigned koff = (unsigned)(c * 2 * MW * 16);
const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
#pragma unroll
for (int t = 0; t < 3; ++t) {
const unsigned long long da = smem_desc(ops[t][0] + koff, MW * 16, 128);
const unsigned long long db = smem_desc(ops[t][1] + koff, MW * 16, 128);
asm volatile(
"{\n.reg .pred q;\nsetp.ne.b32 q, %4, 0;\n"
"tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, q;\n}\n"
:: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
"r"((sl != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
}
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mb[buf]) : "memory");
}
++uses[buf];
__syncthreads();
if (sl + 2 < NSLAB) {
const int nbuf = (sl + 2) % STG;
// reuse of nbuf requires its previous mma group to be complete
if (uses[nbuf] > 0) wait_mbar(mb[nbuf], (unsigned)((uses[nbuf] - 1) & 1));
load_stage(nbuf, sl + 2);
}
}
wait_mbar(mb[(NSLAB - 1) % STG], (unsigned)((uses[(NSLAB - 1) % STG] - 1) & 1));
asm volatile("tcgen05.fence::before_thread_sync;");
__syncthreads();
const int warp = tid >> 5, lane = tid & 31;
const int r = warp * 32 + lane;
for (int c0 = 0; c0 < MW; c0 += 32) {
float v[32];
#pragma unroll
for (int q = 0; q < 4; ++q)
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
"=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
: "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
asm volatile("tcgen05.wait::ld.sync.aligned;");
float* crow = M + (row_i + r) * n + row_k + c0;
#pragma unroll
for (int q = 0; q < 8; ++q) {
float4* cp = (float4*)(crow + q * 4);
float4 c4 = *cp;
c4.x -= v[q * 4]; c4.y -= v[q * 4 + 1];
c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
*cp = c4;
}
}
__syncthreads();
if (tid < 32)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(MW) : "memory");
#endif
}
// b=128 blocked driver with the pipelined tensor-core update; panels are
// pre-split once per column. Requires n % 128 == 0.
void cholesky_pipe128(float* A, int n, long long batch) {
constexpr int B = 128;
constexpr int PSM = B * (B + 1) * sizeof(float);
constexpr int TSM = 2 * B * (B + 1) * sizeof(float);
constexpr int USM = 3 * 4 * 128 * 16 * sizeof(float);
const int N = n / B;
const long long bstride = (long long)n * n;
const long long sstride = (long long)(N - 1) * 8 * 2 * 2048;
static float* scratch = nullptr;
static long long scap = 0;
if (sstride * batch > scap) {
if (scratch) cudaFree(scratch);
cudaMalloc(&scratch, sstride * batch * sizeof(float));
scap = sstride * batch;
}
cudaFuncSetAttribute(potrf_diag_kernel<B, 256>,
cudaFuncAttributeMaxDynamicSharedMemorySize, PSM);
cudaFuncSetAttribute(trsm_tile_128_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, TSM);
cudaFuncSetAttribute(update128_pipe_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, USM);
for (int J = 0; J < N; ++J) {
potrf_diag_kernel<B, 256><<<batch, 256, PSM>>>(A, n, bstride, J);
const int T = N - J - 1;
if (T > 0) {
trsm_tile_128_kernel<<<dim3(T, batch), B, TSM>>>(A, n, bstride, J);
presplit_panel_kernel<<<dim3(T, 8, batch), 128>>>(A, n, bstride, J, scratch, sstride);
update128_pipe_kernel
<<<dim3(T * (T + 1) / 2, batch), 128, USM>>>(A, n, bstride, J, scratch, sstride);
}
}
}
"""
CUDA_SRC = r"""
#include <cudaTypedefs.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
""" + CUDA_KERNELS + r"""
// mode: 0 = fp32 launch-per-column, 1 = tf32 update, 2 = persistent kernel.
at::Tensor batched_cholesky(at::Tensor A, int64_t mode) {
const long long batch = A.size(0);
const int n = A.size(1);
if (mode == 5 && n > 64 && n % 64 == 0) {
at::Tensor L = A.new_empty(A.sizes());
launch_fused_small<64, 256>(A.data_ptr<float>(), L.data_ptr<float>(), n, batch);
return L;
}
if (mode == 12 && n == 256) {
at::Tensor L = A.new_empty(A.sizes());
launch_fused256(A.data_ptr<float>(), L.data_ptr<float>(), batch);
return L;
}
if ((mode == 6 || mode == 7) && n > 64 && n % 64 == 0) {
at::Tensor L = A.new_empty(A.sizes());
if (mode == 7)
launch_fused_fast<256, 2>(A.data_ptr<float>(), L.data_ptr<float>(), n, batch);
else
launch_fused_fast<256, 1>(A.data_ptr<float>(), L.data_ptr<float>(), n, batch);
return L;
}
if (n == 32 || n == 64 || n == 128) {
at::Tensor L = A.new_empty(A.sizes());
const float* src = A.data_ptr<float>();
float* dst = L.data_ptr<float>();
if (n == 32)
potrf32_warp_kernel<<<(batch + 3) / 4, 128>>>(src, dst, batch);
else if (n == 64)
potrf64_rows_kernel<<<batch, 64>>>(src, dst, batch);
else // n == 128: regs scheme beats the smem-rows kernel at every batch
potrf128_regs_kernel<<<batch, 256>>>(src, dst, batch);
return L;
}
TORCH_CHECK(n % 64 == 0, "batched_cholesky: unsupported n=", n);
TORCH_CHECK(batch <= 65535, "batched_cholesky: batch too large");
if (mode == 13) {
TORCH_CHECK(n % 128 == 0, "b128 needs n % 128 == 0");
at::Tensor L = A.new_empty(A.sizes());
const int nblk4 = (n * (n / 4) + 255) / 256;
tril_f4_kernel<<<dim3(nblk4 < 512 ? nblk4 : 512, batch), 256>>>(
A.data_ptr<float>(), L.data_ptr<float>(), n, (long long)n * n);
cholesky_b128regs(L.data_ptr<float>(), n, batch);
return L;
}
at::Tensor L = A.tril(); // blocked path never touches the upper triangle
if (mode == 10) {
TORCH_CHECK(n % 128 == 0, "pipe128 needs n % 128 == 0");
cholesky_pipe128(L.data_ptr<float>(), n, batch);
} else if (mode == 9) {
cholesky_sched5(L.data_ptr<float>(), n, batch);
} else if (mode == 8) {
cholesky_sched(L.data_ptr<float>(), n, batch);
} else if (mode == 4) {
TORCH_CHECK(n % 128 == 0, "tc5b128 needs n % 128 == 0");
cholesky_blocked_128(L.data_ptr<float>(), n, batch);
} else if (mode == 2) {
cholesky_persistent<64>(L.data_ptr<float>(), n, batch);
} else {
cholesky_blocked<64>(L.data_ptr<float>(), n, batch, (int)mode);
}
return L;
}
// Left-looking blocked-128 hybrid: per column one strided-batched cuBLAS
// baddbmm_ (tf32-emulated fp32, K grows with the column) applies ALL prior
// panels at once, then the register panel sweep factors the column. The
// caller must enable torch's allow_tf32 matmul flag around this call.
at::Tensor b128_hybrid(at::Tensor A) {
const long long batch = A.size(0);
const int n = A.size(1);
TORCH_CHECK(n % 128 == 0, "b128t needs n % 128 == 0");
const int nblk = n / 128;
const long long bstride = (long long)n * n;
at::Tensor L = A.new_empty(A.sizes());
const int nblk4 = (n * (n / 4) + 255) / 256;
tril_f4_kernel<<<dim3(nblk4 < 512 ? nblk4 : 512, batch), 256>>>(
A.data_ptr<float>(), L.data_ptr<float>(), n, bstride);
at::Tensor piv = A.new_empty({batch, (long long)nblk, 128, 128});
float* Lp = L.data_ptr<float>();
for (int Q = 0; Q < nblk; ++Q) {
if (Q > 0) {
at::Tensor Lc = L.narrow(1, 128 * Q, n - 128 * Q).narrow(2, 0, 128 * Q);
at::Tensor Lr = L.narrow(1, 128 * Q, 128).narrow(2, 0, 128 * Q);
at::Tensor C = L.narrow(1, 128 * Q, n - 128 * Q).narrow(2, 128 * Q, 128);
C.baddbmm_(Lc, Lr.mT(), 1, -1);
}
const int cons = nblk - Q - 1;
panel512_left_kernel<<<dim3(cons > 0 ? cons : 1, batch), 512>>>(
Lp, piv.data_ptr<float>(), n, bstride, Q);
}
if (nblk > 1)
copy_pivots_kernel<<<dim3(nblk - 1, batch), 256>>>(
Lp, piv.data_ptr<float>(), n, bstride);
return L;
}
TORCH_LIBRARY(codex, m) {
m.def("batched_cholesky(Tensor A, int mode) -> Tensor");
m.impl("batched_cholesky", &batched_cholesky);
m.def("b128_hybrid(Tensor A) -> Tensor");
m.impl("b128_hybrid", &b128_hybrid);
}
"""
BUILD_DIR = Path(__file__).resolve().parent / ".build"
BUILD_DIR.mkdir(exist_ok=True)
_mod = None
_modt = None
try:
load_inline(
name="codex",
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
verbose=True,
is_python_module=False,
no_implicit_headers=True,
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=[
"-O3",
"-gencode=arch=compute_100a,code=sm_100a",
"--use_fast_math",
"--expt-relaxed-constexpr",
"--relocatable-device-code=false",
"-lineinfo",
"-Xptxas=-v",
],
build_directory=str(BUILD_DIR),
)
_mod = torch.ops.codex.batched_cholesky
_modt = torch.ops.codex.b128_hybrid
P("[cholesky] compiled OK")
except Exception as e:
P(f"[cholesky] compilation failed (tail of build log follows)")
P(str(e)[-4000:])
# Blocked right-looking Cholesky from torch ops; the trailing update runs
# cuBLAS TF32 tensor-core kernels, SYRK-aware (only at/below-diagonal
# chunks, in-place). Wins for n >= 8192.
def blocked_cholesky_torch(A: torch.Tensor, b: int) -> torch.Tensor:
saved = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
n = A.size(-1)
L = A.tril()
for s in range(0, n, b):
e = min(s + b, n)
Ljj = torch.linalg.cholesky(L[:, s:e, s:e])
L[:, s:e, s:e] = Ljj
if e < n:
X = torch.linalg.solve_triangular(
Ljj.mT, L[:, e:, s:e], upper=True, left=False)
L[:, e:, s:e] = X
for c in range(e, n, b):
ce = min(c + b, n)
L[:, c:, c:ce].baddbmm_(
X[:, c - e:, :], X[:, c - e:ce - e, :].mT, beta=1, alpha=-1)
return L.tril()
finally:
torch.backends.cuda.matmul.allow_tf32 = saved
# (batch, n) -> fastest measured path on B200 (2026-07-19). "cuda" is our
# kernels, "torch" is batched cuSOLVER, "loop" is per-matrix cuSOLVER,
# "torch_blocked" is blocked_cholesky_torch.
DISPATCH = {
(4096, 32): "cuda",
(1024, 64): "cuda",
(256, 128): "regs128",
(64, 256): "fused256",
(16, 512): "b128t",
(4, 512): "b128t",
(640, 512): "b128t",
(4, 1024): "b128t",
(60, 1024): "b128t",
(2, 2048): "b128t",
(8, 2048): "b128t",
(1, 4096): "torch",
(2, 4096): "loop",
(1, 8192): "torch_blocked",
(1, 16384): "torch_blocked",
(1, 32768): "torch_blocked",
}
def custom_kernel(data: input_t) -> output_t:
batch, n = data.size(0), data.size(-1)
supported = (n in (32, 64, 128) or n % 64 == 0) and batch <= 65535
mode = DISPATCH.get((batch, n), "cuda" if supported else "torch")
if mode == "b128t" and _mod is not None and n % 128 == 0:
# plain tf32 bmm fails the checker (residuals ~2.5); fp32 cuBLAS on
# B200 is multi-pass emulated and both accurate and fast
saved = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = False
try:
return _modt(data)
except Exception as e:
P(f"[cholesky] b128t failed: {e}")
finally:
torch.backends.cuda.matmul.allow_tf32 = saved
if mode in ("cuda", "tf32", "persist", "tc5", "tc5b128", "fused",
"fusedfast", "fusedfast2", "sched", "sched5", "pipe128",
"regs128", "fused256", "b128") and _mod is not None and supported:
try:
return _mod(data, {"cuda": 0, "tf32": 1, "persist": 2, "tc5": 3,
"tc5b128": 4, "fused": 5, "fusedfast": 6,
"fusedfast2": 7, "sched": 8, "sched5": 9,
"pipe128": 10, "regs128": 11, "fused256": 12,
"b128": 13}[mode])
except Exception as e:
P(f"[cholesky] kernel failed: {e}")
if mode == "torch_blocked":
return blocked_cholesky_torch(data, 2048)
if mode == "loop":
out = torch.empty_like(data)
for i in range(batch):
torch.linalg.cholesky(data[i], out=out[i])
return out
return torch.linalg.cholesky(data)
scrolls · 2290 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