submission 912875
icecuber · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 942 lines, June 9 Researcher Reciprocity License v1.0.
champ_submit.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-912875?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:bbe9dba598b007d5f41dad1d04ca84f12420e002cd813d380c7c95dc1087b41e
license declaredunknown
license concludedunknown
authorsicecuber
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float smem[];vector-width = float4
const float4 v = *(const float4 *)(src + t4 * 4);Kernel source
champ_submit.py942 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
import os
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "10.0")
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <ATen/cuda/CUDAContext.h>
#define FULL 0xffffffffu
// n == 32: one warp per matrix, lane i owns row i entirely in registers.
// Both loops are fully unrolled so row[] indices stay compile-time constants;
// a dynamic index would force the array into local memory.
__global__ void potrf32_warp(const float *__restrict__ A,
float *__restrict__ L, int batch) {
constexpr int N = 32;
const int gid = blockIdx.x * blockDim.x + threadIdx.x;
const int mid = gid >> 5;
if (mid >= batch) return;
const int lane = gid & 31;
const float *src = A + (long long)mid * N * N;
float *dst = L + (long long)mid * N * N;
// Stage through shared memory so the global read is contiguous per warp,
// then each lane picks up its own row. LDA = N + 1 keeps those per-lane
// row reads spread across banks.
constexpr int LDA = N + 1;
extern __shared__ float smem[];
float *tile = smem + (threadIdx.x >> 5) * N * LDA;
// float4 on the global side: 8 LDG.128 per matrix instead of 32 LDG.32.
// The SMEM side stays scalar so LDA=33 keeps the 32 lanes on 32 banks
// (r = t4>>3, c = (t4&7)<<2 gives base banks {0..3} x {0,4,..28}).
#pragma unroll
for (int t4 = lane; t4 < N * N / 4; t4 += 32) {
const float4 v = *(const float4 *)(src + t4 * 4);
float *d = tile + (t4 >> 3) * LDA + ((t4 & 7) << 2);
d[0] = v.x; d[1] = v.y; d[2] = v.z; d[3] = v.w;
}
__syncwarp();
float row[N];
#pragma unroll
for (int j = 0; j < N; ++j) row[j] = tile[lane * LDA + j];
#pragma unroll
for (int k = 0; k < N; ++k) {
const float akk = __shfl_sync(FULL, row[k], k);
const float dk = akk > 0.0f ? sqrtf(akk) : 1.0e-20f;
// For lane == k this is akk/sqrt(akk) == sqrt(akk), so the diagonal
// needs no special case. Lanes above k are set to zero.
const float lk = (lane >= k) ? row[k] / dk : 0.0f;
row[k] = lk;
#pragma unroll
for (int j = k + 1; j < N; ++j) {
const float ljk = __shfl_sync(FULL, lk, j);
if (lane >= j) row[j] -= lk * ljk;
}
}
#pragma unroll
for (int j = 0; j < N; ++j) tile[lane * LDA + j] = row[j];
__syncwarp();
#pragma unroll
for (int t4 = lane; t4 < N * N / 4; t4 += 32) {
const float *s = tile + (t4 >> 3) * LDA + ((t4 & 7) << 2);
*(float4 *)(dst + t4 * 4) = make_float4(s[0], s[1], s[2], s[3]);
}
}
// General small n: one CTA per matrix, tile in shared memory.
// LDA = N + 1 so walking down a column strides across banks.
template <int N>
__global__ void potrf_shared(const float *__restrict__ A,
float *__restrict__ L, int batch) {
constexpr int LDA = N + 1;
const int mid = blockIdx.x;
if (mid >= batch) return;
extern __shared__ float sh[];
const float *src = A + (long long)mid * N * N;
float *dst = L + (long long)mid * N * N;
const int tid = threadIdx.x;
const int nt = blockDim.x;
for (int idx = tid; idx < N * N; idx += nt) {
const int r = idx / N, c = idx - r * N;
sh[r * LDA + c] = (r >= c) ? src[idx] : 0.0f;
}
__syncthreads();
for (int k = 0; k < N; ++k) {
if (tid == 0) {
const float d = sh[k * LDA + k];
sh[k * LDA + k] = d > 0.0f ? sqrtf(d) : 1.0e-20f;
}
__syncthreads();
const float inv = 1.0f / sh[k * LDA + k];
for (int i = k + 1 + tid; i < N; i += nt) sh[i * LDA + k] *= inv;
__syncthreads();
const int m = N - k - 1;
for (int t = tid; t < m * m; t += nt) {
const int jj = t / m;
const int ii = t - jj * m;
if (ii >= jj) {
const int i = k + 1 + ii, j = k + 1 + jj;
sh[i * LDA + j] -= sh[i * LDA + k] * sh[j * LDA + k];
}
}
__syncthreads();
}
for (int idx = tid; idx < N * N; idx += nt) {
const int r = idx / N;
dst[idx] = sh[r * LDA + idx - r * N];
}
}
// Blocked right-looking Cholesky, one CTA per matrix, whole matrix in SMEM.
// Per block column: warp 0 factorizes the NB x NB diagonal block in registers
// via shuffles, then one thread per row does the panel TRSM with no interior
// barrier, then the trailing update runs a 4x4 register tile per thread.
// Barriers drop from 3-per-column to 3-per-block-column and the trailing
// update reads 8 SMEM words per 16 FFMAs (vs ~3 per 1 in potrf_shared).
// PACK stores only the lower triangle, row i starting at i(i+1)/2, which halves
// the SMEM footprint (n=256: 263 KB square -> 132 KB, the difference between not
// fitting and fitting in a B200 CTA). Every access below is (row, col<=row) or
// within a row, so one row-base helper covers both layouts.
template <int N, int NB, int TX, bool PACK>
struct Layout {
// Plus an NB x N scratch panel (see potrf_blk) appended after the triangle.
static constexpr size_t tri = PACK ? (size_t)N * (N + 1) / 2
: (size_t)N * (N + 1);
// Plus NB words holding 1/L[j][j] of the current diagonal block.
static constexpr size_t words = tri + (size_t)NB * N + NB;
__device__ __forceinline__ static int base(int i) {
return PACK ? ((i * (i + 1)) >> 1) : (i * (N + 1));
}
};
// TM is the tile height in rows (tile width is fixed at 4, one float4 load).
// TM = 8 halves the v-loads per FFMA; __launch_bounds__ is what keeps ptxas
// inside the 64-register budget a 1024-thread block needs.
template <int N, int NB, int TX, bool PACK, int TM>
__global__ __launch_bounds__(TX) void potrf_blk(const float *__restrict__ A,
float *__restrict__ L, int batch) {
using LO = Layout<N, NB, TX, PACK>;
const int mid = blockIdx.x;
if (mid >= batch) return;
extern __shared__ float sh[];
const float *src = A + (long long)mid * N * N;
float *dst = L + (long long)mid * N * N;
const int tid = threadIdx.x;
// 1/L[j][j] for the current diagonal block, published by the diagonal warp
// so the TRSM multiplies instead of dividing: the TRSM's serial chain is
// NB dependent divides deep and sits on the critical path.
float *dinv = sh + LO::tri + (size_t)NB * N;
// Read the full square from global (coalesced) but keep only tril.
for (int idx = tid; idx < N * N; idx += TX) {
const int r = idx / N, c = idx - r * N;
if (r >= c) sh[LO::base(r) + c] = src[idx];
else if (!PACK) sh[LO::base(r) + c] = 0.0f;
}
__syncthreads();
for (int k = 0; k < N; k += NB) {
// --- diagonal block, warp 0, lane p owns row k+p ---
if (tid < 32) {
const int lane = tid;
float r[NB];
#pragma unroll
// Only p <= lane is in the triangle; under PACK the slots past it
// belong to row k+lane+1, so they must not be read or written.
for (int p = 0; p < NB; ++p)
r[p] = (lane < NB && p <= lane) ? sh[LO::base(k + lane) + k + p]
: 0.0f;
#pragma unroll
for (int c = 0; c < NB; ++c) {
const float acc = __shfl_sync(FULL, r[c], c);
// rsqrt is one MUFU op; sqrt + divide is two dependent ones, and
// lane c gets acc * rsqrt(acc) == sqrt(acc), so the diagonal is
// still exact-in-kind. 1e20 reproduces the old 1/1e-20 guard.
const float dc = acc > 0.0f ? rsqrtf(acc) : 1.0e20f;
const float lc = (lane >= c) ? r[c] * dc : 0.0f;
r[c] = lc;
if (lane == c) dinv[c] = dc;
#pragma unroll
for (int j = c + 1; j < NB; ++j) {
const float ljc = __shfl_sync(FULL, lc, j);
if (lane >= j) r[j] -= lc * ljc;
}
}
if (lane < NB) {
#pragma unroll
for (int p = 0; p < NB; ++p)
if (p <= lane) sh[LO::base(k + lane) + k + p] = r[p];
}
}
__syncthreads();
const int k1 = k + NB;
if (k1 >= N) break;
// --- panel TRSM: row i solves x * Lkk^T = A[i, k:k1] ---
// The solved row is written twice: into the packed triangle, and
// transposed into `pan` (p-major, row-index minor). The trailing update
// then reads four consecutive rows with one 128-bit LDS instead of four
// scattered ones -- under PACK, row i lives at i(i+1)/2, so a warp's
// 32 lanes hit near-random banks and the update was LDS-conflict bound.
float *pan = sh + LO::tri;
for (int i = k1 + tid; i < N; i += TX) {
const int bi = LO::base(i);
float x[NB];
#pragma unroll
for (int p = 0; p < NB; ++p) x[p] = sh[bi + k + p];
#pragma unroll
for (int j = 0; j < NB; ++j) {
const int bj = LO::base(k + j);
float s = x[j];
#pragma unroll
for (int p = 0; p < j; ++p) s -= x[p] * sh[bj + k + p];
x[j] = s * dinv[j];
}
#pragma unroll
for (int p = 0; p < NB; ++p) {
sh[bi + k + p] = x[p];
pan[p * N + (i - k1)] = x[p];
}
}
__syncthreads();
// --- trailing update, lower triangle only, 4x4 register tile ---
// N and NB are both multiples of 4, so m = N - k1 is too and every
// 4-row group is whole: no tail guard, and the loads can be float4.
const int m = N - k1;
const int TI = m / TM;
const int TJ = m >> 2;
// Enumerate only the tiles at or below the diagonal. Striding over the
// full square grid and skipping the upper half wasted half the loop
// trips. Tile row ti owns R*(ti+1) tiles (R = TM/4 tile-columns per
// tile-row of height TM), so tiles before row ti number R*ti(ti+1)/2.
constexpr int R = TM / 4;
const int NT = R * ((TI * (TI + 1)) >> 1);
for (int t = tid; t < NT; t += TX) {
// Invert the cumulative count. Fast-math sqrt is approximate, so
// nudge into range rather than trusting it.
int ti = (int)((sqrtf(1.0f + 8.0f * (float)t / (float)R) - 1.0f) * 0.5f);
while (R * (((ti + 1) * (ti + 2)) >> 1) <= t) ++ti;
while (R * ((ti * (ti + 1)) >> 1) > t) --ti;
const int tj = t - R * ((ti * (ti + 1)) >> 1);
const int i0 = k1 + ti * TM, j0 = k1 + tj * 4;
float acc[TM][4];
#pragma unroll
for (int a = 0; a < TM; ++a)
#pragma unroll
for (int b = 0; b < 4; ++b) acc[a][b] = 0.0f;
#pragma unroll
for (int p = 0; p < NB; ++p) {
const float *row = pan + p * N;
float u[TM];
#pragma unroll
for (int a = 0; a < TM; a += 4) {
const float4 u4 = *(const float4 *)(row + ti * TM + a);
u[a] = u4.x; u[a + 1] = u4.y; u[a + 2] = u4.z; u[a + 3] = u4.w;
}
const float4 v4 = *(const float4 *)(row + tj * 4);
const float v[4] = {v4.x, v4.y, v4.z, v4.w};
#pragma unroll
for (int a = 0; a < TM; ++a)
#pragma unroll
for (int b = 0; b < 4; ++b) acc[a][b] += u[a] * v[b];
}
#pragma unroll
for (int a = 0; a < TM; ++a) {
const int i = i0 + a;
const int bi = LO::base(i);
#pragma unroll
for (int b = 0; b < 4; ++b) {
const int j = j0 + b;
if (j <= i) sh[bi + j] -= acc[a][b];
}
}
}
__syncthreads();
}
for (int idx = tid; idx < N * N; idx += TX) {
const int r = idx / N, c = idx - r * N;
dst[idx] = (r >= c) ? sh[LO::base(r) + c] : 0.0f;
}
}
// Two-level blocked variant. potrf_blk writes every trailing tile back to SMEM
// once per NB-wide column block (N/NB times); the FFMA:writeback ratio there is
// NB:1. Here the right-looking update inside an NBO-wide outer block is confined
// to that block's own columns, and the trailing submatrix is updated once per
// outer block with depth NBO -- NBO:1, so N/NBO passes instead of N/NB.
// pan holds the whole NBO-wide panel, indexed by absolute row.
template <int N, int NB, int NBO, int TX, int TM>
struct Layout2 {
static constexpr size_t tri = (size_t)N * (N + 1) / 2;
static constexpr size_t words = tri + (size_t)NBO * N + NB;
__device__ __forceinline__ static int base(int i) { return (i * (i + 1)) >> 1; }
};
template <int N, int NB, int NBO, int TX, int TM>
__global__ __launch_bounds__(TX) void potrf_blk2(const float *__restrict__ A,
float *__restrict__ L, int batch) {
using LO = Layout2<N, NB, NBO, TX, TM>;
const int mid = blockIdx.x;
if (mid >= batch) return;
extern __shared__ float sh[];
const float *src = A + (long long)mid * N * N;
float *dst = L + (long long)mid * N * N;
const int tid = threadIdx.x;
float *pan = sh + LO::tri;
float *dinv = pan + (size_t)NBO * N;
for (int idx = tid; idx < N * N; idx += TX) {
const int r = idx / N, c = idx - r * N;
if (r >= c) sh[LO::base(r) + c] = src[idx];
}
__syncthreads();
for (int k0 = 0; k0 < N; k0 += NBO) {
const int kend = k0 + NBO; // N is a multiple of NBO
for (int k = k0; k < kend; k += NB) {
// --- diagonal block, warp 0, lane p owns row k+p (as potrf_blk) ---
if (tid < 32) {
const int lane = tid;
float r[NB];
#pragma unroll
for (int p = 0; p < NB; ++p)
r[p] = (lane < NB && p <= lane) ? sh[LO::base(k + lane) + k + p]
: 0.0f;
#pragma unroll
for (int c = 0; c < NB; ++c) {
const float acc = __shfl_sync(FULL, r[c], c);
const float dc = acc > 0.0f ? rsqrtf(acc) : 1.0e20f;
const float lc = (lane >= c) ? r[c] * dc : 0.0f;
r[c] = lc;
if (lane == c) dinv[c] = dc;
#pragma unroll
for (int j = c + 1; j < NB; ++j) {
const float ljc = __shfl_sync(FULL, lc, j);
if (lane >= j) r[j] -= lc * ljc;
}
}
if (lane < NB) {
#pragma unroll
for (int p = 0; p < NB; ++p)
if (p <= lane) sh[LO::base(k + lane) + k + p] = r[p];
}
}
__syncthreads();
const int k1 = k + NB;
if (k1 >= N) break;
// --- panel TRSM, all the way to row N-1: the solved rows are the
// outer block's panel and are consumed by the big update below. ---
const int po = k - k0;
for (int i = k1 + tid; i < N; i += TX) {
const int bi = LO::base(i);
float x[NB];
#pragma unroll
for (int p = 0; p < NB; ++p) x[p] = sh[bi + k + p];
#pragma unroll
for (int j = 0; j < NB; ++j) {
const int bj = LO::base(k + j);
float s = x[j];
#pragma unroll
for (int p = 0; p < j; ++p) s -= x[p] * sh[bj + k + p];
x[j] = s * dinv[j];
}
#pragma unroll
for (int p = 0; p < NB; ++p) {
sh[bi + k + p] = x[p];
pan[(po + p) * N + i] = x[p];
}
}
__syncthreads();
// --- narrow update: only the outer block's own columns, so the
// next inner step sees an up-to-date panel. One thread per element;
// the region is at most (N-k1) x (NBO-NB), a small fraction of the
// work the big update below does. ---
const int W = kend - k1;
if (W > 0) {
// 2x2 register tile: 2 LDS.64 feed 4 FFMAs (1.5 instr/FFMA) where
// the one-thread-per-element form issued 2 LDS + 1 FFMA (3.0),
// and it costs one integer divide per tile instead of per
// element. 2x2 rather than 4x4 because at n=256 the region is
// only (N-k1) x 24, and 4x4 leaves 1024 threads with ~370 tiles.
const int nct = W >> 1; // W is a multiple of NB
const int nt2 = ((N - k1) >> 1) * nct;
for (int t = tid; t < nt2; t += TX) {
const int r = t / nct;
const int i0 = k1 + (r << 1), j0 = k1 + ((t - r * nct) << 1);
if (j0 > i0 + 1) continue;
float acc[2][2] = {{0.0f, 0.0f}, {0.0f, 0.0f}};
#pragma unroll
for (int p = 0; p < NB; ++p) {
const float *row = pan + (po + p) * N;
const float2 u = *(const float2 *)(row + i0);
const float2 v = *(const float2 *)(row + j0);
acc[0][0] += u.x * v.x; acc[0][1] += u.x * v.y;
acc[1][0] += u.y * v.x; acc[1][1] += u.y * v.y;
}
#pragma unroll
for (int a = 0; a < 2; ++a) {
const int bi = LO::base(i0 + a);
#pragma unroll
for (int b = 0; b < 2; ++b)
if (j0 + b <= i0 + a) sh[bi + j0 + b] -= acc[a][b];
}
}
__syncthreads();
}
}
// --- big update: trailing submatrix below/right of the outer block,
// accumulated over all NBO panel columns and written back once. ---
if (kend >= N) break;
const int m = N - kend;
const int TI = m / TM;
constexpr int R = TM / 4;
const int NT = R * ((TI * (TI + 1)) >> 1);
for (int t = tid; t < NT; t += TX) {
int ti = (int)((sqrtf(1.0f + 8.0f * (float)t / (float)R) - 1.0f) * 0.5f);
while (R * (((ti + 1) * (ti + 2)) >> 1) <= t) ++ti;
while (R * ((ti * (ti + 1)) >> 1) > t) --ti;
const int tj = t - R * ((ti * (ti + 1)) >> 1);
const int i0 = kend + ti * TM, j0 = kend + tj * 4;
float acc[TM][4];
#pragma unroll
for (int a = 0; a < TM; ++a)
#pragma unroll
for (int b = 0; b < 4; ++b) acc[a][b] = 0.0f;
#pragma unroll 4
for (int p = 0; p < NBO; ++p) {
const float *row = pan + p * N;
float u[TM];
#pragma unroll
for (int a = 0; a < TM; a += 4) {
const float4 u4 = *(const float4 *)(row + i0 + a);
u[a] = u4.x; u[a + 1] = u4.y; u[a + 2] = u4.z; u[a + 3] = u4.w;
}
const float4 v4 = *(const float4 *)(row + j0);
const float v[4] = {v4.x, v4.y, v4.z, v4.w};
#pragma unroll
for (int a = 0; a < TM; ++a)
#pragma unroll
for (int b = 0; b < 4; ++b) acc[a][b] += u[a] * v[b];
}
#pragma unroll
for (int a = 0; a < TM; ++a) {
const int i = i0 + a;
const int bi = LO::base(i);
#pragma unroll
for (int b = 0; b < 4; ++b) {
const int j = j0 + b;
if (j <= i) sh[bi + j] -= acc[a][b];
}
}
}
__syncthreads();
}
for (int idx = tid; idx < N * N; idx += TX) {
const int r = idx / N, c = idx - r * N;
dst[idx] = (r >= c) ? sh[LO::base(r) + c] : 0.0f;
}
}
template <int N, int NB, int NBO, int TX, int TM>
static bool launch_blk2(const at::Tensor &A, at::Tensor &L, int batch) {
const size_t shmem = Layout2<N, NB, NBO, TX, TM>::words * sizeof(float);
auto kernel = potrf_blk2<N, NB, NBO, TX, TM>;
if (shmem > 48 * 1024) {
if (cudaFuncSetAttribute(kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)shmem) != cudaSuccess) {
cudaGetLastError();
return false;
}
}
kernel<<<batch, TX, shmem>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch);
if (cudaGetLastError() != cudaSuccess) return false;
return true;
}
// Returns false if the device cannot give this CTA the shared memory it needs
// (packed n=256 wants 132 KB, which sm_89 cannot opt into), so the caller can
// fall back to cuSOLVER instead of launching a kernel that would fail.
template <int N, int NB, int TX, bool PACK, int TM>
static bool launch_blk(const at::Tensor &A, at::Tensor &L, int batch) {
const size_t shmem = Layout<N, NB, TX, PACK>::words * sizeof(float);
auto kernel = potrf_blk<N, NB, TX, PACK, TM>;
if (shmem > 48 * 1024) {
if (cudaFuncSetAttribute(kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)shmem) != cudaSuccess) {
cudaGetLastError();
return false;
}
}
kernel<<<batch, TX, shmem>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch);
// A launch that fails on resources (TX * regs > 64K) would otherwise leave L
// uninitialised and silently return garbage.
if (cudaGetLastError() != cudaSuccess) return false;
return true;
}
template <int N>
static void launch_shared(const at::Tensor &A, at::Tensor &L, int batch,
int threads) {
const size_t shmem = (size_t)N * (N + 1) * sizeof(float);
auto kernel = potrf_shared<N>;
if (shmem > 48 * 1024) {
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)shmem);
}
kernel<<<batch, threads, shmem>>>(A.data_ptr<float>(),
L.data_ptr<float>(), batch);
}
// C += alpha * A @ B^T with fp16 inputs and fp32 accumulate/output.
// torch cannot express fp16-in/fp32-out, and materialising an fp16 product
// would cost tens of GB of extra traffic at n=32768. beta=1 accumulates
// straight into the (strided, non-contiguous) view of the trailing matrix.
// fp16 and tf32 share a 10-bit mantissa, so this is accuracy-neutral versus
// the tf32 path at roughly twice the tensor-core rate.
void gemm_nt_f16(at::Tensor C, at::Tensor A, at::Tensor B, double alpha) {
TORCH_CHECK(A.scalar_type() == at::kHalf && B.scalar_type() == at::kHalf);
TORCH_CHECK(C.scalar_type() == at::kFloat);
TORCH_CHECK(A.stride(-1) == 1 && B.stride(-1) == 1 && C.stride(-1) == 1,
"inner dimension must be contiguous");
const int m = (int)A.size(-2), k = (int)A.size(-1), nn = (int)B.size(-2);
TORCH_CHECK((int)B.size(-1) == k && (int)C.size(-2) == m &&
(int)C.size(-1) == nn);
int batch = 1;
long long sa = 0, sb = 0, sc = 0;
if (A.dim() == 3) {
batch = (int)A.size(0);
sa = A.stride(0); sb = B.stride(0); sc = C.stride(0);
}
const float al = (float)alpha, be = 1.0f;
// Row-major C(m x nn) is column-major C^T(nn x m) = B_rm * A_rm^T, which in
// column-major terms is op(B)=T over the k x nn buffer times op(A)=N.
auto st = cublasGemmStridedBatchedEx(
at::cuda::getCurrentCUDABlasHandle(), CUBLAS_OP_T, CUBLAS_OP_N,
nn, m, k, &al,
(const void *)B.data_ptr<at::Half>(), CUDA_R_16F, (int)B.stride(-2), sb,
(const void *)A.data_ptr<at::Half>(), CUDA_R_16F, (int)A.stride(-2), sa,
&be,
(void *)C.data_ptr<float>(), CUDA_R_32F, (int)C.stride(-2), sc,
batch, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasGemmStridedBatchedEx failed");
}
at::Tensor potrf(at::Tensor A) {
TORCH_CHECK(A.dim() == 3 && A.size(1) == A.size(2), "expected batch x n x n");
TORCH_CHECK(A.scalar_type() == at::kFloat, "expected float32");
A = A.contiguous();
const int batch = (int)A.size(0);
const int n = (int)A.size(1);
auto L = at::empty_like(A);
if (n == 32) {
constexpr int TPB = 128;
const int warps = TPB / 32;
const int blocks = (batch + warps - 1) / warps;
const size_t shmem = (size_t)warps * 32 * 33 * sizeof(float);
potrf32_warp<<<blocks, TPB, shmem>>>(A.data_ptr<float>(),
L.data_ptr<float>(), batch);
return L;
}
bool ok;
// Measured and rejected: n=32 through launch_blk<32,8,64,true,4> (18.1 ->
// 36.1 us). At n=32 the blocked kernel is all diagonal block and TRSM with
// a 24x24 trailing matrix -- there is no update to amortise the barriers.
// TM = 8 was measured on B200 and lost: it halves the tile count, and
// both shapes are already short of tiles (256x128 45.2 -> 73.3 us).
// NBO = 64 at n=256 lost (87.2 -> 91.8): narrow-update work grows as NBO^2.
// NBO = 16 wins where the trailing matrix is small (n=64 22.6 -> 21.2,
// n=128 37.0 -> 34.9) but loses at n=256 (88.2 -> 90.0), where the halved
// big-update depth costs more than the halved narrow update.
switch (n) {
case 64: ok = launch_blk2<64, 8, 16, 128, 4>(A, L, batch); break;
case 128: ok = launch_blk2<128, 8, 16, 512, 4>(A, L, batch); break;
case 256: ok = launch_blk2<256, 8, 32, 1024, 4>(A, L, batch); break;
default: TORCH_CHECK(false, "unsupported n");
}
// Empty result tells the caller to use the library path (n=256 needs 132 KB
// of SMEM per CTA, which only Hopper/Blackwell can opt into).
if (!ok) return at::empty({0}, A.options());
return L;
}
"""
CPP_SRC = (
"at::Tensor potrf(at::Tensor A);\n"
"void gemm_nt_f16(at::Tensor C, at::Tensor A, at::Tensor B, double alpha);"
)
_ext = load_inline(
name="cholesky_adapt",
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=["potrf", "gemm_nt_f16"],
extra_cuda_cflags=["-O3", "--expt-relaxed-constexpr", "-use_fast_math"],
extra_ldflags=["-lcublas"],
verbose=False,
)
_CUSTOM_N = (32, 64, 128, 256)
# Blocked right-looking Cholesky for large n. The diagonal blocks and panel
# solves stay in fp32; the trailing updates, which hold essentially all of the
# n^3/3 flops, run on TF32 tensor cores. cuSOLVER does not use TF32 for potrf,
# which is where the headroom comes from.
# Gate on total work, not n: 1x4096 loses to cuSOLVER but 2x4096 wins by 1.4x,
# so batch matters as much as size.
_TF32_MIN_WORK = 1.0e11
_NB = 2048
# Base block for _tri_inv. The fp32 base solve dominates, not the launch count,
# so smaller is better: B200 1x32768 measured 512 -> 37.9ms, 256 -> 37.6,
# 128 -> 29.5, 64 -> 28.8.
_TRI_BASE = int(os.environ.get("CHOL_TRI_BASE", 64))
# Same knob for the mid regime (n in 512/1024). Smaller wins there too:
# B200 640x512 measured 64 -> 1854, 32 -> 1854/1067, 16 -> 1812.
_MID_TRI_BASE = int(os.environ.get("CHOL_MID_TRI_BASE", 16))
def _h(x):
# solve_triangular returns a column-major result and .half() preserves those
# strides, which the cuBLAS wrapper rejects. Cast and transpose in one pass.
return x.to(torch.float16, memory_format=torch.contiguous_format)
def _blk(T, off, cnt, s, step):
# cnt sub-blocks of size s x s, the first at element offset `off` into T's
# storage, successive ones `step` elements apart. Regular strides, so the
# whole set is one batched op.
lead = tuple(T.shape[:-2])
lstr = tuple(T.stride()[:-2])
return T.as_strided(
lead + (cnt, s, s),
lstr + (step, T.stride(-2), 1),
T.storage_offset() + off,
)
def _tri_inv(Lkk, base):
# Inverse of a lower-triangular m x m block, bottom-up 2x2 recursion:
# inv([[A,0],[C,B]]) = [[Ai,0],[-Bi C Ai, Bi]]
# Every level is ONE strided-batched op over all blocks of that level, so the
# whole inverse costs ~9 launches regardless of `base`, and only base^2/m^2 of
# the work is the non-tensor-core fp32 solve. A per-block recursion instead
# emits ~25 small launches per step and measured slower on B200.
# cuSOLVER hands back a column-major L; _blk assumes stride(-1) == 1, and
# reading L^T as lower-triangular silently yields a diagonal-only "inverse"
# (-> indefinite trailing matrix -> NaN), so normalise the layout first.
Lkk = Lkk.contiguous()
m = Lkk.shape[-1]
ld = Lkk.stride(-2)
X = torch.zeros_like(Lkk)
s = min(base, m)
cnt = m // s
D = _blk(Lkk, 0, cnt, s, s * (ld + 1))
eye = torch.eye(s, device=Lkk.device, dtype=Lkk.dtype).expand_as(D)
_blk(X, 0, cnt, s, s * (ld + 1)).copy_(
torch.linalg.solve_triangular(D, eye, upper=False, left=True)
)
while s < m:
cnt = m // (2 * s)
step = 2 * s * (ld + 1)
Ai = _blk(X, 0, cnt, s, step)
Bi = _blk(X, s * (ld + 1), cnt, s, step)
C = _blk(Lkk, s * ld, cnt, s, step)
_blk(X, s * ld, cnt, s, step).copy_(-(Bi @ C @ Ai))
s *= 2
return X
# Left-looking crossover. Right-looking rewrites the whole trailing matrix once
# per step (~8n^3/(3nb) bytes of fp32 read-modify-write, ~45 GB at n=32768) plus
# an n^2 clone; left-looking touches each block column exactly once and pulls its
# update out of an fp16 mirror (~n^3/(3nb) bytes, ~6 GB) with no clone. That only
# pays once there are enough block columns to amortise the mirror: measured on
# B200, n=8192 loses (3910 -> 4130, four columns), n=16384 is a wash (+0.7%),
# n=32768 wins (28000 -> 27100, reproduced twice).
_LEFT_MIN_N = int(os.environ.get("CHOL_LEFT_MIN_N", 32768))
# Must be _TRI_BASE * 2^k for _tri_inv. 1024 measured worse on B200 (1x32768
# 26700 -> 27300): halving nb doubles the number of cuSOLVER diagonal potrf
# calls and C(nb)/nb only improves slightly, while the left-looking traffic it
# would save is already the small term.
_LEFT_NB = int(os.environ.get("CHOL_LEFT_NB", 2048))
def _left_tf32(A, nb):
# Left-looking blocked Cholesky, batch 1, fp16 trailing math.
# Block column k is A[k:, k:e] - Lh[k:, :k] @ Lh[k:e, :k]^T, factored in
# place; the fp32 working buffer is n x nb, not n x n.
n = A.shape[-1]
# Must be zeros, not empty: the panel gemm_nt_f16 accumulates (beta=1) into
# L[e:, k:e], so that region has to start at zero.
L = torch.zeros_like(A)
# empty, not zeros: step k reads only Lh[k:, :k], and every one of those rows
# was written by an earlier step's diagonal block or panel. Skips a 2 GB
# memset at n=32768.
Lh = torch.empty(n, n, device=A.device, dtype=torch.float16)
buf = torch.empty(n, nb, device=A.device, dtype=torch.float32)
for k in range(0, n, nb):
e = min(k + nb, n)
w = e - k
W = buf[: n - k, :w]
W.copy_(A[k:, k:e])
if k:
_ext.gemm_nt_f16(W, Lh[k:, :k], Lh[k:e, :k], -1.0)
Lkk = torch.linalg.cholesky_ex(W[:w], check_errors=False).L
L[k:e, k:e] = Lkk
Lh[k:e, k:e] = Lkk
if e >= n:
break
Linv = _tri_inv(Lkk, _TRI_BASE)
L21 = L[e:, k:e]
_ext.gemm_nt_f16(L21, _h(W[w:]), _h(Linv), 1.0)
Lh[e:, k:e] = L21
return L
def _blocked_tf32(A, nb=None):
n = A.shape[-1]
# Block size tuned per size: small n needs several steps to amortise
# overhead, large n wants fewer, larger GEMMs. Once the trailing updates run
# in fp16 the leftover critical path is the O(n*nb^2) fp32 panel work, so the
# optimum moves down from 4096 to 2048 -- but only at n >= 16384; nb=2048
# measurably regresses n=8192.
if nb is None:
nb = 1024 if n <= 4096 else (2048 if n >= 16384 else 4096)
# fp16 in / fp32 out for the trailing update and the panel GEMM: same 10-bit
# mantissa as tf32, ~2x the tensor-core rate. Below 8192 the cast traffic is
# not amortised.
fp16 = n >= 8192
# A batch of 1 still routes cuSOLVER through its *batched* potrf/trsm, which
# are built for many-small and degrade badly for one-large. Squeezing the
# singleton dim puts the nb x nb diagonal block on the classical path; that
# block is the largest remaining non-tensor-core item at n=32768.
squeeze = A.dim() == 3 and A.shape[0] == 1
if squeeze:
A = A[0]
prev = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
if squeeze and n >= _LEFT_MIN_N:
return _left_tf32(A, _LEFT_NB).unsqueeze(0)
W = A.clone()
L = torch.zeros_like(A)
for k in range(0, n, nb):
e = min(k + nb, n)
Lkk = torch.linalg.cholesky_ex(W[..., k:e, k:e], check_errors=False).L
L[..., k:e, k:e] = Lkk
if e >= n:
break
# Panel: X @ Lkk^T = A21. A direct TRSM here is fp32 and never
# touches tensor cores, and its cost (nb^2 * (n-e)) dominates the
# trailing tf32 GEMMs at large n. Invert the nb x nb block once
# (nb^3/2, ~14x cheaper at the first step) and turn the panel into a
# tf32 GEMM instead.
# With _tri_inv (batched 2x2 recursion) the inverse is cheap enough to
# pay off from n=8192 up: B200 1x8192 5010 -> 3950us.
if n >= 8192:
Linv = _tri_inv(Lkk, _TRI_BASE)
if fp16:
# Accumulate straight into L's panel (still all zeros there,
# so beta=1 is an overwrite). The obvious form -- zeros_like,
# GEMM into it, then copy into L -- costs three extra full
# passes over an (n-e) x nb fp32 buffer per step, ~6.5 GB of
# traffic over the whole factorization at n=32768.
L21 = L[..., e:, k:e]
_ext.gemm_nt_f16(L21, _h(W[..., e:, k:e]), _h(Linv), 1.0)
else:
L21 = W[..., e:, k:e] @ Linv.transpose(-1, -2)
L[..., e:, k:e] = L21
else:
L21 = torch.linalg.solve_triangular(
Lkk.transpose(-1, -2), W[..., e:, k:e], upper=True, left=False
)
L[..., e:, k:e] = L21
# Trailing update by block column, touching only the lower triangle.
# Updating the full square would double the flops.
L21h = _h(L21) if fp16 else None
for jb in range(e, n, nb):
je = min(jb + nb, n)
if fp16:
_ext.gemm_nt_f16(
W[..., jb:, jb:je],
L21h[..., jb - e:, :],
L21h[..., jb - e:je - e, :],
-1.0,
)
else:
W[..., jb:, jb:je] -= (
L21[..., jb - e:, :]
@ L21[..., jb - e:je - e, :].transpose(-1, -2)
)
return L.unsqueeze(0) if squeeze else L
finally:
torch.backends.cuda.matmul.allow_tf32 = prev
def _blocked_mid(A, nb, tri_base=None):
# n in {512, 1024}, high batch. Same right-looking structure as
# _blocked_tf32, but the panel solve is an inverted-diagonal tf32 GEMM
# rather than an fp32 batched TRSM: at these sizes the TRSM, not the
# trailing SYRK, is the bottleneck (nb=256+TRSM measured 4090us at 640x512
# vs 3280us for this). Kept separate so the large-n path is untouched.
n = A.shape[-1]
prev = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
W = A.clone()
L = torch.zeros_like(A)
for k in range(0, n, nb):
e = min(k + nb, n)
Dkk = W[..., k:e, k:e]
Lkk = _ext.potrf(Dkk.contiguous()) if (e - k) in _CUSTOM_N else None
if Lkk is None or not Lkk.numel():
Lkk = torch.linalg.cholesky_ex(Dkk, check_errors=False).L
L[..., k:e, k:e] = Lkk
if e >= n:
break
# Batched 2x2 recursion instead of solve_triangular(Lkk, eye): only
# base^2/nb^2 of the panel-inverse work stays on the non-TC fp32
# solve. B200 640x512 2390 -> 1854, 60x1024 1431 -> 1067.
Linv = _tri_inv(Lkk, tri_base or _MID_TRI_BASE)
L21 = W[..., e:, k:e] @ Linv.transpose(-1, -2)
L[..., e:, k:e] = L21
# Trailing update in fp16 in / fp32 accumulate: same 10-bit mantissa
# as tf32 for ~2x the tensor-core rate, so residuals are unchanged.
# Measured and rejected: doing the panel GEMM in fp16 too (the extra
# W21 cast costs more than the 2x rate saves), and cutting the copy
# traffic here (block-lower-triangle W copy, matmul out= into L) --
# 1-3% worse on B200 even though it cut local 4070 time 17%.
L21h = _h(L21)
for jb in range(e, n, nb):
je = min(jb + nb, n)
_ext.gemm_nt_f16(
W[..., jb:, jb:je],
L21h[..., jb - e:, :],
L21h[..., jb - e:je - e, :],
-1.0,
)
return L
finally:
torch.backends.cuda.matmul.allow_tf32 = prev
def custom_kernel(data: input_t) -> output_t:
n = data.shape[-1]
if data.dim() == 3:
if n in _CUSTOM_N:
out = _ext.potrf(data)
if out.numel():
return out
# SMEM opt-in failed for this n on this device; fall through.
# 60x1024 is below _TF32_MIN_WORK but still beats batched cuSOLVER on
# the blocked tf32 path with a small nb (B200: 2890 -> 2520 us).
# Explicit batch floors rather than a b*n^3 work gate: with the custom
# potrf, fp16 trailing and _tri_inv panel, the low-batch pair now wins
# here too (4x1024 1281 -> 683, 16x512 596 -> 580). The b>=4 floor at
# n=1024 is load-bearing -- 2x1024 lowrank NaNs on this path.
# Batch floors on top of the work gate: five B200 rounds (waves
# 8/10/15/16/17/18) measure the low-batch pair winning here too
# (4x1024 1281 -> 745, 16x512 596 -> 416). Settled; do not re-litigate.
b = data.shape[0]
if n in (512, 1024) and (
b * float(n) ** 3 >= 5e10
or (n == 1024 and b >= 4)
or (n == 512 and b >= 16)
):
nb = 128 if n == 512 else 256
# Low-batch n=512: collapse _tri_inv to a single batched
# solve_triangular (base == nb, no 2x2 recursion). 16x512 is
# dispatch-bound so the ~20 launches saved beat the wider fp32 base
# solve; at n=1024 the same change costs (base=16 -> 745,
# 128 -> 1314, 256 -> 1337), so this is n==512-only by measurement.
tb = nb if (n == 512 and b * float(n) ** 3 < 5e10) else None
return _blocked_mid(data, nb=nb, tri_base=tb)
# Few-large: cuSOLVER's batched potrf is built for many-small and loses
# badly here -- 2x2048 gets 1.7 TF/s where 1x4096 gets 15 on the very same
# library via the classical single-matrix routine. One call per matrix
# keeps every one of them on the classical path.
# b <= 4 only: measured on B200, 8x2048 loses to the batched call
# (5060 -> 5400us) while 2x2048 wins 3450 -> 1349 and 2x4096 6930 -> 3210.
# b == 1 is already on the classical path, so looping only adds an n^2
# copy into the output buffer (1x4096 measured 1528 -> 1603us).
# Measured and rejected on B200: routing 8x2048 and 2x4096 to
# _blocked_batch(256, 1024) -- a blocked tf32 factorization with the
# custom single-CTA kernel on the diagonal, i.e. no cuSOLVER anywhere.
# 8x2048 4910 -> 5520, 2x4096 3210 -> 4660. The 4070 said the opposite
# (1.3-1.4x wins); on the B200 cuSOLVER is fast enough at n>=2048 that a
# Python-level blocked loop cannot pay for its launches. The crossover to
# _blocked_tf32 winning is above 4096, not below it.
# 8x2048 is the worst-throughput entry in the large regime (4.5 TF/s):
# cuSOLVER's batched potrf at n=2048 is slow no matter the batch. Route it
# through the mid blocked path with nb=256, whose diagonal block goes to
# the custom single-CTA potrf, so cuSOLVER is out of the loop entirely.
# b >= 8 only: measured on B200, at b == 2 this path is far worse than one
# classical cuSOLVER call per matrix (2x2048 1348 -> 3730, 2x4096
# 3210 -> 9300) -- nb=256 leaves too little parallelism per launch.
if n == 2048 and data.shape[0] >= 8:
return _blocked_mid(data, nb=256, tri_base=128)
if 2048 <= n <= 4096 and 2 <= data.shape[0] <= 4:
L = torch.empty_like(data)
for i in range(data.shape[0]):
L[i] = torch.linalg.cholesky_ex(data[i], check_errors=False).L
return L
if data.shape[0] * float(n) ** 3 >= _TF32_MIN_WORK:
return _blocked_tf32(data)
return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 942 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