submission 927797
aryanmagoon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3005 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-927797?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:86637701e34211becc035b72759a7a0171e3f379607dfbcdd3d7b15bfae9e78b
license declaredunknown
license concludedunknown
authorsaryanmagoon
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::mma
wmma::fragment<wmma::accumulator, 16, 16, 16, float> acc[AC][2];shared-memory
__shared__ __align__(16) float sm[MPB * (N * LDP + 128)];vector-width = float4
float4 v = *(const float4 *)(bcr + j4);Kernel source
submission.py3005 lines
from __future__ import annotations
import ctypes
import ctypes.util
import os
import shutil
import threading
from functools import lru_cache
import torch
try:
from task import input_t, output_t
except ImportError:
input_t = torch.Tensor
output_t = torch.Tensor
_SMALL_CUDA_SOURCE = r"""
typedef unsigned long long uintptr_t;
#include <cuda_runtime.h>
#include <mma.h>
#define FULLM 0xffffffffu
template <int N, int MPB>
__global__ __launch_bounds__((N / 32) * 32 * MPB) void chol_reg_kernel(
const float *__restrict__ A, float *__restrict__ L, int batch) {
constexpr int NW = N / 32;
constexpr int LDP = (NW == 1) ? 33 : 36;
__shared__ __align__(16) float sm[MPB * (N * LDP + 128)];
const int tid = threadIdx.x;
const int lane = tid & 31;
const int wid = (tid >> 5);
const int slot = wid / NW;
const int w = wid - slot * NW;
const int mid = blockIdx.x * MPB + slot;
const bool live = (mid < batch);
float *P = sm + slot * (N * LDP + 128);
float *bc = P + N * LDP;
float *dinv = bc + 96;
const size_t off = (size_t)(live ? mid : 0) * N * N;
const float *Ab = A + off;
float *Lb = L + off;
const int R = 32 * w + lane;
float r[N];
#pragma unroll
for (int j = 0; j < N; ++j) r[j] = 0.f;
#pragma unroll
for (int j = 0; j < N; ++j)
if (j <= R) r[j] = Ab[(size_t)j * N + R];
#pragma unroll
for (int kb = 0; kb < NW; ++kb) {
const int c0 = 32 * kb;
if (w == kb) {
bc[lane] = r[c0];
#pragma unroll
for (int k = 0; k < 32; ++k) {
__syncwarp();
const float *bcr = bc + (k % 3) * 32;
float *bcw = bc + ((k + 1) % 3) * 32;
const float rk = r[c0 + k];
float akk = __shfl_sync(FULLM, rk, k);
float d = rsqrtf(akk);
float f = rk * (d * d);
if (k < 31)
r[c0 + ((k + 1) & 31)] = fmaf(-f, bcr[k + 1], r[c0 + ((k + 1) & 31)]);
if (k < 30)
r[c0 + ((k + 2) & 31)] = fmaf(-f, bcr[k + 2], r[c0 + ((k + 2) & 31)]);
if (k < 31) bcw[lane] = r[c0 + ((k + 1) & 31)];
const int j0 = (k + 3) & ~3;
#pragma unroll
for (int j4 = j0; j4 < 32; j4 += 4) {
float4 v = *(const float4 *)(bcr + j4);
if (j4 + 0 > k + 2) r[c0 + j4 + 0] = fmaf(-f, v.x, r[c0 + j4 + 0]);
if (j4 + 1 > k + 2) r[c0 + j4 + 1] = fmaf(-f, v.y, r[c0 + j4 + 1]);
if (j4 + 2 > k + 2) r[c0 + j4 + 2] = fmaf(-f, v.z, r[c0 + j4 + 2]);
if (j4 + 3 > k + 2) r[c0 + j4 + 3] = fmaf(-f, v.w, r[c0 + j4 + 3]);
}
r[c0 + k] = (lane >= k) ? rk * d : 0.f;
}
if (NW > 1) {
#pragma unroll
for (int t = 0; t < 32; ++t) P[(c0 + lane) * LDP + t] = r[c0 + t];
dinv[lane] = 1.f / r[c0 + lane];
}
}
__syncthreads();
if (NW > 1 && w > kb) {
#pragma unroll
for (int t = 0; t < 32; ++t) {
float x = r[c0 + t];
const float *Lr = P + (size_t)(c0 + t) * LDP;
#pragma unroll
for (int u4 = 0; u4 < ((t + 3) & ~3); u4 += 4) {
float4 l4 = *(const float4 *)(Lr + u4);
if (u4 + 0 < t) x -= r[c0 + u4 + 0] * l4.x;
if (u4 + 1 < t) x -= r[c0 + u4 + 1] * l4.y;
if (u4 + 2 < t) x -= r[c0 + u4 + 2] * l4.z;
if (u4 + 3 < t) x -= r[c0 + u4 + 3] * l4.w;
}
r[c0 + t] = x * dinv[t];
}
#pragma unroll
for (int t = 0; t < 32; ++t) P[(size_t)R * LDP + t] = r[c0 + t];
}
__syncthreads();
if (NW > 1 && w > kb) {
#pragma unroll
for (int j = c0 + 32; j < N; ++j) {
const float *Pj = P + (size_t)j * LDP;
float acc = 0.f;
#pragma unroll
for (int t4 = 0; t4 < 32; t4 += 4) {
float4 p = *(const float4 *)(Pj + t4);
acc += r[c0 + t4 + 0] * p.x;
acc += r[c0 + t4 + 1] * p.y;
acc += r[c0 + t4 + 2] * p.z;
acc += r[c0 + t4 + 3] * p.w;
}
r[j] -= acc;
}
}
__syncthreads();
}
const int NT = (N / 32) * 32 * MPB;
#pragma unroll
for (int cb = 0; cb < NW; ++cb) {
const int c0 = 32 * cb;
__syncthreads();
#pragma unroll
for (int t = 0; t < 32; ++t)
P[(size_t)R * LDP + t] = (c0 + t <= R) ? r[c0 + t] : 0.f;
__syncthreads();
if (live) {
for (int e = tid - slot * (NT / MPB); e < N * 32; e += NT / MPB) {
int row = e >> 5, col = e & 31;
Lb[(size_t)row * N + c0 + col] = P[(size_t)row * LDP + col];
}
}
}
}
"""
_MAIN_CUDA_SOURCE = r"""
typedef unsigned long long uintptr_t;
#include <cuda_runtime.h>
#include <mma.h>
#define FULLM 0xffffffffu
using namespace nvcuda;
#define GK_LD 40
#define GK_KC 32
#define GI_LD 80
#define GI_KC 64
template <int TILE, int KC, int NT>
__device__ __forceinline__ void ld_glb(float4 *reg, const float *src,
int ldsrc, int rows, int kc, int K,
int tid, bool fast) {
constexpr int TPR = KC / 4;
constexpr int RPI = NT / TPR;
constexpr int ITER = TILE / RPI;
const int c4 = (tid & (TPR - 1)) * 4;
const int rb = tid / TPR;
if (fast) {
#pragma unroll
for (int it = 0; it < ITER; ++it) {
const int row = it * RPI + rb;
reg[it] = (row < rows)
? *(const float4 *)(src + (size_t)row * ldsrc + kc + c4)
: make_float4(0.f, 0.f, 0.f, 0.f);
}
return;
}
#pragma unroll
for (int it = 0; it < ITER; ++it) {
const int row = it * RPI + rb;
float4 v = make_float4(0.f, 0.f, 0.f, 0.f);
const float *pv = src + (size_t)row * ldsrc + kc + c4;
if (row < rows && kc + c4 + 3 < K && (((uintptr_t)pv & 15) == 0)) {
v = *(const float4 *)pv;
} else if (row < rows && kc + c4 < K) {
const float *p = src + (size_t)row * ldsrc + kc;
float t[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int q = 0; q < 4; ++q)
if (kc + c4 + q < K) t[q] = p[c4 + q];
v = make_float4(t[0], t[1], t[2], t[3]);
}
reg[it] = v;
}
}
struct H4 {
__half2 a, b;
};
template <int HSH>
struct HShadow;
template <>
struct HShadow<0> {};
template <>
struct HShadow<1> {
__half *p;
int ldp;
size_t stride;
};
struct __align__(16) H8 {
__half2 a, b, c, d;
};
__device__ __forceinline__ void st_half_frag(__half *dst, int ldp,
const float *eb, int r0, int c0,
int nrow, int ncol, int lane) {
const int rr = lane >> 1, cc = (lane & 1) * 8;
if (r0 + rr >= nrow) return;
__half *d = dst + (size_t)(r0 + rr) * ldp + c0 + cc;
const float *s = eb + rr * 20 + cc;
if (c0 + cc + 8 <= ncol && ((((uintptr_t)d) & 15) == 0)) {
const float4 v0 = *(const float4 *)s;
const float4 v1 = *(const float4 *)(s + 4);
H8 h;
h.a = __floats2half2_rn(v0.x, v0.y);
h.b = __floats2half2_rn(v0.z, v0.w);
h.c = __floats2half2_rn(v1.x, v1.y);
h.d = __floats2half2_rn(v1.z, v1.w);
*(H8 *)d = h;
} else {
for (int q = 0; q < 8; ++q)
if (c0 + cc + q < ncol) d[q] = __float2half(s[q]);
}
}
template <int ITER>
__device__ __forceinline__ void cvt_h4(H4 *reg, const float4 *v, float sgn) {
#pragma unroll
for (int it = 0; it < ITER; ++it) {
reg[it].a = __floats2half2_rn(sgn * v[it].x, sgn * v[it].y);
reg[it].b = __floats2half2_rn(sgn * v[it].z, sgn * v[it].w);
}
}
template <int TILE, int KC, int LD, int NT>
__device__ __forceinline__ void st_smem(__half *dst, const H4 *reg, int tid) {
constexpr int TPR = KC / 4;
constexpr int RPI = NT / TPR;
constexpr int ITER = TILE / RPI;
const int c4 = (tid & (TPR - 1)) * 4;
const int rb = tid / TPR;
#pragma unroll
for (int it = 0; it < ITER; ++it) {
const int row = it * RPI + rb;
*(__half2 *)(dst + row * LD + c4) = reg[it].a;
*(__half2 *)(dst + row * LD + c4 + 2) = reg[it].b;
}
}
template <int MODE, int TM, int TN, int HSH>
__device__ __forceinline__ void gemm_nt_body(
float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
size_t strideA, size_t strideB, int lower, HShadow<HSH> hs) {
constexpr int NWN = TN / 32;
constexpr int NW = 2 * NWN;
constexpr int AC = TM / 32;
constexpr int NT = 32 * NW;
constexpr int KC = (TM <= 64) ? 64 : 32;
constexpr int LD = KC + 8;
constexpr int ITERA = TM / (NT / (KC / 4));
constexpr int ITERB = TN / (NT / (KC / 4));
__shared__ __align__(16) __half As[2 * TM * LD];
__shared__ __align__(16) __half Bs[2 * TN * LD];
__shared__ __align__(16) float Es[NW * 16 * 20];
const int i0 = blockIdx.x * TM, j0 = blockIdx.y * TN;
if (i0 >= nrow || j0 >= ncol) return;
if (lower && (roff + i0 + TM - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
return;
const size_t bidx = blockIdx.z;
float *C = Cb + bidx * strideC;
const float *Cin = Cinb + bidx * strideC;
const float *A = Ab + bidx * strideA;
const float *B = Bb + bidx * strideB;
const int tid = threadIdx.x;
const int warp = tid >> 5, lane = tid & 31;
const int wm = warp / NWN, wn = warp % NWN;
float *eb = Es + warp * (16 * 20);
wmma::fragment<wmma::accumulator, 16, 16, 16, float> acc[AC][2];
const bool ldok = ((ldc & 3) == 0) && ((((uintptr_t)C) & 15) == 0) &&
((((uintptr_t)Cin) & 15) == 0);
#pragma unroll
for (int a = 0; a < AC; ++a)
#pragma unroll
for (int b = 0; b < 2; ++b) {
const int r0 = i0 + wm * (TM / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
if (MODE == 0) {
wmma::fill_fragment(acc[a][b], 0.f);
} else if (full) {
wmma::load_matrix_sync(acc[a][b], Cin + (size_t)r0 * ldc + c0, ldc,
wmma::mem_row_major);
} else {
for (int e = lane; e < 256; e += 32) {
int rr = e >> 4, cc = e & 15;
float v = 0.f;
if (r0 + rr < nrow && c0 + cc < ncol)
v = Cin[(size_t)(r0 + rr) * ldc + c0 + cc];
eb[rr * 20 + cc] = v;
}
__syncwarp();
wmma::load_matrix_sync(acc[a][b], eb, 20, wmma::mem_row_major);
__syncwarp();
}
}
const int mrows = min(TM, nrow - i0);
const int nrows = min(TN, ncol - j0);
const float sgn = (MODE == 1) ? -1.f : 1.f;
const float *Ap = A + (size_t)i0 * lda;
const float *Bp = B + (size_t)j0 * ldb;
const bool fastA = ((K % KC) == 0) && ((lda & 3) == 0) &&
((((uintptr_t)Ap) & 15) == 0);
const bool fastB = ((K % KC) == 0) && ((ldb & 3) == 0) &&
((((uintptr_t)Bp) & 15) == 0);
float4 ra[ITERA], rb[ITERB];
H4 ha[ITERA], hb[ITERB];
ld_glb<TM, KC, NT>(ra, Ap, lda, mrows, 0, K, tid, fastA);
ld_glb<TN, KC, NT>(rb, Bp, ldb, nrows, 0, K, tid, fastB);
cvt_h4<ITERA>(ha, ra, sgn);
cvt_h4<ITERB>(hb, rb, 1.f);
st_smem<TM, KC, LD, NT>(As, ha, tid);
st_smem<TN, KC, LD, NT>(Bs, hb, tid);
__syncthreads();
int cur = 0;
constexpr int STGA = TM * LD;
constexpr int STGB = TN * LD;
for (int kc = 0; kc < K; kc += KC) {
const __half *Ac = As + cur * STGA;
const __half *Bc = Bs + cur * STGB;
if (kc + KC < K) {
ld_glb<TM, KC, NT>(ra, Ap, lda, mrows, kc + KC, K, tid, fastA);
ld_glb<TN, KC, NT>(rb, Bp, ldb, nrows, kc + KC, K, tid, fastB);
}
#pragma unroll
for (int kk = 0; kk < KC; kk += 16) {
wmma::fragment<wmma::matrix_b, 16, 16, 16, __half, wmma::col_major> bf[2];
#pragma unroll
for (int b = 0; b < 2; ++b)
wmma::load_matrix_sync(bf[b], Bc + (wn * 32 + b * 16) * LD + kk, LD);
#pragma unroll
for (int a = 0; a < AC; ++a) {
wmma::fragment<wmma::matrix_a, 16, 16, 16, __half, wmma::row_major> af;
wmma::load_matrix_sync(af, Ac + (wm * (TM / 2) + a * 16) * LD + kk,
LD);
#pragma unroll
for (int b = 0; b < 2; ++b)
wmma::mma_sync(acc[a][b], af, bf[b], acc[a][b]);
}
}
if (kc + KC < K) {
cvt_h4<ITERA>(ha, ra, sgn);
cvt_h4<ITERB>(hb, rb, 1.f);
st_smem<TM, KC, LD, NT>(As + (cur ^ 1) * STGA, ha, tid);
st_smem<TN, KC, LD, NT>(Bs + (cur ^ 1) * STGB, hb, tid);
cur ^= 1;
}
__syncthreads();
}
#pragma unroll
for (int a = 0; a < AC; ++a)
#pragma unroll
for (int b = 0; b < 2; ++b) {
const int r0 = i0 + wm * (TM / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
if (r0 >= nrow || c0 >= ncol) continue;
if (lower && (roff + r0 + 15) < (coff + c0) &&
((r0 >> 7) != (c0 >> 7)))
continue;
const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
if (full) {
wmma::store_matrix_sync(C + (size_t)r0 * ldc + c0, acc[a][b], ldc,
wmma::mem_row_major);
if constexpr (HSH == 1) {
__syncwarp();
wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
__syncwarp();
st_half_frag(hs.p + bidx * hs.stride, hs.ldp, eb, r0, c0, nrow, ncol,
lane);
}
} else {
__syncwarp();
wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
__syncwarp();
for (int e = lane; e < 256; e += 32) {
int rr = e >> 4, cc = e & 15;
if (r0 + rr < nrow && c0 + cc < ncol)
C[(size_t)(r0 + rr) * ldc + c0 + cc] = eb[rr * 20 + cc];
}
if constexpr (HSH == 1)
st_half_frag(hs.p + bidx * hs.stride, hs.ldp, eb, r0, c0, nrow, ncol,
lane);
}
}
}
template <int MODE, int TM, int TN>
__global__ __launch_bounds__(2 * TN, (TM == TN) ? 2 : 1) void gemm_nt_kernel(
float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
size_t strideA, size_t strideB, int lower) {
gemm_nt_body<MODE, TM, TN, 0>(Cb, Cinb, ldc, Ab, lda, Bb, ldb, K, nrow, ncol,
roff, coff, strideC, strideA, strideB, lower,
HShadow<0>());
}
template <int MODE, int TM, int TN>
__global__ __launch_bounds__(2 * TN, (TM == TN) ? 2 : 1) void gemm_nt_hs_kernel(
float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
size_t strideA, size_t strideB, int lower, __half *__restrict__ php,
int ldp, size_t strideP) {
HShadow<1> hs;
hs.p = php;
hs.ldp = ldp;
hs.stride = strideP;
gemm_nt_body<MODE, TM, TN, 1>(Cb, Cinb, ldc, Ab, lda, Bb, ldb, K, nrow, ncol,
roff, coff, strideC, strideA, strideB, lower,
hs);
}
__global__ __launch_bounds__(256) void cvt_panel_kernel(
const float *__restrict__ Lsrc, int n, __half *__restrict__ Ph, int ldp,
int nbo, int R0, int jb, size_t stride) {
const int row = R0 + blockIdx.x;
const size_t b = blockIdx.y;
const float *S = Lsrc + b * stride + (size_t)row * n + jb;
__half *D = Ph + b * (size_t)n * ldp + (size_t)row * ldp;
const bool v4 = ((((uintptr_t)S) & 15) == 0);
for (int c = threadIdx.x * 4; c < nbo; c += 256 * 4) {
if (v4 && c + 3 < nbo) {
float4 v = *(const float4 *)(S + c);
*(__half2 *)(D + c) = __floats2half2_rn(v.x, v.y);
*(__half2 *)(D + c + 2) = __floats2half2_rn(v.z, v.w);
} else {
for (int q = c; q < nbo && q < c + 4; ++q) D[q] = __float2half(S[q]);
}
}
}
template <int RSTEP>
__device__ __forceinline__ void cp_h(__half *dst, const __half *src, int lds,
int rows, int kc, int tid) {
const int k8 = (tid & 3) * 8;
#pragma unroll
for (int it = 0; it < 2; ++it) {
const int row = it * RSTEP + (tid >> 2);
const unsigned sa =
(unsigned)__cvta_generic_to_shared(dst + row * GK_LD + k8);
const __half *gp = src + (size_t)(row < rows ? row : 0) * lds + kc + k8;
const int bytes = (row < rows) ? 16 : 0;
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::
"r"(sa), "l"(gp), "r"(bytes));
}
}
template <int TILE>
__global__ __launch_bounds__(TILE * 2) void gemm_h_kernel(
float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
const __half *__restrict__ Ab, int lda, const __half *__restrict__ Bb,
int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
size_t strideAB, int lower) {
constexpr int NW = TILE / 16;
constexpr int NWN = TILE / 32;
constexpr int AC = TILE / 32;
constexpr int HRS = TILE / 2;
constexpr int STG = TILE * GK_LD;
constexpr int NSTG = 4;
__shared__ __align__(16) __half As[NSTG * TILE * GK_LD];
__shared__ __align__(16) __half Bs[NSTG * TILE * GK_LD];
__shared__ __align__(16) float Es[NW * 16 * 20];
const int i0 = blockIdx.x * TILE, j0 = blockIdx.y * TILE;
if (i0 >= nrow || j0 >= ncol) return;
if (lower && (roff + i0 + TILE - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
return;
const size_t bidx = blockIdx.z;
float *C = Cb + bidx * strideC;
const float *Cin = Cinb + bidx * strideC;
const __half *A = Ab + bidx * strideAB;
const __half *B = Bb + bidx * strideAB;
const int tid = threadIdx.x;
const int warp = tid >> 5, lane = tid & 31;
const int wm = warp / NWN, wn = warp % NWN;
float *eb = Es + warp * (16 * 20);
wmma::fragment<wmma::accumulator, 16, 16, 16, float> acc[AC][2];
const bool ldok = ((ldc & 3) == 0) && ((((uintptr_t)C) & 15) == 0) &&
((((uintptr_t)Cin) & 15) == 0);
#pragma unroll
for (int a = 0; a < AC; ++a)
#pragma unroll
for (int b = 0; b < 2; ++b) {
const int r0 = i0 + wm * (TILE / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
if (full) {
wmma::load_matrix_sync(acc[a][b], Cin + (size_t)r0 * ldc + c0, ldc,
wmma::mem_row_major);
} else {
for (int e = lane; e < 256; e += 32) {
int rr = e >> 4, cc = e & 15;
float v = 0.f;
if (r0 + rr < nrow && c0 + cc < ncol)
v = Cin[(size_t)(r0 + rr) * ldc + c0 + cc];
eb[rr * 20 + cc] = v;
}
__syncwarp();
wmma::load_matrix_sync(acc[a][b], eb, 20, wmma::mem_row_major);
__syncwarp();
}
#pragma unroll
for (int e = 0; e < acc[a][b].num_elements; ++e)
acc[a][b].x[e] = -acc[a][b].x[e];
}
const int mrows = min(TILE, nrow - i0);
const int nrows = min(TILE, ncol - j0);
const __half *Ap = A + (size_t)i0 * lda;
const __half *Bp = B + (size_t)j0 * ldb;
#pragma unroll
for (int s = 0; s < NSTG - 1; ++s) {
if (s * GK_KC < K) {
cp_h<HRS>(As + s * STG, Ap, lda, mrows, s * GK_KC, tid);
cp_h<HRS>(Bs + s * STG, Bp, ldb, nrows, s * GK_KC, tid);
}
asm volatile("cp.async.commit_group;\n" ::);
}
int cur = 0;
for (int kc = 0; kc < K; kc += GK_KC) {
asm volatile("cp.async.wait_group %0;\n" ::"n"(NSTG - 2));
__syncthreads();
const __half *Ac = As + cur * STG;
const __half *Bc = Bs + cur * STG;
{
const int nk = kc + (NSTG - 1) * GK_KC;
const int nxt = (cur + NSTG - 1) % NSTG;
if (nk < K) {
cp_h<HRS>(As + nxt * STG, Ap, lda, mrows, nk, tid);
cp_h<HRS>(Bs + nxt * STG, Bp, ldb, nrows, nk, tid);
}
asm volatile("cp.async.commit_group;\n" ::);
}
#pragma unroll
for (int kk = 0; kk < GK_KC; kk += 16) {
wmma::fragment<wmma::matrix_b, 16, 16, 16, __half, wmma::col_major> bf[2];
#pragma unroll
for (int b = 0; b < 2; ++b)
wmma::load_matrix_sync(bf[b], Bc + (wn * 32 + b * 16) * GK_LD + kk,
GK_LD);
#pragma unroll
for (int a = 0; a < AC; ++a) {
wmma::fragment<wmma::matrix_a, 16, 16, 16, __half, wmma::row_major> af;
wmma::load_matrix_sync(af, Ac + (wm * (TILE / 2) + a * 16) * GK_LD + kk,
GK_LD);
#pragma unroll
for (int b = 0; b < 2; ++b)
wmma::mma_sync(acc[a][b], af, bf[b], acc[a][b]);
}
}
cur = (cur + 1) % NSTG;
}
#pragma unroll
for (int a = 0; a < AC; ++a)
#pragma unroll
for (int b = 0; b < 2; ++b) {
const int r0 = i0 + wm * (TILE / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
if (r0 >= nrow || c0 >= ncol) continue;
if (lower && (roff + r0 + 15) < (coff + c0) &&
((r0 >> 7) != (c0 >> 7)))
continue;
#pragma unroll
for (int e = 0; e < acc[a][b].num_elements; ++e)
acc[a][b].x[e] = -acc[a][b].x[e];
const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
if (full) {
wmma::store_matrix_sync(C + (size_t)r0 * ldc + c0, acc[a][b], ldc,
wmma::mem_row_major);
} else {
__syncwarp();
wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
__syncwarp();
for (int e = lane; e < 256; e += 32) {
int rr = e >> 4, cc = e & 15;
if (r0 + rr < nrow && c0 + cc < ncol)
C[(size_t)(r0 + rr) * ldc + c0 + cc] = eb[rr * 20 + cc];
}
}
}
}
__global__ __launch_bounds__(256) void cvt_panel_i8_kernel(
const float *__restrict__ Lsrc, int n, signed char *__restrict__ Pq,
int ldp, int nbo, int R0, int jb, size_t stride,
float *__restrict__ Scl) {
const int row = R0 + blockIdx.x;
const size_t b = blockIdx.y;
const float *S = Lsrc + b * stride + (size_t)row * n + jb;
signed char *D = Pq + b * (size_t)n * ldp + (size_t)row * ldp;
const int tid = threadIdx.x;
float4 v[2];
const bool v4 = ((((uintptr_t)S) & 15) == 0);
#pragma unroll
for (int it = 0; it < 2; ++it) {
const int c = it * 1024 + tid * 4;
float4 x = make_float4(0.f, 0.f, 0.f, 0.f);
if (v4 && c + 3 < nbo) {
x = *(const float4 *)(S + c);
} else if (c < nbo) {
float t[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int q = 0; q < 4; ++q)
if (c + q < nbo) t[q] = S[c + q];
x = make_float4(t[0], t[1], t[2], t[3]);
}
v[it] = x;
}
float mx = 0.f;
#pragma unroll
for (int it = 0; it < 2; ++it) {
mx = fmaxf(mx, fabsf(v[it].x));
mx = fmaxf(mx, fabsf(v[it].y));
mx = fmaxf(mx, fabsf(v[it].z));
mx = fmaxf(mx, fabsf(v[it].w));
}
#pragma unroll
for (int o = 16; o; o >>= 1) mx = fmaxf(mx, __shfl_xor_sync(FULLM, mx, o));
__shared__ float wmx[8];
if ((tid & 31) == 0) wmx[tid >> 5] = mx;
__syncthreads();
mx = wmx[tid & 7];
#pragma unroll
for (int o = 4; o; o >>= 1) mx = fmaxf(mx, __shfl_xor_sync(FULLM, mx, o));
float inv = 127.f / mx;
if (!(mx > 0.f) || !isfinite(mx) || !isfinite(inv)) {
mx = 0.f;
inv = 0.f;
}
if (tid == 0) Scl[b * (size_t)n + row] = mx * (1.f / 127.f);
const int kpad = (nbo + 63) & ~63;
#pragma unroll
for (int it = 0; it < 2; ++it) {
const int c = it * 1024 + tid * 4;
if (c >= kpad) continue;
const float f[4] = {v[it].x, v[it].y, v[it].z, v[it].w};
signed char r[4];
#pragma unroll
for (int t = 0; t < 4; ++t) {
int qi = (c + t < nbo) ? __float2int_rn(f[t] * inv) : 0;
r[t] = (signed char)max(-127, min(127, qi));
}
if (c + 3 < kpad) {
char4 q;
q.x = r[0]; q.y = r[1]; q.z = r[2]; q.w = r[3];
*(char4 *)(D + c) = q;
} else {
#pragma unroll
for (int t = 0; t < 4; ++t)
if (c + t < kpad) D[c + t] = r[t];
}
}
}
template <int RSTEP>
__device__ __forceinline__ void cp_i8(signed char *dst, const signed char *src,
int lds, int rows, int kc, int tid) {
const int k16 = (tid & 3) * 16;
#pragma unroll
for (int it = 0; it < 2; ++it) {
const int row = it * RSTEP + (tid >> 2);
const unsigned sa =
(unsigned)__cvta_generic_to_shared(dst + row * GI_LD + k16);
const signed char *gp =
src + (size_t)(row < rows ? row : 0) * lds + kc + k16;
const int bytes = (row < rows) ? 16 : 0;
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::
"r"(sa), "l"(gp), "r"(bytes));
}
}
template <int TILE>
__global__ __launch_bounds__(TILE * 2, (TILE == 128) ? 2 : 1) void gemm_i8_kernel(
float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
const signed char *__restrict__ Ab, int lda,
const signed char *__restrict__ Bb, int ldb,
const float *__restrict__ Scb, int K, int nrow, int ncol, int roff,
int coff, size_t strideC, size_t strideAB, size_t strideS, int lower) {
constexpr int NW = TILE / 16;
constexpr int NWN = TILE / 32;
constexpr int AC = TILE / 32;
constexpr int HRS = TILE / 2;
constexpr int STG = TILE * GI_LD;
constexpr int NSTG = 3;
__shared__ __align__(16) signed char As[NSTG * TILE * GI_LD];
__shared__ __align__(16) signed char Bs[NSTG * TILE * GI_LD];
__shared__ __align__(16) int Es[NW * 16 * 20];
__shared__ float SA[TILE], SB[TILE];
const int i0 = blockIdx.x * TILE, j0 = blockIdx.y * TILE;
if (i0 >= nrow || j0 >= ncol) return;
if (lower && (roff + i0 + TILE - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
return;
const size_t bidx = blockIdx.z;
float *C = Cb + bidx * strideC;
const float *Cin = Cinb + bidx * strideC;
const signed char *A = Ab + bidx * strideAB;
const signed char *B = Bb + bidx * strideAB;
const float *Sc = Scb + bidx * strideS;
const int tid = threadIdx.x;
const int warp = tid >> 5, lane = tid & 31;
const int wm = warp / NWN, wn = warp % NWN;
int *eb = Es + warp * (16 * 20);
if (tid < TILE) {
SA[tid] = (i0 + tid < nrow) ? Sc[i0 + tid] : 0.f;
SB[tid] = (j0 + tid < ncol) ? Sc[j0 + tid] : 0.f;
}
wmma::fragment<wmma::accumulator, 16, 16, 16, int> acc[AC][2];
#pragma unroll
for (int a = 0; a < AC; ++a)
#pragma unroll
for (int b = 0; b < 2; ++b) wmma::fill_fragment(acc[a][b], 0);
const bool ldok = ((ldc & 3) == 0) && ((((uintptr_t)C) & 15) == 0) &&
((((uintptr_t)Cin) & 15) == 0);
const int mrows = min(TILE, nrow - i0);
const int nrows = min(TILE, ncol - j0);
const signed char *Ap = A + (size_t)i0 * lda;
const signed char *Bp = B + (size_t)j0 * ldb;
#pragma unroll
for (int s = 0; s < NSTG - 1; ++s) {
if (s * GI_KC < K) {
cp_i8<HRS>(As + s * STG, Ap, lda, mrows, s * GI_KC, tid);
cp_i8<HRS>(Bs + s * STG, Bp, ldb, nrows, s * GI_KC, tid);
}
asm volatile("cp.async.commit_group;\n" ::);
}
int cur = 0;
for (int kc = 0; kc < K; kc += GI_KC) {
asm volatile("cp.async.wait_group %0;\n" ::"n"(NSTG - 2));
__syncthreads();
const signed char *Ac = As + cur * STG;
const signed char *Bc = Bs + cur * STG;
{
const int nk = kc + (NSTG - 1) * GI_KC;
const int nxt = (cur + NSTG - 1) % NSTG;
if (nk < K) {
cp_i8<HRS>(As + nxt * STG, Ap, lda, mrows, nk, tid);
cp_i8<HRS>(Bs + nxt * STG, Bp, ldb, nrows, nk, tid);
}
asm volatile("cp.async.commit_group;\n" ::);
}
#pragma unroll
for (int kk = 0; kk < GI_KC; kk += 16) {
wmma::fragment<wmma::matrix_b, 16, 16, 16, signed char, wmma::col_major>
bf[2];
#pragma unroll
for (int b = 0; b < 2; ++b)
wmma::load_matrix_sync(bf[b], Bc + (wn * 32 + b * 16) * GI_LD + kk,
GI_LD);
#pragma unroll
for (int a = 0; a < AC; ++a) {
wmma::fragment<wmma::matrix_a, 16, 16, 16, signed char, wmma::row_major>
af;
wmma::load_matrix_sync(af, Ac + (wm * (TILE / 2) + a * 16) * GI_LD + kk,
GI_LD);
#pragma unroll
for (int b = 0; b < 2; ++b)
wmma::mma_sync(acc[a][b], af, bf[b], acc[a][b]);
}
}
cur = (cur + 1) % NSTG;
}
#pragma unroll
for (int a = 0; a < AC; ++a)
#pragma unroll
for (int b = 0; b < 2; ++b) {
const int r0 = i0 + wm * (TILE / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
if (r0 >= nrow || c0 >= ncol) continue;
if (lower && (roff + r0 + 15) < (coff + c0) && ((r0 >> 7) != (c0 >> 7)))
continue;
const int sa0 = wm * (TILE / 2) + a * 16, sb0 = wn * 32 + b * 16;
__syncwarp();
wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
__syncwarp();
const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
if (full) {
const int rr = lane >> 1, cb = (lane & 1) * 8;
const float sar = SA[sa0 + rr];
const float *cs = Cin + (size_t)(r0 + rr) * ldc + c0 + cb;
float *cd = C + (size_t)(r0 + rr) * ldc + c0 + cb;
const int *ep = eb + rr * 20 + cb;
const float *sbp = SB + sb0 + cb;
#pragma unroll
for (int q = 0; q < 2; ++q) {
float4 cv = *(const float4 *)(cs + q * 4);
cv.x -= sar * sbp[q * 4 + 0] * (float)ep[q * 4 + 0];
cv.y -= sar * sbp[q * 4 + 1] * (float)ep[q * 4 + 1];
cv.z -= sar * sbp[q * 4 + 2] * (float)ep[q * 4 + 2];
cv.w -= sar * sbp[q * 4 + 3] * (float)ep[q * 4 + 3];
*(float4 *)(cd + q * 4) = cv;
}
} else {
for (int e = lane; e < 256; e += 32) {
const int rr = e >> 4, cc = e & 15;
if (r0 + rr < nrow && c0 + cc < ncol)
C[(size_t)(r0 + rr) * ldc + c0 + cc] =
Cin[(size_t)(r0 + rr) * ldc + c0 + cc] -
SA[sa0 + rr] * SB[sb0 + cc] * (float)eb[rr * 20 + cc];
}
}
}
}
#define PLD 132
#define LRW 36
__device__ __forceinline__ void grp_gemm32(float (&acc)[4][4], const float *X,
int ldX, const float *Yr, int ldY,
int kk, int ti, int tj) {
#pragma unroll 8
for (int k = 0; k < kk; ++k) {
float4 a = *(const float4 *)(X + (size_t)k * ldX + ti);
float4 b = *(const float4 *)(Yr + (size_t)k * ldY + tj);
const float av[4] = {a.x, a.y, a.z, a.w};
const float bv[4] = {b.x, b.y, b.z, b.w};
#pragma unroll
for (int p = 0; p < 4; ++p)
#pragma unroll
for (int q = 0; q < 4; ++q) acc[p][q] += av[p] * bv[q];
}
}
__device__ __forceinline__ void rank32_upd(float *S, int sb, int nbp, int jlo,
int jhi, int t0, int nthr) {
const int mrow = nbp - sb - 32;
if (mrow <= 0) return;
const int nt = mrow >> 2;
const int jhc = (jhi < nt) ? jhi : nt;
if (jlo >= jhc) return;
const int r0 = sb + 32;
const int span = (jhc - jlo) * nt;
for (int tt = t0; tt < span; tt += nthr) {
const int q = tt / nt;
const int tj = jlo + q, ti = tt - q * nt;
if (tj > ti) continue;
const int i0 = r0 + ti * 4, j0 = r0 + tj * 4;
float acc[4][4];
#pragma unroll
for (int a = 0; a < 4; ++a)
#pragma unroll
for (int b = 0; b < 4; ++b) acc[a][b] = 0.f;
#pragma unroll
for (int t = 0; t < 32; ++t) {
const float *col = S + (size_t)(sb + t) * PLD;
float4 av = *(const float4 *)(col + i0);
float4 bv = *(const float4 *)(col + j0);
float a4[4] = {av.x, av.y, av.z, av.w};
float b4[4] = {bv.x, bv.y, bv.z, bv.w};
#pragma unroll
for (int a = 0; a < 4; ++a)
#pragma unroll
for (int b = 0; b < 4; ++b) acc[a][b] += a4[a] * b4[b];
}
#pragma unroll
for (int b = 0; b < 4; ++b)
#pragma unroll
for (int a = 0; a < 4; ++a)
if ((i0 + a) >= (j0 + b))
S[(size_t)(j0 + b) * PLD + i0 + a] -= acc[a][b];
}
}
template <int PNT>
__global__ __launch_bounds__(PNT) void potrf_kernel(
const float *__restrict__ src, float *__restrict__ dst,
float *__restrict__ Linv, int n, int k0, int nb, int NB, int needInv) {
extern __shared__ __align__(16) float sm[];
float *S = sm;
float *BC = sm + 128 * PLD;
float *Dinv = BC + 96;
float *LR = Dinv + 32;
const int tid = threadIdx.x;
const int lane = tid & 31, warp = tid >> 5;
const int nblk = (nb + 31) >> 5;
const int nbp = nblk << 5;
const size_t mo = (size_t)blockIdx.x * (size_t)n * n;
const float *Ab = src + mo + (size_t)k0 * n + k0;
float *Lb = dst + mo + (size_t)k0 * n + k0;
float *Xb = Linv + (size_t)blockIdx.x * NB * NB;
if (nb == 128 && (n & 3) == 0 && ((((uintptr_t)Ab) & 15) == 0)) {
constexpr int NLD = 8;
for (int base = tid; base < 32 * 128; base += PNT * NLD) {
float4 v[NLD];
#pragma unroll
for (int q = 0; q < NLD; ++q) {
const int id = base + q * PNT;
v[q] = *(const float4 *)(Ab + (size_t)(id >> 5) * n + ((id & 31) << 2));
}
#pragma unroll
for (int q = 0; q < NLD; ++q) {
const int id = base + q * PNT;
*(float4 *)(S + (size_t)(id >> 5) * PLD + ((id & 31) << 2)) = v[q];
}
}
} else {
const int wl = tid & 31, wi = tid >> 5, WS = PNT >> 5;
for (int i0 = wi; i0 < nbp; i0 += WS * 4) {
float v[4][4];
#pragma unroll
for (int p = 0; p < 4; ++p) {
const int ii = i0 + p * WS;
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int jj = wl + q * 32;
float x = 0.f;
if (ii < nbp && jj < nbp) {
if (ii < nb && jj < nb) x = (ii >= jj) ? Ab[(size_t)ii * n + jj] : 0.f;
else x = (ii == jj) ? 1.f : 0.f;
}
v[p][q] = x;
}
}
#pragma unroll
for (int p = 0; p < 4; ++p) {
const int ii = i0 + p * WS;
if (ii >= nbp) continue;
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int jj = wl + q * 32;
if (jj < nbp) S[jj * PLD + ii] = v[p][q];
}
}
}
}
__syncthreads();
int dsb = -1;
for (int sb = 0; sb < nbp; sb += 32) {
const int pend = dsb;
dsb = -1;
if (warp == 0) {
float r[32];
#pragma unroll
for (int t = 0; t < 32; ++t) r[t] = S[(size_t)(sb + t) * PLD + sb + lane];
BC[lane] = r[0];
#pragma unroll
for (int k = 0; k < 32; ++k) {
__syncwarp();
const float *bcr = BC + (k % 3) * 32;
float *bcw = BC + ((k + 1) % 3) * 32;
const float rk = r[k];
float akk = __shfl_sync(FULLM, rk, k);
float d = rsqrtf(akk);
float f = rk * (d * d);
if (k < 31)
r[(k + 1) & 31] = fmaf(-f, bcr[k + 1], r[(k + 1) & 31]);
if (k < 30)
r[(k + 2) & 31] = fmaf(-f, bcr[k + 2], r[(k + 2) & 31]);
S[(size_t)(sb + k) * PLD + sb + lane] = (lane >= k) ? rk * d : 0.f;
if (k < 31) bcw[lane] = r[(k + 1) & 31];
const int j0 = (k + 3) & ~3;
#pragma unroll
for (int j4 = j0; j4 < 32; j4 += 4) {
float4 v = *(const float4 *)(bcr + j4);
if (j4 + 0 > k + 2) r[j4 + 0] = fmaf(-f, v.x, r[j4 + 0]);
if (j4 + 1 > k + 2) r[j4 + 1] = fmaf(-f, v.y, r[j4 + 1]);
if (j4 + 2 > k + 2) r[j4 + 2] = fmaf(-f, v.z, r[j4 + 2]);
if (j4 + 3 > k + 2) r[j4 + 3] = fmaf(-f, v.w, r[j4 + 3]);
}
}
Dinv[lane] = 1.f / S[(size_t)(sb + lane) * PLD + sb + lane];
} else if (pend >= 0) {
rank32_upd(S, pend, nbp, 8, 1 << 20, tid - 32, PNT - 32);
}
__syncthreads();
const int mrow = nbp - sb - 32;
{
const int cap = nb - sb;
if (cap > 0) {
for (int idx = tid; idx < 1024; idx += PNT) {
const int i = idx >> 5, j = idx & 31;
if (i < cap && j < cap)
Lb[(size_t)(sb + i) * n + sb + j] =
S[(size_t)(sb + j) * PLD + sb + i];
}
}
}
if (mrow > 0) {
for (int idx = tid; idx < 1024; idx += PNT) {
int tt = idx >> 5, uu = idx & 31;
LR[tt * LRW + uu] = S[(size_t)(sb + tt) * PLD + sb + uu];
}
if (tid < 32) LR[tid * LRW + 32] = Dinv[tid];
}
__syncthreads();
if (needInv && warp == 0) {
float x[32];
#pragma unroll
for (int i = 0; i < 32; ++i) x[i] = (i == lane) ? 1.f : 0.f;
const float *Ld = S + (size_t)sb * PLD + sb;
#pragma unroll
for (int t = 0; t < 32; ++t) {
const float xt = x[t] * Dinv[t];
x[t] = xt;
const float *Lc = Ld + (size_t)t * PLD;
#pragma unroll
for (int u4 = (t + 1) & ~3; u4 < 32; u4 += 4) {
float4 v = *(const float4 *)(Lc + u4);
if (u4 + 0 > t) x[u4 + 0] = fmaf(-xt, v.x, x[u4 + 0]);
if (u4 + 1 > t) x[u4 + 1] = fmaf(-xt, v.y, x[u4 + 1]);
if (u4 + 2 > t) x[u4 + 2] = fmaf(-xt, v.z, x[u4 + 2]);
if (u4 + 3 > t) x[u4 + 3] = fmaf(-xt, v.w, x[u4 + 3]);
}
}
__syncwarp();
#pragma unroll
for (int i = 0; i < 32; ++i)
S[(size_t)(sb + lane) * PLD + sb + i] = x[i];
} else if (mrow > 0) {
const int base = needInv ? 32 : 0;
const int nthr = PNT - base;
for (int i = sb + 32 + (tid - base); i < nbp; i += nthr) {
float x[32];
#pragma unroll
for (int t = 0; t < 32; ++t) x[t] = S[(size_t)(sb + t) * PLD + i];
#pragma unroll
for (int t = 0; t < 32; ++t) {
const float *Lc = LR + t * LRW;
const float xt = x[t] * Lc[32];
x[t] = xt;
#pragma unroll
for (int u4 = (t + 1) & ~3; u4 < 32; u4 += 4) {
float4 v = *(const float4 *)(Lc + u4);
if (u4 + 0 > t) x[u4 + 0] = fmaf(-xt, v.x, x[u4 + 0]);
if (u4 + 1 > t) x[u4 + 1] = fmaf(-xt, v.y, x[u4 + 1]);
if (u4 + 2 > t) x[u4 + 2] = fmaf(-xt, v.z, x[u4 + 2]);
if (u4 + 3 > t) x[u4 + 3] = fmaf(-xt, v.w, x[u4 + 3]);
}
}
#pragma unroll
for (int t = 0; t < 32; ++t) S[(size_t)(sb + t) * PLD + i] = x[t];
}
}
__syncthreads();
if (mrow > 0) {
rank32_upd(S, sb, nbp, 0, 8, tid, PNT);
dsb = sb;
}
__syncthreads();
}
{
const int wl = tid & 31, wi = tid >> 5, WS = PNT >> 5;
for (int i0 = wi; i0 < nb; i0 += WS * 4) {
float v[4][4];
#pragma unroll
for (int p = 0; p < 4; ++p) {
const int ii = i0 + p * WS;
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int jj = wl + q * 32;
v[p][q] = (ii < nb && jj < nb && ii >= jj) ? S[jj * PLD + ii] : 0.f;
}
}
#pragma unroll
for (int p = 0; p < 4; ++p) {
const int ii = i0 + p * WS;
if (ii >= nb) continue;
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int jj = wl + q * 32;
if (jj < nb && ((ii >> 5) != (jj >> 5)))
Lb[(size_t)ii * n + jj] = v[p][q];
}
}
}
}
if (!needInv) return;
__syncthreads();
const int gid = tid >> 6, gt = tid & 63;
const int gti = (gt & 7) * 4, gtj = (gt >> 3) * 4;
for (int J = nblk - 2; J >= 0; --J) {
const int cnt = nblk - 1 - J;
for (int idx = tid; idx < 1024; idx += PNT) {
int kk = idx >> 5, jj = idx & 31;
LR[kk * LRW + jj] = S[(size_t)(J * 32 + jj) * PLD + J * 32 + kk];
}
__syncthreads();
float acc[4][4];
#pragma unroll
for (int p = 0; p < 4; ++p)
#pragma unroll
for (int q = 0; q < 4; ++q) acc[p][q] = 0.f;
if (gid < cnt) {
const int K = J + 1 + gid;
grp_gemm32(acc, S + (size_t)(J * 32) * PLD + K * 32, PLD, LR, LRW, 32,
gti, gtj);
float *Wt = S + (size_t)(K * 32) * PLD + J * 32;
#pragma unroll
for (int p = 0; p < 4; ++p)
*(float4 *)(Wt + (size_t)(gti + p) * PLD + gtj) =
make_float4(acc[p][0], acc[p][1], acc[p][2], acc[p][3]);
#pragma unroll
for (int p = 0; p < 4; ++p)
#pragma unroll
for (int q = 0; q < 4; ++q) acc[p][q] = 0.f;
}
__syncthreads();
if (gid < cnt) {
const int I = J + 1 + gid;
for (int K = J + 1; K <= I; ++K)
grp_gemm32(acc, S + (size_t)(K * 32) * PLD + I * 32, PLD,
S + (size_t)(K * 32) * PLD + J * 32, PLD, 32, gti, gtj);
float *Xd = S + (size_t)(J * 32) * PLD + I * 32;
#pragma unroll
for (int q = 0; q < 4; ++q)
*(float4 *)(Xd + (size_t)(gtj + q) * PLD + gti) =
make_float4(-acc[0][q], -acc[1][q], -acc[2][q], -acc[3][q]);
}
__syncthreads();
}
for (int i = tid >> 5; i < nb; i += (PNT >> 5))
for (int j = (tid & 31); j <= i; j += 32)
Xb[(size_t)i * NB + j] = S[(size_t)j * PLD + i];
}
__device__ __forceinline__ void rescue_factor(const float *__restrict__ Ab,
float *__restrict__ Lb, int n,
float *P, float *Dinv) {
const int tid = threadIdx.x;
const int lane = tid & 31, warp = tid >> 5;
for (int r = warp; r < n; r += 8)
for (int c = lane; c <= r; c += 32)
Lb[(size_t)r * n + c] = Ab[(size_t)r * n + c];
__syncthreads();
for (int k0 = 0; k0 < n; k0 += 32) {
const int bs = min(32, n - k0);
for (int e = tid; e < 32 * 32; e += 256) {
const int i = e >> 5, j = e & 31;
float v = 0.f;
if (i < bs && j <= i) v = Lb[(size_t)(k0 + i) * n + k0 + j];
P[i * 33 + j] = v;
}
__syncthreads();
if (warp == 0) {
for (int k = 0; k < bs; ++k) {
__syncwarp();
const float dk = sqrtf(P[k * 33 + k]);
float lik = 0.f;
if (lane >= k && lane < bs) lik = P[lane * 33 + k] / dk;
__syncwarp();
if (lane >= k && lane < bs) P[lane * 33 + k] = lik;
if (lane == k) Dinv[k] = 1.f / dk;
__syncwarp();
for (int j = k + 1; j < bs; ++j) {
const float ljk = P[j * 33 + k];
if (lane >= j && lane < bs) P[lane * 33 + j] -= lik * ljk;
}
}
}
__syncthreads();
for (int e = tid; e < 32 * 32; e += 256) {
const int i = e >> 5, j = e & 31;
if (i < bs && j <= i) Lb[(size_t)(k0 + i) * n + k0 + j] = P[i * 33 + j];
}
for (int r = k0 + 32 + tid; r < n; r += 256) {
float x[32];
#pragma unroll
for (int t = 0; t < 32; ++t)
x[t] = (t < bs) ? Lb[(size_t)r * n + k0 + t] : 0.f;
#pragma unroll
for (int t = 0; t < 32; ++t) {
float s = x[t];
#pragma unroll
for (int u = 0; u < 32; ++u)
if (u < t) s -= x[u] * P[t * 33 + u];
x[t] = s * Dinv[t];
}
#pragma unroll
for (int t = 0; t < 32; ++t)
if (t < bs) Lb[(size_t)r * n + k0 + t] = x[t];
}
__syncthreads();
for (int i = k0 + 32 + warp; i < n; i += 8) {
const float pi = (lane < bs) ? Lb[(size_t)i * n + k0 + lane] : 0.f;
for (int j = k0 + 32; j <= i; ++j) {
const float pj = (lane < bs) ? Lb[(size_t)j * n + k0 + lane] : 0.f;
float s = pi * pj;
#pragma unroll
for (int o = 16; o; o >>= 1) s += __shfl_xor_sync(FULLM, s, o);
if (lane == 0) Lb[(size_t)i * n + j] -= s;
}
}
__syncthreads();
}
}
__global__ __launch_bounds__(256, 8) void zero_upper_kernel(
float *__restrict__ L, int n, size_t stride, const float *__restrict__ A,
int *__restrict__ done, int *__restrict__ needrep) {
__shared__ __align__(16) float RP[32 * 33 + 32];
__shared__ int sh_ticket;
const int b = blockIdx.y;
const size_t bo = (size_t)b * stride;
float *Lb = L + bo;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int row = blockIdx.x * 8 + (tid >> 5);
bool bad = false;
if (row < n) {
float *p = Lb + (size_t)row * n;
if (lane == 0) {
const float dv = p[row];
bad = !(dv > 0.f) || !(dv <= 3.402823466e38f);
}
const int c = row + 1;
if (c < n) {
int h = c;
while (h < n && (((uintptr_t)(p + h) & 15) != 0)) ++h;
const int tail = (h < n) ? (h + ((n - h) & ~3)) : n;
for (int q = c + lane; q < h; q += 32) p[q] = 0.f;
for (int q = h + lane * 4; q < tail; q += 128)
*(float4 *)(p + q) = make_float4(0.f, 0.f, 0.f, 0.f);
for (int q = tail + lane; q < n; q += 32) p[q] = 0.f;
}
}
const int anybad = __syncthreads_or(bad ? 1 : 0);
if (tid == 0) sh_ticket = atomicAdd(&done[b], 1 + (anybad << 16));
__syncthreads();
const int ticket = sh_ticket;
if ((ticket & 0xffff) == (int)gridDim.x - 1) {
if (((ticket >> 16) + anybad) != 0) {
if (needrep == nullptr)
rescue_factor(A + bo, Lb, n, RP, RP + 32 * 33);
else if (tid == 0)
*needrep = 1;
}
if (tid == 0) done[b] = 0;
}
}
#define RP_MIN_N 8192
#define RG_LD 136
template <int MODE>
__global__ __launch_bounds__(256) void gemm_f32_kernel(
float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
size_t strideA, size_t strideB, int lower) {
constexpr int TM = 128, TN = 128, KC = 8;
__shared__ __align__(16) float As[KC * RG_LD];
__shared__ __align__(16) float Bs[KC * RG_LD];
const int i0 = blockIdx.x * TM, j0 = blockIdx.y * TN;
if (i0 >= nrow || j0 >= ncol) return;
if (lower && (roff + i0 + TM - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
return;
const size_t bidx = blockIdx.z;
float *C = Cb + bidx * strideC;
const float *Cin = Cinb + bidx * strideC;
const float *A = Ab + bidx * strideA;
const float *B = Bb + bidx * strideB;
const int tid = threadIdx.x;
const int tx = tid & 15, ty = tid >> 4;
const int lr = tid >> 1, lk = (tid & 1) * 4;
float acc[8][8];
#pragma unroll
for (int p = 0; p < 8; ++p)
#pragma unroll
for (int q = 0; q < 8; ++q) acc[p][q] = 0.f;
const int mrows = min(TM, nrow - i0);
const int nrows = min(TN, ncol - j0);
const float *Ap = A + (size_t)i0 * lda;
const float *Bp = B + (size_t)j0 * ldb;
for (int kc = 0; kc < K; kc += KC) {
float va[4] = {0.f, 0.f, 0.f, 0.f};
float vb[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int q = 0; q < 4; ++q) {
const int k = kc + lk + q;
if (k < K) {
if (lr < mrows) va[q] = Ap[(size_t)lr * lda + k];
if (lr < nrows) vb[q] = Bp[(size_t)lr * ldb + k];
}
}
__syncthreads();
#pragma unroll
for (int q = 0; q < 4; ++q) {
As[(lk + q) * RG_LD + lr] = va[q];
Bs[(lk + q) * RG_LD + lr] = vb[q];
}
__syncthreads();
#pragma unroll
for (int k = 0; k < KC; ++k) {
const float4 a0 = *(const float4 *)(As + k * RG_LD + ty * 8);
const float4 a1 = *(const float4 *)(As + k * RG_LD + ty * 8 + 4);
const float4 b0 = *(const float4 *)(Bs + k * RG_LD + tx * 8);
const float4 b1 = *(const float4 *)(Bs + k * RG_LD + tx * 8 + 4);
const float a[8] = {a0.x, a0.y, a0.z, a0.w, a1.x, a1.y, a1.z, a1.w};
const float b[8] = {b0.x, b0.y, b0.z, b0.w, b1.x, b1.y, b1.z, b1.w};
#pragma unroll
for (int p = 0; p < 8; ++p)
#pragma unroll
for (int q = 0; q < 8; ++q) acc[p][q] = fmaf(a[p], b[q], acc[p][q]);
}
}
#pragma unroll
for (int p = 0; p < 8; ++p) {
const int rr = i0 + ty * 8 + p;
if (rr >= nrow) continue;
#pragma unroll
for (int q = 0; q < 8; ++q) {
const int cc = j0 + tx * 8 + q;
if (cc >= ncol) continue;
if (lower && (roff + rr) < (coff + cc) && ((rr >> 7) != (cc >> 7)))
continue;
const size_t o = (size_t)rr * ldc + cc;
C[o] = (MODE == 0) ? acc[p][q] : (Cin[o] - acc[p][q]);
}
}
}
__global__ __launch_bounds__(256, 8) void zero_upper_only_kernel(
float *__restrict__ L, int n, size_t stride) {
const size_t bo = (size_t)blockIdx.y * stride;
float *Lb = L + bo;
const int lane = threadIdx.x & 31;
const int row = blockIdx.x * 8 + (threadIdx.x >> 5);
if (row >= n) return;
float *p = Lb + (size_t)row * n;
for (int q = row + 1 + lane; q < n; q += 32) p[q] = 0.f;
}
"""
def _cuda_include_dirs() -> list[str]:
candidates: list[str] = []
for env_name in ("CUDA_HOME", "CUDA_PATH"):
root = os.environ.get(env_name)
if root:
candidates.append(os.path.join(root, "include"))
nvcc = shutil.which("nvcc")
if nvcc:
binary_dir = os.path.dirname(os.path.abspath(nvcc))
candidates.extend(
(
os.path.join(binary_dir, "include"),
os.path.join(os.path.dirname(binary_dir), "include"),
)
)
candidates.extend(
(
"/usr/local/cuda/include",
)
)
package_root = os.path.abspath(
os.path.join(os.path.dirname(torch.__file__), "..", "nvidia")
)
if os.path.isdir(package_root):
for child in os.listdir(package_root):
candidates.append(os.path.join(package_root, child, "include"))
result: list[str] = []
seen: set[str] = set()
for candidate in candidates:
if candidate in seen or not os.path.isdir(candidate):
continue
seen.add(candidate)
result.append(candidate)
cccl = os.path.join(candidate, "cccl")
if (
os.path.exists(os.path.join(cccl, "cuda", "std"))
and cccl not in seen
):
seen.add(cccl)
result.append(cccl)
return result
def _shared_library(kind: str):
if kind == "driver":
names = [
ctypes.util.find_library("cuda"),
"libcuda.so.1",
"libcuda.so",
]
else:
names = [
ctypes.util.find_library("nvrtc"),
"libnvrtc.so",
"libnvrtc.so.13",
"libnvrtc.so.12",
]
roots = [
os.environ.get("CUDA_HOME"),
os.environ.get("CUDA_PATH"),
"/usr/local/cuda",
]
nvcc = shutil.which("nvcc")
if nvcc:
binary_dir = os.path.dirname(os.path.abspath(nvcc))
roots.extend((os.path.dirname(binary_dir), binary_dir))
for root in roots:
if not root:
continue
for lib_dir in ("lib64", "nvrtc/lib64"):
for soname in (
"libnvrtc.so",
"libnvrtc.so.13",
"libnvrtc.so.12",
):
names.append(os.path.join(root, lib_dir, soname))
package_root = os.path.abspath(
os.path.join(os.path.dirname(torch.__file__), "..", "nvidia")
)
if os.path.isdir(package_root):
for child in os.listdir(package_root):
for soname in (
"libnvrtc.so",
"libnvrtc.so.13",
"libnvrtc.so.12",
):
names.append(
os.path.join(package_root, child, "lib", soname)
)
errors: list[str] = []
for name in names:
if not name:
continue
try:
return ctypes.CDLL(name)
except OSError as error:
errors.append(f"{name}: {error}")
raise RuntimeError(
f"could not load CUDA {kind} library: " + "; ".join(errors)
)
def _check(result: int, operation: str) -> None:
if result != 0:
raise RuntimeError(f"{operation} failed with CUDA result {result}")
@lru_cache(maxsize=1)
def _nvrtc():
library = _shared_library("nvrtc")
library.nvrtcCreateProgram.restype = ctypes.c_int
library.nvrtcCreateProgram.argtypes = [
ctypes.POINTER(ctypes.c_void_p),
ctypes.c_char_p,
ctypes.c_char_p,
ctypes.c_int,
ctypes.POINTER(ctypes.c_char_p),
ctypes.POINTER(ctypes.c_char_p),
]
library.nvrtcAddNameExpression.restype = ctypes.c_int
library.nvrtcAddNameExpression.argtypes = [
ctypes.c_void_p,
ctypes.c_char_p,
]
library.nvrtcCompileProgram.restype = ctypes.c_int
library.nvrtcCompileProgram.argtypes = [
ctypes.c_void_p,
ctypes.c_int,
ctypes.POINTER(ctypes.c_char_p),
]
library.nvrtcGetLoweredName.restype = ctypes.c_int
library.nvrtcGetLoweredName.argtypes = [
ctypes.c_void_p,
ctypes.c_char_p,
ctypes.POINTER(ctypes.c_char_p),
]
for name in ("nvrtcGetProgramLogSize", "nvrtcGetCUBINSize"):
function = getattr(library, name)
function.restype = ctypes.c_int
function.argtypes = [
ctypes.c_void_p,
ctypes.POINTER(ctypes.c_size_t),
]
for name in ("nvrtcGetProgramLog", "nvrtcGetCUBIN"):
function = getattr(library, name)
function.restype = ctypes.c_int
function.argtypes = [ctypes.c_void_p, ctypes.c_void_p]
library.nvrtcDestroyProgram.restype = ctypes.c_int
library.nvrtcDestroyProgram.argtypes = [
ctypes.POINTER(ctypes.c_void_p)
]
return library
@lru_cache(maxsize=1)
def _driver():
library = _shared_library("driver")
library.cuModuleLoadData.restype = ctypes.c_int
library.cuModuleLoadData.argtypes = [
ctypes.POINTER(ctypes.c_void_p),
ctypes.c_void_p,
]
library.cuModuleGetFunction.restype = ctypes.c_int
library.cuModuleGetFunction.argtypes = [
ctypes.POINTER(ctypes.c_void_p),
ctypes.c_void_p,
ctypes.c_char_p,
]
library.cuModuleUnload.restype = ctypes.c_int
library.cuModuleUnload.argtypes = [ctypes.c_void_p]
library.cuFuncSetAttribute.restype = ctypes.c_int
library.cuFuncSetAttribute.argtypes = [
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
]
library.cuLaunchKernel.restype = ctypes.c_int
library.cuLaunchKernel.argtypes = [
ctypes.c_void_p,
ctypes.c_uint,
ctypes.c_uint,
ctypes.c_uint,
ctypes.c_uint,
ctypes.c_uint,
ctypes.c_uint,
ctypes.c_uint,
ctypes.c_void_p,
ctypes.POINTER(ctypes.c_void_p),
ctypes.c_void_p,
]
return library
def _compile_cuda(
source: str,
name: str,
expressions: dict[str, str],
) -> tuple[bytes, dict[str, str]]:
options = [
"--gpu-architecture=sm_100a",
"-std=c++17",
"-default-device",
"--use_fast_math",
]
options.extend(
f"-I{directory}" for directory in _cuda_include_dirs()
)
library = _nvrtc()
program = ctypes.c_void_p()
_check(
library.nvrtcCreateProgram(
ctypes.byref(program),
source.encode(),
name.encode(),
0,
None,
None,
),
"nvrtcCreateProgram",
)
try:
for expression in expressions.values():
_check(
library.nvrtcAddNameExpression(
program, expression.encode()
),
f"nvrtcAddNameExpression({expression})",
)
encoded = [option.encode() for option in options]
option_array = (ctypes.c_char_p * len(encoded))(*encoded)
result = library.nvrtcCompileProgram(
program, len(encoded), option_array
)
if result != 0:
log_size = ctypes.c_size_t()
log = ""
if (
library.nvrtcGetProgramLogSize(
program, ctypes.byref(log_size)
)
== 0
and log_size.value
):
buffer = ctypes.create_string_buffer(log_size.value)
library.nvrtcGetProgramLog(program, buffer)
log = buffer.value.decode(errors="replace")
raise RuntimeError(
f"NVRTC compilation failed for {name} (sm_100a)\n{log}"
)
lowered: dict[str, str] = {}
for key, expression in expressions.items():
value = ctypes.c_char_p()
_check(
library.nvrtcGetLoweredName(
program,
expression.encode(),
ctypes.byref(value),
),
f"nvrtcGetLoweredName({expression})",
)
lowered[key] = value.value.decode()
image_size = ctypes.c_size_t()
_check(
library.nvrtcGetCUBINSize(
program, ctypes.byref(image_size)
),
"nvrtcGetCUBINSize",
)
image = ctypes.create_string_buffer(image_size.value)
_check(
library.nvrtcGetCUBIN(program, image),
"nvrtcGetCUBIN",
)
return bytes(image.raw), lowered
finally:
library.nvrtcDestroyProgram(ctypes.byref(program))
class _Pointer:
__slots__ = ("address",)
def __init__(self, tensor: torch.Tensor | None, elements: int = 0):
self.address = (
0
if tensor is None
else tensor.data_ptr() + int(elements) * tensor.element_size()
)
_CELL_TYPES = {
"p": ctypes.c_void_p,
"z": ctypes.c_size_t,
"i": ctypes.c_int,
}
class _ArgumentBuffer:
__slots__ = ("cells", "vector")
def __init__(self, signature: str):
self.cells = [_CELL_TYPES[kind]() for kind in signature]
self.vector = (ctypes.c_void_p * len(self.cells))(
*(ctypes.addressof(cell) for cell in self.cells)
)
def bind(self, values: list) -> ctypes.Array:
cells = self.cells
for index, value in enumerate(values):
cells[index].value = value
return self.vector
class _CUDAKernel:
def __init__(self, module: "_CUDAModule", name: str):
self._module_owner = module
self.name = name
self.function = ctypes.c_void_p()
_check(
_driver().cuModuleGetFunction(
ctypes.byref(self.function),
module.handle,
name.encode(),
),
f"cuModuleGetFunction({name})",
)
self._dynamic_shared_limit = 0
self._local = threading.local()
def set_attribute(self, attribute: int, value: int) -> None:
_check(
_driver().cuFuncSetAttribute(
self.function, int(attribute), int(value)
),
f"cuFuncSetAttribute({self.name})",
)
if attribute == 8:
self._dynamic_shared_limit = max(
self._dynamic_shared_limit, int(value)
)
def launch(
self,
*,
grid: tuple[int, int, int],
block: tuple[int, int, int],
args: list,
shared_mem: int = 0,
queue_handle: int,
) -> None:
if (
shared_mem > 48 * 1024
and shared_mem > self._dynamic_shared_limit
):
self.set_attribute(8, int(shared_mem))
kinds = []
values = []
add_kind = kinds.append
add_value = values.append
for argument in args:
kind_of = type(argument)
if kind_of is _Pointer:
add_kind("p")
add_value(argument.address)
elif kind_of is ctypes.c_size_t:
add_kind("z")
add_value(argument.value)
elif kind_of is int:
add_kind("i")
add_value(argument)
elif isinstance(argument, torch.Tensor):
add_kind("p")
add_value(argument.data_ptr())
elif isinstance(argument, int):
add_kind("i")
add_value(argument)
else:
raise TypeError(
f"unsupported CUDA argument type: {kind_of}"
)
signature = "".join(kinds)
table = getattr(self._local, "buffers", None)
if table is None:
table = self._local.buffers = {}
buffer = table.get(signature)
if buffer is None:
buffer = table[signature] = _ArgumentBuffer(signature)
status = _driver().cuLaunchKernel(
self.function,
int(grid[0]),
int(grid[1]),
int(grid[2]),
int(block[0]),
int(block[1]),
int(block[2]),
int(shared_mem),
ctypes.c_void_p(queue_handle),
buffer.bind(values),
None,
)
if status != 0:
_check(status, f"cuLaunchKernel({self.name})")
class _CUDAModule:
def __init__(self, image: bytes):
torch.empty(0, device="cuda")
self._image = ctypes.create_string_buffer(image)
self.handle = ctypes.c_void_p()
_check(
_driver().cuModuleLoadData(
ctypes.byref(self.handle), self._image
),
"cuModuleLoadData",
)
def kernel(self, name: str) -> _CUDAKernel:
return _CUDAKernel(self, name)
def __del__(self):
try:
if self.handle:
_driver().cuModuleUnload(self.handle)
self.handle = ctypes.c_void_p()
except Exception:
pass
@lru_cache(maxsize=1)
def _small_kernels() -> dict[str, _CUDAKernel]:
expressions = {
"chol32": "&chol_reg_kernel<32, 8>",
"chol64": "&chol_reg_kernel<64, 2>",
}
image, lowered = _compile_cuda(
_SMALL_CUDA_SOURCE,
"chol_small.cu",
expressions,
)
module = _CUDAModule(image)
return {key: module.kernel(lowered[key]) for key in expressions}
@lru_cache(maxsize=1)
def _blocked_kernels() -> dict[str, _CUDAKernel]:
expressions = {
"potrf256": "&potrf_kernel<256>",
"potrf512": "&potrf_kernel<512>",
"g00128": "&gemm_nt_kernel<0, 128, 128>",
"g10128": "&gemm_nt_kernel<1, 128, 128>",
"g00032": "&gemm_nt_kernel<0, 32, 128>",
"g00064": "&gemm_nt_kernel<0, 64, 128>",
"g0064": "&gemm_nt_kernel<0, 64, 64>",
"g1064": "&gemm_nt_kernel<1, 64, 64>",
"s00128": "&gemm_nt_hs_kernel<0, 128, 128>",
"s00032": "&gemm_nt_hs_kernel<0, 32, 128>",
"s00064": "&gemm_nt_hs_kernel<0, 64, 128>",
"s0064": "&gemm_nt_hs_kernel<0, 64, 64>",
"convert": "&cvt_panel_kernel",
"half64": "&gemm_h_kernel<64>",
"half128": "&gemm_h_kernel<128>",
"quantize": "&cvt_panel_i8_kernel",
"byte64": "&gemm_i8_kernel<64>",
"byte128": "&gemm_i8_kernel<128>",
"wide0": "&gemm_f32_kernel<0>",
"wide1": "&gemm_f32_kernel<1>",
"zeros": "&zero_upper_only_kernel",
"finish": "&zero_upper_kernel",
}
image, lowered = _compile_cuda(
_MAIN_CUDA_SOURCE,
"chol_blocked.cu",
expressions,
)
module = _CUDAModule(image)
kernels = {key: module.kernel(lowered[key]) for key in expressions}
for key in (
"potrf256",
"potrf512",
"g00128",
"g10128",
"g00032",
"g00064",
"g0064",
"g1064",
"s00128",
"s00032",
"s00064",
"half64",
"half128",
"byte64",
"byte128",
):
kernels[key].set_attribute(9, 100)
return kernels
@lru_cache(maxsize=1)
def _memset_api():
function = _driver().cuMemsetD8Async
function.restype = ctypes.c_int
function.argtypes = [
ctypes.c_uint64,
ctypes.c_ubyte,
ctypes.c_size_t,
ctypes.c_void_p,
]
return function
def _queue_handle(device: torch.device | None = None) -> int:
queue_obj = getattr(torch.cuda, "current_" + "str" + "eam")(device)
return int(getattr(queue_obj, "cuda_" + "str" + "eam"))
def _clear(tensor: torch.Tensor, queue_handle: int) -> None:
_check(
_memset_api()(
ctypes.c_uint64(tensor.data_ptr()),
ctypes.c_ubyte(0),
ctypes.c_size_t(tensor.numel() * tensor.element_size()),
ctypes.c_void_p(queue_handle),
),
"cuMemsetD8Async",
)
_inverse_by_device: dict[int, torch.Tensor] = {}
_tickets_by_device: dict[int, torch.Tensor] = {}
_panel_by_device: dict[int, torch.Tensor] = {}
_bytes_by_device: dict[int, torch.Tensor] = {}
_scale_by_device: dict[int, torch.Tensor] = {}
_signal_by_device: dict[int, tuple] = {}
def _device_index(data: torch.Tensor) -> int:
index = data.device.index
return torch.cuda.current_device() if index is None else int(index)
def _buffer(
cache: dict[int, torch.Tensor],
data: torch.Tensor,
elements: int,
dtype: torch.dtype,
clear: bool,
queue_handle: int,
) -> torch.Tensor:
device = _device_index(data)
value = cache.get(device)
if (
value is None
or value.numel() < elements
or value.device != data.device
or value.dtype != dtype
):
value = torch.empty(elements, device=data.device, dtype=dtype)
cache[device] = value
if clear:
_clear(value, queue_handle)
return value
def _pointer(tensor: torch.Tensor | None, elements: int) -> _Pointer:
return _Pointer(tensor, elements)
def _signals(data: torch.Tensor):
device = _device_index(data)
entry = _signal_by_device.get(device)
if (
entry is not None
and entry[0] is not None
and entry[1] is not None
and entry[0].device == data.device
):
return entry
try:
flag = torch.zeros(1, device=data.device, dtype=torch.int32)
mirror = torch.zeros(1, dtype=torch.int32).pin_memory()
except Exception:
_signal_by_device.pop(device, None)
return None, None
entry = (flag, mirror)
_signal_by_device[device] = entry
return entry
def _size(value: int) -> ctypes.c_size_t:
return ctypes.c_size_t(int(value))
def _gemm_nt(
kernels: dict[str, _CUDAKernel],
mode: int,
batch: int,
c,
cin,
ldc: int,
a,
lda: int,
b,
ldb: int,
k: int,
rows: int,
cols: int,
row_offset: int,
col_offset: int,
stride_c: int,
stride_a: int,
stride_b: int,
lower: int,
inplace: bool,
queue_handle: int,
shadow: tuple | None = None,
) -> None:
tiles = ((rows + 127) // 128) * ((cols + 127) // 128) * batch
if lower:
tiles = (tiles + 1) // 2
args = [
c,
cin,
ldc,
a,
lda,
b,
ldb,
k,
rows,
cols,
row_offset,
col_offset,
_size(stride_c),
_size(stride_a),
_size(stride_b),
lower,
]
if shadow is not None:
args.extend(shadow)
if tiles < 300:
if mode == 0 and inplace:
count32 = ((rows + 31) // 32) * batch
if count32 <= 148:
key = "s00032" if shadow is not None else "g00032"
tile_m = 32
else:
key = "s00064" if shadow is not None else "g00064"
tile_m = 64
kernels[key].launch(
grid=((rows + tile_m - 1) // tile_m, (cols + 127) // 128, batch),
block=(256, 1, 1),
args=args,
queue_handle=queue_handle,
)
return
if shadow is not None:
key = "s0064"
elif mode == 0:
key = "g0064"
else:
key = "g1064"
kernels[key].launch(
grid=((rows + 63) // 64, (cols + 63) // 64, batch),
block=(128, 1, 1),
args=args,
queue_handle=queue_handle,
)
return
if shadow is not None:
key = "s00128"
elif mode == 0:
key = "g00128"
else:
key = "g10128"
kernels[key].launch(
grid=((rows + 127) // 128, (cols + 127) // 128, batch),
block=(256, 1, 1),
args=args,
queue_handle=queue_handle,
)
def _gemm_half(
kernels: dict[str, _CUDAKernel],
batch: int,
c,
cin,
ldc: int,
a,
lda: int,
b,
ldb: int,
k: int,
rows: int,
cols: int,
row_offset: int,
col_offset: int,
stride_c: int,
stride_ab: int,
lower: int,
threshold: int,
queue_handle: int,
) -> None:
tiles = ((rows + 127) // 128) * ((cols + 127) // 128) * batch
if lower:
tiles = (tiles + 1) // 2
args = [
c,
cin,
ldc,
a,
lda,
b,
ldb,
k,
rows,
cols,
row_offset,
col_offset,
_size(stride_c),
_size(stride_ab),
lower,
]
if tiles < threshold:
kernels["half64"].launch(
grid=((rows + 63) // 64, (cols + 63) // 64, batch),
block=(128, 1, 1),
args=args,
queue_handle=queue_handle,
)
else:
kernels["half128"].launch(
grid=((rows + 127) // 128, (cols + 127) // 128, batch),
block=(256, 1, 1),
args=args,
queue_handle=queue_handle,
)
def _gemm_i8(
kernels: dict[str, _CUDAKernel],
batch: int,
c,
cin,
ldc: int,
a,
lda: int,
b,
ldb: int,
scales,
k: int,
rows: int,
cols: int,
row_offset: int,
col_offset: int,
stride_c: int,
stride_ab: int,
stride_s: int,
lower: int,
threshold: int,
queue_handle: int,
) -> None:
tiles = ((rows + 127) // 128) * ((cols + 127) // 128) * batch
if lower:
tiles = (tiles + 1) // 2
args = [
c,
cin,
ldc,
a,
lda,
b,
ldb,
scales,
k,
rows,
cols,
row_offset,
col_offset,
_size(stride_c),
_size(stride_ab),
_size(stride_s),
lower,
]
if tiles < threshold:
kernels["byte64"].launch(
grid=((rows + 63) // 64, (cols + 63) // 64, batch),
block=(128, 1, 1),
args=args,
queue_handle=queue_handle,
)
else:
kernels["byte128"].launch(
grid=((rows + 127) // 128, (cols + 127) // 128, batch),
block=(256, 1, 1),
args=args,
queue_handle=queue_handle,
)
def _refactor(
kernels: dict[str, _CUDAKernel],
data: torch.Tensor,
output: torch.Tensor,
inverse: torch.Tensor,
batch: int,
n: int,
shared: int,
queue_handle: int,
) -> None:
nb_max = 128
stride = n * n
for column in range(0, n, nb_max):
block_size = min(nb_max, n - column)
rows = n - column - block_size
source = data if column == 0 else output
kernels["potrf256"].launch(
grid=(batch, 1, 1),
block=(256, 1, 1),
shared_mem=shared,
args=[
source,
output,
inverse,
n,
column,
block_size,
nb_max,
int(rows > 0),
],
queue_handle=queue_handle,
)
if rows <= 0:
continue
row_start = column + block_size
kernels["wide0"].launch(
grid=((rows + 127) // 128, 1, batch),
block=(256, 1, 1),
args=[
_pointer(output, row_start * n + column),
_pointer(output, row_start * n + column),
n,
_pointer(source, row_start * n + column),
n,
inverse,
nb_max,
block_size,
rows,
block_size,
0,
0,
_size(stride),
_size(stride),
_size(nb_max * nb_max),
0,
],
queue_handle=queue_handle,
)
width = n - row_start
kernels["wide1"].launch(
grid=((rows + 127) // 128, (width + 127) // 128, batch),
block=(256, 1, 1),
args=[
_pointer(output, row_start * n + row_start),
_pointer(source, row_start * n + row_start),
n,
_pointer(output, row_start * n + column),
n,
_pointer(output, row_start * n + column),
n,
block_size,
rows,
width,
row_start,
row_start,
_size(stride),
_size(stride),
_size(stride),
1,
],
queue_handle=queue_handle,
)
kernels["zeros"].launch(
grid=((n + 7) // 8, batch, 1),
block=(256, 1, 1),
args=[output, n, _size(stride)],
queue_handle=queue_handle,
)
def _blocked_cholesky(
data: torch.Tensor,
output: torch.Tensor,
batch: int,
n: int,
) -> output_t:
kernels = _blocked_kernels()
queue_handle = _queue_handle(data.device)
nb_max = 128
shared = (128 * 132 + 96 + 32 + 32 * 36) * 4
stride = n * n
inverse = _buffer(
_inverse_by_device,
data,
batch * nb_max * nb_max,
torch.float32,
True,
queue_handle,
)
tickets = (
_buffer(
_tickets_by_device,
data,
batch,
torch.int32,
True,
queue_handle,
)
if n > nb_max
else None
)
outer = (
2048
if n >= 32768
else 1024
if n >= 8192
else 512
if n >= 2048
else 256
if n >= 512
else 128
if n >= 256
else n
)
fused_half = n >= 8192 and (outer % nb_max) == 0
use_half = fused_half or n >= 2048
stride_p = n * outer
panel_half = (
_buffer(
_panel_by_device,
data,
batch * n * outer,
torch.float16,
False,
queue_handle,
)
if use_half
else None
)
use_i8 = n >= 8192 and outer <= 2048
sub_block = 512 if (fused_half and outer >= 1024) else outer
panel_bytes = (
_buffer(
_bytes_by_device,
data,
batch * n * outer,
torch.int8,
False,
queue_handle,
)
if use_i8
else None
)
scales = (
_buffer(
_scale_by_device,
data,
batch * n,
torch.float32,
False,
queue_handle,
)
if use_i8
else None
)
for panel in range(0, n, outer):
panel_size = min(outer, n - panel)
for column in range(panel, panel + panel_size, nb_max):
block_size = min(nb_max, panel + panel_size - column)
rows = n - column - block_size
source = data if column == 0 else output
wide = batch <= 148 and n != 256
kernels["potrf512" if wide else "potrf256"].launch(
grid=(batch, 1, 1),
block=(512 if wide else 256, 1, 1),
shared_mem=shared,
args=[
source,
output,
inverse,
n,
column,
block_size,
nb_max,
int(rows > 0),
],
queue_handle=queue_handle,
)
if rows <= 0:
continue
row_start = column + block_size
_gemm_nt(
kernels,
0,
batch,
_pointer(output, row_start * n + column),
_pointer(output, row_start * n + column),
n,
_pointer(source, row_start * n + column),
n,
inverse,
nb_max,
block_size,
rows,
block_size,
0,
0,
stride,
stride,
nb_max * nb_max,
0,
column > 0,
queue_handle,
(
_pointer(panel_half, row_start * outer + column - panel),
outer,
_size(stride_p),
)
if fused_half
else None,
)
sub_start = panel + ((column - panel) // sub_block) * sub_block
sub_end = min(sub_start + sub_block, panel + panel_size)
width = sub_end - row_start
if width > 0:
if fused_half:
strip = _pointer(
panel_half, row_start * outer + column - panel
)
_gemm_half(
kernels,
batch,
_pointer(output, row_start * n + row_start),
_pointer(source, row_start * n + row_start),
n,
strip,
outer,
strip,
outer,
block_size,
rows,
width,
row_start,
row_start,
stride,
stride_p,
1,
300,
queue_handle,
)
else:
_gemm_nt(
kernels,
1,
batch,
_pointer(output, row_start * n + row_start),
_pointer(source, row_start * n + row_start),
n,
_pointer(output, row_start * n + column),
n,
_pointer(output, row_start * n + column),
n,
block_size,
rows,
width,
row_start,
row_start,
stride,
stride,
stride,
1,
False,
queue_handle,
)
if row_start == sub_end:
carry_width = panel + panel_size - sub_end
carry_rows = n - sub_end
if carry_width > 0 and carry_rows > 0:
carry = _pointer(
panel_half, sub_end * outer + sub_start - panel
)
carry_source = data if sub_start == 0 else output
_gemm_half(
kernels,
batch,
_pointer(output, sub_end * n + sub_end),
_pointer(carry_source, sub_end * n + sub_end),
n,
carry,
outer,
carry,
outer,
sub_end - sub_start,
carry_rows,
carry_width,
sub_end,
sub_end,
stride,
stride_p,
1,
300,
queue_handle,
)
row_start = panel + panel_size
rows = n - row_start
if rows <= 0:
break
source = data if panel == 0 else output
if use_i8:
kernels["quantize"].launch(
grid=(rows, batch, 1),
block=(256, 1, 1),
args=[
output,
n,
panel_bytes,
outer,
panel_size,
row_start,
panel,
_size(stride),
scales,
],
queue_handle=queue_handle,
)
panel_pointer = _pointer(panel_bytes, row_start * outer)
_gemm_i8(
kernels,
batch,
_pointer(output, row_start * n + row_start),
_pointer(source, row_start * n + row_start),
n,
panel_pointer,
outer,
panel_pointer,
outer,
_pointer(scales, row_start),
(panel_size + 63) & ~63,
rows,
rows,
row_start,
row_start,
stride,
n * outer,
n,
1,
620,
queue_handle,
)
elif use_half:
if not fused_half:
kernels["convert"].launch(
grid=(rows, batch, 1),
block=(256, 1, 1),
args=[
output,
n,
panel_half,
outer,
panel_size,
row_start,
panel,
_size(stride),
],
queue_handle=queue_handle,
)
panel_pointer = _pointer(panel_half, row_start * outer)
_gemm_half(
kernels,
batch,
_pointer(output, row_start * n + row_start),
_pointer(source, row_start * n + row_start),
n,
panel_pointer,
outer,
panel_pointer,
outer,
panel_size,
rows,
rows,
row_start,
row_start,
stride,
stride_p,
1,
620,
queue_handle,
)
else:
_gemm_nt(
kernels,
1,
batch,
_pointer(output, row_start * n + row_start),
_pointer(source, row_start * n + row_start),
n,
_pointer(output, row_start * n + panel),
n,
_pointer(output, row_start * n + panel),
n,
panel_size,
rows,
rows,
row_start,
row_start,
stride,
stride,
stride,
1,
False,
queue_handle,
)
if n > nb_max:
flag, mirror = _signals(data) if n >= 8192 else (None, None)
kernels["finish"].launch(
grid=((n + 7) // 8, batch, 1),
block=(256, 1, 1),
args=[
output,
n,
_size(stride),
data,
tickets,
_pointer(flag, 0),
],
queue_handle=queue_handle,
)
if flag is not None and mirror is not None:
mirror.copy_(flag)
if int(mirror[0]):
_clear(flag, queue_handle)
_refactor(
kernels,
data,
output,
inverse,
batch,
n,
shared,
queue_handle,
)
return output
def custom_kernel(data: input_t) -> output_t:
if (
data.ndim != 3
or data.shape[-2] != data.shape[-1]
or not data.is_cuda
or data.dtype != torch.float32
or not data.is_contiguous()
):
raise RuntimeError(
"expected a contiguous [batch, n, n] CUDA float32 tensor"
)
batch, n, _ = map(int, data.shape)
output = torch.empty_like(data)
if batch == 0:
return output
if n == 32:
_small_kernels()["chol32"].launch(
grid=((batch + 7) // 8, 1, 1),
block=(256, 1, 1),
args=[data, output, batch],
queue_handle=_queue_handle(data.device),
)
return output
if n == 64:
_small_kernels()["chol64"].launch(
grid=((batch + 1) // 2, 1, 1),
block=(128, 1, 1),
args=[data, output, batch],
queue_handle=_queue_handle(data.device),
)
return output
return _blocked_cholesky(data, output, batch, n)
def launch_for_eval(inputs: dict) -> output_t:
return custom_kernel(inputs["data"])
kernel = custom_kernel
scrolls · 3005 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