Skip to content
KernelIndex
Search⌘K

submission 911904

patwrall · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 2290 lines, June 9 Researcher Reciprocity License v1.0.

submission_1b26b37.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-911904?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
858.5µs
#90 of 337
2026-07-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c6e76cea8aa4c6a06b3b8429046395bb2956c1a7f4f3abe7eefb496e776424a2
license declaredunknown
license concludedunknown
authorspatwrall
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

async-copy__device__ __forceinline__ void cp_async16(void* dst, const void* src) {
mbarrierasm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
mma"mma.sync.aligned.m16n8k8.row.col.f32.tf32.tf32.f32 "
persistent-kernelpersistent_cholesky_kernel(float* __restrict__ A, int n, long long bstride,
shared-memoryextern __shared__ float smem[];
tcgen05asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
vector-width = float4for (int q = 0; q < 32; q += 4) *(float4*)&r[q] = *(const float4*)&src[q];

Kernel source

submission_1b26b37.py2290 lines
import sys
from pathlib import Path

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

P = lambda *a: print(*a, file=sys.stderr)
CPP_SRC = r"""
"""

# Pure device code, no torch dependency — extracted verbatim by the local
# sm_86 test harness. Keep torch/ATen includes out of this string.
CUDA_KERNELS = r"""
// Factor one B x B SPD tile (dst <- L, upper zeroed); one CTA, unblocked
// Cholesky staged through smem (B*(B+1) floats, padded). src may equal dst.
template <int B, int NT>
__device__ void potrf_tile(const float* __restrict__ src, float* __restrict__ dst,
                           int lda, float* smem) {
    const int tid = threadIdx.x;
    auto s = [&](int i, int j) -> float& { return smem[i * (B + 1) + j]; };

    for (int idx = tid; idx < B * B; idx += NT)
        s(idx / B, idx % B) = src[(idx / B) * lda + (idx % B)];
    __syncthreads();

    for (int k = 0; k < B; ++k) {
        if (tid == 0) s(k, k) = sqrtf(s(k, k));
        __syncthreads();
        const float dinv = 1.0f / s(k, k);
        for (int i = k + 1 + tid; i < B; i += NT) s(i, k) *= dinv;
        __syncthreads();
        for (int i = k + 1 + tid; i < B; i += NT) {
            const float lik = s(i, k);
            for (int j = k + 1; j <= i; ++j) s(i, j) -= lik * s(j, k);
        }
        __syncthreads();
    }

    for (int idx = tid; idx < B * B; idx += NT) {
        const int i = idx / B, j = idx % B;
        dst[i * lda + j] = (j <= i) ? s(i, j) : 0.0f;
    }
    __syncthreads();
}

// rho regime, n == B: one CTA per matrix, whole matrix is a single tile,
// output written out of place (no clone pass).
template <int B, int NT>
__global__ void potrf_batched_kernel(const float* __restrict__ A, float* __restrict__ L,
                                     long long batch_stride) {
    extern __shared__ float smem[];
    potrf_tile<B, NT>(A + blockIdx.x * batch_stride, L + blockIdx.x * batch_stride, B, smem);
}

template <int B, int NT>
void launch_potrf_batched(const float* A, float* L, long long batch) {
    constexpr int SMEM = B * (B + 1) * sizeof(float);
    cudaFuncSetAttribute(potrf_batched_kernel<B, NT>,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    potrf_batched_kernel<B, NT>
        <<<batch, NT, SMEM>>>(A, L, (long long)B * B);
}

// n == 32 special case: one warp per matrix, each lane owns a row in
// registers, active column published through smem; no block syncs.
__global__ void potrf32_warp_kernel(const float* __restrict__ A, float* __restrict__ L,
                                    long long batch) {
    __shared__ float col[4][32];
    const int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
    const long long m = (long long)blockIdx.x * 4 + warp;
    if (m >= batch) return;
    const float* src = A + m * 32 * 32 + lane * 32;
    float* dst = L + m * 32 * 32 + lane * 32;

    float r[32];
    #pragma unroll
    for (int q = 0; q < 32; q += 4) *(float4*)&r[q] = *(const float4*)&src[q];

    #pragma unroll
    for (int k = 0; k < 32; ++k) {
        const float dk = sqrtf(__shfl_sync(0xffffffffu, r[k], k));
        if (lane == k)     r[k] = dk;
        else if (lane > k) r[k] /= dk;
        col[warp][lane] = r[k];
        __syncwarp();
        #pragma unroll
        for (int j = 0; j < 32; ++j)
            if (j > k && lane >= j) r[j] -= r[k] * col[warp][j];
        __syncwarp();
    }

    #pragma unroll
    for (int j = 0; j < 32; ++j)
        if (j > lane) r[j] = 0.0f;
    #pragma unroll
    for (int q = 0; q < 32; q += 4) *(float4*)&dst[q] = *(float4*)&r[q];
}

// n == 64: one 2-warp CTA per matrix, each thread owns a row in registers.
// The unscaled column is published pre-scaling (1/d^2 folded into the
// update FMA) so each k step needs a single block sync; colbuf is double
// buffered to avoid a second one.
__global__ void __launch_bounds__(64)
potrf64_rows_kernel(const float* __restrict__ A, float* __restrict__ L, long long batch) {
    __shared__ float colbuf[2][64];
    const int i = threadIdx.x;
    const float* src = A + blockIdx.x * 64 * 64 + i * 64;
    float* dst = L + blockIdx.x * 64 * 64 + i * 64;
    float r[64];
    #pragma unroll
    for (int q = 0; q < 64; q += 4) *(float4*)&r[q] = *(const float4*)&src[q];

    #pragma unroll
    for (int k = 0; k < 64; ++k) {
        colbuf[k & 1][i] = r[k];
        __syncthreads();
        const float dkk = colbuf[k & 1][k];
        if (i > k) {
            const float t = r[k] * (1.0f / dkk);
            #pragma unroll
            for (int j = 0; j < 64; ++j)
                if (j > k && j <= i) r[j] -= t * colbuf[k & 1][j];
            r[k] *= rsqrtf(dkk);
        } else if (i == k) {
            r[k] = sqrtf(dkk);
        }
    }
    #pragma unroll
    for (int j = 0; j < 64; ++j)
        if (j > i) r[j] = 0.0f;
    #pragma unroll
    for (int q = 0; q < 64; q += 4) *(float4*)&dst[q] = *(float4*)&r[q];
}

// n == 128: one 128-thread CTA per matrix, rows resident in smem (rows in
// registers would spill); same single-sync-per-step structure.
__global__ void __launch_bounds__(128)
potrf128_rows_kernel(const float* __restrict__ A, float* __restrict__ L, long long batch) {
    extern __shared__ float s[];  // 128 x 129
    const int i = threadIdx.x;
    const float* src = A + blockIdx.x * 128 * 128;
    float* dst = L + blockIdx.x * 128 * 128;
    for (int idx = i; idx < 128 * 128; idx += 128)
        s[(idx / 128) * 129 + idx % 128] = src[idx];
    __syncthreads();

    // elimination keeps unscaled values (diag holds d^2); scaling by
    // rsqrt(d^2) happens per column in the store pass, so nothing mutates
    // a column while other threads still read it
    float* row = s + i * 129;
    for (int k = 0; k < 128; ++k) {
        if (i > k) {
            const float t = row[k] * (1.0f / s[k * 129 + k]);
            for (int j = k + 1; j <= i; ++j) row[j] -= t * s[j * 129 + k];
        }
        __syncthreads();
    }
    for (int idx = i; idx < 128 * 128; idx += 128) {
        const int rr = idx / 128, cc = idx % 128;
        float v = 0.0f;
        if (cc < rr)       v = s[rr * 129 + cc] * rsqrtf(s[cc * 129 + cc]);
        else if (cc == rr) v = sqrtf(s[cc * 129 + cc]);
        dst[idx] = v;
    }
}

// n == 128, register scheme: one 256-thread CTA per matrix, threads i and
// i+128 own the two 64-column halves of row i in registers. Unscaled column
// k is published through double-buffered smem (all rows, so colbuf[i] also
// serves as each row's own a_ik); one block sync per step like the n=64
// scheme, rsqrt scaling folded in after publication.
// Rank-4 variant of the sweep: publish 4 RAW columns per barrier, factor
// the 4x4 pivot block in registers (uniform work, redundant per thread),
// then back-substitute the four combined coefficients so each j costs four
// FMAs against raw columns. Quarter the barriers of rank-1; deeper raw-
// column cancellation than rank-2, so accuracy must be re-validated on any
// new shape it is applied to.
// Column ownership is INTERLEAVED by 4-wide group: the h-half of a row owns
// groups h, h+2, ... (register slot r[4q+u] holds column 8q+4h+u), so both
// halves keep equal live work at every step — column-half ownership left
// half-0 idle for the whole second half of the sweep (30% barrier stall).
template <int ROWS>
__device__ __forceinline__ void panel_sweep128_r4(float (&r)[64], float (*colbuf)[4][ROWS],
                                                  int i, int h, bool active) {
    #pragma unroll
    for (int k = 0; k < 128; k += 4) {
        const int lq = (k >> 3) * 4;                  // r[] base of pivot group
        const bool own = active && (h == ((k >> 2) & 1));
        const int p = (k >> 2) & 1;
        if (own) {
            #pragma unroll
            for (int b = 0; b < 4; ++b) colbuf[p][b][i] = r[lq + b];
        }
        __syncthreads();
        const float M00 = colbuf[p][0][k],     M10 = colbuf[p][0][k + 1],
                    M20 = colbuf[p][0][k + 2], M30 = colbuf[p][0][k + 3],
                    M11 = colbuf[p][1][k + 1], M21 = colbuf[p][1][k + 2],
                    M31 = colbuf[p][1][k + 3], M22 = colbuf[p][2][k + 2],
                    M32 = colbuf[p][2][k + 3], M33 = colbuf[p][3][k + 3];
        const float rd0 = 1.0f / M00;
        const float s10 = M10 * rd0, s20 = M20 * rd0, s30 = M30 * rd0;
        const float d1 = M11 - s10 * M10;
        const float rd1 = 1.0f / d1;
        const float m21 = M21 - s20 * M10, m31 = M31 - s30 * M10;
        const float s21 = m21 * rd1, s31 = m31 * rd1;
        const float d2 = M22 - s20 * M20 - s21 * m21;
        const float rd2 = 1.0f / d2;
        const float m32 = M32 - s30 * M20 - s31 * m21;
        const float s32 = m32 * rd2;
        const float d3 = M33 - s30 * M30 - s31 * m31 - s32 * m32;
        const float e0 = colbuf[p][0][i], e1 = colbuf[p][1][i],
                    e2 = colbuf[p][2][i], e3 = colbuf[p][3][i];
        const float t0 = e0 * rd0;
        const float e1c = e1 - s10 * e0;
        const float t1 = e1c * rd1;
        const float e2c = e2 - s20 * e0 - s21 * e1c;
        const float t2 = e2c * rd2;
        const float e3c = e3 - s30 * e0 - s31 * e1c - s32 * e2c;
        const float t3 = e3c * (1.0f / d3);
        if (active && i > k + 3) {
            const float a3 = t3;
            const float a2 = t2 - a3 * s32;
            const float a1 = t1 - a2 * s21 - a3 * s31;
            const float a0 = t0 - a1 * s10 - a2 * s20 - a3 * s30;
            if (h == 0) {
                #pragma unroll
                for (int q = 0; q < 16; ++q) {
                    if (q * 8 + 3 > k + 3) {
                        const float4 v0 = *(const float4*)&colbuf[p][0][q * 8];
                        const float4 v1 = *(const float4*)&colbuf[p][1][q * 8];
                        const float4 v2 = *(const float4*)&colbuf[p][2][q * 8];
                        const float4 v3 = *(const float4*)&colbuf[p][3][q * 8];
                        const float w0[4] = {v0.x, v0.y, v0.z, v0.w};
                        const float w1[4] = {v1.x, v1.y, v1.z, v1.w};
                        const float w2[4] = {v2.x, v2.y, v2.z, v2.w};
                        const float w3[4] = {v3.x, v3.y, v3.z, v3.w};
                        #pragma unroll
                        for (int u = 0; u < 4; ++u) {
                            const int j = q * 8 + u;
                            if (j > k + 3 && j <= i)
                                r[q * 4 + u] -= a0 * w0[u] + a1 * w1[u] + a2 * w2[u] + a3 * w3[u];
                        }
                    }
                }
            } else {
                #pragma unroll
                for (int q = 0; q < 16; ++q) {
                    if (q * 8 + 4 + 3 > k + 3) {
                        const float4 v0 = *(const float4*)&colbuf[p][0][q * 8 + 4];
                        const float4 v1 = *(const float4*)&colbuf[p][1][q * 8 + 4];
                        const float4 v2 = *(const float4*)&colbuf[p][2][q * 8 + 4];
                        const float4 v3 = *(const float4*)&colbuf[p][3][q * 8 + 4];
                        const float w0[4] = {v0.x, v0.y, v0.z, v0.w};
                        const float w1[4] = {v1.x, v1.y, v1.z, v1.w};
                        const float w2[4] = {v2.x, v2.y, v2.z, v2.w};
                        const float w3[4] = {v3.x, v3.y, v3.z, v3.w};
                        #pragma unroll
                        for (int u = 0; u < 4; ++u) {
                            const int j = q * 8 + 4 + u;
                            if (j > k + 3 && j <= i)
                                r[q * 4 + u] -= a0 * w0[u] + a1 * w1[u] + a2 * w2[u] + a3 * w3[u];
                        }
                    }
                }
            }
        }
        if (own) {
            if (i == k)          r[lq] = sqrtf(M00);
            else if (i > k)      r[lq] = e0 * rsqrtf(M00);
            if (i == k + 1)      r[lq + 1] = sqrtf(d1);
            else if (i > k + 1)  r[lq + 1] = e1c * rsqrtf(d1);
            if (i == k + 2)      r[lq + 2] = sqrtf(d2);
            else if (i > k + 2)  r[lq + 2] = e2c * rsqrtf(d2);
            if (i == k + 3)      r[lq + 3] = sqrtf(d3);
            else if (i > k + 3)  r[lq + 3] = e3c * rsqrtf(d3);
        }
    }
}

// interleaved-group column index of register slot r[4q+u] for half h
__device__ __forceinline__ int ilv_col(int q, int h, int u) { return q * 8 + h * 4 + u; }

__global__ void __launch_bounds__(256, 2)
potrf128_regs_kernel(const float* __restrict__ A, float* __restrict__ L, long long batch) {
    __shared__ float colbuf[2][4][128];  // [parity][column][row]
    const int i = threadIdx.x & 127;
    const int h = threadIdx.x >> 7;
    const float* src = A + blockIdx.x * 128 * 128 + i * 128;
    float* dst = L + blockIdx.x * 128 * 128 + i * 128;
    float r[64];
    #pragma unroll
    for (int q = 0; q < 16; ++q) *(float4*)&r[q * 4] = *(const float4*)&src[ilv_col(q, h, 0)];

    panel_sweep128_r4<128>(r, colbuf, i, h, true);
    #pragma unroll
    for (int q = 0; q < 16; ++q) {
        #pragma unroll
        for (int u = 0; u < 4; ++u)
            if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
        *(float4*)&dst[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
    }
}

// n == 256: one 512-thread CTA per matrix. Phase A sweeps the whole 256x128
// left panel (L11 + L21) held in registers (two 64-col halves per row);
// L21 is then staged to smem ([129] pad -> conflict-free column reads) for
// an in-CTA SYRK into registers, and phase D sweeps the trailing 128x128
// with threads 0..255 while the rest only participate in barriers.
__global__ void __launch_bounds__(512, 1)
fused256_kernel(const float* __restrict__ A, float* __restrict__ Lg, long long bstride) {
    extern __shared__ float smem[];  // sL[128][129] | colbuf 2*4*256
    float (*sL)[129] = (float (*)[129])smem;
    float* colb = smem + 128 * 129;
    const float* src = A + blockIdx.x * bstride;
    float* dst = Lg + blockIdx.x * bstride;
    float r[64];

    const int i = threadIdx.x & 255;
    const int h = (threadIdx.x >> 8) & 1;
    #pragma unroll
    for (int q = 0; q < 16; ++q)
        *(float4*)&r[q * 4] = *(const float4*)&src[i * 256 + ilv_col(q, h, 0)];
    panel_sweep128_r4<256>(r, (float (*)[4][256])colb, i, h, true);

    if (i < 128) {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            #pragma unroll
            for (int u = 0; u < 4; ++u)
                if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
    } else {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            #pragma unroll
            for (int u = 0; u < 4; ++u) sL[i - 128][ilv_col(q, h, u)] = r[q * 4 + u];
    }
    #pragma unroll
    for (int q = 0; q < 16; ++q)
        *(float4*)&dst[i * 256 + ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
    // zero the upper-right 128x128 block
    for (int idx = threadIdx.x; idx < 128 * 32; idx += 512) {
        const int rr = idx >> 5, cc = (idx & 31) * 4;
        *(float4*)&dst[rr * 256 + 128 + cc] = make_float4(0.f, 0.f, 0.f, 0.f);
    }
    __syncthreads();

    // SYRK: 4 rows (lane-stride 1 -> conflict-free column reads of sL) x
    // 8 cols (constant per warp -> broadcast) per thread
    const int rb = threadIdx.x & 31, cb = (threadIdx.x >> 5) * 8;
    float acc[4][8] = {};
    for (int k = 0; k < 128; ++k) {
        float a4[4], b8[8];
        #pragma unroll
        for (int u = 0; u < 4; ++u) a4[u] = sL[rb + 32 * u][k];
        #pragma unroll
        for (int v = 0; v < 8; ++v) b8[v] = sL[cb + v][k];
        #pragma unroll
        for (int u = 0; u < 4; ++u)
            #pragma unroll
            for (int v = 0; v < 8; ++v) acc[u][v] += a4[u] * b8[v];
    }
    float c22[4][8];
    #pragma unroll
    for (int u = 0; u < 4; ++u)
        #pragma unroll
        for (int v = 0; v < 8; v += 4) {
            const float4 a4 = *(const float4*)&src[(128 + rb + 32 * u) * 256 + 128 + cb + v];
            c22[u][v] = a4.x - acc[u][v];     c22[u][v + 1] = a4.y - acc[u][v + 1];
            c22[u][v + 2] = a4.z - acc[u][v + 2]; c22[u][v + 3] = a4.w - acc[u][v + 3];
        }
    __syncthreads();  // all SYRK reads of sL done before overwrite
    #pragma unroll
    for (int u = 0; u < 4; ++u)
        #pragma unroll
        for (int v = 0; v < 8; ++v) sL[rb + 32 * u][cb + v] = c22[u][v];
    __syncthreads();

    // phase D: trailing 128x128 out of smem
    const bool act = threadIdx.x < 256;
    const int i2 = threadIdx.x & 127;
    const int h2 = (threadIdx.x >> 7) & 1;
    if (act) {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            #pragma unroll
            for (int u = 0; u < 4; ++u) r[q * 4 + u] = sL[i2][ilv_col(q, h2, u)];
    }
    __syncthreads();
    panel_sweep128_r4<128>(r, (float (*)[4][128])colb, i2, h2, act);
    if (act) {
        #pragma unroll
        for (int q = 0; q < 16; ++q) {
            #pragma unroll
            for (int u = 0; u < 4; ++u)
                if (ilv_col(q, h2, u) > i2) r[q * 4 + u] = 0.0f;
            *(float4*)&dst[(128 + i2) * 256 + 128 + ilv_col(q, h2, 0)] = *(float4*)&r[q * 4];
        }
    }
}

void launch_fused256(const float* A, float* L, long long batch) {
    constexpr int SMEM = (128 * 129 + 2 * 4 * 256) * sizeof(float);
    cudaFuncSetAttribute(fused256_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    fused256_kernel<<<batch, 512, SMEM>>>(A, L, (long long)256 * 256);
}

// Blocked-128 panel step for n in {512, 1024, ...}: one CTA per consumer
// 128-row block (blockIdx.x = c), each CTA holds pivot block P plus its
// consumer block in registers and runs the fused potrf+trsm sweep; the
// pivot factorization is recomputed redundantly per CTA (cheap) so no
// cross-CTA sync is needed. CTA c==0 stores the pivot rows.
__global__ void __launch_bounds__(512, 1)
panel512_kernel(float* __restrict__ Lg, int n, long long bstride, int P) {
    __shared__ float colbuf[2][4][256];
    float* M = Lg + blockIdx.y * bstride;
    const int i = threadIdx.x & 255;
    const int h = (threadIdx.x >> 8) & 1;
    const int nblk = n / 128;
    const int cons = nblk - P - 1;
    const int grow = (i < 128) ? 128 * P + i
                               : 128 * (P + 1 + (int)blockIdx.x) + (i - 128);
    float* rowp = M + (long long)grow * n + 128 * P;
    float r[64];
    const bool have = (i < 128) || ((int)blockIdx.x < cons);
    if (have) {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            *(float4*)&r[q * 4] = *(const float4*)&rowp[ilv_col(q, h, 0)];
    }
    panel_sweep128_r4<256>(r, colbuf, i, h, have);
    // Pivot rows are stored ONLY when this launch has a single CTA (last
    // column): with consumers in flight, sibling CTAs still read the raw
    // pivot block from global, so the store moves to the update launch.
    if (i < 128) {
        if (blockIdx.x == 0 && cons == 0) {
            #pragma unroll
            for (int q = 0; q < 16; ++q) {
                #pragma unroll
                for (int u = 0; u < 4; ++u)
                    if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
                *(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
            }
        }
    } else if (have) {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            *(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
    }
}

// Left-looking panel step: identical sweep to panel512_kernel, but the
// factored pivot block goes to the side buffer `piv` (one 128x128 slot per
// (matrix, column)) instead of L — in left-looking form nothing reads the
// pivot block after its own column, so sibling CTAs can keep reading the
// raw pivot from L race-free; copy_pivots_kernel patches L at the end.
__global__ void __launch_bounds__(512, 1)
panel512_left_kernel(float* __restrict__ Lg, float* __restrict__ piv,
                     int n, long long bstride, int Q) {
    __shared__ float colbuf[2][4][256];
    float* M = Lg + blockIdx.y * bstride;
    const int i = threadIdx.x & 255;
    const int h = (threadIdx.x >> 8) & 1;
    const int nblk = n / 128;
    const int cons = nblk - Q - 1;
    const int grow = (i < 128) ? 128 * Q + i
                               : 128 * (Q + 1 + (int)blockIdx.x) + (i - 128);
    float* rowp = M + (long long)grow * n + 128 * Q;
    float r[64];
    const bool have = (i < 128) || ((int)blockIdx.x < cons);
    if (have) {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            *(float4*)&r[q * 4] = *(const float4*)&rowp[ilv_col(q, h, 0)];
    }
    panel_sweep128_r4<256>(r, colbuf, i, h, have);
    if (i < 128) {
        if (blockIdx.x == 0) {
            #pragma unroll
            for (int q = 0; q < 16; ++q)
                #pragma unroll
                for (int u = 0; u < 4; ++u)
                    if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
            float* pdst = (cons == 0)
                ? rowp
                : piv + ((long long)blockIdx.y * nblk + Q) * 128 * 128 + i * 128;
            #pragma unroll
            for (int q = 0; q < 16; ++q)
                *(float4*)&pdst[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
        }
    } else if (have) {
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            *(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
    }
}

// Patch the stashed pivot blocks (already tril-zeroed) into L; also
// overwrites the above-diagonal garbage baddbmm left inside pivot blocks.
__global__ void copy_pivots_kernel(float* __restrict__ Lg, const float* __restrict__ piv,
                                   int n, long long bstride) {
    const int nblk = n / 128, Q = blockIdx.x;
    const float* src = piv + ((long long)blockIdx.y * nblk + Q) * 128 * 128;
    float* dst = Lg + blockIdx.y * bstride + (long long)(128 * Q) * n + 128 * Q;
    for (int idx = threadIdx.x; idx < 128 * 32; idx += 256) {
        const int rr = idx >> 5, cc = (idx & 31) * 4;
        *(float4*)&dst[(long long)rr * n + cc] = *(const float4*)&src[rr * 128 + cc];
    }
}

__device__ __forceinline__ void pair_decode(int p, int J, int& I, int& K);

// Trailing update for the blocked-128 path: C(I,K) -= L(I,P) * L(K,P)^T for
// 128-tile pair p (I >= K > P); operands staged transposed ([k][row]) in
// 64-wide k chunks; 256 threads, 8x8 outputs each. Diagonal tiles skip
// above-diagonal stores to preserve the tril zeros.
__global__ void __launch_bounds__(256, 2)
update128_regs_kernel(float* __restrict__ Lg, int n, long long bstride, int P) {
    extern __shared__ float smem[];
    float (*sIT)[132] = (float (*)[132])smem;
    float (*sKT)[132] = (float (*)[132])(smem + 64 * 132);
    float* M = Lg + blockIdx.y * bstride;
    const int cons = n / 128 - P - 1;
    if ((int)blockIdx.x == cons * (cons + 1) / 2) {
        // extra CTA: factor + store the pivot block (raw in global — panel
        // CTAs never wrote it, and no update CTA touches block column P)
        __shared__ float cb[2][4][128];
        const int i = threadIdx.x & 127;
        const int h = (threadIdx.x >> 7) & 1;
        float* rowp = M + (long long)(128 * P + i) * n + 128 * P;
        float r[64];
        #pragma unroll
        for (int q = 0; q < 16; ++q)
            *(float4*)&r[q * 4] = *(const float4*)&rowp[ilv_col(q, h, 0)];
        panel_sweep128_r4<128>(r, cb, i, h, true);
        #pragma unroll
        for (int q = 0; q < 16; ++q) {
            #pragma unroll
            for (int u = 0; u < 4; ++u)
                if (ilv_col(q, h, u) > i) r[q * 4 + u] = 0.0f;
            *(float4*)&rowp[ilv_col(q, h, 0)] = *(float4*)&r[q * 4];
        }
        return;
    }
    int I, K;
    pair_decode(blockIdx.x, P, I, K);
    const float* LI = M + (long long)I * 128 * n + 128 * P;
    const float* LK = M + (long long)K * 128 * n + 128 * P;
    float* C = M + (long long)I * 128 * n + 128 * K;

    const int tx = threadIdx.x & 15, ty = threadIdx.x >> 4;
    float acc[8][8] = {};
    for (int kc = 0; kc < 128; kc += 64) {
        __syncthreads();
        for (int idx = threadIdx.x; idx < 64 * 128; idx += 256) {
            const int rr = idx >> 6, kk = idx & 63;
            sIT[kk][rr] = LI[rr * n + kc + kk];
            sKT[kk][rr] = LK[rr * n + kc + kk];
        }
        __syncthreads();
        #pragma unroll 8
        for (int k = 0; k < 64; ++k) {
            float a8[8], b8[8];
            #pragma unroll
            for (int q = 0; q < 8; q += 4) {
                *(float4*)&a8[q] = *(const float4*)&sIT[k][ty * 8 + q];
                *(float4*)&b8[q] = *(const float4*)&sKT[k][tx * 8 + q];
            }
            #pragma unroll
            for (int u = 0; u < 8; ++u)
                #pragma unroll
                for (int v = 0; v < 8; ++v) acc[u][v] += a8[u] * b8[v];
        }
    }
    #pragma unroll
    for (int u = 0; u < 8; ++u) {
        const int rr = ty * 8 + u;
        #pragma unroll
        for (int v = 0; v < 8; v += 4) {
            const int cc = tx * 8 + v;
            if (I == K && cc + 3 > rr) {
                #pragma unroll
                for (int w = 0; w < 4; ++w)
                    if (cc + w <= rr) C[rr * n + cc + w] -= acc[u][v + w];
            } else {
                float4* cp = (float4*)&C[rr * n + cc];
                float4 c4 = *cp;
                c4.x -= acc[u][v];     c4.y -= acc[u][v + 1];
                c4.z -= acc[u][v + 2]; c4.w -= acc[u][v + 3];
                *cp = c4;
            }
        }
    }
}

// float4 tril copy (out-of-place): the at::tril clone runs ~4x below
// memory speed and showed up at 9% of the 8x2048 profile
__global__ void tril_f4_kernel(const float* __restrict__ A, float* __restrict__ L,
                               int n, long long bstride) {
    const float* src = A + blockIdx.y * bstride;
    float* dst = L + blockIdx.y * bstride;
    const int nq = n >> 2;
    for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < n * nq;
         idx += gridDim.x * blockDim.x) {
        const int i = idx / nq, c = (idx % nq) * 4;
        float4 v = *(const float4*)&src[(long long)i * n + c];
        if (c + 3 > i) {
            if (c + 0 > i) v.x = 0.0f;
            if (c + 1 > i) v.y = 0.0f;
            if (c + 2 > i) v.z = 0.0f;
            v.w = 0.0f;
        }
        *(float4*)&dst[(long long)i * n + c] = v;
    }
}

// Right-looking blocked-128 factorization from the register panel sweep;
// per column: fused potrf+trsm panel CTAs, then the 128-tile update.
void cholesky_b128regs(float* L, int n, long long batch) {
    constexpr int USM = 2 * 64 * 132 * sizeof(float);
    cudaFuncSetAttribute(update128_regs_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, USM);
    const int nblk = n / 128;
    const long long bstride = (long long)n * n;
    for (int P = 0; P < nblk; ++P) {
        const int cons = nblk - P - 1;
        panel512_kernel<<<dim3(cons > 0 ? cons : 1, batch), 512>>>(L, n, bstride, P);
        if (cons > 0)  // +1 CTA factors and stores the pivot block
            update128_regs_kernel
                <<<dim3(cons * (cons + 1) / 2 + 1, batch), 256, USM>>>(L, n, bstride, P);
    }
}

// Factor the diagonal tile of block-column J; one CTA per matrix.
template <int B, int NT>
__global__ void potrf_diag_kernel(float* __restrict__ A, int n, long long bstride, int J) {
    extern __shared__ float smem[];
    float* tile = A + blockIdx.x * bstride + (long long)J * B * (n + 1);
    potrf_tile<B, NT>(tile, tile, n, smem);
}

// Solve L_IJ * L_JJ^T = A_IJ for one tile below the diagonal; smem holds
// the diagonal tile and the tile being solved, one thread per row.
template <int B, int NT>
__device__ void trsm_tile_dev(float* __restrict__ M, int n, int J, int tile_idx, float* smem) {
    float (*sL)[B + 1] = (float (*)[B + 1])smem;
    float (*sX)[B + 1] = (float (*)[B + 1])(smem + B * (B + 1));
    const float* diag = M + (long long)J * B * (n + 1);
    float* tile = M + (long long)(J + 1 + tile_idx) * B * n + (long long)J * B;

    for (int idx = threadIdx.x; idx < B * B; idx += NT) {
        sL[idx / B][idx % B] = diag[(idx / B) * n + idx % B];
        sX[idx / B][idx % B] = tile[(idx / B) * n + idx % B];
    }
    __syncthreads();

    const int r = threadIdx.x;
    if (r < B) {
        for (int k = 0; k < B; ++k) {
            float acc = sX[r][k];
            for (int j = 0; j < k; ++j) acc -= sX[r][j] * sL[k][j];
            sX[r][k] = acc / sL[k][k];
        }
        for (int c = 0; c < B; ++c) tile[r * n + c] = sX[r][c];
    }
    __syncthreads();
}

template <int B>
__global__ void trsm_tile_kernel(float* __restrict__ A, int n, long long bstride, int J) {
    __shared__ float smem[2 * B * (B + 1)];
    trsm_tile_dev<B, B>(A + blockIdx.y * bstride, n, J, blockIdx.x, smem);
}

// Trailing update C_IK -= L_IJ * L_KJ^T for tile pair p (I >= K, I == K
// covers SYRK); 16x16 threads, 4x4 register tile per thread. Operands
// staged transposed ([k][row]) so each k step is two float4 loads.
template <int B>
__device__ void update_tile_dev(float* __restrict__ M, int n, int J, int p, float* smem) {
    float (*sI)[B + 4] = (float (*)[B + 4])smem;
    float (*sK)[B + 4] = (float (*)[B + 4])(smem + B * (B + 4));
    int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
    while ((Ip + 1) * (Ip + 2) / 2 <= p) ++Ip;
    while (Ip * (Ip + 1) / 2 > p) --Ip;
    const int I = J + 1 + Ip, K = J + 1 + (p - Ip * (Ip + 1) / 2);

    const float* LI = M + (long long)I * B * n + (long long)J * B;
    const float* LK = M + (long long)K * B * n + (long long)J * B;
    float* C = M + (long long)I * B * n + (long long)K * B;

    for (int idx = threadIdx.x; idx < B * B; idx += 256) {
        const int r = idx / B, k = idx % B;
        sI[k][r] = LI[r * n + k];
        sK[k][r] = LK[r * n + k];
    }
    __syncthreads();

    const int tx = threadIdx.x % 16, ty = threadIdx.x / 16;
    float acc[4][4] = {};
    #pragma unroll 8
    for (int k = 0; k < B; ++k) {
        const float4 a = *(const float4*)&sI[k][ty * 4];
        const float4 b = *(const float4*)&sK[k][tx * 4];
        const float ar[4] = {a.x, a.y, a.z, a.w};
        const float br[4] = {b.x, b.y, b.z, b.w};
        #pragma unroll
        for (int i = 0; i < 4; ++i)
            #pragma unroll
            for (int j = 0; j < 4; ++j) acc[i][j] += ar[i] * br[j];
    }

    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        float4* cp = (float4*)&C[(ty * 4 + i) * n + tx * 4];
        float4 c = *cp;
        c.x -= acc[i][0]; c.y -= acc[i][1]; c.z -= acc[i][2]; c.w -= acc[i][3];
        *cp = c;
    }
    __syncthreads();
}

template <int B>
__global__ void update_tile_kernel(float* __restrict__ A, int n, long long bstride, int J) {
    __shared__ float smem[2 * B * (B + 4)];
    update_tile_dev<B>(A + blockIdx.y * bstride, n, J, blockIdx.x, smem);
}

// Whole factorization fused into one kernel, one CTA per matrix: no
// per-column launches, matrix stays hot in L2, only block-level syncs.
// Composes the existing tile device functions sequentially per CTA;
// batch supplies the parallelism. Also writes the tril copy itself.
template <int B, int NT>
__global__ void fused_small_kernel(const float* __restrict__ A, float* __restrict__ Lg,
                                   int n, long long bstride) {
    extern __shared__ float scratch[];
    const float* src = A + blockIdx.x * bstride;
    float* L = Lg + blockIdx.x * bstride;
    for (int idx = threadIdx.x; idx < n * n; idx += NT) {
        const int i = idx / n, j = idx % n;
        L[idx] = (j <= i) ? src[idx] : 0.0f;
    }
    __syncthreads();
    const int N = n / B;
    for (int J = 0; J < N; ++J) {
        float* diag = L + (long long)J * B * (n + 1);
        potrf_tile<B, NT>(diag, diag, n, scratch);
        const int T = N - J - 1;
        for (int t = 0; t < T; ++t) trsm_tile_dev<B, NT>(L, n, J, t, scratch);
        const int pairs = T * (T + 1) / 2;
        for (int p = 0; p < pairs; ++p) update_tile_dev<B>(L, n, J, p, scratch);
    }
}

template <int B, int NT>
void launch_fused_small(const float* A, float* L, int n, long long batch) {
    constexpr int SMEM = 2 * B * (B + 4) * sizeof(float);
    cudaFuncSetAttribute(fused_small_kernel<B, NT>,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    fused_small_kernel<B, NT><<<batch, NT, SMEM>>>(A, L, n, (long long)n * n);
}

// In-CTA n=64 diagonal potrf: threads 0..63 own rows in registers, one
// sync per step (all CTA threads participate in syncs). Result goes to
// global and to sDiag (64x68, float4-aligned rows) for the following trsm.
__device__ __forceinline__ void potrf64_regs_dev(float* __restrict__ tile, int lda,
                                 float* __restrict__ sDiag, float* __restrict__ colbuf,
                                 float (&r)[64]) {
    const int i = threadIdx.x;
    if (i < 64) {
        #pragma unroll
        for (int q = 0; q < 64; q += 4) *(float4*)&r[q] = *(const float4*)&tile[i * lda + q];
    }
    #pragma unroll
    for (int k = 0; k < 64; ++k) {
        if (i < 64) colbuf[(k & 1) * 64 + i] = r[k];
        __syncthreads();
        if (i < 64) {
            const float dkk = colbuf[(k & 1) * 64 + k];
            if (i > k) {
                const float t = r[k] * (1.0f / dkk);
                #pragma unroll
                for (int j = 0; j < 64; ++j)
                    if (j > k && j <= i) r[j] -= t * colbuf[(k & 1) * 64 + j];
                r[k] *= rsqrtf(dkk);
            } else if (i == k) {
                r[k] = sqrtf(dkk);
            }
        }
    }
    if (i < 64) {
        #pragma unroll
        for (int j = 0; j < 64; ++j)
            if (j > i) r[j] = 0.0f;
        #pragma unroll
        for (int q = 0; q < 64; q += 4) {
            *(float4*)&tile[i * lda + q] = *(float4*)&r[q];
            *(float4*)&sDiag[i * 68 + q] = *(float4*)&r[q];
        }
    }
    __syncthreads();
}

// In-CTA trsm for up to four 64-tiles at once: 256 threads, one row per
// thread in registers, diagonal tile read from smem; no internal syncs.
__device__ __forceinline__ void trsm4_regs_dev(float* __restrict__ Lg, int n, int J, int t0, int cnt,
                               const float* __restrict__ sDiag, float (&x)[64]) {
    const int tl = threadIdx.x >> 6, r = threadIdx.x & 63;
    if (tl >= cnt) return;
    float* rowp = Lg + (long long)(J + 1 + t0 + tl) * 64 * n + (long long)J * 64 + (long long)r * n;
    #pragma unroll
    for (int q = 0; q < 64; q += 4) *(float4*)&x[q] = *(const float4*)&rowp[q];
    #pragma unroll
    for (int k = 0; k < 64; ++k) {
        float acc = x[k];
        #pragma unroll
        for (int j = 0; j < 64; ++j)
            if (j < k) acc -= x[j] * sDiag[k * 68 + j];
        x[k] = acc * (1.0f / sDiag[k * 68 + k]);
    }
    #pragma unroll
    for (int q = 0; q < 64; q += 4) *(float4*)&rowp[q] = *(float4*)&x[q];
}


// pair decode shared by the update helpers
__device__ __forceinline__ void pair_decode(int p, int J, int& I, int& K) {
    int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
    while ((Ip + 1) * (Ip + 2) <= 2 * p) ++Ip;
    while (Ip * (Ip + 1) > 2 * p) --Ip;
    I = J + 1 + Ip;
    K = J + 1 + (p - Ip * (Ip + 1) / 2);
}

// prefetch pair p's two 64x64 panels into the shared register array
__device__ __forceinline__ void update_prefetch_dev(const float* __restrict__ M, int n,
                                                    int J, int p, float (&rr)[64]) {
    int I, K;
    pair_decode(p, J, I, K);
    const float* LI = M + (long long)I * 64 * n + (long long)J * 64;
    const float* LK = M + (long long)K * 64 * n + (long long)J * 64;
    #pragma unroll
    for (int j = 0; j < 16; ++j) {
        const int idx = threadIdx.x + j * 256;
        const int r = idx / 64, k = idx % 64;
        rr[j] = LI[r * n + k];
        rr[j + 16] = LK[r * n + k];
    }
}

// spill the prefetched registers into the transposed staging layout
__device__ __forceinline__ void update_stage_dev(const float (&rr)[64], float* smem) {
    float (*sI)[68] = (float (*)[68])smem;
    float (*sK)[68] = (float (*)[68])(smem + 64 * 68);
    #pragma unroll
    for (int j = 0; j < 16; ++j) {
        const int idx = threadIdx.x + j * 256;
        const int r = idx / 64, k = idx % 64;
        sI[k][r] = rr[j];
        sK[k][r] = rr[j + 16];
    }
}

// compute pair p's update from already-staged smem operands
__device__ __forceinline__ void update_compute_dev(float* __restrict__ M, int n, int J,
                                                   int p, float* smem) {
    float (*sI)[68] = (float (*)[68])smem;
    float (*sK)[68] = (float (*)[68])(smem + 64 * 68);
    int I, K;
    pair_decode(p, J, I, K);
    float* C = M + (long long)I * 64 * n + (long long)K * 64;

    const int tx = threadIdx.x % 16, ty = threadIdx.x / 16;
    float acc[4][4] = {};
    #pragma unroll 8
    for (int k = 0; k < 64; ++k) {
        const float4 a = *(const float4*)&sI[k][ty * 4];
        const float4 b = *(const float4*)&sK[k][tx * 4];
        const float ar[4] = {a.x, a.y, a.z, a.w};
        const float br[4] = {b.x, b.y, b.z, b.w};
        #pragma unroll
        for (int i = 0; i < 4; ++i)
            #pragma unroll
            for (int j = 0; j < 4; ++j) acc[i][j] += ar[i] * br[j];
    }
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        float4* cp = (float4*)&C[(ty * 4 + i) * n + tx * 4];
        float4 c = *cp;
        c.x -= acc[i][0]; c.y -= acc[i][1]; c.z -= acc[i][2]; c.w -= acc[i][3];
        *cp = c;
    }
}

// Fused-fast: one CTA per matrix, composing the register-resident potrf
// and 4-wide register trsm with the register-tiled update; only the
// per-step syncs those primitives need. For n in {256, 512, 1024}.
template <int NT, int MINB>
__global__ void __launch_bounds__(256, MINB)
fused_fast_kernel(const float* __restrict__ A, float* __restrict__ Lg,
                  int n, long long bstride) {
    extern __shared__ float scratch[];  // update scratch | sDiag | colbuf
    float* sDiag = scratch + 2 * 64 * 68;
    float* colbuf = sDiag + 64 * 68;
    const float* src = A + blockIdx.x * bstride;
    float* L = Lg + blockIdx.x * bstride;
    for (int idx = threadIdx.x; idx < n * n; idx += NT) {
        const int i = idx / n, j = idx % n;
        L[idx] = (j <= i) ? src[idx] : 0.0f;
    }
    __syncthreads();
    const int N = n / 64;
    float rr[64];
    for (int J = 0; J < N; ++J) {
        potrf64_regs_dev(L + (long long)J * 64 * (n + 1), n, sDiag, colbuf, rr);
        const int T = N - J - 1;
        for (int g = 0; g < T; g += 4)
            trsm4_regs_dev(L, n, J, g, (T - g < 4) ? (T - g) : 4, sDiag, rr);
        __syncthreads();
        const int pairs = T * (T + 1) / 2;
        for (int p = 0; p < pairs; ++p) update_tile_dev<64>(L, n, J, p, scratch);
    }
}

template <int NT, int MINB>
void launch_fused_fast(const float* A, float* L, int n, long long batch) {
    constexpr int SMEM = (2 * 64 * 68 + 64 * 68 + 128) * sizeof(float);
    cudaFuncSetAttribute(fused_fast_kernel<NT, MINB>,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    fused_fast_kernel<NT, MINB><<<batch, NT, SMEM>>>(A, L, n, (long long)n * n);
}

// Grid-wide software barrier for the persistent kernel: generation counter
// plus arrive count. Valid because the grid never exceeds resident capacity.
__device__ void grid_barrier(int* count, volatile int* gen, int nctas) {
    __syncthreads();
    if (threadIdx.x == 0) {
        const int g = *gen;
        if (atomicAdd(count, 1) == nctas - 1) {
            *count = 0;
            __threadfence();
            *gen = g + 1;
        } else {
            while (*gen == g) __nanosleep(64);
        }
        __threadfence();
    }
    __syncthreads();
}

// Whole factorization in one launch: persistent CTAs sweep the potrf, trsm
// and update phases of each block-column, separated by grid barriers.
// Kills the per-column launch spine that dominates low-batch mid sizes.
template <int B>
__global__ void __launch_bounds__(256)
persistent_cholesky_kernel(float* __restrict__ A, int n, long long bstride,
                           int batch, int nctas, int* bar) {
    __shared__ float smem[2 * B * (B + 4)];
    const int N = n / B;
    for (int J = 0; J < N; ++J) {
        for (int it = blockIdx.x; it < batch; it += nctas) {
            float* tile = A + it * bstride + (long long)J * B * (n + 1);
            potrf_tile<B, 256>(tile, tile, n, smem);
        }
        grid_barrier(bar, bar + 1, nctas);
        const int T = N - J - 1;
        if (T == 0) break;
        for (int it = blockIdx.x; it < T * batch; it += nctas)
            trsm_tile_dev<B, 256>(A + (it / T) * bstride, n, J, it % T, smem);
        grid_barrier(bar, bar + 1, nctas);
        const int pairs = T * (T + 1) / 2;
        for (int it = blockIdx.x; it < pairs * batch; it += nctas)
            update_tile_dev<B>(A + (it / pairs) * bstride, n, J, it % pairs, smem);
        grid_barrier(bar, bar + 1, nctas);
    }
}

// Launch config: as many CTAs as can be simultaneously resident. The
// barrier buffer is allocated once; count self-resets and gen only needs
// monotonicity, so no zeroing between calls.
template <int B>
void cholesky_persistent(float* A, int n, long long batch) {
    static int nctas = 0;
    static int* bar = nullptr;
    if (nctas == 0) {
        int dev, sms, per_sm;
        cudaGetDevice(&dev);
        cudaDeviceGetAttribute(&sms, cudaDevAttrMultiProcessorCount, dev);
        cudaOccupancyMaxActiveBlocksPerMultiprocessor(
            &per_sm, persistent_cholesky_kernel<B>, 256, 0);
        nctas = sms * per_sm;
        cudaMalloc(&bar, 2 * sizeof(int));
        cudaMemset(bar, 0, 2 * sizeof(int));
    }
    persistent_cholesky_kernel<B><<<nctas, 256>>>(A, n, (long long)n * n, batch, nctas, bar);
}

// Round-to-nearest fp32 -> tf32 conversion.
__device__ __forceinline__ float to_tf32(float x) {
    unsigned u;
    asm("cvt.rna.tf32.f32 %0, %1;" : "=r"(u) : "f"(x));
    return __uint_as_float(u);
}

// One m16n8k8 TF32 tensor-core MMA, fp32 accumulate.
__device__ __forceinline__ void mma_tf32(float d[4], const unsigned a[4], const unsigned b[2]) {
    asm volatile(
        "mma.sync.aligned.m16n8k8.row.col.f32.tf32.tf32.f32 "
        "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
        : "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3])
        : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]));
}

// 16-byte async global->shared copy plus commit/wait wrappers.
__device__ __forceinline__ void cp_async16(void* dst, const void* src) {
    const unsigned d = (unsigned)__cvta_generic_to_shared(dst);
    asm volatile("cp.async.ca.shared.global [%0], [%1], 16;\n" :: "r"(d), "l"(src));
}
__device__ __forceinline__ void cp_commit() { asm volatile("cp.async.commit_group;\n"); }
template <int N>
__device__ __forceinline__ void cp_wait() { asm volatile("cp.async.wait_group %0;\n" :: "n"(N)); }

// TF32 trailing update over 128x128 supertiles (2x2 groups of 64-tiles).
// 8 warps per CTA in a 2x4 grid, each warp computes a 64x32 slice with
// m16n8k8 MMAs. 16-wide k-slabs are cp.async double-buffered so the next
// slab loads in while the current one feeds the tensor cores. Edge
// subtiles (beyond T or above the diagonal) are masked.
template <int B>
__global__ void __launch_bounds__(256, 1)
update_supertile_tf32_kernel(float* __restrict__ A, int n, long long bstride, int J) {
    constexpr int KW = 16, PAD = 20;
    __shared__ float sI[2][2 * B][PAD], sK[2][2 * B][PAD];

    const int T = n / B - J - 1;
    const int p = blockIdx.x;
    int Sp = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
    while ((Sp + 1) * (Sp + 2) / 2 <= p) ++Sp;
    while (Sp * (Sp + 1) / 2 > p) --Sp;
    const int si = Sp, sk = p - Sp * (Sp + 1) / 2;

    float* M = A + blockIdx.y * bstride;
    const long long panel_col = (long long)J * B;
    const long long row_i = (long long)(J + 1 + 2 * si) * B;
    const long long row_k = (long long)(J + 1 + 2 * sk) * B;

    auto stage = [&](int buf, int ks) {
        for (int idx = threadIdx.x; idx < 2 * B * (KW / 4); idx += 256) {
            const int r = idx / (KW / 4), q = (idx % (KW / 4)) * 4;
            if (2 * si + r / B < T)
                cp_async16(&sI[buf][r][q], M + (row_i + r) * n + panel_col + ks + q);
            else
                *(float4*)&sI[buf][r][q] = make_float4(0.f, 0.f, 0.f, 0.f);
            if (2 * sk + r / B < T)
                cp_async16(&sK[buf][r][q], M + (row_k + r) * n + panel_col + ks + q);
            else
                *(float4*)&sK[buf][r][q] = make_float4(0.f, 0.f, 0.f, 0.f);
        }
    };

    const int warp = threadIdx.x >> 5, lane = threadIdx.x & 31;
    const int wm = warp >> 2, wn = warp & 3;
    const int g = lane >> 2, t = lane & 3;
    float acc[4][4][4] = {};

    stage(0, 0);
    cp_commit();
    constexpr int NSLAB = B / KW;
    for (int s = 0; s < NSLAB; ++s) {
        if (s + 1 < NSLAB) {
            stage((s + 1) & 1, (s + 1) * KW);
            cp_commit();
            cp_wait<1>();
        } else {
            cp_wait<0>();
        }
        __syncthreads();
        const int buf = s & 1;
        #pragma unroll
        for (int koff = 0; koff < KW; koff += 8) {
            unsigned a[4][4], b[4][2];
            #pragma unroll
            for (int mt = 0; mt < 4; ++mt) {
                const int r0 = wm * 64 + mt * 16 + g;
                a[mt][0] = __float_as_uint(to_tf32(sI[buf][r0][koff + t]));
                a[mt][1] = __float_as_uint(to_tf32(sI[buf][r0 + 8][koff + t]));
                a[mt][2] = __float_as_uint(to_tf32(sI[buf][r0][koff + t + 4]));
                a[mt][3] = __float_as_uint(to_tf32(sI[buf][r0 + 8][koff + t + 4]));
            }
            #pragma unroll
            for (int nt = 0; nt < 4; ++nt) {
                const int c0 = wn * 32 + nt * 8 + g;
                b[nt][0] = __float_as_uint(to_tf32(sK[buf][c0][koff + t]));
                b[nt][1] = __float_as_uint(to_tf32(sK[buf][c0][koff + t + 4]));
            }
            #pragma unroll
            for (int mt = 0; mt < 4; ++mt)
                #pragma unroll
                for (int nt = 0; nt < 4; ++nt) mma_tf32(acc[mt][nt], a[mt], b[nt]);
        }
        __syncthreads();
    }

    #pragma unroll
    for (int mt = 0; mt < 4; ++mt) {
        #pragma unroll
        for (int nt = 0; nt < 4; ++nt) {
            #pragma unroll
            for (int e = 0; e < 4; ++e) {
                const int r = wm * 64 + mt * 16 + g + (e >= 2 ? 8 : 0);
                const int c = wn * 32 + nt * 8 + t * 2 + (e & 1);
                const int rt = 2 * si + r / B, ct = 2 * sk + c / B;
                if (rt < T && ct < T && ct <= rt)
                    M[(row_i + r) * n + row_k + c] -= acc[mt][nt][e];
            }
        }
    }
}

#if __CUDA_ARCH__ >= 1000
// 64-bit shared-memory matrix descriptor (no swizzle, K-major,
// core-matrix-tiled storage).
__device__ __forceinline__ unsigned long long smem_desc(unsigned addr, unsigned lbo, unsigned sbo) {
    return (unsigned long long)((addr & 0x3FFFFu) >> 4)
         | ((unsigned long long)((lbo >> 4) & 0x3FFFu) << 16)
         | ((unsigned long long)((sbo >> 4) & 0x3FFFu) << 32)
         | (1ULL << 46);
}
#endif

// Split-TF32 trailing update on 5th-gen tensor cores: same 128x128
// supertile scheme as update_supertile_tf32_kernel, but the 128x128x64
// product runs as K=8 tcgen05 MMA triples (hi*hi + hi*lo + lo*hi, ~fp32
// accuracy) accumulating in tensor memory. Operands are staged in
// core-matrix-tiled smem (8x4 tf32 blocks of 128 contiguous bytes).
template <int B>
__global__ void __launch_bounds__(128)
update_supertile_tc5_kernel(float* __restrict__ A, int n, long long bstride, int J) {
#if __CUDA_ARCH__ >= 1000
    constexpr int M = 2 * B;
    extern __shared__ __align__(1024) float dsm[];
    float* sA = dsm;
    float* sB = dsm + M * 16;
    __shared__ __align__(8) unsigned long long mbar;
    __shared__ unsigned taddr_slot;

    const int T = n / B - J - 1;
    const int p = blockIdx.x;
    int Sp = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
    while ((Sp + 1) * (Sp + 2) / 2 <= p) ++Sp;
    while (Sp * (Sp + 1) / 2 > p) --Sp;
    const int si = Sp, sk = p - Sp * (Sp + 1) / 2;

    float* Mm = A + blockIdx.y * bstride;
    const long long panel_col = (long long)J * B;
    const long long row_i = (long long)(J + 1 + 2 * si) * B;
    const long long row_k = (long long)(J + 1 + 2 * sk) * B;

    const int tid = threadIdx.x;
    const unsigned mbar_a = (unsigned)__cvta_generic_to_shared(&mbar);
    if (tid == 0) {
        asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
        asm volatile("fence.mbarrier_init.release.cluster;");
    }
    if (tid < 32) {
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
                     :: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(M) : "memory");
        asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
    }

    // K is processed in 16-wide slabs so smem stays at 32KB and several
    // CTAs stay resident per SM (TMEM caps residency at 4); hi/lo split
    // panels give ~fp32 accuracy via 3-term MMA (hi*hi + hi*lo + lo*hi)
    constexpr int KS = 16, NSLAB = 64 / KS, PT = M * KS / 128;
    float* sAlo = dsm + 2 * M * KS;
    float* sBlo = dsm + 3 * M * KS;
    __syncthreads();
    const unsigned taddr = taddr_slot;

    const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
                         | ((unsigned)(M >> 3) << 17) | ((unsigned)(M >> 4) << 24);

    // threads prefetch the next slab into registers while the current
    // slab's MMAs run, then wait for the MMAs before rewriting smem
    float ra[PT], rb[PT];
    auto load_regs = [&](int ks) {
        #pragma unroll
        for (int j = 0; j < PT; ++j) {
            const int idx = tid + j * 128;
            const int m = idx / KS, k = idx % KS;
            ra[j] = (2 * si + m / B < T) ? Mm[(row_i + m) * n + panel_col + ks + k] : 0.0f;
            rb[j] = (2 * sk + m / B < T) ? Mm[(row_k + m) * n + panel_col + ks + k] : 0.0f;
        }
    };
    auto wait_mbar = [&](unsigned phase) {
        unsigned done = 0;
        while (!done)
            asm volatile(
                "{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
                "selp.u32 %0, 1, 0, p;\n}\n"
                : "=r"(done) : "r"(mbar_a), "r"(phase));
    };

    load_regs(0);
    for (int s = 0; s < NSLAB; ++s) {
        if (s > 0) {
            wait_mbar((unsigned)((s - 1) & 1));
            asm volatile("tcgen05.fence::before_thread_sync;");
            __syncthreads();
        }
        #pragma unroll
        for (int j = 0; j < PT; ++j) {
            const int idx = tid + j * 128;
            const int m = idx / KS, k = idx % KS;
            const int off = (k / 4) * (M * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
            const float ah = to_tf32(ra[j]), bh = to_tf32(rb[j]);
            sA[off] = ah;
            sB[off] = bh;
            sAlo[off] = to_tf32(ra[j] - ah);
            sBlo[off] = to_tf32(rb[j] - bh);
        }
        asm volatile("fence.proxy.async.shared::cta;");  // generic->async visibility
        __syncthreads();
        if (tid == 0) {
            asm volatile("tcgen05.fence::after_thread_sync;");
            const unsigned aa = (unsigned)__cvta_generic_to_shared(sA);
            const unsigned bb = (unsigned)__cvta_generic_to_shared(sB);
            const unsigned al = (unsigned)__cvta_generic_to_shared(sAlo);
            const unsigned bl = (unsigned)__cvta_generic_to_shared(sBlo);
            #pragma unroll
            for (int c = 0; c < KS / 8; ++c) {
                const unsigned koff = (unsigned)(c * 2 * M * 16);
                const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
                #pragma unroll
                for (int t = 0; t < 3; ++t) {
                    const unsigned long long da = smem_desc(ops[t][0] + koff, M * 16, 128);
                    const unsigned long long db = smem_desc(ops[t][1] + koff, M * 16, 128);
                    asm volatile(
                        "{\n.reg .pred p;\nsetp.ne.b32 p, %4, 0;\n"
                        "tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, p;\n}\n"
                        :: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
                           "r"((s != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
                }
            }
            asm volatile(
                "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                :: "r"(mbar_a) : "memory");
        }
        if (s + 1 < NSLAB) load_regs((s + 1) * KS);
    }
    wait_mbar((unsigned)((NSLAB - 1) & 1));
    asm volatile("tcgen05.fence::before_thread_sync;");
    __syncthreads();

    // each warp reads its 32-row slice of the accumulator, four x8 loads
    // per wait, and subtracts from C with edge masking
    const int warp = tid >> 5, lane = tid & 31;
    const int r = warp * 32 + lane;
    const int rt = 2 * si + r / B;
    const bool interior = (2 * si + 1 < T) && (2 * sk + 1 < T);
    for (int c0 = 0; c0 < M; c0 += 32) {
        float v[32];
        #pragma unroll
        for (int q = 0; q < 4; ++q)
            asm volatile(
                "tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
                : "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
                  "=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
                : "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        float* crow = Mm + (row_i + r) * n + row_k + c0;
        if (interior && si != sk) {
            #pragma unroll
            for (int q = 0; q < 8; ++q) {
                float4* cp = (float4*)(crow + q * 4);
                float4 c4 = *cp;
                c4.x -= v[q * 4];     c4.y -= v[q * 4 + 1];
                c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
                *cp = c4;
            }
        } else if (interior) {
            // diagonal supertile: subtile above the diagonal is masked out
            const int climit = (r < B) ? B : M;
            #pragma unroll
            for (int q = 0; q < 8; ++q) {
                if (c0 + q * 4 < climit) {
                    float4* cp = (float4*)(crow + q * 4);
                    float4 c4 = *cp;
                    c4.x -= v[q * 4];     c4.y -= v[q * 4 + 1];
                    c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
                    *cp = c4;
                }
            }
        } else if (rt < T) {
            #pragma unroll
            for (int e = 0; e < 32; ++e) {
                const int c = c0 + e;
                const int ct = 2 * sk + c / B;
                if (ct < T && ct <= rt)
                    Mm[(row_i + r) * n + row_k + c] -= v[e];
            }
        }
    }
    __syncthreads();
    if (tid < 32)
        asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
                     :: "r"(taddr), "r"(M) : "memory");
#endif
}

// b=128 split-TF32 trailing update: one CTA per 128x128 tile pair (I >= K)
// over a K=128 panel — twice the FLOP per pair of the b=64 supertile path
// and, since n % 128 == 0, no edge masking anywhere. Same slab/MMA/ld
// machinery as update_supertile_tc5_kernel.
__global__ void __launch_bounds__(128)
update_tile_tc5_128_kernel(float* __restrict__ A, int n, long long bstride, int J) {
#if __CUDA_ARCH__ >= 1000
    constexpr int M = 128, KS = 16, NSLAB = 128 / KS, PT = M * KS / 128;
    extern __shared__ __align__(1024) float dsm[];
    float* sA = dsm;
    float* sB = dsm + M * KS;
    float* sAlo = dsm + 2 * M * KS;
    float* sBlo = dsm + 3 * M * KS;
    __shared__ __align__(8) unsigned long long mbar;
    __shared__ unsigned taddr_slot;

    const int p = blockIdx.x;
    int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
    while ((Ip + 1) * (Ip + 2) / 2 <= p) ++Ip;
    while (Ip * (Ip + 1) / 2 > p) --Ip;
    const int I = J + 1 + Ip, K = J + 1 + (p - Ip * (Ip + 1) / 2);

    float* Mm = A + blockIdx.y * bstride;
    const long long panel_col = (long long)J * M;
    const long long row_i = (long long)I * M;
    const long long row_k = (long long)K * M;

    const int tid = threadIdx.x;
    const unsigned mbar_a = (unsigned)__cvta_generic_to_shared(&mbar);
    if (tid == 0) {
        asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
        asm volatile("fence.mbarrier_init.release.cluster;");
    }
    if (tid < 32) {
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
                     :: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(M) : "memory");
        asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
    }
    __syncthreads();
    const unsigned taddr = taddr_slot;
    const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
                         | ((unsigned)(M >> 3) << 17) | ((unsigned)(M >> 4) << 24);

    float ra[PT], rb[PT];
    auto load_regs = [&](int ks) {
        #pragma unroll
        for (int j = 0; j < PT; ++j) {
            const int idx = tid + j * 128;
            const int m = idx / KS, k = idx % KS;
            ra[j] = Mm[(row_i + m) * n + panel_col + ks + k];
            rb[j] = Mm[(row_k + m) * n + panel_col + ks + k];
        }
    };
    auto wait_mbar = [&](unsigned phase) {
        unsigned done = 0;
        while (!done)
            asm volatile(
                "{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
                "selp.u32 %0, 1, 0, p;\n}\n"
                : "=r"(done) : "r"(mbar_a), "r"(phase));
    };

    load_regs(0);
    for (int s = 0; s < NSLAB; ++s) {
        if (s > 0) {
            wait_mbar((unsigned)((s - 1) & 1));
            asm volatile("tcgen05.fence::before_thread_sync;");
            __syncthreads();
        }
        #pragma unroll
        for (int j = 0; j < PT; ++j) {
            const int idx = tid + j * 128;
            const int m = idx / KS, k = idx % KS;
            const int off = (k / 4) * (M * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
            const float ah = to_tf32(ra[j]), bh = to_tf32(rb[j]);
            sA[off] = ah;
            sB[off] = bh;
            sAlo[off] = to_tf32(ra[j] - ah);
            sBlo[off] = to_tf32(rb[j] - bh);
        }
        asm volatile("fence.proxy.async.shared::cta;");
        __syncthreads();
        if (tid == 0) {
            asm volatile("tcgen05.fence::after_thread_sync;");
            const unsigned aa = (unsigned)__cvta_generic_to_shared(sA);
            const unsigned bb = (unsigned)__cvta_generic_to_shared(sB);
            const unsigned al = (unsigned)__cvta_generic_to_shared(sAlo);
            const unsigned bl = (unsigned)__cvta_generic_to_shared(sBlo);
            #pragma unroll
            for (int c = 0; c < KS / 8; ++c) {
                const unsigned koff = (unsigned)(c * 2 * M * 16);
                const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
                #pragma unroll
                for (int t = 0; t < 3; ++t) {
                    const unsigned long long da = smem_desc(ops[t][0] + koff, M * 16, 128);
                    const unsigned long long db = smem_desc(ops[t][1] + koff, M * 16, 128);
                    asm volatile(
                        "{\n.reg .pred q;\nsetp.ne.b32 q, %4, 0;\n"
                        "tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, q;\n}\n"
                        :: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
                           "r"((s != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
                }
            }
            asm volatile(
                "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                :: "r"(mbar_a) : "memory");
        }
        if (s + 1 < NSLAB) load_regs((s + 1) * KS);
    }
    wait_mbar((unsigned)((NSLAB - 1) & 1));
    asm volatile("tcgen05.fence::before_thread_sync;");
    __syncthreads();

    const int warp = tid >> 5, lane = tid & 31;
    const int r = warp * 32 + lane;
    for (int c0 = 0; c0 < M; c0 += 32) {
        float v[32];
        #pragma unroll
        for (int q = 0; q < 4; ++q)
            asm volatile(
                "tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
                : "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
                  "=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
                : "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        float* crow = Mm + (row_i + r) * n + row_k + c0;
        #pragma unroll
        for (int q = 0; q < 8; ++q) {
            float4* cp = (float4*)(crow + q * 4);
            float4 c4 = *cp;
            c4.x -= v[q * 4];     c4.y -= v[q * 4 + 1];
            c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
            *cp = c4;
        }
    }
    __syncthreads();
    if (tid < 32)
        asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
                     :: "r"(taddr), "r"(M) : "memory");
#endif
}

// 128-wide trsm needs 132KB smem, so dynamic (B200 only).
__global__ void trsm_tile_128_kernel(float* __restrict__ A, int n, long long bstride, int J) {
    extern __shared__ float smem[];
    trsm_tile_dev<128, 128>(A + blockIdx.y * bstride, n, J, blockIdx.x, smem);
}

// b=128 blocked driver: fp32 potrf/trsm panels (B200 has the smem for the
// 128-wide trsm), tc5 trailing update. Requires n % 128 == 0.
void cholesky_blocked_128(float* A, int n, long long batch) {
    constexpr int B = 128;
    constexpr int PSMEM = B * (B + 1) * sizeof(float);
    constexpr int TSMEM = 2 * B * (B + 1) * sizeof(float);
    constexpr int USMEM = 56 * 1024;
    const int N = n / B;
    const long long bstride = (long long)n * n;
    cudaFuncSetAttribute(potrf_diag_kernel<B, 256>,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, PSMEM);
    cudaFuncSetAttribute(trsm_tile_128_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, TSMEM);
    cudaFuncSetAttribute(update_tile_tc5_128_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, USMEM);
    for (int J = 0; J < N; ++J) {
        potrf_diag_kernel<B, 256><<<batch, 256, PSMEM>>>(A, n, bstride, J);
        const int T = N - J - 1;
        if (T > 0) {
            trsm_tile_128_kernel<<<dim3(T, batch), B, TSMEM>>>(A, n, bstride, J);
            update_tile_tc5_128_kernel
                <<<dim3(T * (T + 1) / 2, batch), 128, USMEM>>>(A, n, bstride, J);
        }
    }
}

// Right-looking blocked Cholesky, one launch trio per block-column.
// Expects the strict upper triangle of every matrix to already be zero.
template <int B>
void cholesky_blocked(float* A, int n, long long batch, int mode) {
    constexpr int SMEM = B * (B + 1) * sizeof(float);
    const int N = n / B;
    const long long bstride = (long long)n * n;
    for (int J = 0; J < N; ++J) {
        potrf_diag_kernel<B, 128><<<batch, 128, SMEM>>>(A, n, bstride, J);
        const int T = N - J - 1;
        if (T > 0) {
            trsm_tile_kernel<B><<<dim3(T, batch), B>>>(A, n, bstride, J);
            const int S = (T + 1) / 2;
            if (mode == 3) {
                // padded above the 32KB actually used so at most 4 CTAs
                // are resident per SM (4 x 128 TMEM columns = capacity)
                constexpr int SMEM5 = 56 * 1024;
                cudaFuncSetAttribute(update_supertile_tc5_kernel<B>,
                                     cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM5);
                update_supertile_tc5_kernel<B>
                    <<<dim3(S * (S + 1) / 2, batch), 128, SMEM5>>>(A, n, bstride, J);
            } else if (mode == 1) {
                update_supertile_tf32_kernel<B>
                    <<<dim3(S * (S + 1) / 2, batch), 256>>>(A, n, bstride, J);
            } else {
                update_tile_kernel<B>
                    <<<dim3(T * (T + 1) / 2, batch), 256>>>(A, n, bstride, J);
            }
        }
    }
}
// Dynamic dataflow scheduler: persistent CTAs claim nodes from a global
// ticket in topological order (per column: potrf, trsms, update pairs;
// round-robin over matrices) and poll per-tile dependency state. No phase
// barriers, so column J+1 panel work overlaps column J's update tail.
// ucount[I][K] counts columns applied to tile (I,K); readiness is:
// potrf(J): ucount[J][J]==J; trsm(I,J): potrf done && ucount[I][J]==J;
// update(I,K,J): both trsms done && ucount[I][K]==J.
__global__ void __launch_bounds__(256)
sched_cholesky_kernel(float* __restrict__ A, int n, long long bstride, int batch,
                      int* __restrict__ ticket, int* __restrict__ potrf_done,
                      int* __restrict__ trsm_done, int* __restrict__ ucount,
                      long long total_per_m) {
    extern __shared__ float scratch[];
    __shared__ int cur;
    const int N = n / 64;
    const int tid = threadIdx.x;
    for (;;) {
        if (tid == 0) cur = atomicAdd(ticket, 1);
        __syncthreads();
        const long long t = cur;
        __syncthreads();
        if (t >= total_per_m * batch) return;
        const int m = (int)(t % batch);
        long long off = t / batch;
        int J = 0, T, P;
        for (;; ++J) {
            T = N - J - 1;
            P = T * (T + 1) / 2;
            if (off < 1 + T + P) break;
            off -= 1 + T + P;
        }
        float* M = A + (long long)m * bstride;
        volatile int* pd = (volatile int*)(potrf_done + (long long)m * N);
        volatile int* td = (volatile int*)(trsm_done + (long long)m * N * N);
        volatile int* uc = (volatile int*)(ucount + (long long)m * N * N);
        if (off == 0) {
            if (tid == 0) while (uc[J * N + J] < J) __nanosleep(64);
            __syncthreads();
            float* diag = M + (long long)J * 64 * (n + 1);
            potrf_tile<64, 256>(diag, diag, n, scratch);
            if (tid == 0) { __threadfence(); pd[J] = 1; }
        } else if (off <= T) {
            const int I = J + 1 + (int)(off - 1);
            if (tid == 0) while (!pd[J] || uc[I * N + J] < J) __nanosleep(64);
            __syncthreads();
            trsm_tile_dev<64, 256>(M, n, J, I - J - 1, scratch);
            if (tid == 0) { __threadfence(); td[J * N + I] = 1; }
        } else {
            const int p = (int)(off - 1 - T);
            int I, K;
            pair_decode(p, J, I, K);
            if (tid == 0)
                while (!td[J * N + I] || !td[J * N + K] || uc[I * N + K] < J) __nanosleep(64);
            __syncthreads();
            update_tile_dev<64>(M, n, J, p, scratch);
            if (tid == 0) { __threadfence(); atomicAdd((int*)&uc[I * N + K], 1); }
        }
        __syncthreads();
    }
}

// Scheduler state buffers are cached and grown as needed; one memset per
// call resets ticket + flags.
void cholesky_sched(float* A, int n, long long batch) {
    const int N = n / 64;
    long long total = 0;
    for (int J = 0; J < N; ++J) {
        const long long T = N - J - 1;
        total += 1 + T + T * (T + 1) / 2;
    }
    static int* buf = nullptr;
    static long long cap = 0;
    const long long need = 1 + batch * (long long)N + 2 * batch * (long long)N * N;
    if (need > cap) {
        if (buf) cudaFree(buf);
        cudaMalloc(&buf, need * sizeof(int));
        cap = need;
    }
    cudaMemset(buf, 0, need * sizeof(int));
    int* ticket = buf;
    int* pd = buf + 1;
    int* td = pd + batch * N;
    int* uc = td + batch * (long long)N * N;
    const long long work = total * batch;
    const int grid = (int)((work < 444) ? work : 444);
    constexpr int SMEM = 2 * 64 * 68 * sizeof(float);
    sched_cholesky_kernel<<<grid, 256, SMEM>>>(A, n, (long long)n * n, (int)batch,
                                               ticket, pd, td, uc, total);
}


// tc5 split-tf32 supertile update as a scheduler node: the TMEM
// accumulator and mbarrier are CTA-persistent (allocated once by the
// scheduler); `commits` tracks mbarrier phase parity across nodes.
__device__ void sched_tc5_update_dev(float* __restrict__ M, int n, int J, int sp, int T,
                                     float* dsm, unsigned taddr, unsigned mbar_a,
                                     int& commits) {
#if __CUDA_ARCH__ >= 1000
    constexpr int B = 64, MW = 128, KS = 16;
    int Sp = (int)((sqrt(8.0 * sp + 1.0) - 1.0) / 2.0);
    while ((Sp + 1) * (Sp + 2) <= 2 * sp) ++Sp;
    while (Sp * (Sp + 1) > 2 * sp) --Sp;
    const int si = Sp, sk = sp - Sp * (Sp + 1) / 2;
    const long long panel_col = (long long)J * B;
    const long long row_i = (long long)(J + 1 + 2 * si) * B;
    const long long row_k = (long long)(J + 1 + 2 * sk) * B;
    const int tid = threadIdx.x;
    float* sA = dsm;
    float* sB = dsm + MW * KS;
    float* sAlo = dsm + 2 * MW * KS;
    float* sBlo = dsm + 3 * MW * KS;
    const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
                         | ((unsigned)(MW >> 3) << 17) | ((unsigned)(MW >> 4) << 24);

    for (int ks = 0; ks < B; ks += KS) {
        for (int idx = tid; idx < MW * KS; idx += 128) {
            const int m = idx / KS, k = idx % KS;
            const int off = (k / 4) * (MW * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
            const float av = (2 * si + m / B < T) ? M[(row_i + m) * n + panel_col + ks + k] : 0.0f;
            const float bv = (2 * sk + m / B < T) ? M[(row_k + m) * n + panel_col + ks + k] : 0.0f;
            const float ah = to_tf32(av), bh = to_tf32(bv);
            sA[off] = ah;
            sB[off] = bh;
            sAlo[off] = to_tf32(av - ah);
            sBlo[off] = to_tf32(bv - bh);
        }
        asm volatile("fence.proxy.async.shared::cta;");
        __syncthreads();
        if (tid == 0) {
            asm volatile("tcgen05.fence::after_thread_sync;");
            const unsigned aa = (unsigned)__cvta_generic_to_shared(sA);
            const unsigned bb = (unsigned)__cvta_generic_to_shared(sB);
            const unsigned al = (unsigned)__cvta_generic_to_shared(sAlo);
            const unsigned bl = (unsigned)__cvta_generic_to_shared(sBlo);
            #pragma unroll
            for (int c = 0; c < KS / 8; ++c) {
                const unsigned koff = (unsigned)(c * 2 * MW * 16);
                const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
                #pragma unroll
                for (int t = 0; t < 3; ++t) {
                    const unsigned long long da = smem_desc(ops[t][0] + koff, MW * 16, 128);
                    const unsigned long long db = smem_desc(ops[t][1] + koff, MW * 16, 128);
                    asm volatile(
                        "{\n.reg .pred q;\nsetp.ne.b32 q, %4, 0;\n"
                        "tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, q;\n}\n"
                        :: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
                           "r"((ks != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
                }
            }
            asm volatile(
                "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                :: "r"(mbar_a) : "memory");
        }
        ++commits;
        {
            unsigned done = 0;
            const unsigned phase = (unsigned)((commits - 1) & 1);
            while (!done)
                asm volatile(
                    "{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
                    "selp.u32 %0, 1, 0, p;\n}\n"
                    : "=r"(done) : "r"(mbar_a), "r"(phase));
        }
        asm volatile("tcgen05.fence::before_thread_sync;");
        __syncthreads();
    }

    const int warp = tid >> 5, lane = tid & 31;
    const int r = warp * 32 + lane;
    const int rt = 2 * si + r / B;
    const bool interior = (2 * si + 1 < T) && (2 * sk + 1 < T);
    for (int c0 = 0; c0 < MW; c0 += 32) {
        float v[32];
        #pragma unroll
        for (int q = 0; q < 4; ++q)
            asm volatile(
                "tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
                : "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
                  "=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
                : "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        float* crow = M + (row_i + r) * n + row_k + c0;
        if (interior && si != sk) {
            #pragma unroll
            for (int q = 0; q < 8; ++q) {
                float4* cp = (float4*)(crow + q * 4);
                float4 c4 = *cp;
                c4.x -= v[q * 4];     c4.y -= v[q * 4 + 1];
                c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
                *cp = c4;
            }
        } else if (interior) {
            const int climit = (r < B) ? B : MW;
            #pragma unroll
            for (int q = 0; q < 8; ++q) {
                if (c0 + q * 4 < climit) {
                    float4* cp = (float4*)(crow + q * 4);
                    float4 c4 = *cp;
                    c4.x -= v[q * 4];     c4.y -= v[q * 4 + 1];
                    c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
                    *cp = c4;
                }
            }
        } else if (rt < T) {
            #pragma unroll
            for (int e = 0; e < 32; ++e) {
                const int c = c0 + e;
                const int ct = 2 * sk + c / B;
                if (ct < T && ct <= rt)
                    M[(row_i + r) * n + row_k + c] -= v[e];
            }
        }
    }
    __syncthreads();
#endif
}

// Dataflow scheduler with tc5 supertile update nodes: same ticket/poll
// design as sched_cholesky_kernel, but updates are 128x128 supertiles on
// the tensor cores, and each supertile completion bumps all its pair
// counters. TMEM and the mbarrier live for the whole CTA.
__global__ void __launch_bounds__(128)
sched5_cholesky_kernel(float* __restrict__ A, int n, long long bstride, int batch,
                       int* __restrict__ ticket, int* __restrict__ potrf_done,
                       int* __restrict__ trsm_done, int* __restrict__ ucount,
                       long long total_per_m) {
#if __CUDA_ARCH__ >= 1000
    extern __shared__ __align__(1024) float dsm[];
    __shared__ __align__(8) unsigned long long mbar;
    __shared__ unsigned taddr_slot;
    __shared__ int cur;
    const int N = n / 64;
    const int tid = threadIdx.x;
    const unsigned mbar_a = (unsigned)__cvta_generic_to_shared(&mbar);
    if (tid == 0) {
        asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(mbar_a));
        asm volatile("fence.mbarrier_init.release.cluster;");
    }
    if (tid < 32) {
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
                     :: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(128) : "memory");
        asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
    }
    __syncthreads();
    const unsigned taddr = taddr_slot;
    int commits = 0;
    for (;;) {
        if (tid == 0) cur = atomicAdd(ticket, 1);
        __syncthreads();
        const long long t = cur;
        __syncthreads();
        if (t >= total_per_m * batch) break;
        const int m = (int)(t % batch);
        long long off = t / batch;
        int J = 0, T, S;
        for (;; ++J) {
            T = N - J - 1;
            S = (T + 1) / 2;
            const long long cnt = 1 + T + (long long)S * (S + 1) / 2;
            if (off < cnt) break;
            off -= cnt;
        }
        float* M = A + (long long)m * bstride;
        volatile int* pd = (volatile int*)(potrf_done + (long long)m * N);
        volatile int* td = (volatile int*)(trsm_done + (long long)m * N * N);
        volatile int* uc = (volatile int*)(ucount + (long long)m * N * N);
        if (off == 0) {
            if (tid == 0) while (uc[J * N + J] < J) __nanosleep(64);
            __syncthreads();
            float* diag = M + (long long)J * 64 * (n + 1);
            potrf_tile<64, 128>(diag, diag, n, dsm);
            if (tid == 0) { __threadfence(); pd[J] = 1; }
        } else if (off <= T) {
            const int I = J + 1 + (int)(off - 1);
            if (tid == 0) while (!pd[J] || uc[I * N + J] < J) __nanosleep(64);
            __syncthreads();
            trsm_tile_dev<64, 128>(M, n, J, I - J - 1, dsm);
            if (tid == 0) { __threadfence(); td[J * N + I] = 1; }
        } else {
            const int sp = (int)(off - 1 - T);
            int Sp = (int)((sqrt(8.0 * sp + 1.0) - 1.0) / 2.0);
            while ((Sp + 1) * (Sp + 2) <= 2 * sp) ++Sp;
            while (Sp * (Sp + 1) > 2 * sp) --Sp;
            const int si = Sp, sk = sp - Sp * (Sp + 1) / 2;
            if (tid == 0) {
                for (int d1 = 0; d1 < 2; ++d1)
                    for (int d2 = 0; d2 < 2; ++d2) {
                        const int It = 2 * si + d1, Kt = 2 * sk + d2;
                        if (It < T && Kt < T && Kt <= It) {
                            const int I = J + 1 + It, K = J + 1 + Kt;
                            while (!td[J * N + I] || !td[J * N + K] || uc[I * N + K] < J)
                                __nanosleep(64);
                        }
                    }
            }
            __syncthreads();
            sched_tc5_update_dev(M, n, J, sp, T, dsm, taddr, mbar_a, commits);
            if (tid == 0) {
                __threadfence();
                for (int d1 = 0; d1 < 2; ++d1)
                    for (int d2 = 0; d2 < 2; ++d2) {
                        const int It = 2 * si + d1, Kt = 2 * sk + d2;
                        if (It < T && Kt < T && Kt <= It)
                            atomicAdd((int*)&uc[(J + 1 + It) * N + (J + 1 + Kt)], 1);
                    }
            }
        }
        __syncthreads();
    }
    __syncthreads();
    if (tid < 32)
        asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
                     :: "r"(taddr), "r"(128) : "memory");
#endif
}

// 56KB smem pad keeps at most 4 CTAs per SM resident, matching the 512
// TMEM columns so no CTA ever blocks in tcgen05.alloc.
void cholesky_sched5(float* A, int n, long long batch) {
    const int N = n / 64;
    long long total = 0;
    for (int J = 0; J < N; ++J) {
        const long long T = N - J - 1, S = (T + 1) / 2;
        total += 1 + T + S * (S + 1) / 2;
    }
    static int* buf = nullptr;
    static long long cap = 0;
    const long long need = 1 + batch * (long long)N + 2 * batch * (long long)N * N;
    if (need > cap) {
        if (buf) cudaFree(buf);
        cudaMalloc(&buf, need * sizeof(int));
        cap = need;
    }
    cudaMemset(buf, 0, need * sizeof(int));
    int* ticket = buf;
    int* pd = buf + 1;
    int* td = pd + batch * N;
    int* uc = td + batch * (long long)N * N;
    const long long work = total * batch;
    const int grid = (int)((work < 592) ? work : 592);
    constexpr int SMEM = 56 * 1024;
    cudaFuncSetAttribute(sched5_cholesky_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    sched5_cholesky_kernel<<<grid, 128, SMEM>>>(A, n, (long long)n * n, (int)batch,
                                                ticket, pd, td, uc, total);
}


// Split one 128-wide block-column panel into tf32 hi/lo pairs stored
// slab-major in core-matrix-tiled order so pipeline stages can be filled
// with plain contiguous async copies. One CTA per (trailing rowblock,
// 16-wide slab): scratch offset = ((rb * 8 + s) * 2 + {hi,lo}) * 2048.
__global__ void presplit_panel_kernel(const float* __restrict__ L, int n, long long bstride,
                                      int J, float* __restrict__ scratch, long long sstride) {
    const int rb = blockIdx.x, sl = blockIdx.y;
    const float* M = L + blockIdx.z * bstride;
    float* out = scratch + blockIdx.z * sstride + ((long long)(rb * 8 + sl) * 2) * 2048;
    const long long row0 = (long long)(J + 1) * 128 + (long long)rb * 128;
    const long long col0 = (long long)J * 128 + sl * 16;
    for (int idx = threadIdx.x; idx < 128 * 16; idx += 128) {
        const int m = idx / 16, k = idx % 16;
        const int off = (k / 4) * (128 * 4) + (m / 8) * 32 + (m % 8) * 4 + (k % 4);
        const float v = M[(row0 + m) * n + col0 + k];
        const float h = to_tf32(v);
        out[off] = h;
        out[2048 + off] = to_tf32(v - h);
    }
}

// Warp-specialized-style pipelined trailing update for b=128 tiles:
// 3-stage cp.async pipeline over pre-split slabs; the tensor pipe never
// drains (mbarrier wait is only for buffer reuse two slabs back). One CTA
// per 128x128 tile pair, K = 128 in eight 16-wide slabs, split-tf32.
__global__ void __launch_bounds__(128)
update128_pipe_kernel(float* __restrict__ A, int n, long long bstride, int J,
                      const float* __restrict__ scratch, long long sstride) {
#if __CUDA_ARCH__ >= 1000
    constexpr int MW = 128, KS = 16, NSLAB = 8, STG = 3;
    constexpr int CH = MW * KS;  // floats per hi or lo chunk (2048)
    extern __shared__ __align__(1024) float dsm[];  // STG stages x 4 chunks
    __shared__ __align__(8) unsigned long long mbar[STG];
    __shared__ unsigned taddr_slot;

    const int p = blockIdx.x;
    int Ip = (int)((sqrt(8.0 * p + 1.0) - 1.0) / 2.0);
    while ((Ip + 1) * (Ip + 2) <= 2 * p) ++Ip;
    while (Ip * (Ip + 1) > 2 * p) --Ip;
    const int Kp = p - Ip * (Ip + 1) / 2;
    float* M = A + blockIdx.y * bstride;
    const float* scr = scratch + blockIdx.y * sstride;
    const long long row_i = (long long)(J + 1 + Ip) * MW;
    const long long row_k = (long long)(J + 1 + Kp) * MW;

    const int tid = threadIdx.x;
    unsigned mb[STG];
    if (tid == 0) {
        #pragma unroll
        for (int i = 0; i < STG; ++i)
            asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;"
                         :: "r"((unsigned)__cvta_generic_to_shared(&mbar[i])));
        asm volatile("fence.mbarrier_init.release.cluster;");
    }
    #pragma unroll
    for (int i = 0; i < STG; ++i) mb[i] = (unsigned)__cvta_generic_to_shared(&mbar[i]);
    if (tid < 32) {
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
                     :: "r"((unsigned)__cvta_generic_to_shared(&taddr_slot)), "r"(MW) : "memory");
        asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
    }
    __syncthreads();
    const unsigned taddr = taddr_slot;
    const unsigned idesc = (1u << 4) | (2u << 7) | (2u << 10)
                         | ((unsigned)(MW >> 3) << 17) | ((unsigned)(MW >> 4) << 24);

    // stage loader: copy 4 pre-tiled chunks (Ahi, Alo, Khi, Klo) for slab sl
    auto load_stage = [&](int buf, int sl) {
        float* dst = dsm + (long long)buf * 4 * CH;
        const float* si = scr + ((long long)(Ip * 8 + sl) * 2) * CH;
        const float* sk = scr + ((long long)(Kp * 8 + sl) * 2) * CH;
        for (int q = tid; q < CH / 4; q += 128) {
            cp_async16(dst + q * 4, si + q * 4);
            cp_async16(dst + CH + q * 4, si + CH + q * 4);
            cp_async16(dst + 2 * CH + q * 4, sk + q * 4);
            cp_async16(dst + 3 * CH + q * 4, sk + CH + q * 4);
        }
        cp_commit();
    };
    auto wait_mbar = [&](unsigned a, unsigned phase) {
        unsigned done = 0;
        while (!done)
            asm volatile(
                "{\n.reg .pred p;\nmbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n"
                "selp.u32 %0, 1, 0, p;\n}\n"
                : "=r"(done) : "r"(a), "r"(phase));
    };

    int uses[STG] = {0, 0, 0};
    load_stage(0, 0);
    load_stage(1, 1);
    for (int sl = 0; sl < NSLAB; ++sl) {
        const int buf = sl % STG;
        // loads for slab sl were issued 2 stages ago (or in the prologue);
        // allow one newer group to stay in flight — except at the tail,
        // where the newest group IS this slab's and must complete
        if (sl + 2 < NSLAB) cp_wait<1>(); else cp_wait<0>();
        asm volatile("fence.proxy.async.shared::cta;");
        __syncthreads();
        if (tid == 0) {
            asm volatile("tcgen05.fence::after_thread_sync;");
            const unsigned aa = (unsigned)__cvta_generic_to_shared(dsm + (long long)buf * 4 * CH);
            const unsigned al = aa + CH * 4;
            const unsigned bb = aa + 2 * CH * 4;
            const unsigned bl = aa + 3 * CH * 4;
            #pragma unroll
            for (int c = 0; c < KS / 8; ++c) {
                const unsigned koff = (unsigned)(c * 2 * MW * 16);
                const unsigned ops[3][2] = {{aa, bb}, {aa, bl}, {al, bb}};
                #pragma unroll
                for (int t = 0; t < 3; ++t) {
                    const unsigned long long da = smem_desc(ops[t][0] + koff, MW * 16, 128);
                    const unsigned long long db = smem_desc(ops[t][1] + koff, MW * 16, 128);
                    asm volatile(
                        "{\n.reg .pred q;\nsetp.ne.b32 q, %4, 0;\n"
                        "tcgen05.mma.cta_group::1.kind::tf32 [%0], %1, %2, %3, q;\n}\n"
                        :: "r"(taddr), "l"(da), "l"(db), "r"(idesc),
                           "r"((sl != 0 || c != 0 || t != 0) ? 1 : 0) : "memory");
                }
            }
            asm volatile(
                "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                :: "r"(mb[buf]) : "memory");
        }
        ++uses[buf];
        __syncthreads();
        if (sl + 2 < NSLAB) {
            const int nbuf = (sl + 2) % STG;
            // reuse of nbuf requires its previous mma group to be complete
            if (uses[nbuf] > 0) wait_mbar(mb[nbuf], (unsigned)((uses[nbuf] - 1) & 1));
            load_stage(nbuf, sl + 2);
        }
    }
    wait_mbar(mb[(NSLAB - 1) % STG], (unsigned)((uses[(NSLAB - 1) % STG] - 1) & 1));
    asm volatile("tcgen05.fence::before_thread_sync;");
    __syncthreads();

    const int warp = tid >> 5, lane = tid & 31;
    const int r = warp * 32 + lane;
    for (int c0 = 0; c0 < MW; c0 += 32) {
        float v[32];
        #pragma unroll
        for (int q = 0; q < 4; ++q)
            asm volatile(
                "tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
                : "=f"(v[q * 8]), "=f"(v[q * 8 + 1]), "=f"(v[q * 8 + 2]), "=f"(v[q * 8 + 3]),
                  "=f"(v[q * 8 + 4]), "=f"(v[q * 8 + 5]), "=f"(v[q * 8 + 6]), "=f"(v[q * 8 + 7])
                : "r"(taddr + ((unsigned)(warp * 32) << 16) + (unsigned)(c0 + q * 8)));
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        float* crow = M + (row_i + r) * n + row_k + c0;
        #pragma unroll
        for (int q = 0; q < 8; ++q) {
            float4* cp = (float4*)(crow + q * 4);
            float4 c4 = *cp;
            c4.x -= v[q * 4];     c4.y -= v[q * 4 + 1];
            c4.z -= v[q * 4 + 2]; c4.w -= v[q * 4 + 3];
            *cp = c4;
        }
    }
    __syncthreads();
    if (tid < 32)
        asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
                     :: "r"(taddr), "r"(MW) : "memory");
#endif
}

// b=128 blocked driver with the pipelined tensor-core update; panels are
// pre-split once per column. Requires n % 128 == 0.
void cholesky_pipe128(float* A, int n, long long batch) {
    constexpr int B = 128;
    constexpr int PSM = B * (B + 1) * sizeof(float);
    constexpr int TSM = 2 * B * (B + 1) * sizeof(float);
    constexpr int USM = 3 * 4 * 128 * 16 * sizeof(float);
    const int N = n / B;
    const long long bstride = (long long)n * n;
    const long long sstride = (long long)(N - 1) * 8 * 2 * 2048;
    static float* scratch = nullptr;
    static long long scap = 0;
    if (sstride * batch > scap) {
        if (scratch) cudaFree(scratch);
        cudaMalloc(&scratch, sstride * batch * sizeof(float));
        scap = sstride * batch;
    }
    cudaFuncSetAttribute(potrf_diag_kernel<B, 256>,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, PSM);
    cudaFuncSetAttribute(trsm_tile_128_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, TSM);
    cudaFuncSetAttribute(update128_pipe_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, USM);
    for (int J = 0; J < N; ++J) {
        potrf_diag_kernel<B, 256><<<batch, 256, PSM>>>(A, n, bstride, J);
        const int T = N - J - 1;
        if (T > 0) {
            trsm_tile_128_kernel<<<dim3(T, batch), B, TSM>>>(A, n, bstride, J);
            presplit_panel_kernel<<<dim3(T, 8, batch), 128>>>(A, n, bstride, J, scratch, sstride);
            update128_pipe_kernel
                <<<dim3(T * (T + 1) / 2, batch), 128, USM>>>(A, n, bstride, J, scratch, sstride);
        }
    }
}

"""

CUDA_SRC = r"""
#include <cudaTypedefs.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
""" + CUDA_KERNELS + r"""

// mode: 0 = fp32 launch-per-column, 1 = tf32 update, 2 = persistent kernel.
at::Tensor batched_cholesky(at::Tensor A, int64_t mode) {
    const long long batch = A.size(0);
    const int n = A.size(1);
    if (mode == 5 && n > 64 && n % 64 == 0) {
        at::Tensor L = A.new_empty(A.sizes());
        launch_fused_small<64, 256>(A.data_ptr<float>(), L.data_ptr<float>(), n, batch);
        return L;
    }
    if (mode == 12 && n == 256) {
        at::Tensor L = A.new_empty(A.sizes());
        launch_fused256(A.data_ptr<float>(), L.data_ptr<float>(), batch);
        return L;
    }
    if ((mode == 6 || mode == 7) && n > 64 && n % 64 == 0) {
        at::Tensor L = A.new_empty(A.sizes());
        if (mode == 7)
            launch_fused_fast<256, 2>(A.data_ptr<float>(), L.data_ptr<float>(), n, batch);
        else
            launch_fused_fast<256, 1>(A.data_ptr<float>(), L.data_ptr<float>(), n, batch);
        return L;
    }
    if (n == 32 || n == 64 || n == 128) {
        at::Tensor L = A.new_empty(A.sizes());
        const float* src = A.data_ptr<float>();
        float* dst = L.data_ptr<float>();
        if (n == 32)
            potrf32_warp_kernel<<<(batch + 3) / 4, 128>>>(src, dst, batch);
        else if (n == 64)
            potrf64_rows_kernel<<<batch, 64>>>(src, dst, batch);
        else  // n == 128: regs scheme beats the smem-rows kernel at every batch
            potrf128_regs_kernel<<<batch, 256>>>(src, dst, batch);
        return L;
    }
    TORCH_CHECK(n % 64 == 0, "batched_cholesky: unsupported n=", n);
    TORCH_CHECK(batch <= 65535, "batched_cholesky: batch too large");
    if (mode == 13) {
        TORCH_CHECK(n % 128 == 0, "b128 needs n % 128 == 0");
        at::Tensor L = A.new_empty(A.sizes());
        const int nblk4 = (n * (n / 4) + 255) / 256;
        tril_f4_kernel<<<dim3(nblk4 < 512 ? nblk4 : 512, batch), 256>>>(
            A.data_ptr<float>(), L.data_ptr<float>(), n, (long long)n * n);
        cholesky_b128regs(L.data_ptr<float>(), n, batch);
        return L;
    }
    at::Tensor L = A.tril();  // blocked path never touches the upper triangle
    if (mode == 10) {
        TORCH_CHECK(n % 128 == 0, "pipe128 needs n % 128 == 0");
        cholesky_pipe128(L.data_ptr<float>(), n, batch);
    } else if (mode == 9) {
        cholesky_sched5(L.data_ptr<float>(), n, batch);
    } else if (mode == 8) {
        cholesky_sched(L.data_ptr<float>(), n, batch);
    } else if (mode == 4) {
        TORCH_CHECK(n % 128 == 0, "tc5b128 needs n % 128 == 0");
        cholesky_blocked_128(L.data_ptr<float>(), n, batch);
    } else if (mode == 2) {
        cholesky_persistent<64>(L.data_ptr<float>(), n, batch);
    } else {
        cholesky_blocked<64>(L.data_ptr<float>(), n, batch, (int)mode);
    }
    return L;
}

// Left-looking blocked-128 hybrid: per column one strided-batched cuBLAS
// baddbmm_ (tf32-emulated fp32, K grows with the column) applies ALL prior
// panels at once, then the register panel sweep factors the column. The
// caller must enable torch's allow_tf32 matmul flag around this call.
at::Tensor b128_hybrid(at::Tensor A) {
    const long long batch = A.size(0);
    const int n = A.size(1);
    TORCH_CHECK(n % 128 == 0, "b128t needs n % 128 == 0");
    const int nblk = n / 128;
    const long long bstride = (long long)n * n;
    at::Tensor L = A.new_empty(A.sizes());
    const int nblk4 = (n * (n / 4) + 255) / 256;
    tril_f4_kernel<<<dim3(nblk4 < 512 ? nblk4 : 512, batch), 256>>>(
        A.data_ptr<float>(), L.data_ptr<float>(), n, bstride);
    at::Tensor piv = A.new_empty({batch, (long long)nblk, 128, 128});
    float* Lp = L.data_ptr<float>();
    for (int Q = 0; Q < nblk; ++Q) {
        if (Q > 0) {
            at::Tensor Lc = L.narrow(1, 128 * Q, n - 128 * Q).narrow(2, 0, 128 * Q);
            at::Tensor Lr = L.narrow(1, 128 * Q, 128).narrow(2, 0, 128 * Q);
            at::Tensor C  = L.narrow(1, 128 * Q, n - 128 * Q).narrow(2, 128 * Q, 128);
            C.baddbmm_(Lc, Lr.mT(), 1, -1);
        }
        const int cons = nblk - Q - 1;
        panel512_left_kernel<<<dim3(cons > 0 ? cons : 1, batch), 512>>>(
            Lp, piv.data_ptr<float>(), n, bstride, Q);
    }
    if (nblk > 1)
        copy_pivots_kernel<<<dim3(nblk - 1, batch), 256>>>(
            Lp, piv.data_ptr<float>(), n, bstride);
    return L;
}

TORCH_LIBRARY(codex, m) {
    m.def("batched_cholesky(Tensor A, int mode) -> Tensor");
    m.impl("batched_cholesky", &batched_cholesky);
    m.def("b128_hybrid(Tensor A) -> Tensor");
    m.impl("b128_hybrid", &b128_hybrid);
}
"""
BUILD_DIR = Path(__file__).resolve().parent / ".build"
BUILD_DIR.mkdir(exist_ok=True)

_mod = None
_modt = None
try:
    load_inline(
        name="codex",
        cpp_sources=CPP_SRC,
        cuda_sources=CUDA_SRC,
        verbose=True,
        is_python_module=False,
        no_implicit_headers=True,
        extra_cflags=["-O3", "-std=c++17"],
        extra_cuda_cflags=[
            "-O3",
            "-gencode=arch=compute_100a,code=sm_100a",
            "--use_fast_math",
            "--expt-relaxed-constexpr",
            "--relocatable-device-code=false",
            "-lineinfo",
            "-Xptxas=-v",
        ],
        build_directory=str(BUILD_DIR),
    )
    _mod = torch.ops.codex.batched_cholesky
    _modt = torch.ops.codex.b128_hybrid
    P("[cholesky] compiled OK")
except Exception as e:
    P(f"[cholesky] compilation failed (tail of build log follows)")
    P(str(e)[-4000:])

# Blocked right-looking Cholesky from torch ops; the trailing update runs
# cuBLAS TF32 tensor-core kernels, SYRK-aware (only at/below-diagonal
# chunks, in-place). Wins for n >= 8192.
def blocked_cholesky_torch(A: torch.Tensor, b: int) -> torch.Tensor:
    saved = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        n = A.size(-1)
        L = A.tril()
        for s in range(0, n, b):
            e = min(s + b, n)
            Ljj = torch.linalg.cholesky(L[:, s:e, s:e])
            L[:, s:e, s:e] = Ljj
            if e < n:
                X = torch.linalg.solve_triangular(
                    Ljj.mT, L[:, e:, s:e], upper=True, left=False)
                L[:, e:, s:e] = X
                for c in range(e, n, b):
                    ce = min(c + b, n)
                    L[:, c:, c:ce].baddbmm_(
                        X[:, c - e:, :], X[:, c - e:ce - e, :].mT, beta=1, alpha=-1)
        return L.tril()
    finally:
        torch.backends.cuda.matmul.allow_tf32 = saved


# (batch, n) -> fastest measured path on B200 (2026-07-19). "cuda" is our
# kernels, "torch" is batched cuSOLVER, "loop" is per-matrix cuSOLVER,
# "torch_blocked" is blocked_cholesky_torch.
DISPATCH = {
    (4096, 32): "cuda",
    (1024, 64): "cuda",
    (256, 128): "regs128",
    (64, 256): "fused256",
    (16, 512): "b128t",
    (4, 512): "b128t",
    (640, 512): "b128t",
    (4, 1024): "b128t",
    (60, 1024): "b128t",
    (2, 2048): "b128t",
    (8, 2048): "b128t",
    (1, 4096): "torch",
    (2, 4096): "loop",
    (1, 8192): "torch_blocked",
    (1, 16384): "torch_blocked",
    (1, 32768): "torch_blocked",
}

def custom_kernel(data: input_t) -> output_t:
    batch, n = data.size(0), data.size(-1)
    supported = (n in (32, 64, 128) or n % 64 == 0) and batch <= 65535
    mode = DISPATCH.get((batch, n), "cuda" if supported else "torch")
    if mode == "b128t" and _mod is not None and n % 128 == 0:
        # plain tf32 bmm fails the checker (residuals ~2.5); fp32 cuBLAS on
        # B200 is multi-pass emulated and both accurate and fast
        saved = torch.backends.cuda.matmul.allow_tf32
        torch.backends.cuda.matmul.allow_tf32 = False
        try:
            return _modt(data)
        except Exception as e:
            P(f"[cholesky] b128t failed: {e}")
        finally:
            torch.backends.cuda.matmul.allow_tf32 = saved
    if mode in ("cuda", "tf32", "persist", "tc5", "tc5b128", "fused",
                "fusedfast", "fusedfast2", "sched", "sched5", "pipe128",
                "regs128", "fused256", "b128") and _mod is not None and supported:
        try:
            return _mod(data, {"cuda": 0, "tf32": 1, "persist": 2, "tc5": 3,
                               "tc5b128": 4, "fused": 5, "fusedfast": 6,
                               "fusedfast2": 7, "sched": 8, "sched5": 9,
                               "pipe128": 10, "regs128": 11, "fused256": 12,
                               "b128": 13}[mode])
        except Exception as e:
            P(f"[cholesky] kernel failed: {e}")
    if mode == "torch_blocked":
        return blocked_cholesky_torch(data, 2048)
    if mode == "loop":
        out = torch.empty_like(data)
        for i in range(batch):
            torch.linalg.cholesky(data[i], out=out[i])
        return out
    return torch.linalg.cholesky(data)
scrolls · 2290 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON