submission 930470
diacl_54609 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2361 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-930470?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:aeb6bc565e5c665787be4e2d55ffde71fd280cc4dbbc20c406f11fab48380c6b
license declaredunknown
license concludedunknown
authorsdiacl_54609
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mbarrier
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;"mma
using bf3_acc_t = nvcuda::wmma::fragment<nvcuda::wmma::accumulator,shared-memory
__shared__ __align__(16) float stage[WPB * N * LDP];vector-width = float4
const float4* g4 = reinterpret_cast<const float4*>(A + (size_t)b0 * N * N);Kernel source
submission.py2361 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200
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 <cusolverDn.h>
#include <cublas_v2.h>
#include <cuda_bf16.h>
#include <mma.h>
template <int N>
__global__ void __launch_bounds__(128) chol_warp_reg_kernel(
const float* __restrict__ A, float* __restrict__ L, int B) {
constexpr int WPB = 4;
constexpr int LDP = N + 1;
__shared__ __align__(16) float stage[WPB * N * LDP];
__shared__ __align__(16) float cbuf[WPB][2][N];
const int tid = threadIdx.x;
const int warp = tid >> 5, lane = tid & 31;
const int b0 = blockIdx.x * WPB;
const int nmat = (B - b0 < WPB) ? (B - b0) : WPB;
const float4* g4 = reinterpret_cast<const float4*>(A + (size_t)b0 * N * N);
const int quads = nmat * (N * N / 4);
for (int i = tid; i < quads; i += 128) {
float4 v = g4[i];
int m = i >> 8;
int q = i & 255;
int r = q >> 3, c = (q & 7) << 2;
float* dst = stage + (m * N + r) * LDP + c;
dst[0] = v.x; dst[1] = v.y; dst[2] = v.z; dst[3] = v.w;
}
__syncthreads();
if (warp >= nmat) return;
float* st = stage + warp * N * LDP;
float row[N];
#pragma unroll
for (int c = 0; c < N; ++c) row[c] = st[lane * LDP + c];
#pragma unroll
for (int j = 0; j < N; ++j) {
float* cb = cbuf[warp][j & 1];
cb[lane] = row[j];
__syncwarp();
float inv = rsqrtf(cb[j]);
float lij = row[j] * inv;
float li2 = lij * inv;
row[j] = lij;
#pragma unroll
for (int g = (j + 1) >> 2; g < N / 4; ++g) {
float4 c4 = *reinterpret_cast<const float4*>(cb + g * 4);
if (g * 4 > j) row[g * 4 ] -= li2 * c4.x;
if (g * 4 + 1 > j) row[g * 4 + 1] -= li2 * c4.y;
if (g * 4 + 2 > j) row[g * 4 + 2] -= li2 * c4.z;
if (g * 4 + 3 > j) row[g * 4 + 3] -= li2 * c4.w;
}
}
#pragma unroll
for (int c = 0; c < N; ++c)
st[lane * LDP + c] = (c <= lane) ? row[c] : 0.f;
__syncwarp();
float4* o4 = reinterpret_cast<float4*>(L + (size_t)(b0 + warp) * N * N);
#pragma unroll
for (int k = 0; k < N * N / 4 / 32; ++k) {
int q = k * 32 + lane;
int r = q >> 3, c = (q & 7) << 2;
const float* s = st + r * LDP + c;
o4[q] = make_float4(s[0], s[1], s[2], s[3]);
}
}
template <int N>
__global__ void __launch_bounds__(64) chol_rowreg_kernel(
const float* __restrict__ A, float* __restrict__ L, int B) {
constexpr int MPB = 2;
constexpr int LDP = N + 1;
__shared__ __align__(16) float stage[MPB * N * LDP];
__shared__ __align__(16) float cbuf[MPB][2][N];
const int tid = threadIdx.x;
const int lane = tid & 31;
const int m = tid >> 5;
const int b0 = blockIdx.x * MPB;
const int nmat = (B - b0 < MPB) ? (B - b0) : MPB;
const float4* g4 = reinterpret_cast<const float4*>(A + (size_t)b0 * N * N);
const int quads = nmat * (N * N / 4);
for (int i = tid; i < quads; i += 64) {
float4 v = g4[i];
int mm = i >> 10;
int q = i & 1023;
int r = q >> 4, c = (q & 15) << 2;
float* dst = stage + (mm * N + r) * LDP + c;
dst[0] = v.x; dst[1] = v.y; dst[2] = v.z; dst[3] = v.w;
}
__syncthreads();
if (m < nmat) {
float* st = stage + m * N * LDP;
float r0[N], r1[N];
#pragma unroll
for (int c = 0; c < N; ++c) {
r0[c] = st[lane * LDP + c];
r1[c] = st[(lane + 32) * LDP + c];
}
#pragma unroll
for (int j = 0; j < N; ++j) {
float* cb = cbuf[m][j & 1];
cb[lane] = r0[j];
cb[lane + 32] = r1[j];
__syncwarp();
float inv = rsqrtf(cb[j]);
float l0 = r0[j] * inv;
float l1 = r1[j] * inv;
float s0 = l0 * inv;
float s1 = l1 * inv;
r0[j] = l0;
r1[j] = l1;
#pragma unroll
for (int g = (j + 1) >> 2; g < N / 4; ++g) {
float4 c4 = *reinterpret_cast<const float4*>(cb + g * 4);
if (g * 4 > j) { r0[g * 4 ] -= s0 * c4.x;
r1[g * 4 ] -= s1 * c4.x; }
if (g * 4 + 1 > j) { r0[g * 4 + 1] -= s0 * c4.y;
r1[g * 4 + 1] -= s1 * c4.y; }
if (g * 4 + 2 > j) { r0[g * 4 + 2] -= s0 * c4.z;
r1[g * 4 + 2] -= s1 * c4.z; }
if (g * 4 + 3 > j) { r0[g * 4 + 3] -= s0 * c4.w;
r1[g * 4 + 3] -= s1 * c4.w; }
}
}
#pragma unroll
for (int c = 0; c < N; ++c) {
st[lane * LDP + c] = (c <= lane) ? r0[c] : 0.f;
st[(lane + 32) * LDP + c] = (c <= lane + 32) ? r1[c] : 0.f;
}
}
__syncthreads();
float4* o4 = reinterpret_cast<float4*>(L + (size_t)b0 * N * N);
for (int i = tid; i < quads; i += 64) {
int mm = i >> 10;
int q = i & 1023;
int r = q >> 4, c = (q & 15) << 2;
const float* s = stage + (mm * N + r) * LDP + c;
o4[i] = make_float4(s[0], s[1], s[2], s[3]);
}
}
torch::Tensor chol_small(torch::Tensor A) {
TORCH_CHECK(A.is_cuda() && A.dtype() == torch::kFloat32 && A.dim() == 3 &&
A.is_contiguous(), "chol_small: bad input");
const int B = A.size(0);
const int n = A.size(1);
auto L = torch::empty_like(A);
if (n == 32) {
const int wpb = 4;
chol_warp_reg_kernel<32><<<(B + wpb - 1) / wpb, wpb * 32>>>(
A.data_ptr<float>(), L.data_ptr<float>(), B);
} else if (n == 64) {
const int mpb = 2;
chol_rowreg_kernel<64><<<(B + mpb - 1) / mpb, 64>>>(
A.data_ptr<float>(), L.data_ptr<float>(), B);
} else {
TORCH_CHECK(false, "chol_small: unsupported n");
}
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, cudaGetErrorString(err));
return L;
}
#define NB 64
#define LDX 65
__device__ __forceinline__ int ld_acquire(const int* p) {
int v;
asm volatile("ld.acquire.gpu.b32 %0, [%1];" : "=r"(v) : "l"(p) : "memory");
return v;
}
__device__ __forceinline__ void st_release(int* p, int v) {
asm volatile("st.release.gpu.b32 [%0], %1;" :: "l"(p), "r"(v) : "memory");
}
__device__ __forceinline__ void wait_ge(const int* p, int target) {
int ns = 8;
while (ld_acquire(p) < target) {
__nanosleep(ns);
if (ns < 32) ns <<= 1;
}
}
__device__ __forceinline__ int next_task(
unsigned long long* p, unsigned int epoch) {
const unsigned long long base =
static_cast<unsigned long long>(epoch) << 32;
while (true) {
unsigned long long cur = atomicAdd(p, 0ULL);
if (static_cast<unsigned int>(cur >> 32) != epoch) {
if (atomicCAS(p, cur, base) != cur) continue;
}
unsigned long long got = atomicAdd(p, 1ULL);
if (static_cast<unsigned int>(got >> 32) == epoch)
return static_cast<int>(got);
}
}
template <int TNB>
__device__ __forceinline__ void load_tile(float* dst, const float* src,
int n, int tid) {
constexpr int LDT = TNB + 1;
constexpr int QROW = TNB / 4;
#pragma unroll
for (int q = 0; q < TNB * TNB / 4 / 256; ++q) {
int idx = (q << 8) + tid;
int r = idx / QROW, c = (idx % QROW) << 2;
float4 v = *reinterpret_cast<const float4*>(src + (size_t)r * n + c);
float* d = dst + r * LDT + c;
d[0] = v.x; d[1] = v.y; d[2] = v.z; d[3] = v.w;
}
}
template <int TNB>
__device__ __forceinline__ void tile_ldg(float4 (&v)[TNB * TNB / 1024],
const float* src, int n, int tid) {
constexpr int QROW = TNB / 4;
#pragma unroll
for (int q = 0; q < TNB * TNB / 1024; ++q) {
int idx = (q << 8) + tid;
int r = idx / QROW, c = (idx % QROW) << 2;
v[q] = *reinterpret_cast<const float4*>(src + (size_t)r * n + c);
}
}
template <int TNB, bool BF1 = false>
__device__ __forceinline__ void tile_sts(__nv_bfloat16* hi,
__nv_bfloat16* lo,
const float4 (&v)[TNB * TNB / 1024],
int tid) {
constexpr int BDT = TNB + 8;
constexpr int QROW = TNB / 4;
#pragma unroll
for (int q = 0; q < TNB * TNB / 1024; ++q) {
int idx = (q << 8) + tid;
int r = idx / QROW, c = (idx % QROW) << 2;
__nv_bfloat162 h01 = __floats2bfloat162_rn(v[q].x, v[q].y);
__nv_bfloat162 h23 = __floats2bfloat162_rn(v[q].z, v[q].w);
*reinterpret_cast<uint2*>(hi + r * BDT + c) = make_uint2(
reinterpret_cast<const unsigned int&>(h01),
reinterpret_cast<const unsigned int&>(h23));
if constexpr (!BF1) {
__nv_bfloat162 l01 = __floats2bfloat162_rn(
v[q].x - __bfloat162float(h01.x),
v[q].y - __bfloat162float(h01.y));
__nv_bfloat162 l23 = __floats2bfloat162_rn(
v[q].z - __bfloat162float(h23.x),
v[q].w - __bfloat162float(h23.y));
*reinterpret_cast<uint2*>(lo + r * BDT + c) = make_uint2(
reinterpret_cast<const unsigned int&>(l01),
reinterpret_cast<const unsigned int&>(l23));
}
}
}
template <int TNB, bool BF1 = false>
__device__ __forceinline__ void load_tile_bf3(
__nv_bfloat16* hi, __nv_bfloat16* lo, const float* src, int n,
int tid) {
float4 v[TNB * TNB / 1024];
tile_ldg<TNB>(v, src, n, tid);
tile_sts<TNB, BF1>(hi, lo, v, tid);
}
using bf3_acc_t = nvcuda::wmma::fragment<nvcuda::wmma::accumulator,
16, 16, 16, float>;
template <int TNB, bool BF1 = false, bool BF2 = false>
__device__ __forceinline__ void bf3_mma_slab(
bf3_acc_t& c0, bf3_acc_t& c1,
const __nv_bfloat16* hiA, const __nv_bfloat16* loA,
const __nv_bfloat16* hiB, const __nv_bfloat16* loB) {
namespace wmma = nvcuda::wmma;
constexpr int BDT = TNB + 8;
const int warp = threadIdx.x >> 5;
if constexpr (TNB == 32) {
const int bc = warp & 1;
const int br = (warp >> 1) & 1;
const int kk = (warp >> 2) << 4;
wmma::fragment<wmma::matrix_b, 16, 16, 16, __nv_bfloat16,
wmma::col_major> bh, bl;
wmma::fragment<wmma::matrix_a, 16, 16, 16, __nv_bfloat16,
wmma::row_major> ah, al;
wmma::load_matrix_sync(bh, hiB + (bc * 16) * BDT + kk, BDT);
wmma::load_matrix_sync(ah, hiA + (br * 16) * BDT + kk, BDT);
wmma::mma_sync(c0, ah, bh, c0);
if constexpr (!BF1) {
if constexpr (!BF2)
wmma::load_matrix_sync(
bl, loB + (bc * 16) * BDT + kk, BDT);
wmma::load_matrix_sync(al, loA + (br * 16) * BDT + kk, BDT);
wmma::mma_sync(c0, al, bh, c0);
if constexpr (!BF2)
wmma::mma_sync(c0, ah, bl, c0);
}
return;
}
const int bc = warp & 3;
const int br0 = warp >> 2;
const int br1 = br0 + 2;
#pragma unroll
for (int kk = 0; kk < TNB; kk += 16) {
wmma::fragment<wmma::matrix_b, 16, 16, 16, __nv_bfloat16,
wmma::col_major> bh, bl;
wmma::fragment<wmma::matrix_a, 16, 16, 16, __nv_bfloat16,
wmma::row_major> ah, al;
wmma::load_matrix_sync(bh, hiB + (bc * 16) * BDT + kk, BDT);
wmma::load_matrix_sync(ah, hiA + (br0 * 16) * BDT + kk, BDT);
wmma::mma_sync(c0, ah, bh, c0);
if constexpr (!BF1) {
if constexpr (!BF2)
wmma::load_matrix_sync(
bl, loB + (bc * 16) * BDT + kk, BDT);
wmma::load_matrix_sync(al, loA + (br0 * 16) * BDT + kk, BDT);
wmma::mma_sync(c0, al, bh, c0);
if constexpr (!BF2)
wmma::mma_sync(c0, ah, bl, c0);
}
wmma::load_matrix_sync(ah, hiA + (br1 * 16) * BDT + kk, BDT);
wmma::mma_sync(c1, ah, bh, c1);
if constexpr (!BF1) {
wmma::load_matrix_sync(al, loA + (br1 * 16) * BDT + kk, BDT);
wmma::mma_sync(c1, al, bh, c1);
if constexpr (!BF2)
wmma::mma_sync(c1, ah, bl, c1);
}
}
}
template <int TNB, typename AccT>
__device__ __forceinline__ void bf3_fold(
float (&acc)[TNB / 16][TNB / 16], AccT& c0, AccT& c1,
float* stage, float* stage2, int tr, int tc) {
namespace wmma = nvcuda::wmma;
constexpr int SDF = TNB + 4;
const int warp = threadIdx.x >> 5;
if constexpr (TNB == 32) {
const int bc = warp & 1;
const int br = (warp >> 1) & 1;
float* dst = (warp >> 2) ? stage2 : stage;
wmma::store_matrix_sync(dst + (br * 16) * SDF + bc * 16, c0, SDF,
wmma::mem_row_major);
__syncthreads();
#pragma unroll
for (int u = 0; u < 2; ++u)
#pragma unroll
for (int v = 0; v < 2; ++v)
acc[u][v] -= stage[(tr * 2 + u) * SDF + tc * 2 + v]
+ stage2[(tr * 2 + u) * SDF + tc * 2 + v];
return;
}
const int bc = warp & 3;
const int br0 = warp >> 2;
const int br1 = br0 + 2;
wmma::store_matrix_sync(stage + (br0 * 16) * SDF + bc * 16, c0, SDF,
wmma::mem_row_major);
wmma::store_matrix_sync(stage + (br1 * 16) * SDF + bc * 16, c1, SDF,
wmma::mem_row_major);
__syncthreads();
#pragma unroll
for (int u = 0; u < 4; ++u)
#pragma unroll
for (int v = 0; v < 4; ++v)
acc[u][v] -= stage[(tr * 4 + u) * SDF + tc * 4 + v];
}
template <int TNB>
__device__ __forceinline__ void stage_acc(float* dst, const float* At, int n,
float (&acc)[TNB / 16][TNB / 16],
int tr, int tc) {
constexpr int LDT = TNB + 1;
constexpr int R = TNB / 16;
if constexpr (R == 4) {
#pragma unroll
for (int u = 0; u < 4; ++u) {
float4 c4 = *reinterpret_cast<const float4*>(
At + (size_t)(tr * 4 + u) * n + tc * 4);
float* d = dst + (tr * 4 + u) * LDT + tc * 4;
d[0] = c4.x + acc[u][0]; d[1] = c4.y + acc[u][1];
d[2] = c4.z + acc[u][2]; d[3] = c4.w + acc[u][3];
}
} else {
#pragma unroll
for (int u = 0; u < R; ++u) {
const float* s = At + (size_t)(tr * R + u) * n + tc * R;
float* d = dst + (tr * R + u) * LDT + tc * R;
#pragma unroll
for (int v = 0; v < R; ++v) d[v] = s[v] + acc[u][v];
}
}
}
template <int TNB>
__device__ __forceinline__ void stage_acc_sm(
float* dst, const float* Asm, float (&acc)[TNB / 16][TNB / 16],
int tr, int tc) {
constexpr int LDT = TNB + 1;
constexpr int R = TNB / 16;
#pragma unroll
for (int u = 0; u < R; ++u) {
const float* s = Asm + (tr * R + u) * LDT + tc * R;
float* d = dst + (tr * R + u) * LDT + tc * R;
#pragma unroll
for (int v = 0; v < R; ++v) d[v] = s[v] + acc[u][v];
}
}
template <int TNB, int LDT = TNB + 1>
__device__ __forceinline__ void potrf_tile_rank4(float* tileL, int tid) {
constexpr int NS = 256 / TNB;
const int r = tid & (TNB - 1), q = tid / TNB;
for (int c = 0; c < TNB; c += 4) {
const float* K = tileL + c * LDT + c;
float inv0 = rsqrtf(K[0]);
float l10 = K[LDT] * inv0;
float l20 = K[2 * LDT] * inv0;
float l30 = K[3 * LDT] * inv0;
float inv1 = rsqrtf(K[LDT + 1] - l10 * l10);
float l21 = (K[2 * LDT + 1] - l20 * l10) * inv1;
float l31 = (K[3 * LDT + 1] - l30 * l10) * inv1;
float inv2 = rsqrtf(K[2 * LDT + 2] - l20 * l20 - l21 * l21);
float l32 = (K[3 * LDT + 2] - l30 * l20 - l31 * l21) * inv2;
float inv3 = rsqrtf(K[3 * LDT + 3] - l30 * l30 - l31 * l31
- l32 * l32);
float x0 = 0.f, x1 = 0.f, x2 = 0.f, x3 = 0.f;
if (r >= c) {
const float* rowp = tileL + r * LDT + c;
x0 = rowp[0] * inv0;
x1 = (rowp[1] - x0 * l10) * inv1;
x2 = (rowp[2] - x0 * l20 - x1 * l21) * inv2;
x3 = (rowp[3] - x0 * l30 - x1 * l31 - x2 * l32) * inv3;
}
__syncthreads();
if (q == 0 && r >= c) {
float* rowp = tileL + r * LDT + c;
rowp[0] = x0; rowp[1] = x1; rowp[2] = x2; rowp[3] = x3;
}
__syncthreads();
for (int m = c + 4 + q; m <= r; m += NS) {
const float* pm = tileL + m * LDT + c;
tileL[r * LDT + m] -= x0 * pm[0] + x1 * pm[1] +
x2 * pm[2] + x3 * pm[3];
}
__syncthreads();
}
}
template <int TNB, int LDL = TNB + 1, int LDC = LDL>
__device__ __forceinline__ void trsm_tile_rank4(const float* tileL,
float* tileC, float* rec,
int tid) {
constexpr int NS = 256 / TNB;
if (tid < TNB) rec[tid] = 1.0f / tileL[tid * LDL + tid];
__syncthreads();
const int r = tid & (TNB - 1), q = tid / TNB;
for (int c = 0; c < TNB; c += 4) {
float inv0 = rec[c], inv1 = rec[c + 1];
float inv2 = rec[c + 2], inv3 = rec[c + 3];
float l10 = tileL[(c + 1) * LDL + c];
float l20 = tileL[(c + 2) * LDL + c];
float l21 = tileL[(c + 2) * LDL + c + 1];
float l30 = tileL[(c + 3) * LDL + c];
float l31 = tileL[(c + 3) * LDL + c + 1];
float l32 = tileL[(c + 3) * LDL + c + 2];
float x0 = tileC[r * LDC + c] * inv0;
float x1 = (tileC[r * LDC + c + 1] - x0 * l10) * inv1;
float x2 = (tileC[r * LDC + c + 2] - x0 * l20 - x1 * l21) * inv2;
float x3 = (tileC[r * LDC + c + 3] - x0 * l30 - x1 * l31
- x2 * l32) * inv3;
__syncthreads();
if (q == 0) {
tileC[r * LDC + c ] = x0;
tileC[r * LDC + c + 1] = x1;
tileC[r * LDC + c + 2] = x2;
tileC[r * LDC + c + 3] = x3;
}
for (int m = c + 4 + q; m < TNB; m += NS)
tileC[r * LDC + m] -=
x0 * tileL[m * LDL + c] + x1 * tileL[m * LDL + c + 1] +
x2 * tileL[m * LDL + c + 2] + x3 * tileL[m * LDL + c + 3];
__syncthreads();
}
}
__device__ __forceinline__ void potrf32_warp(float* tileL,
float (&cb2)[2][32], int tid) {
constexpr int LDT = 33;
if (tid < 32) {
float row[32];
#pragma unroll
for (int c = 0; c < 32; ++c) row[c] = tileL[tid * LDT + c];
#pragma unroll
for (int j = 0; j < 32; ++j) {
float* cb = cb2[j & 1];
cb[tid] = row[j];
__syncwarp();
float inv = rsqrtf(cb[j]);
float lij = row[j] * inv;
float li2 = lij * inv;
row[j] = lij;
#pragma unroll
for (int g = (j + 1) >> 2; g < 8; ++g) {
float4 c4 = *reinterpret_cast<const float4*>(cb + g * 4);
if (g * 4 > j) row[g * 4 ] -= li2 * c4.x;
if (g * 4 + 1 > j) row[g * 4 + 1] -= li2 * c4.y;
if (g * 4 + 2 > j) row[g * 4 + 2] -= li2 * c4.z;
if (g * 4 + 3 > j) row[g * 4 + 3] -= li2 * c4.w;
}
}
#pragma unroll
for (int c = 0; c < 32; ++c) tileL[tid * LDT + c] = row[c];
}
__syncthreads();
}
__device__ __forceinline__ void trsm32_warp(const float* tileL, float* tileC,
int tid) {
constexpr int LDT = 33;
if (tid < 32) {
float x[32];
#pragma unroll
for (int c = 0; c < 32; ++c) x[c] = tileC[tid * LDT + c];
#pragma unroll
for (int j = 0; j < 32; ++j) {
float xj = x[j] * (1.0f / tileL[j * LDT + j]);
x[j] = xj;
#pragma unroll
for (int c = j + 1; c < 32; ++c)
x[c] -= xj * tileL[c * LDT + j];
}
#pragma unroll
for (int c = 0; c < 32; ++c) tileC[tid * LDT + c] = x[c];
}
__syncthreads();
}
template <int TNB, int MINB = 2, bool BF1 = false, bool BF2 = false,
bool PACK = false>
__global__ void __launch_bounds__(256, MINB) chol_persist_kernel(
const float* __restrict__ Asrc, float* __restrict__ W,
const long long* __restrict__ tasks, int ntasks,
unsigned long long* __restrict__ counter,
int* __restrict__ state, int n, int t,
int bstride, int epoch, __nv_bfloat16* __restrict__ Hout) {
constexpr int LDT = TNB + 1;
constexpr int SMW = TNB * (TNB + 8);
__shared__ __align__(32) float smA[SMW];
__shared__ __align__(32) float smB[SMW];
float* tileL = smA;
float* tileC = smB;
__shared__ float colbuf[TNB];
__shared__ float cbw[2][32];
__shared__ float sL32[TNB == 32 ? 32 * 33 : 1];
__shared__ __align__(32) float smC[TNB == 32 ? SMW : 1];
__shared__ __align__(32) float smD[TNB == 32 ? SMW : 1];
__shared__ long long shTask;
const int tid = threadIdx.x;
while (true) {
if (tid == 0) {
int id = next_task(counter, epoch);
shTask = (id < ntasks) ? tasks[id] : -1;
}
__syncthreads();
const long long tk = shTask;
if (tk < 0) return;
const int type = (int)(tk >> 52) & 7;
const int b = (int)(tk >> 30) & 0x3FFFFF;
const int j = (int)(tk >> 10) & 1023;
const int i = (int)tk & 1023;
int* stb = state + (size_t)b * bstride;
const float* Ab = Asrc + (size_t)b * n * n;
float* Wb = W + (size_t)b * n * n;
__nv_bfloat16* Hb = nullptr;
if constexpr (PACK) Hb = Hout + (size_t)b * n * n;
const int tr = tid >> 4, tc = tid & 15;
int* flag;
if (type == 0) {
float acc[TNB / 16][TNB / 16] = {};
{
bf3_acc_t c0, c1;
nvcuda::wmma::fill_fragment(c0, 0.0f);
nvcuda::wmma::fill_fragment(c1, 0.0f);
__nv_bfloat16* hiA = reinterpret_cast<__nv_bfloat16*>(smA);
const float* row = Wb + (size_t)(j * TNB) * n;
for (int k = 0; k < j; ++k) {
if (tid == 0) wait_ge(&stb[j * t + k], epoch);
__syncthreads();
load_tile_bf3<TNB, BF1>(
hiA, hiA + TNB * (TNB + 8),
row + k * TNB, n, tid);
__syncthreads();
bf3_mma_slab<TNB, BF1, BF2>(
c0, c1, hiA,
hiA + TNB * (TNB + 8), hiA,
hiA + TNB * (TNB + 8));
}
__syncthreads();
bf3_fold<TNB>(acc, c0, c1, smA, smB, tr, tc);
}
const float* At = Ab + (size_t)(j * TNB) * n + j * TNB;
float* Wt = Wb + (size_t)(j * TNB) * n + j * TNB;
__syncthreads();
stage_acc<TNB>(tileL, At, n, acc, tr, tc);
__syncthreads();
if constexpr (TNB == 32)
potrf32_warp(tileL, cbw, tid);
else
potrf_tile_rank4<TNB>(tileL, tid);
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
Wt[(size_t)r * n + c] = (c <= r) ? tileL[r * LDT + c] : 0.f;
}
flag = &stb[j * t + j];
} else if (type == 2) {
if constexpr (TNB != 32) {
flag = &stb[i * t + j];
} else {
float accP[TNB / 16][TNB / 16] = {};
float accD[TNB / 16][TNB / 16] = {};
bf3_acc_t c0, c1;
nvcuda::wmma::fill_fragment(c0, 0.0f);
nvcuda::wmma::fill_fragment(c1, 0.0f);
__nv_bfloat16* hiA = reinterpret_cast<__nv_bfloat16*>(smA);
__nv_bfloat16* hiB = reinterpret_cast<__nv_bfloat16*>(smB);
const float* rowA = Wb + (size_t)(i * TNB) * n;
const float* rowB = Wb + (size_t)(j * TNB) * n;
load_tile<TNB>(smC, Ab + (size_t)(i * TNB) * n + j * TNB, n,
tid);
load_tile<TNB>(smD, Ab + (size_t)(i * TNB) * n + i * TNB, n,
tid);
for (int k = 0; k < j; ++k) {
if (tid == 0) {
wait_ge(&stb[i * t + k], epoch);
wait_ge(&stb[j * t + k], epoch);
}
__syncthreads();
load_tile_bf3<TNB, BF1>(
hiA, hiA + TNB * (TNB + 8),
rowA + k * TNB, n, tid);
load_tile_bf3<TNB, BF1 || BF2>(
hiB, hiB + TNB * (TNB + 8),
rowB + k * TNB, n, tid);
__syncthreads();
bf3_mma_slab<TNB, BF1, BF2>(
c0, c1, hiA, hiA + TNB * (TNB + 8),
hiB, hiB + TNB * (TNB + 8));
bf3_mma_slab<TNB, BF1, BF2>(
c1, c0, hiA, hiA + TNB * (TNB + 8),
hiA, hiA + TNB * (TNB + 8));
}
__syncthreads();
bf3_fold<TNB>(accP, c0, c1, smA, smB, tr, tc);
if (tid == 0) wait_ge(&stb[j * t + j], epoch);
__syncthreads();
load_tile<TNB>(sL32, Wb + (size_t)(j * TNB) * n + j * TNB,
n, tid);
{
float* Wt = Wb + (size_t)(i * TNB) * n + j * TNB;
stage_acc_sm<TNB>(tileC, smC, accP, tr, tc);
__syncthreads();
trsm32_warp(sL32, tileC, tid);
float* Wu = Wb + (size_t)(j * TNB) * n + i * TNB;
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
Wt[(size_t)r * n + c] = tileC[r * LDT + c];
Wu[(size_t)r * n + c] = 0.f;
}
__syncthreads();
if (tid == 0) {
st_release(&stb[i * t + j], epoch);
}
}
for (int q = 0; q < (TNB * TNB + 255) / 256; ++q) {
int idx = q * 256 + tid;
int r = idx / TNB, c = idx % TNB;
float v = tileC[r * LDT + c];
__nv_bfloat16 h = __float2bfloat16(v);
hiA[r * (TNB + 8) + c] = h;
if constexpr (!BF1)
hiA[TNB * (TNB + 8) + r * (TNB + 8) + c] =
__float2bfloat16(v - __bfloat162float(h));
}
__syncthreads();
bf3_mma_slab<TNB, BF1, BF2>(
c1, c0, hiA, hiA + TNB * (TNB + 8),
hiA, hiA + TNB * (TNB + 8));
__syncthreads();
bf3_fold<TNB>(accD, c1, c0, smA, smB, tr, tc);
float* Wt = Wb + (size_t)(i * TNB) * n + i * TNB;
__syncthreads();
stage_acc_sm<TNB>(tileL, smD, accD, tr, tc);
__syncthreads();
potrf32_warp(tileL, cbw, tid);
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
Wt[(size_t)r * n + c] = (c <= r) ? tileL[r * LDT + c] : 0.f;
}
flag = &stb[i * t + i];
}
} else {
float acc[TNB / 16][TNB / 16] = {};
{
bf3_acc_t c0, c1;
nvcuda::wmma::fill_fragment(c0, 0.0f);
nvcuda::wmma::fill_fragment(c1, 0.0f);
__nv_bfloat16* hiA = reinterpret_cast<__nv_bfloat16*>(smA);
__nv_bfloat16* hiB = reinterpret_cast<__nv_bfloat16*>(smB);
const float* rowA = Wb + (size_t)(i * TNB) * n;
const float* rowB = Wb + (size_t)(j * TNB) * n;
for (int k = 0; k < j; ++k) {
if (tid == 0) {
wait_ge(&stb[i * t + k], epoch);
wait_ge(&stb[j * t + k], epoch);
}
__syncthreads();
load_tile_bf3<TNB, BF1>(
hiA, hiA + TNB * (TNB + 8),
rowA + k * TNB, n, tid);
load_tile_bf3<TNB, BF1 || BF2>(
hiB, hiB + TNB * (TNB + 8),
rowB + k * TNB, n, tid);
__syncthreads();
bf3_mma_slab<TNB, BF1, BF2>(
c0, c1, hiA,
hiA + TNB * (TNB + 8), hiB,
hiB + TNB * (TNB + 8));
}
__syncthreads();
bf3_fold<TNB>(acc, c0, c1, smA, smB, tr, tc);
}
if (tid == 0) wait_ge(&stb[j * t + j], epoch);
__syncthreads();
load_tile<TNB>(tileL, Wb + (size_t)(j * TNB) * n + j * TNB,
n, tid);
const float* At = Ab + (size_t)(i * TNB) * n + j * TNB;
float* Wt = Wb + (size_t)(i * TNB) * n + j * TNB;
stage_acc<TNB>(tileC, At, n, acc, tr, tc);
__syncthreads();
if constexpr (TNB == 32)
trsm32_warp(tileL, tileC, tid);
else
trsm_tile_rank4<TNB>(tileL, tileC, colbuf, tid);
float* Wu = Wb + (size_t)(j * TNB) * n + i * TNB;
__nv_bfloat16* Ht = nullptr;
if constexpr (PACK)
if (i >= t)
Ht = Hb + (size_t)(i * TNB) * n + j * TNB;
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
float v = tileC[r * LDT + c];
Wt[(size_t)r * n + c] = v;
if constexpr (PACK)
if (Ht) Ht[(size_t)r * n + c] = __float2bfloat16(v);
Wu[(size_t)r * n + c] = 0.f;
}
flag = &stb[i * t + j];
}
__syncthreads();
if (tid == 0) {
st_release(flag, epoch);
}
}
}
__device__ __forceinline__ unsigned int ring_smem_u32(const void* p) {
return (unsigned int)__cvta_generic_to_shared(p);
}
__device__ __forceinline__ unsigned int ring_mapa(unsigned int a,
unsigned int rank) {
unsigned int r;
asm volatile("mapa.shared::cluster.u32 %0, %1, %2;" : "=r"(r)
: "r"(a), "r"(rank));
return r;
}
#define RING_R 8
#define RING_LF 1056
#define RING_PU 1280
template <bool BF1 = false, int MINB = 2, bool PACK = false>
__global__ void __launch_bounds__(256, MINB) chol_ring_kernel(
const float* __restrict__ Asrc, float* __restrict__ W,
const long long* __restrict__ tasks, int ntasks,
unsigned long long* __restrict__ counter,
int* __restrict__ state, int n, int t,
int nring, int bstride, int epoch,
__nv_bfloat16* __restrict__ Hout) {
constexpr int TNB = 32;
constexpr int LDT = 33;
constexpr int SMW = 32 * 40;
constexpr int RING_PU_USED = BF1 ? RING_PU / 2 : RING_PU;
__shared__ __align__(32) float smA[SMW];
__shared__ __align__(32) float smB[SMW];
__shared__ __align__(32) float smH[SMW];
__shared__ __align__(16) float tileP[32 * 33];
__shared__ __align__(16) float sL[32 * 33];
__shared__ __align__(16) float inbox[RING_LF + RING_PU_USED];
__shared__ __align__(8) unsigned long long rbar;
__shared__ __align__(16) float cbw[2][32];
__shared__ long long shTask;
const int tid = threadIdx.x;
const int tr = tid >> 4, tc2 = tid & 15;
unsigned int rank;
asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));
const int cid = blockIdx.x / RING_R;
if (cid < nring) {
const int b = cid;
int* stb = state + (size_t)b * bstride;
const float* Ab = Asrc + (size_t)b * n * n;
float* Wb = W + (size_t)b * n * n;
__nv_bfloat16* hA = reinterpret_cast<__nv_bfloat16*>(smA);
__nv_bfloat16* hB = reinterpret_cast<__nv_bfloat16*>(smB);
__nv_bfloat16* hH = reinterpret_cast<__nv_bfloat16*>(smH);
float* inL = inbox;
__nv_bfloat16* inP =
reinterpret_cast<__nv_bfloat16*>(inbox + RING_LF);
if (tid == 0) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;"
:: "r"(ring_smem_u32(&rbar)));
asm volatile("fence.mbarrier_init.release.cluster;");
}
__syncthreads();
asm volatile("barrier.cluster.arrive.aligned;");
asm volatile("barrier.cluster.wait.aligned;");
const unsigned int nxt = (rank + 1) % RING_R;
const unsigned int nxtInbox = ring_mapa(ring_smem_u32(inbox), nxt);
const unsigned int nxtBar = ring_mapa(ring_smem_u32(&rbar), nxt);
const unsigned int myBar = ring_smem_u32(&rbar);
unsigned int phase = 0;
for (int j = (int)rank; j < t; j += RING_R) {
bf3_acc_t fP, fD;
nvcuda::wmma::fill_fragment(fP, 0.0f);
nvcuda::wmma::fill_fragment(fD, 0.0f);
for (int k = 0; k + 2 <= j; ++k) {
if (tid == 0) {
wait_ge(&stb[j * t + k], epoch);
if (k + 3 <= j)
wait_ge(&stb[(j - 1) * t + k], epoch);
}
__syncthreads();
__nv_bfloat16* dA = (k + 2 == j) ? hH : hA;
load_tile_bf3<TNB, BF1>(
dA, dA + SMW,
Wb + (size_t)(j * TNB) * n + k * TNB,
n, tid);
if (k + 3 <= j)
load_tile_bf3<TNB, BF1>(
hB, hB + SMW,
Wb + (size_t)((j - 1) * TNB) * n +
k * TNB, n, tid);
__syncthreads();
bf3_mma_slab<TNB, BF1>(
fD, fP, dA, dA + SMW, dA, dA + SMW);
if (k + 3 <= j)
bf3_mma_slab<TNB, BF1>(
fP, fD, dA, dA + SMW, hB, hB + SMW);
}
if (j > 0) {
if (tid == 0)
asm volatile(
"mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;"
:: "r"(myBar),
"r"((j == 1) ? RING_LF * 4
: (RING_LF + RING_PU_USED) * 4));
asm volatile(
"{\n\t.reg .pred p;\n\trw%=:\n\t"
"mbarrier.try_wait.parity.shared::cta.b64 p, [%0], %1;\n\t"
"@!p bra rw%=;\n\t}\n" :: "r"(myBar), "r"(phase));
phase ^= 1;
__syncthreads();
if (j >= 2)
bf3_mma_slab<TNB, BF1>(
fP, fD, hH, hH + SMW, inP, inP + SMW);
}
__syncthreads();
if (j > 0) {
{
float a32[TNB / 16][TNB / 16] = {};
bf3_fold<TNB>(a32, fP, fD, smA, smB, tr, tc2);
__syncthreads();
stage_acc<TNB>(tileP,
Ab + (size_t)(j * TNB) * n +
(j - 1) * TNB, n, a32, tr, tc2);
__syncthreads();
trsm32_warp(inL, tileP, tid);
}
for (int q = 0; q < (TNB * TNB + 255) / 256; ++q) {
int idx = q * 256 + tid;
int r = idx / TNB, c = idx % TNB;
float v = tileP[r * LDT + c];
__nv_bfloat16 h = __float2bfloat16(v);
hA[r * (TNB + 8) + c] = h;
if constexpr (!BF1)
hA[SMW + r * (TNB + 8) + c] =
__float2bfloat16(v - __bfloat162float(h));
}
__syncthreads();
bf3_mma_slab<TNB, BF1>(
fD, fP, hA, hA + SMW, hA, hA + SMW);
}
__syncthreads();
{
float a32[TNB / 16][TNB / 16] = {};
bf3_fold<TNB>(a32, fD, fP, smB, inbox, tr, tc2);
__syncthreads();
stage_acc<TNB>(sL,
Ab + (size_t)(j * TNB) * n + j * TNB, n,
a32, tr, tc2);
__syncthreads();
potrf32_warp(sL, cbw, tid);
}
if (j + 1 < t) {
for (int i = tid; i < RING_LF; i += 256)
asm volatile(
"st.async.shared::cluster.mbarrier::complete_tx::bytes.u32 [%0], %1, [%2];"
:: "r"(nxtInbox + 4u * i),
"r"(__float_as_uint(sL[i])), "r"(nxtBar));
if (j > 0) {
const unsigned int* hu =
reinterpret_cast<const unsigned int*>(hA);
for (int i = tid; i < RING_PU_USED; i += 256)
asm volatile(
"st.async.shared::cluster.mbarrier::complete_tx::bytes.u32 [%0], %1, [%2];"
:: "r"(nxtInbox + 4u * (RING_LF + i)),
"r"(hu[i]), "r"(nxtBar));
}
}
if (j > 0) {
float* Wt = Wb + (size_t)(j * TNB) * n + (j - 1) * TNB;
float* Wu = Wb + (size_t)((j - 1) * TNB) * n + j * TNB;
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
Wt[(size_t)r * n + c] = tileP[r * LDT + c];
Wu[(size_t)r * n + c] = 0.f;
}
}
{
float* Wt = Wb + (size_t)(j * TNB) * n + j * TNB;
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
Wt[(size_t)r * n + c] =
(c <= r) ? sL[r * LDT + c] : 0.f;
}
}
__syncthreads();
if (tid == 0) {
if (j > 0)
st_release(&stb[j * t + (j - 1)], epoch);
st_release(&stb[j * t + j], epoch);
}
}
asm volatile("barrier.cluster.arrive.aligned;");
asm volatile("barrier.cluster.wait.aligned;");
} else {
while (true) {
if (tid == 0) {
int id = next_task(counter, epoch);
shTask = (id < ntasks) ? tasks[id] : -1;
}
__syncthreads();
const long long tk = shTask;
if (tk < 0) break;
const int b = (int)(tk >> 30) & 0x3FFFFF;
const int j = (int)(tk >> 10) & 1023;
const int i = (int)tk & 1023;
int* stb = state + (size_t)b * bstride;
const float* Ab = Asrc + (size_t)b * n * n;
float* Wb = W + (size_t)b * n * n;
__nv_bfloat16* Hb = nullptr;
if constexpr (PACK) Hb = Hout + (size_t)b * n * n;
float acc[TNB / 16][TNB / 16] = {};
{
bf3_acc_t c0, c1;
nvcuda::wmma::fill_fragment(c0, 0.0f);
nvcuda::wmma::fill_fragment(c1, 0.0f);
__nv_bfloat16* hiA = reinterpret_cast<__nv_bfloat16*>(smA);
__nv_bfloat16* hiB = reinterpret_cast<__nv_bfloat16*>(smB);
const float* rowA = Wb + (size_t)(i * TNB) * n;
const float* rowB = Wb + (size_t)(j * TNB) * n;
for (int k = 0; k < j; ++k) {
if (tid == 0) {
wait_ge(&stb[i * t + k], epoch);
wait_ge(&stb[j * t + k], epoch);
}
__syncthreads();
load_tile_bf3<TNB, BF1>(
hiA, hiA + SMW, rowA + k * TNB, n, tid);
load_tile_bf3<TNB, BF1>(
hiB, hiB + SMW, rowB + k * TNB, n, tid);
__syncthreads();
bf3_mma_slab<TNB, BF1>(
c0, c1, hiA, hiA + SMW, hiB, hiB + SMW);
}
__syncthreads();
bf3_fold<TNB>(acc, c0, c1, smA, smB, tr, tc2);
}
if (tid == 0) wait_ge(&stb[j * t + j], epoch);
__syncthreads();
load_tile<TNB>(sL, Wb + (size_t)(j * TNB) * n + j * TNB, n,
tid);
stage_acc<TNB>(tileP,
Ab + (size_t)(i * TNB) * n + j * TNB, n, acc,
tr, tc2);
__syncthreads();
trsm32_warp(sL, tileP, tid);
float* Wt = Wb + (size_t)(i * TNB) * n + j * TNB;
float* Wu = Wb + (size_t)(j * TNB) * n + i * TNB;
__nv_bfloat16* Ht = nullptr;
if constexpr (PACK)
if (i >= t)
Ht = Hb + (size_t)(i * TNB) * n + j * TNB;
for (int q = tid; q < TNB * TNB; q += 256) {
int r = q / TNB, c = q % TNB;
float v = tileP[r * LDT + c];
Wt[(size_t)r * n + c] = v;
if constexpr (PACK)
if (Ht) Ht[(size_t)r * n + c] = __float2bfloat16(v);
Wu[(size_t)r * n + c] = 0.f;
}
__syncthreads();
if (tid == 0) {
st_release(&stb[i * t + j], epoch);
}
}
}
}
int chol_ring(torch::Tensor A, torch::Tensor L, torch::Tensor tasks,
torch::Tensor state, torch::Tensor counter,
int64_t epochc) {
const int n = A.size(2);
const int B = A.size(0);
const int t = n / 32;
static int bps = -1;
static int nsm = 0;
if (bps < 0) {
int dev = 0;
cudaGetDevice(&dev);
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, dev);
nsm = prop.multiProcessorCount;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&bps, chol_ring_kernel<false>, 256, 0);
if (bps < 1) bps = 1;
}
int blocks = nsm * bps;
blocks -= blocks % RING_R;
if (B * RING_R + RING_R > blocks) return 1;
const float* Ap = A.data_ptr<float>();
float* Wp = L.data_ptr<float>();
const long long* tp =
reinterpret_cast<const long long*>(tasks.data_ptr<int64_t>());
unsigned long long* cp =
reinterpret_cast<unsigned long long*>(counter.data_ptr<int>());
int* sp = state.data_ptr<int>();
int n_ = n, t_ = t, nt_ = (int)tasks.size(0), nring = B;
int bstride = t * t, epoch = (int)epochc;
__nv_bfloat16* Hout = nullptr;
void* args[] = {
&Ap, &Wp, &tp, &nt_, &cp, &sp, &n_, &t_, &nring, &bstride,
&epoch, &Hout
};
cudaLaunchConfig_t cfg = {};
cfg.gridDim = dim3(blocks, 1, 1);
cfg.blockDim = dim3(256, 1, 1);
cudaLaunchAttribute attr[1];
attr[0].id = cudaLaunchAttributeClusterDimension;
attr[0].val.clusterDim.x = RING_R;
attr[0].val.clusterDim.y = 1;
attr[0].val.clusterDim.z = 1;
cfg.attrs = attr;
cfg.numAttrs = 1;
cudaError_t err = cudaLaunchKernelExC(
&cfg, (void*)chol_ring_kernel<false>, args);
if (err != cudaSuccess) {
(void)cudaGetLastError();
return 1;
}
return 0;
}
template <int MINB>
static int chol_ring_panel_bf1_impl(
torch::Tensor W, torch::Tensor hi, int64_t offc, int64_t nbc,
torch::Tensor tasks, torch::Tensor state, torch::Tensor counter,
int64_t epochc) {
const int n = W.size(2);
const int B = W.size(0);
const int off = (int)offc;
const int nb = (int)nbc;
const int t = nb / 32;
const int trows = (n - off) / 32;
static int bps = -1;
static int nsm = 0;
if (bps < 0) {
int dev = 0;
cudaGetDevice(&dev);
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, dev);
nsm = prop.multiProcessorCount;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&bps, chol_ring_kernel<true, MINB, true>, 256, 0);
if (bps < 1) bps = 1;
}
int blocks = nsm * bps;
blocks -= blocks % RING_R;
if (B * RING_R + RING_R > blocks) return 1;
int epoch = (int)epochc;
if (epoch == 0) {
cudaMemsetAsync(state.data_ptr<int>(), 0,
(size_t)state.numel() * sizeof(int));
cudaMemsetAsync(counter.data_ptr<int>(), 0, 2 * sizeof(int));
epoch = 1;
}
float* base = W.data_ptr<float>() + (size_t)off * n + off;
const float* Ap = base;
float* Wp = base;
const long long* tp =
reinterpret_cast<const long long*>(tasks.data_ptr<int64_t>());
unsigned long long* cp =
reinterpret_cast<unsigned long long*>(counter.data_ptr<int>());
int* sp = state.data_ptr<int>();
int n_ = n, t_ = t, nt_ = (int)tasks.size(0), nring = B;
int bstride = trows * t;
__nv_bfloat16* Hout =
reinterpret_cast<__nv_bfloat16*>(hi.data_ptr()) +
(size_t)off * n + off;
void* args[] = {
&Ap, &Wp, &tp, &nt_, &cp, &sp, &n_, &t_, &nring, &bstride,
&epoch, &Hout
};
cudaLaunchConfig_t cfg = {};
cfg.gridDim = dim3(blocks, 1, 1);
cfg.blockDim = dim3(256, 1, 1);
cudaLaunchAttribute attr[1];
attr[0].id = cudaLaunchAttributeClusterDimension;
attr[0].val.clusterDim.x = RING_R;
attr[0].val.clusterDim.y = 1;
attr[0].val.clusterDim.z = 1;
cfg.attrs = attr;
cfg.numAttrs = 1;
cudaError_t err = cudaLaunchKernelExC(
&cfg, (void*)chol_ring_kernel<true, MINB, true>, args);
if (err != cudaSuccess) {
(void)cudaGetLastError();
return 1;
}
return 0;
}
int chol_ring_panel_bf1(torch::Tensor W, torch::Tensor hi,
int64_t off, int64_t nb,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t minb,
int64_t epoch) {
return minb == 3
? chol_ring_panel_bf1_impl<3>(
W, hi, off, nb, tasks, state, counter, epoch)
: chol_ring_panel_bf1_impl<2>(
W, hi, off, nb, tasks, state, counter, epoch);
}
template <int TY>
__global__ void tril_transpose_kernel(const float* __restrict__ W,
float* __restrict__ L, int n) {
__shared__ float tile[32][33];
const size_t base = (size_t)blockIdx.z * n * (size_t)n;
const int rb = blockIdx.y << 5;
const int cb = blockIdx.x << 5;
if (cb > rb) {
for (int dy = threadIdx.y; dy < 32; dy += TY) {
float4 z = make_float4(0.f, 0.f, 0.f, 0.f);
*reinterpret_cast<float4*>(
L + base + (size_t)(rb + dy) * n + cb + (threadIdx.x << 2)) = z;
}
return;
}
for (int dy = threadIdx.y; dy < 32; dy += TY) {
float4 v = *reinterpret_cast<const float4*>(
W + base + (size_t)(cb + dy) * n + rb + (threadIdx.x << 2));
tile[threadIdx.x * 4 ][dy] = v.x;
tile[threadIdx.x * 4 + 1][dy] = v.y;
tile[threadIdx.x * 4 + 2][dy] = v.z;
tile[threadIdx.x * 4 + 3][dy] = v.w;
}
__syncthreads();
for (int dy = threadIdx.y; dy < 32; dy += TY) {
int r = rb + dy;
int c0 = cb + (threadIdx.x << 2);
float4 v;
v.x = (c0 <= r) ? tile[dy][threadIdx.x * 4 ] : 0.f;
v.y = (c0 + 1 <= r) ? tile[dy][threadIdx.x * 4 + 1] : 0.f;
v.z = (c0 + 2 <= r) ? tile[dy][threadIdx.x * 4 + 2] : 0.f;
v.w = (c0 + 3 <= r) ? tile[dy][threadIdx.x * 4 + 3] : 0.f;
*reinterpret_cast<float4*>(L + base + (size_t)r * n + c0) = v;
}
}
torch::Tensor chol_tril_t(torch::Tensor W) {
auto L = torch::empty_like(W);
const int B = W.size(0);
const int n = W.size(2);
dim3 grid(n / 32, n / 32, B);
if (n == 4096) {
tril_transpose_kernel<16><<<grid, dim3(8, 16)>>>(
W.data_ptr<float>(), L.data_ptr<float>(), n);
} else {
tril_transpose_kernel<8><<<grid, dim3(8, 8)>>>(
W.data_ptr<float>(), L.data_ptr<float>(), n);
}
return L;
}
static cusolverDnHandle_t cusolver_handle() {
static cusolverDnHandle_t h = nullptr;
if (!h) {
cusolverDnCreate(&h);
cusolverDnSetMathMode(h, CUSOLVER_FP32_EMULATED_BF16X9_MATH);
}
return h;
}
torch::Tensor chol_cusolver_low(torch::Tensor A) {
auto W = A.clone();
const int B = A.size(0);
const int64_t n = A.size(2);
cusolverDnHandle_t h = cusolver_handle();
static cusolverDnParams_t p = nullptr;
if (!p) cusolverDnCreateParams(&p);
size_t dws = 0, hws = 0;
cusolverDnXpotrf_bufferSize(h, p, CUBLAS_FILL_MODE_LOWER, n,
CUDA_R_32F, W.data_ptr<float>(), n,
CUDA_R_32F, &dws, &hws);
auto ws = torch::empty({(int64_t)dws}, A.options().dtype(torch::kUInt8));
std::vector<char> hbuf(hws ? hws : 1);
auto info = torch::empty({1}, A.options().dtype(torch::kInt32));
for (int b = 0; b < B; ++b) {
cusolverDnXpotrf(h, p, CUBLAS_FILL_MODE_LOWER, n, CUDA_R_32F,
W.data_ptr<float>() + (size_t)b * n * n, n,
CUDA_R_32F, ws.data_ptr(), dws, hbuf.data(), hws,
info.data_ptr<int>());
}
return W;
}
static cublasHandle_t cublas_emu() {
static cublasHandle_t h = nullptr;
if (!h) {
cublasCreate(&h);
cublasSetMathMode(h, (cublasMath_t)CUBLAS_FP32_EMULATED_BF16X9_MATH);
cublasSetEmulationSpecialValuesSupport(
h, CUDA_EMULATION_SPECIAL_VALUES_SUPPORT_NONE);
}
return h;
}
static cublasHandle_t cublas_fp32() {
static cublasHandle_t h = nullptr;
if (!h) cublasCreate(&h);
return h;
}
static cublasHandle_t cublas_bf1() {
static cublasHandle_t h = nullptr;
if (!h) cublasCreate(&h);
return h;
}
static cublasHandle_t cublas_tf32() {
static cublasHandle_t h = nullptr;
if (!h) {
cublasCreate(&h);
cublasSetMathMode(h, CUBLAS_TF32_TENSOR_OP_MATH);
}
return h;
}
__global__ void fill_ptrs_kernel(float** out, float* base, size_t stride,
int B) {
int b = blockIdx.x * blockDim.x + threadIdx.x;
if (b < B) out[b] = base + (size_t)b * stride;
}
__global__ void zero_upper_kernel(float* L, int quads, int log2n) {
const int q = blockIdx.x * blockDim.x + threadIdx.x;
if (q >= quads) return;
const int n = 1 << log2n;
const size_t base = (size_t)blockIdx.y * n * (size_t)n;
const int off = q << 2;
const int r = off >> log2n;
const int c = off & (n - 1);
if (c > r) {
*reinterpret_cast<float4*>(L + base + off) =
make_float4(0.f, 0.f, 0.f, 0.f);
} else if (c + 3 > r) {
float* p = L + base + off;
if (c > r) p[0] = 0.f;
if (c + 1 > r) p[1] = 0.f;
if (c + 2 > r) p[2] = 0.f;
if (c + 3 > r) p[3] = 0.f;
}
}
template <int CTNB = NB, int MINB = 2, bool BF1 = false,
bool PACK = false>
static int coop_potrf_diag(float* diag, int n, int tsub,
torch::Tensor& tasks, torch::Tensor& state,
torch::Tensor& counter, int bstride,
__nv_bfloat16* Hout = nullptr,
int epoch = 0) {
static int blocks_per_sm = -1;
static int nsm = 0;
if (blocks_per_sm < 0) {
int dev = 0;
cudaGetDevice(&dev);
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, dev);
nsm = prop.multiProcessorCount;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&blocks_per_sm,
chol_persist_kernel<CTNB, MINB, BF1, false, PACK>, 256, 0);
if (blocks_per_sm < 1) blocks_per_sm = 1;
if (!prop.cooperativeLaunch) blocks_per_sm = 0;
}
if (blocks_per_sm == 0) return 1;
int ntasks = (int)tasks.size(0);
int pblocks = nsm * blocks_per_sm;
if (pblocks > ntasks) pblocks = ntasks;
if (pblocks < 2) pblocks = 2;
if (epoch == 0) {
cudaMemsetAsync(state.data_ptr<int>(), 0,
(size_t)state.numel() * sizeof(int));
cudaMemsetAsync(counter.data_ptr<int>(), 0, 2 * sizeof(int));
epoch = 1;
}
const float* Ap = diag;
float* Wp = diag;
const long long* tp = reinterpret_cast<const long long*>(
tasks.data_ptr<int64_t>());
unsigned long long* cp =
reinterpret_cast<unsigned long long*>(counter.data_ptr<int>());
int* sp = state.data_ptr<int>();
int n_ = n, t_ = tsub, nt_ = ntasks, bs_ = bstride;
void* args[] = {
&Ap, &Wp, &tp, &nt_, &cp, &sp, &n_, &t_, &bs_, &epoch, &Hout
};
cudaError_t err = cudaLaunchKernel(
(void*)chol_persist_kernel<CTNB, MINB, BF1, false, PACK>,
dim3(pblocks), dim3(256),
args, 0, 0);
if (err != cudaSuccess) {
(void)cudaGetLastError();
return 1;
}
return 0;
}
int chol_blocked(torch::Tensor A, torch::Tensor L, int64_t nbc,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter) {
const int B = A.size(0);
const int n = A.size(2);
const int nb = (int)nbc;
const int tc = n / nb;
const int tsub = nb / NB;
const size_t mstride = (size_t)n * n;
int log2n = 0;
while ((1 << log2n) < n) ++log2n;
torch::Tensor parr;
float** pa = nullptr;
float** pb = nullptr;
if (B > 1) {
parr = torch::empty({2 * B}, A.options().dtype(torch::kInt64));
pa = reinterpret_cast<float**>(parr.data_ptr<int64_t>());
pb = pa + B;
}
cudaMemcpy(L.data_ptr<float>(), A.data_ptr<float>(),
(size_t)B * mstride * sizeof(float),
cudaMemcpyDeviceToDevice);
float* W = L.data_ptr<float>();
const float one = 1.f, minus_one = -1.f;
const int pgrid = (B + 127) / 128;
for (int k = 0; k < tc; ++k) {
const int off = k * nb;
const int mfull = n - off;
float* diag = W + (size_t)off * n + off;
if (k > 0) {
if (B == 1) {
cublasSgemm(
cublas_emu(), CUBLAS_OP_T, CUBLAS_OP_N,
nb, mfull, off, &minus_one,
W + (size_t)off * n, n,
W + (size_t)off * n, n, &one, diag, n);
} else {
cublasSgemmStridedBatched(
cublas_emu(), CUBLAS_OP_T, CUBLAS_OP_N,
nb, mfull, off, &minus_one,
W + (size_t)off * n, n, mstride,
W + (size_t)off * n, n, mstride,
&one, diag, n, mstride, B);
}
}
if (coop_potrf_diag(diag, n, tsub, tasks, state, counter,
tsub * tsub))
return 1;
const int m2 = mfull - nb;
if (m2 > 0) {
float* below = W + (size_t)(off + nb) * n + off;
if (B == 1) {
cublasStrsm(cublas_fp32(), CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER, CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT, nb, m2, &one,
diag, n, below, n);
} else {
fill_ptrs_kernel<<<pgrid, 128>>>(pa, diag, mstride, B);
fill_ptrs_kernel<<<pgrid, 128>>>(
pb, below, mstride, B);
cublasStrsmBatched(cublas_fp32(), CUBLAS_SIDE_LEFT,
CUBLAS_FILL_MODE_UPPER, CUBLAS_OP_T,
CUBLAS_DIAG_NON_UNIT, nb, m2, &one,
(const float**)pa, n, pb, n, B);
}
}
}
const int quads = n * n / 4;
dim3 cgrid((quads + 255) / 256, B);
zero_upper_kernel<<<cgrid, 256>>>(L.data_ptr<float>(), quads, log2n);
return 0;
}
__global__ void copy_panel_kernel(const float* __restrict__ A,
float* __restrict__ W,
int B, int n, int r0, int c0,
int w, int rows) {
const int q = blockIdx.x * blockDim.x + threadIdx.x;
const int wq = w >> 2;
const int per_batch = rows * wq;
if (q >= B * per_batch) return;
const int b = q / per_batch;
const int qb = q - b * per_batch;
const size_t base = (size_t)b * n * n;
const size_t r = r0 + qb / wq;
const int c = c0 + ((qb % wq) << 2);
*reinterpret_cast<float4*>(W + base + r * n + c) =
*reinterpret_cast<const float4*>(A + base + r * n + c);
}
void copy_panel(torch::Tensor A, torch::Tensor W,
int64_t r0, int64_t c0, int64_t w) {
const int B = A.size(0);
const int n = A.size(-1);
const int rows = n - (int)r0;
if (rows <= 0) return;
const int total = B * rows * ((int)w >> 2);
copy_panel_kernel<<<(total + 255) / 256, 256>>>(
A.data_ptr<float>(), W.data_ptr<float>(), B, n,
(int)r0, (int)c0, (int)w, rows);
}
int chol_potrf_panel_bf1(torch::Tensor W, torch::Tensor hi,
int64_t offc, int64_t nbc,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t epochc) {
const int n = W.size(-1);
const int off = (int)offc;
const int nb = (int)nbc;
const int tsub = nb / NB;
const int trows = (n - off) / NB;
float* diag = W.data_ptr<float>() + (size_t)off * n + off;
__nv_bfloat16* Hout =
reinterpret_cast<__nv_bfloat16*>(hi.data_ptr()) +
(size_t)off * n + off;
return coop_potrf_diag<NB, 2, true, true>(
diag, n, tsub, tasks, state, counter, trows * tsub, Hout,
(int)epochc);
}
int chol_potrf_panel32_bf1_m3(torch::Tensor W, torch::Tensor hi,
int64_t offc, int64_t nbc,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t epochc) {
const int n = W.size(-1);
const int off = (int)offc;
const int nb = (int)nbc;
const int tsub = nb / 32;
const int trows = (n - off) / 32;
float* diag = W.data_ptr<float>() + (size_t)off * n + off;
__nv_bfloat16* Hout =
reinterpret_cast<__nv_bfloat16*>(hi.data_ptr()) +
(size_t)off * n + off;
return coop_potrf_diag<32, 3, true, true>(
diag, n, tsub, tasks, state, counter, trows * tsub, Hout,
(int)epochc);
}
int chol_potrf_panel32(torch::Tensor W, int64_t offc, int64_t nbc,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t minbc,
int64_t epochc) {
const int n = W.size(-1);
const int off = (int)offc;
const int nb = (int)nbc;
const int tsub = nb / 32;
const int trows = (n - off) / 32;
float* diag = W.data_ptr<float>() + (size_t)off * n + off;
if (minbc == 3)
return coop_potrf_diag<32, 3, false>(
diag, n, tsub, tasks, state, counter, trows * tsub, nullptr,
(int)epochc);
return coop_potrf_diag<32, 2, false>(
diag, n, tsub, tasks, state, counter, trows * tsub, nullptr,
(int)epochc);
}
void chol_gemm_tf32(torch::Tensor W, int64_t offc, int64_t nbc) {
const int B = W.size(0);
const int n = W.size(2);
const int off = (int)offc;
const int nb = (int)nbc;
if (off == 0) return;
const size_t mstride = (size_t)n * n;
const float one = 1.f, minus_one = -1.f;
float* Wp = W.data_ptr<float>();
float* diag = Wp + (size_t)off * n + off;
const int rows = n - off;
if (B == 1) {
cublasSgemm(cublas_tf32(), CUBLAS_OP_T, CUBLAS_OP_N,
nb, rows, off, &minus_one,
Wp + (size_t)off * n, n,
Wp + (size_t)off * n, n, &one, diag, n);
} else {
cublasSgemmStridedBatched(
cublas_tf32(), CUBLAS_OP_T, CUBLAS_OP_N,
nb, rows, off, &minus_one,
Wp + (size_t)off * n, n, mstride,
Wp + (size_t)off * n, n, mstride,
&one, diag, n, mstride, B);
}
}
void chol_gemm_bf1(torch::Tensor W, torch::Tensor hi,
int64_t offc, int64_t nbc) {
const int B = W.size(0);
const int n = W.size(2);
const int off = (int)offc;
const int nb = (int)nbc;
if (off == 0) return;
const int rows = n - off;
const __nv_bfloat16* H =
reinterpret_cast<const __nv_bfloat16*>(hi.data_ptr()) +
(size_t)off * n;
float* C = W.data_ptr<float>() + (size_t)off * n + off;
const float minus_one = -1.f, one = 1.f;
const long long stride = (long long)n * n;
cublasStatus_t st;
if (B == 1) {
st = cublasGemmEx(
cublas_bf1(), CUBLAS_OP_T, CUBLAS_OP_N,
nb, rows, off, &minus_one,
H, CUDA_R_16BF, n,
H, CUDA_R_16BF, n,
&one, C, CUDA_R_32F, n,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
} else {
st = cublasGemmStridedBatchedEx(
cublas_bf1(), CUBLAS_OP_T, CUBLAS_OP_N,
nb, rows, off, &minus_one,
H, CUDA_R_16BF, n, stride,
H, CUDA_R_16BF, n, stride,
&one, C, CUDA_R_32F, n, stride, B,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
}
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "BF16x1 panel GEMM failed");
}
__device__ __forceinline__ void tile_sm_bf3_32(
__nv_bfloat16* hi, __nv_bfloat16* lo, const float* src, int tid) {
constexpr int TNB = 32, LDT = 33, BDT = 40;
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int idx = tid + q * 256;
const int r = idx >> 5, c = idx & 31;
const float v = src[r * LDT + c];
const __nv_bfloat16 h = __float2bfloat16(v);
hi[r * BDT + c] = h;
lo[r * BDT + c] =
__float2bfloat16(v - __bfloat162float(h));
}
}
template <int T>
__global__ void __launch_bounds__(256) chol_res32_tc_kernel(
const float* __restrict__ A, float* __restrict__ L) {
constexpr int TN = 32, LDT = 33, BDT = 40;
constexpr int N = T * TN;
constexpr int NTILES = T * (T + 1) / 2;
constexpr int TS = TN * LDT;
constexpr int WSF = TN * BDT;
extern __shared__ __align__(16) float sm[];
float* workA = sm + NTILES * TS;
float* workB = workA + WSF;
__nv_bfloat16* ah = reinterpret_cast<__nv_bfloat16*>(workA);
__nv_bfloat16* al = ah + WSF;
__nv_bfloat16* bh = reinterpret_cast<__nv_bfloat16*>(workB);
__nv_bfloat16* bl = bh + WSF;
__shared__ __align__(16) float cbw[2][32];
const int tid = threadIdx.x;
const int tr = tid >> 4, tc = tid & 15;
const float* Ab = A + (size_t)blockIdx.x * N * N;
float* Lb = L + (size_t)blockIdx.x * N * N;
#define TILE32(i, j) \
(sm + ((i) * ((i) + 1) / 2 + (j)) * TS)
for (int i = 0; i < T; ++i)
for (int j = 0; j <= i; ++j)
load_tile<TN>(
TILE32(i, j),
Ab + (size_t)(i * TN) * N + j * TN, N, tid);
__syncthreads();
for (int j = 0; j < T; ++j) {
{
float acc[2][2] = {};
bf3_acc_t c0, c1;
nvcuda::wmma::fill_fragment(c0, 0.0f);
nvcuda::wmma::fill_fragment(c1, 0.0f);
for (int k = 0; k < j; ++k) {
tile_sm_bf3_32(
ah, al, TILE32(j, k), tid);
__syncthreads();
bf3_mma_slab<32>(
c0, c1, ah, al, ah, al);
__syncthreads();
}
bf3_fold<32>(
acc, c0, c1, workA, workB, tr, tc);
stage_acc_sm<32>(
TILE32(j, j), TILE32(j, j), acc, tr, tc);
__syncthreads();
potrf32_warp(TILE32(j, j), cbw, tid);
__syncthreads();
}
for (int i = j + 1; i < T; ++i) {
float acc[2][2] = {};
bf3_acc_t c0, c1;
nvcuda::wmma::fill_fragment(c0, 0.0f);
nvcuda::wmma::fill_fragment(c1, 0.0f);
for (int k = 0; k < j; ++k) {
tile_sm_bf3_32(
ah, al, TILE32(i, k), tid);
tile_sm_bf3_32(
bh, bl, TILE32(j, k), tid);
__syncthreads();
bf3_mma_slab<32>(
c0, c1, ah, al, bh, bl);
__syncthreads();
}
bf3_fold<32>(
acc, c0, c1, workA, workB, tr, tc);
stage_acc_sm<32>(
TILE32(i, j), TILE32(i, j), acc, tr, tc);
__syncthreads();
trsm32_warp(TILE32(j, j), TILE32(i, j), tid);
__syncthreads();
}
}
for (int i = 0; i < T; ++i) {
for (int j = 0; j < T; ++j) {
float* dst = Lb + (size_t)(i * TN) * N + j * TN;
if (j > i) {
for (int q = tid; q < TN * TN / 4; q += 256) {
const int r = q >> 3, c = (q & 7) << 2;
*reinterpret_cast<float4*>(
dst + (size_t)r * N + c) =
make_float4(0.f, 0.f, 0.f, 0.f);
}
} else {
const float* src = TILE32(i, j);
for (int q = tid; q < TN * TN / 4; q += 256) {
const int r = q >> 3, c = (q & 7) << 2;
float4 v = make_float4(
src[r * LDT + c],
src[r * LDT + c + 1],
src[r * LDT + c + 2],
src[r * LDT + c + 3]);
if (i == j) {
if (c > r) v.x = 0.f;
if (c + 1 > r) v.y = 0.f;
if (c + 2 > r) v.z = 0.f;
if (c + 3 > r) v.w = 0.f;
}
*reinterpret_cast<float4*>(
dst + (size_t)r * N + c) = v;
}
}
}
}
#undef TILE32
}
torch::Tensor chol_res32_tc(torch::Tensor A) {
auto L = torch::empty_like(A);
const int B = A.size(0);
const int n = A.size(2);
if (n == 128) {
constexpr int nf = (10 * 32 * 33 + 2 * 32 * 40);
constexpr size_t smem = nf * sizeof(float);
static bool attr = false;
if (!attr) {
cudaFuncSetAttribute(
chol_res32_tc_kernel<4>,
cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
attr = true;
}
chol_res32_tc_kernel<4><<<B, 256, smem>>>(
A.data_ptr<float>(), L.data_ptr<float>());
} else {
TORCH_CHECK(false, "chol_res32_tc: unsupported n");
}
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, cudaGetErrorString(err));
return L;
}
int chol_persist(torch::Tensor A, torch::Tensor L, torch::Tensor tasks,
torch::Tensor state, torch::Tensor counter,
int64_t modec, int64_t epochc) {
const int mode = (int)modec;
if (mode != 1 && mode != 2 && mode != 4 &&
mode != 6 && mode != 7) return 1;
const int n = A.size(2);
const int t = n / ((mode == 2 || mode == 4) ? 32 : NB);
void* kernel = (mode == 6) ? ((t <= 16)
? (void*)chol_persist_kernel<NB, 3, true>
: (void*)chol_persist_kernel<NB, 2, true>)
: (mode == 7) ? ((t <= 16)
? (void*)chol_persist_kernel<NB, 3, false, true>
: (void*)chol_persist_kernel<NB, 2, false, true>)
: (mode == 4) ? (void*)chol_persist_kernel<32, 2>
: (mode == 2) ? (void*)chol_persist_kernel<32, 3>
: (t <= 16) ? (void*)chol_persist_kernel<NB, 3>
: (void*)chol_persist_kernel<NB>;
const int kidx = (mode == 6) ? ((t <= 16) ? 4 : 5)
: (mode == 7) ? ((t <= 16) ? 6 : 7)
: (mode == 4) ? 3
: (mode == 2) ? 0 : (t <= 16) ? 1 : 2;
static int blocks_per_sm[8] = {
-1, -1, -1, -1, -1, -1, -1, -1
};
static int nsm = 0;
if (blocks_per_sm[kidx] < 0) {
int dev = 0;
cudaGetDevice(&dev);
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, dev);
nsm = prop.multiProcessorCount;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&blocks_per_sm[kidx], kernel, 256, 0);
if (blocks_per_sm[kidx] < 1) blocks_per_sm[kidx] = 1;
if (!prop.cooperativeLaunch) blocks_per_sm[kidx] = 0;
}
if (blocks_per_sm[kidx] == 0) return 1;
int ntasks = (int)tasks.size(0);
int blocks = nsm * blocks_per_sm[kidx];
if (blocks > ntasks) blocks = ntasks;
if (blocks < 2) blocks = 2;
const float* Ap = A.data_ptr<float>();
float* W = L.data_ptr<float>();
const long long* tp =
reinterpret_cast<const long long*>(tasks.data_ptr<int64_t>());
unsigned long long* cp =
reinterpret_cast<unsigned long long*>(counter.data_ptr<int>());
int* sp = state.data_ptr<int>();
int n_ = n, t_ = t, nt_ = ntasks, bs_ = t * t;
int epoch = (int)epochc;
__nv_bfloat16* Hout = nullptr;
void* args[] = {
&Ap, &W, &tp, &nt_, &cp, &sp, &n_, &t_, &bs_, &epoch, &Hout
};
cudaError_t err = cudaLaunchKernel(
kernel, dim3(blocks), dim3(256), args, 0, 0);
if (err != cudaSuccess) {
(void)cudaGetLastError();
return 1;
}
return 0;
}
"""
CPP_SRC = """
torch::Tensor chol_small(torch::Tensor A);
int chol_persist(torch::Tensor A, torch::Tensor L, torch::Tensor tasks,
torch::Tensor state, torch::Tensor counter, int64_t mode,
int64_t epoch);
torch::Tensor chol_cusolver_low(torch::Tensor A);
torch::Tensor chol_tril_t(torch::Tensor W);
torch::Tensor chol_res32_tc(torch::Tensor A);
int chol_blocked(torch::Tensor A, torch::Tensor L, int64_t nb,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter);
void copy_panel(torch::Tensor A, torch::Tensor W,
int64_t r0, int64_t c0, int64_t w);
int chol_potrf_panel_bf1(torch::Tensor W, torch::Tensor hi,
int64_t off, int64_t nb,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t epoch);
int chol_potrf_panel32_bf1_m3(torch::Tensor W, torch::Tensor hi,
int64_t off, int64_t nb,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t epoch);
int chol_potrf_panel32(torch::Tensor W, int64_t off, int64_t nb,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t minb,
int64_t epoch);
void chol_gemm_tf32(torch::Tensor W, int64_t off, int64_t nb);
void chol_gemm_bf1(torch::Tensor W, torch::Tensor hi,
int64_t off, int64_t nb);
int chol_ring(torch::Tensor A, torch::Tensor L, torch::Tensor tasks,
torch::Tensor state, torch::Tensor counter, int64_t epoch);
int chol_ring_panel_bf1(torch::Tensor W, torch::Tensor hi,
int64_t off, int64_t nb,
torch::Tensor tasks, torch::Tensor state,
torch::Tensor counter, int64_t minb,
int64_t epoch);
"""
_module = load_inline(
name="cholesky_kernels",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["chol_small", "chol_persist", "chol_cusolver_low",
"chol_tril_t", "chol_res32_tc", "chol_blocked",
"copy_panel", "chol_potrf_panel_bf1",
"chol_potrf_panel32_bf1_m3", "chol_potrf_panel32",
"chol_gemm_tf32", "chol_gemm_bf1", "chol_ring",
"chol_ring_panel_bf1"],
extra_cuda_cflags=["-O3", "--use_fast_math",
"-gencode=arch=compute_100a,code=sm_100a"],
extra_ldflags=["-lcusolver", "-lcublas"],
verbose=False,
)
_NB = 64
_PATH_B_SIZES = frozenset((256, 512, 1024, 2048, 4096, 8192))
def _path_b_mode(B: int, n: int) -> int:
t = n // 64
return 2 if B * t * t <= 4096 else 1
_task_cache = {}
_workspace_cache = {}
def _cached_epoch_workspace(key, nstate: int, device, epochs: int = 1):
ckey = (key, str(device))
entry = _workspace_cache.get(ckey)
if entry is None:
pad = nstate & 1
workspace = torch.zeros(
nstate + pad + 2, dtype=torch.int32, device=device
)
entry = [workspace[:nstate], workspace[nstate + pad:], 0]
_workspace_cache[ckey] = entry
start = entry[2] + 1
entry[2] += epochs
if entry[2] >= 0x7FFFFFFF:
entry[0].zero_()
entry[1].zero_()
start = 1
entry[2] = epochs
return entry[0], entry[1], start
def _tasks_for(B: int, t: int, device) -> torch.Tensor:
key = (B, t)
v = _task_cache.get(key)
if v is None:
import numpy as np
boff = np.arange(B, dtype=np.int64) << 30
segs = []
for j in range(t):
ii = np.arange(j, t, dtype=np.int64)
seg = (np.int64(1) << 52) | (j << 10) | ii
seg[0] = (j << 10) | j
segs.append((seg[:, None] | boff[None, :]).ravel())
v = torch.from_numpy(np.concatenate(segs)).to(device)
_task_cache[key] = v
return v
def _tasks_row(B: int, t: int, device) -> torch.Tensor:
key = ("row", B, t)
v = _task_cache.get(key)
if v is None:
import numpy as np
boff = np.arange(B, dtype=np.int64) << 30
segs = [(np.int64(0) << 52) | boff]
for j in range(t - 1):
ii = np.arange(j + 1, t, dtype=np.int64)
seg = (np.int64(1) << 52) | (j << 10) | ii
seg[0] = (np.int64(2) << 52) | (j << 10) | (j + 1)
segs.append((seg[:, None] | boff[None, :]).ravel())
v = torch.from_numpy(np.concatenate(segs)).to(device)
_task_cache[key] = v
return v
def _tasks_panels(B: int, t: int, device) -> torch.Tensor:
key = ("panels", B, t)
v = _task_cache.get(key)
if v is None:
import numpy as np
boff = np.arange(B, dtype=np.int64) << 30
segs = []
for j in range(t - 2):
ii = np.arange(j + 2, t, dtype=np.int64)
seg = (np.int64(1) << 52) | (j << 10) | ii
segs.append((seg[:, None] | boff[None, :]).ravel())
v = torch.from_numpy(np.concatenate(segs)).to(device)
_task_cache[key] = v
return v
def _tasks_panels_rect(
B: int, trows: int, tcols: int, device
) -> torch.Tensor:
key = ("panels_rect", B, trows, tcols)
v = _task_cache.get(key)
if v is None:
import numpy as np
boff = np.arange(B, dtype=np.int64) << 30
segs = []
for j in range(tcols):
ii = np.concatenate((
np.arange(j + 2, tcols, dtype=np.int64),
np.arange(tcols, trows, dtype=np.int64),
))
if ii.size:
seg = (np.int64(1) << 52) | (j << 10) | ii
segs.append((seg[:, None] | boff[None, :]).ravel())
v = torch.from_numpy(np.concatenate(segs)).to(device)
_task_cache[key] = v
return v
def _tasks_rect_row32(
B: int, trows: int, tcols: int, device
) -> torch.Tensor:
key = ("rect_row32", B, trows, tcols)
v = _task_cache.get(key)
if v is None:
import numpy as np
boff = np.arange(B, dtype=np.int64) << 30
segs = [boff.copy()]
for j in range(tcols):
ii = np.arange(j + 1, trows, dtype=np.int64)
if ii.size == 0:
continue
seg = (np.int64(1) << 52) | (j << 10) | ii
if j + 1 < tcols:
seg[0] = (
(np.int64(2) << 52) | (j << 10) | (j + 1)
)
segs.append((seg[:, None] | boff[None, :]).ravel())
v = torch.from_numpy(np.concatenate(segs)).to(device)
_task_cache[key] = v
return v
def _tasks_rect_diag_first(
B: int, trows: int, tcols: int, device, row32: bool
) -> torch.Tensor:
key = ("rect_diag_first", row32, B, trows, tcols)
v = _task_cache.get(key)
if v is None:
import numpy as np
boff = np.arange(B, dtype=np.int64) << 30
segs = []
if row32:
segs.append(boff.copy())
for j in range(tcols - 1):
ii = np.arange(j + 1, tcols, dtype=np.int64)
seg = (np.int64(1) << 52) | (j << 10) | ii
seg[0] = (
(np.int64(2) << 52) | (j << 10) | (j + 1)
)
segs.append((seg[:, None] | boff[None, :]).ravel())
else:
for j in range(tcols):
ii = np.arange(j, tcols, dtype=np.int64)
seg = (np.int64(1) << 52) | (j << 10) | ii
seg[0] = (j << 10) | j
segs.append((seg[:, None] | boff[None, :]).ravel())
if trows > tcols:
ii = np.arange(tcols, trows, dtype=np.int64)
for j in range(tcols):
seg = (np.int64(1) << 52) | (j << 10) | ii
segs.append((seg[:, None] | boff[None, :]).ravel())
v = torch.from_numpy(np.concatenate(segs)).to(device)
_task_cache[key] = v
return v
def _path_b(
A: torch.Tensor, fast_bf1: bool = False, fast_bf2: bool = False
):
B, n = A.size(0), A.size(-1)
mode = _path_b_mode(B, n)
if fast_bf1:
assert mode == 1
mode = 6
if fast_bf2:
assert mode == 1 and not fast_bf1
mode = 7
t = n // (
32 if mode == 2 else _NB
)
use_row = mode == 2 and B <= 16
if mode == 2 and use_row and t == 32:
tasks = _tasks_panels(B, t, A.device)
nstate = B * t * t
state, counter, epoch = _cached_epoch_workspace(
("ring", B, t), nstate, A.device
)
L = torch.empty_like(A)
if _module.chol_ring(
A, L, tasks, state, counter, epoch
) == 0:
return L
tasks = (_tasks_row if use_row else _tasks_for)(B, t, A.device)
kmode = (
4
if mode == 2 and use_row and t <= 64 and (B, n) != (2, 2048)
else mode
)
nstate = B * t * t
state, counter, epoch = _cached_epoch_workspace(
("persist", B, t), nstate, A.device,
2 if kmode in (6, 7) else 1
)
L = torch.empty_like(A)
if _module.chol_persist(
A, L, tasks, state, counter, kmode, epoch
) != 0:
if kmode in (6, 7):
state.zero_()
counter.zero_()
if kmode in (6, 7) and _module.chol_persist(
A, L, tasks, state, counter, 1, epoch + 1
) == 0:
return L
return None
return L
def _path_c_blocked_bf1(
A: torch.Tensor, nb: int = 512, panel32: bool = False,
tail32_at: int | None = None, diag_first: bool = False,
ring_panel: bool = False, ring_max_remaining: int | None = None,
ring_minb: int = 2
):
B = A.size(0)
n = A.size(-1)
panel_init = (B, n) == (8, 2048) or n >= 32768
W = torch.empty_like(A) if panel_init else A.clone()
hi = torch.empty_like(A, dtype=torch.bfloat16)
any_panel32 = panel32 or tail32_at is not None
state_unit = 32 if any_panel32 else _NB
max_tsub = max(nb // (32 if panel32 else _NB),
256 // 32 if tail32_at is not None else 0)
nstate = B * (n // state_unit) * max_tsub
off = 0
while off < n:
use_panel32 = panel32 or (
tail32_at is not None and off >= tail32_at
)
step_nb = (
256 if tail32_at is not None and off >= tail32_at else nb
)
step_nb = min(step_nb, n - off)
panel_unit = 32 if use_panel32 else _NB
tsub = step_nb // panel_unit
use_ring = (
ring_panel and use_panel32
and (
ring_max_remaining is None
or n - off <= ring_max_remaining
)
)
if panel_init:
_module.copy_panel(A, W, off, off, step_nb)
if off > 0:
_module.chol_gemm_bf1(W, hi, off, step_nb)
trows = (n - off) // panel_unit
tasks = (
_tasks_panels_rect(
B, trows, tsub, A.device
)
if use_ring else
_tasks_rect_diag_first(
B, trows, tsub, A.device, use_panel32
)
if diag_first else
_tasks_rect_row32(B, trows, tsub, A.device)
)
state, counter, epoch = _cached_epoch_workspace(
("blocked_bf1", B, n, state_unit, max_tsub),
nstate, A.device
)
if use_ring:
status = _module.chol_ring_panel_bf1(
W, hi, off, step_nb, tasks, state, counter, ring_minb,
epoch
)
elif use_panel32:
status = _module.chol_potrf_panel32_bf1_m3(
W, hi, off, step_nb, tasks, state, counter, epoch
)
else:
status = _module.chol_potrf_panel_bf1(
W, hi, off, step_nb, tasks, state, counter, epoch
)
if status != 0:
return None
off += step_nb
return W
def _path_c_blocked(A: torch.Tensor, nb: int):
B, t = A.size(0), nb // _NB
tasks = _tasks_for(B, t, A.device)
state = torch.empty(B * t * t, dtype=torch.int32, device=A.device)
counter = torch.empty(2, dtype=torch.int32, device=A.device)
L = torch.empty_like(A)
if _module.chol_blocked(A, L, nb, tasks, state, counter) != 0:
return None
return L
def _path_b_blocked_tf32_panel32(
A: torch.Tensor, nb: int, minb: int = 3
):
B, n = A.size(0), A.size(-1)
W = A.clone()
tsub = nb // 32
nstate = B * (n // 32) * tsub
off = 0
while off < n:
step_nb = min(nb, n - off)
if off:
_module.chol_gemm_tf32(W, off, step_nb)
tsub_step = step_nb // 32
tasks = _tasks_rect_row32(
B, (n - off) // 32, tsub_step, A.device
)
state, counter, epoch = _cached_epoch_workspace(
("blocked_tf32", B, n, tsub), nstate, A.device
)
if _module.chol_potrf_panel32(
W, off, step_nb, tasks, state, counter, minb, epoch
) != 0:
return None
off += step_nb
return W
def custom_kernel(data: input_t) -> output_t:
A = data
if (
A.is_cuda
and A.dtype == torch.float32
and A.dim() == 3
and A.is_contiguous()
):
n = A.size(-1)
if n in (32, 64):
return _module.chol_small(A)
if n == 128:
return _module.chol_res32_tc(A)
if n in _PATH_B_SIZES:
B = A.size(0)
use_blocked_bf1 = (B, n) in (
(8, 2048),
(2, 4096),
(1, 8192),
)
use_persist_bf1 = (B, n) == (60, 1024)
use_persist_bf2 = (B, n) == (640, 512)
L = (_path_b_blocked_tf32_panel32(A, 1024, 2)
if (B, n) == (1, 4096) else
_path_c_blocked_bf1(
A, panel32=True, ring_panel=True, ring_minb=3
)
if (B, n) == (1, 8192) else
_path_c_blocked_bf1(
A, panel32=True
)
if use_blocked_bf1 else _path_b(
A, use_persist_bf1, use_persist_bf2
))
if L is not None:
return L
if (
A.is_cuda
and A.dtype == torch.float32
and A.dim() == 3
and A.is_contiguous()
and A.size(0) <= 4
and A.size(-1) >= 2048
):
if A.size(-1) >= 16384:
if A.size(0) == 1:
L = (
_path_c_blocked_bf1(
A, panel32=False,
tail32_at=20480, diag_first=True,
ring_panel=True,
ring_max_remaining=4096, ring_minb=2
)
if A.size(-1) >= 32768 else
_path_c_blocked_bf1(
A, panel32=True, diag_first=True,
ring_panel=True, ring_max_remaining=4096,
ring_minb=2
)
)
if L is not None:
return L
L = _path_c_blocked(A, 256)
if L is not None:
return L
W = _module.chol_cusolver_low(A)
return _module.chol_tril_t(W)
return torch.linalg.cholesky_ex(A, check_errors=False).L
scrolls · 2361 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