Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
359.2µs
#16 of 337
2026-07-29

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.

mbarrierasm volatile("mbarrier.init.shared::cta.b64 [%0], 1;"
mmausing bf3_acc_t = nvcuda::wmma::fragment<nvcuda::wmma::accumulator,
shared-memory__shared__ __align__(16) float stage[WPB * N * LDP];
vector-width = float4const 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