Skip to content
KernelIndex
Search⌘K

submission 912875

icecuber · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bbe9dba598b007d5f41dad1d04ca84f12420e002cd813d380c7c95dc1087b41e
license declaredunknown
license concludedunknown
authorsicecuber
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float smem[];
vector-width = float4const float4 v = *(const float4 *)(src + t4 * 4);

Kernel source

champ_submit.py942 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

import os

os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "10.0")

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <ATen/cuda/CUDAContext.h>

#define FULL 0xffffffffu

// n == 32: one warp per matrix, lane i owns row i entirely in registers.
// Both loops are fully unrolled so row[] indices stay compile-time constants;
// a dynamic index would force the array into local memory.
__global__ void potrf32_warp(const float *__restrict__ A,
                             float *__restrict__ L, int batch) {
    constexpr int N = 32;
    const int gid = blockIdx.x * blockDim.x + threadIdx.x;
    const int mid = gid >> 5;
    if (mid >= batch) return;
    const int lane = gid & 31;

    const float *src = A + (long long)mid * N * N;
    float *dst = L + (long long)mid * N * N;

    // Stage through shared memory so the global read is contiguous per warp,
    // then each lane picks up its own row. LDA = N + 1 keeps those per-lane
    // row reads spread across banks.
    constexpr int LDA = N + 1;
    extern __shared__ float smem[];
    float *tile = smem + (threadIdx.x >> 5) * N * LDA;

    // float4 on the global side: 8 LDG.128 per matrix instead of 32 LDG.32.
    // The SMEM side stays scalar so LDA=33 keeps the 32 lanes on 32 banks
    // (r = t4>>3, c = (t4&7)<<2 gives base banks {0..3} x {0,4,..28}).
#pragma unroll
    for (int t4 = lane; t4 < N * N / 4; t4 += 32) {
        const float4 v = *(const float4 *)(src + t4 * 4);
        float *d = tile + (t4 >> 3) * LDA + ((t4 & 7) << 2);
        d[0] = v.x; d[1] = v.y; d[2] = v.z; d[3] = v.w;
    }
    __syncwarp();

    float row[N];
#pragma unroll
    for (int j = 0; j < N; ++j) row[j] = tile[lane * LDA + j];

#pragma unroll
    for (int k = 0; k < N; ++k) {
        const float akk = __shfl_sync(FULL, row[k], k);
        const float dk = akk > 0.0f ? sqrtf(akk) : 1.0e-20f;
        // For lane == k this is akk/sqrt(akk) == sqrt(akk), so the diagonal
        // needs no special case. Lanes above k are set to zero.
        const float lk = (lane >= k) ? row[k] / dk : 0.0f;
        row[k] = lk;
#pragma unroll
        for (int j = k + 1; j < N; ++j) {
            const float ljk = __shfl_sync(FULL, lk, j);
            if (lane >= j) row[j] -= lk * ljk;
        }
    }

#pragma unroll
    for (int j = 0; j < N; ++j) tile[lane * LDA + j] = row[j];
    __syncwarp();
#pragma unroll
    for (int t4 = lane; t4 < N * N / 4; t4 += 32) {
        const float *s = tile + (t4 >> 3) * LDA + ((t4 & 7) << 2);
        *(float4 *)(dst + t4 * 4) = make_float4(s[0], s[1], s[2], s[3]);
    }
}

// General small n: one CTA per matrix, tile in shared memory.
// LDA = N + 1 so walking down a column strides across banks.
template <int N>
__global__ void potrf_shared(const float *__restrict__ A,
                             float *__restrict__ L, int batch) {
    constexpr int LDA = N + 1;
    const int mid = blockIdx.x;
    if (mid >= batch) return;

    extern __shared__ float sh[];
    const float *src = A + (long long)mid * N * N;
    float *dst = L + (long long)mid * N * N;
    const int tid = threadIdx.x;
    const int nt = blockDim.x;

    for (int idx = tid; idx < N * N; idx += nt) {
        const int r = idx / N, c = idx - r * N;
        sh[r * LDA + c] = (r >= c) ? src[idx] : 0.0f;
    }
    __syncthreads();

    for (int k = 0; k < N; ++k) {
        if (tid == 0) {
            const float d = sh[k * LDA + k];
            sh[k * LDA + k] = d > 0.0f ? sqrtf(d) : 1.0e-20f;
        }
        __syncthreads();

        const float inv = 1.0f / sh[k * LDA + k];
        for (int i = k + 1 + tid; i < N; i += nt) sh[i * LDA + k] *= inv;
        __syncthreads();

        const int m = N - k - 1;
        for (int t = tid; t < m * m; t += nt) {
            const int jj = t / m;
            const int ii = t - jj * m;
            if (ii >= jj) {
                const int i = k + 1 + ii, j = k + 1 + jj;
                sh[i * LDA + j] -= sh[i * LDA + k] * sh[j * LDA + k];
            }
        }
        __syncthreads();
    }

    for (int idx = tid; idx < N * N; idx += nt) {
        const int r = idx / N;
        dst[idx] = sh[r * LDA + idx - r * N];
    }
}

// Blocked right-looking Cholesky, one CTA per matrix, whole matrix in SMEM.
// Per block column: warp 0 factorizes the NB x NB diagonal block in registers
// via shuffles, then one thread per row does the panel TRSM with no interior
// barrier, then the trailing update runs a 4x4 register tile per thread.
// Barriers drop from 3-per-column to 3-per-block-column and the trailing
// update reads 8 SMEM words per 16 FFMAs (vs ~3 per 1 in potrf_shared).
// PACK stores only the lower triangle, row i starting at i(i+1)/2, which halves
// the SMEM footprint (n=256: 263 KB square -> 132 KB, the difference between not
// fitting and fitting in a B200 CTA). Every access below is (row, col<=row) or
// within a row, so one row-base helper covers both layouts.
template <int N, int NB, int TX, bool PACK>
struct Layout {
    // Plus an NB x N scratch panel (see potrf_blk) appended after the triangle.
    static constexpr size_t tri = PACK ? (size_t)N * (N + 1) / 2
                                       : (size_t)N * (N + 1);
    // Plus NB words holding 1/L[j][j] of the current diagonal block.
    static constexpr size_t words = tri + (size_t)NB * N + NB;
    __device__ __forceinline__ static int base(int i) {
        return PACK ? ((i * (i + 1)) >> 1) : (i * (N + 1));
    }
};

// TM is the tile height in rows (tile width is fixed at 4, one float4 load).
// TM = 8 halves the v-loads per FFMA; __launch_bounds__ is what keeps ptxas
// inside the 64-register budget a 1024-thread block needs.
template <int N, int NB, int TX, bool PACK, int TM>
__global__ __launch_bounds__(TX) void potrf_blk(const float *__restrict__ A,
                          float *__restrict__ L, int batch) {
    using LO = Layout<N, NB, TX, PACK>;
    const int mid = blockIdx.x;
    if (mid >= batch) return;

    extern __shared__ float sh[];
    const float *src = A + (long long)mid * N * N;
    float *dst = L + (long long)mid * N * N;
    const int tid = threadIdx.x;

    // 1/L[j][j] for the current diagonal block, published by the diagonal warp
    // so the TRSM multiplies instead of dividing: the TRSM's serial chain is
    // NB dependent divides deep and sits on the critical path.
    float *dinv = sh + LO::tri + (size_t)NB * N;

    // Read the full square from global (coalesced) but keep only tril.
    for (int idx = tid; idx < N * N; idx += TX) {
        const int r = idx / N, c = idx - r * N;
        if (r >= c) sh[LO::base(r) + c] = src[idx];
        else if (!PACK) sh[LO::base(r) + c] = 0.0f;
    }
    __syncthreads();

    for (int k = 0; k < N; k += NB) {
        // --- diagonal block, warp 0, lane p owns row k+p ---
        if (tid < 32) {
            const int lane = tid;
            float r[NB];
#pragma unroll
            // Only p <= lane is in the triangle; under PACK the slots past it
            // belong to row k+lane+1, so they must not be read or written.
            for (int p = 0; p < NB; ++p)
                r[p] = (lane < NB && p <= lane) ? sh[LO::base(k + lane) + k + p]
                                                : 0.0f;
#pragma unroll
            for (int c = 0; c < NB; ++c) {
                const float acc = __shfl_sync(FULL, r[c], c);
                // rsqrt is one MUFU op; sqrt + divide is two dependent ones, and
                // lane c gets acc * rsqrt(acc) == sqrt(acc), so the diagonal is
                // still exact-in-kind. 1e20 reproduces the old 1/1e-20 guard.
                const float dc = acc > 0.0f ? rsqrtf(acc) : 1.0e20f;
                const float lc = (lane >= c) ? r[c] * dc : 0.0f;
                r[c] = lc;
                if (lane == c) dinv[c] = dc;
#pragma unroll
                for (int j = c + 1; j < NB; ++j) {
                    const float ljc = __shfl_sync(FULL, lc, j);
                    if (lane >= j) r[j] -= lc * ljc;
                }
            }
            if (lane < NB) {
#pragma unroll
                for (int p = 0; p < NB; ++p)
                    if (p <= lane) sh[LO::base(k + lane) + k + p] = r[p];
            }
        }
        __syncthreads();
        const int k1 = k + NB;
        if (k1 >= N) break;

        // --- panel TRSM: row i solves x * Lkk^T = A[i, k:k1] ---
        // The solved row is written twice: into the packed triangle, and
        // transposed into `pan` (p-major, row-index minor). The trailing update
        // then reads four consecutive rows with one 128-bit LDS instead of four
        // scattered ones -- under PACK, row i lives at i(i+1)/2, so a warp's
        // 32 lanes hit near-random banks and the update was LDS-conflict bound.
        float *pan = sh + LO::tri;
        for (int i = k1 + tid; i < N; i += TX) {
            const int bi = LO::base(i);
            float x[NB];
#pragma unroll
            for (int p = 0; p < NB; ++p) x[p] = sh[bi + k + p];
#pragma unroll
            for (int j = 0; j < NB; ++j) {
                const int bj = LO::base(k + j);
                float s = x[j];
#pragma unroll
                for (int p = 0; p < j; ++p) s -= x[p] * sh[bj + k + p];
                x[j] = s * dinv[j];
            }
#pragma unroll
            for (int p = 0; p < NB; ++p) {
                sh[bi + k + p] = x[p];
                pan[p * N + (i - k1)] = x[p];
            }
        }
        __syncthreads();

        // --- trailing update, lower triangle only, 4x4 register tile ---
        // N and NB are both multiples of 4, so m = N - k1 is too and every
        // 4-row group is whole: no tail guard, and the loads can be float4.
        const int m = N - k1;
        const int TI = m / TM;
        const int TJ = m >> 2;
        // Enumerate only the tiles at or below the diagonal. Striding over the
        // full square grid and skipping the upper half wasted half the loop
        // trips. Tile row ti owns R*(ti+1) tiles (R = TM/4 tile-columns per
        // tile-row of height TM), so tiles before row ti number R*ti(ti+1)/2.
        constexpr int R = TM / 4;
        const int NT = R * ((TI * (TI + 1)) >> 1);
        for (int t = tid; t < NT; t += TX) {
            // Invert the cumulative count. Fast-math sqrt is approximate, so
            // nudge into range rather than trusting it.
            int ti = (int)((sqrtf(1.0f + 8.0f * (float)t / (float)R) - 1.0f) * 0.5f);
            while (R * (((ti + 1) * (ti + 2)) >> 1) <= t) ++ti;
            while (R * ((ti * (ti + 1)) >> 1) > t) --ti;
            const int tj = t - R * ((ti * (ti + 1)) >> 1);
            const int i0 = k1 + ti * TM, j0 = k1 + tj * 4;
            float acc[TM][4];
#pragma unroll
            for (int a = 0; a < TM; ++a)
#pragma unroll
                for (int b = 0; b < 4; ++b) acc[a][b] = 0.0f;
#pragma unroll
            for (int p = 0; p < NB; ++p) {
                const float *row = pan + p * N;
                float u[TM];
#pragma unroll
                for (int a = 0; a < TM; a += 4) {
                    const float4 u4 = *(const float4 *)(row + ti * TM + a);
                    u[a] = u4.x; u[a + 1] = u4.y; u[a + 2] = u4.z; u[a + 3] = u4.w;
                }
                const float4 v4 = *(const float4 *)(row + tj * 4);
                const float v[4] = {v4.x, v4.y, v4.z, v4.w};
#pragma unroll
                for (int a = 0; a < TM; ++a)
#pragma unroll
                    for (int b = 0; b < 4; ++b) acc[a][b] += u[a] * v[b];
            }
#pragma unroll
            for (int a = 0; a < TM; ++a) {
                const int i = i0 + a;
                const int bi = LO::base(i);
#pragma unroll
                for (int b = 0; b < 4; ++b) {
                    const int j = j0 + b;
                    if (j <= i) sh[bi + j] -= acc[a][b];
                }
            }
        }
        __syncthreads();
    }

    for (int idx = tid; idx < N * N; idx += TX) {
        const int r = idx / N, c = idx - r * N;
        dst[idx] = (r >= c) ? sh[LO::base(r) + c] : 0.0f;
    }
}

// Two-level blocked variant. potrf_blk writes every trailing tile back to SMEM
// once per NB-wide column block (N/NB times); the FFMA:writeback ratio there is
// NB:1. Here the right-looking update inside an NBO-wide outer block is confined
// to that block's own columns, and the trailing submatrix is updated once per
// outer block with depth NBO -- NBO:1, so N/NBO passes instead of N/NB.
// pan holds the whole NBO-wide panel, indexed by absolute row.
template <int N, int NB, int NBO, int TX, int TM>
struct Layout2 {
    static constexpr size_t tri = (size_t)N * (N + 1) / 2;
    static constexpr size_t words = tri + (size_t)NBO * N + NB;
    __device__ __forceinline__ static int base(int i) { return (i * (i + 1)) >> 1; }
};

template <int N, int NB, int NBO, int TX, int TM>
__global__ __launch_bounds__(TX) void potrf_blk2(const float *__restrict__ A,
                          float *__restrict__ L, int batch) {
    using LO = Layout2<N, NB, NBO, TX, TM>;
    const int mid = blockIdx.x;
    if (mid >= batch) return;

    extern __shared__ float sh[];
    const float *src = A + (long long)mid * N * N;
    float *dst = L + (long long)mid * N * N;
    const int tid = threadIdx.x;
    float *pan = sh + LO::tri;
    float *dinv = pan + (size_t)NBO * N;

    for (int idx = tid; idx < N * N; idx += TX) {
        const int r = idx / N, c = idx - r * N;
        if (r >= c) sh[LO::base(r) + c] = src[idx];
    }
    __syncthreads();

    for (int k0 = 0; k0 < N; k0 += NBO) {
        const int kend = k0 + NBO;  // N is a multiple of NBO
        for (int k = k0; k < kend; k += NB) {
            // --- diagonal block, warp 0, lane p owns row k+p (as potrf_blk) ---
            if (tid < 32) {
                const int lane = tid;
                float r[NB];
#pragma unroll
                for (int p = 0; p < NB; ++p)
                    r[p] = (lane < NB && p <= lane) ? sh[LO::base(k + lane) + k + p]
                                                    : 0.0f;
#pragma unroll
                for (int c = 0; c < NB; ++c) {
                    const float acc = __shfl_sync(FULL, r[c], c);
                    const float dc = acc > 0.0f ? rsqrtf(acc) : 1.0e20f;
                    const float lc = (lane >= c) ? r[c] * dc : 0.0f;
                    r[c] = lc;
                    if (lane == c) dinv[c] = dc;
#pragma unroll
                    for (int j = c + 1; j < NB; ++j) {
                        const float ljc = __shfl_sync(FULL, lc, j);
                        if (lane >= j) r[j] -= lc * ljc;
                    }
                }
                if (lane < NB) {
#pragma unroll
                    for (int p = 0; p < NB; ++p)
                        if (p <= lane) sh[LO::base(k + lane) + k + p] = r[p];
                }
            }
            __syncthreads();
            const int k1 = k + NB;
            if (k1 >= N) break;

            // --- panel TRSM, all the way to row N-1: the solved rows are the
            // outer block's panel and are consumed by the big update below. ---
            const int po = k - k0;
            for (int i = k1 + tid; i < N; i += TX) {
                const int bi = LO::base(i);
                float x[NB];
#pragma unroll
                for (int p = 0; p < NB; ++p) x[p] = sh[bi + k + p];
#pragma unroll
                for (int j = 0; j < NB; ++j) {
                    const int bj = LO::base(k + j);
                    float s = x[j];
#pragma unroll
                    for (int p = 0; p < j; ++p) s -= x[p] * sh[bj + k + p];
                    x[j] = s * dinv[j];
                }
#pragma unroll
                for (int p = 0; p < NB; ++p) {
                    sh[bi + k + p] = x[p];
                    pan[(po + p) * N + i] = x[p];
                }
            }
            __syncthreads();

            // --- narrow update: only the outer block's own columns, so the
            // next inner step sees an up-to-date panel. One thread per element;
            // the region is at most (N-k1) x (NBO-NB), a small fraction of the
            // work the big update below does. ---
            const int W = kend - k1;
            if (W > 0) {
                // 2x2 register tile: 2 LDS.64 feed 4 FFMAs (1.5 instr/FFMA) where
                // the one-thread-per-element form issued 2 LDS + 1 FFMA (3.0),
                // and it costs one integer divide per tile instead of per
                // element. 2x2 rather than 4x4 because at n=256 the region is
                // only (N-k1) x 24, and 4x4 leaves 1024 threads with ~370 tiles.
                const int nct = W >> 1;              // W is a multiple of NB
                const int nt2 = ((N - k1) >> 1) * nct;
                for (int t = tid; t < nt2; t += TX) {
                    const int r = t / nct;
                    const int i0 = k1 + (r << 1), j0 = k1 + ((t - r * nct) << 1);
                    if (j0 > i0 + 1) continue;
                    float acc[2][2] = {{0.0f, 0.0f}, {0.0f, 0.0f}};
#pragma unroll
                    for (int p = 0; p < NB; ++p) {
                        const float *row = pan + (po + p) * N;
                        const float2 u = *(const float2 *)(row + i0);
                        const float2 v = *(const float2 *)(row + j0);
                        acc[0][0] += u.x * v.x; acc[0][1] += u.x * v.y;
                        acc[1][0] += u.y * v.x; acc[1][1] += u.y * v.y;
                    }
#pragma unroll
                    for (int a = 0; a < 2; ++a) {
                        const int bi = LO::base(i0 + a);
#pragma unroll
                        for (int b = 0; b < 2; ++b)
                            if (j0 + b <= i0 + a) sh[bi + j0 + b] -= acc[a][b];
                    }
                }
                __syncthreads();
            }
        }

        // --- big update: trailing submatrix below/right of the outer block,
        // accumulated over all NBO panel columns and written back once. ---
        if (kend >= N) break;
        const int m = N - kend;
        const int TI = m / TM;
        constexpr int R = TM / 4;
        const int NT = R * ((TI * (TI + 1)) >> 1);
        for (int t = tid; t < NT; t += TX) {
            int ti = (int)((sqrtf(1.0f + 8.0f * (float)t / (float)R) - 1.0f) * 0.5f);
            while (R * (((ti + 1) * (ti + 2)) >> 1) <= t) ++ti;
            while (R * ((ti * (ti + 1)) >> 1) > t) --ti;
            const int tj = t - R * ((ti * (ti + 1)) >> 1);
            const int i0 = kend + ti * TM, j0 = kend + tj * 4;
            float acc[TM][4];
#pragma unroll
            for (int a = 0; a < TM; ++a)
#pragma unroll
                for (int b = 0; b < 4; ++b) acc[a][b] = 0.0f;
#pragma unroll 4
            for (int p = 0; p < NBO; ++p) {
                const float *row = pan + p * N;
                float u[TM];
#pragma unroll
                for (int a = 0; a < TM; a += 4) {
                    const float4 u4 = *(const float4 *)(row + i0 + a);
                    u[a] = u4.x; u[a + 1] = u4.y; u[a + 2] = u4.z; u[a + 3] = u4.w;
                }
                const float4 v4 = *(const float4 *)(row + j0);
                const float v[4] = {v4.x, v4.y, v4.z, v4.w};
#pragma unroll
                for (int a = 0; a < TM; ++a)
#pragma unroll
                    for (int b = 0; b < 4; ++b) acc[a][b] += u[a] * v[b];
            }
#pragma unroll
            for (int a = 0; a < TM; ++a) {
                const int i = i0 + a;
                const int bi = LO::base(i);
#pragma unroll
                for (int b = 0; b < 4; ++b) {
                    const int j = j0 + b;
                    if (j <= i) sh[bi + j] -= acc[a][b];
                }
            }
        }
        __syncthreads();
    }

    for (int idx = tid; idx < N * N; idx += TX) {
        const int r = idx / N, c = idx - r * N;
        dst[idx] = (r >= c) ? sh[LO::base(r) + c] : 0.0f;
    }
}

template <int N, int NB, int NBO, int TX, int TM>
static bool launch_blk2(const at::Tensor &A, at::Tensor &L, int batch) {
    const size_t shmem = Layout2<N, NB, NBO, TX, TM>::words * sizeof(float);
    auto kernel = potrf_blk2<N, NB, NBO, TX, TM>;
    if (shmem > 48 * 1024) {
        if (cudaFuncSetAttribute(kernel,
                                 cudaFuncAttributeMaxDynamicSharedMemorySize,
                                 (int)shmem) != cudaSuccess) {
            cudaGetLastError();
            return false;
        }
    }
    kernel<<<batch, TX, shmem>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch);
    if (cudaGetLastError() != cudaSuccess) return false;
    return true;
}

// Returns false if the device cannot give this CTA the shared memory it needs
// (packed n=256 wants 132 KB, which sm_89 cannot opt into), so the caller can
// fall back to cuSOLVER instead of launching a kernel that would fail.
template <int N, int NB, int TX, bool PACK, int TM>
static bool launch_blk(const at::Tensor &A, at::Tensor &L, int batch) {
    const size_t shmem = Layout<N, NB, TX, PACK>::words * sizeof(float);
    auto kernel = potrf_blk<N, NB, TX, PACK, TM>;
    if (shmem > 48 * 1024) {
        if (cudaFuncSetAttribute(kernel,
                                 cudaFuncAttributeMaxDynamicSharedMemorySize,
                                 (int)shmem) != cudaSuccess) {
            cudaGetLastError();
            return false;
        }
    }
    kernel<<<batch, TX, shmem>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch);
    // A launch that fails on resources (TX * regs > 64K) would otherwise leave L
    // uninitialised and silently return garbage.
    if (cudaGetLastError() != cudaSuccess) return false;
    return true;
}

template <int N>
static void launch_shared(const at::Tensor &A, at::Tensor &L, int batch,
                          int threads) {
    const size_t shmem = (size_t)N * (N + 1) * sizeof(float);
    auto kernel = potrf_shared<N>;
    if (shmem > 48 * 1024) {
        cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
                             (int)shmem);
    }
    kernel<<<batch, threads, shmem>>>(A.data_ptr<float>(),
                                      L.data_ptr<float>(), batch);
}

// C += alpha * A @ B^T with fp16 inputs and fp32 accumulate/output.
// torch cannot express fp16-in/fp32-out, and materialising an fp16 product
// would cost tens of GB of extra traffic at n=32768. beta=1 accumulates
// straight into the (strided, non-contiguous) view of the trailing matrix.
// fp16 and tf32 share a 10-bit mantissa, so this is accuracy-neutral versus
// the tf32 path at roughly twice the tensor-core rate.
void gemm_nt_f16(at::Tensor C, at::Tensor A, at::Tensor B, double alpha) {
    TORCH_CHECK(A.scalar_type() == at::kHalf && B.scalar_type() == at::kHalf);
    TORCH_CHECK(C.scalar_type() == at::kFloat);
    TORCH_CHECK(A.stride(-1) == 1 && B.stride(-1) == 1 && C.stride(-1) == 1,
                "inner dimension must be contiguous");
    const int m = (int)A.size(-2), k = (int)A.size(-1), nn = (int)B.size(-2);
    TORCH_CHECK((int)B.size(-1) == k && (int)C.size(-2) == m &&
                (int)C.size(-1) == nn);
    int batch = 1;
    long long sa = 0, sb = 0, sc = 0;
    if (A.dim() == 3) {
        batch = (int)A.size(0);
        sa = A.stride(0); sb = B.stride(0); sc = C.stride(0);
    }
    const float al = (float)alpha, be = 1.0f;
    // Row-major C(m x nn) is column-major C^T(nn x m) = B_rm * A_rm^T, which in
    // column-major terms is op(B)=T over the k x nn buffer times op(A)=N.
    auto st = cublasGemmStridedBatchedEx(
        at::cuda::getCurrentCUDABlasHandle(), CUBLAS_OP_T, CUBLAS_OP_N,
        nn, m, k, &al,
        (const void *)B.data_ptr<at::Half>(), CUDA_R_16F, (int)B.stride(-2), sb,
        (const void *)A.data_ptr<at::Half>(), CUDA_R_16F, (int)A.stride(-2), sa,
        &be,
        (void *)C.data_ptr<float>(), CUDA_R_32F, (int)C.stride(-2), sc,
        batch, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP);
    TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasGemmStridedBatchedEx failed");
}

at::Tensor potrf(at::Tensor A) {
    TORCH_CHECK(A.dim() == 3 && A.size(1) == A.size(2), "expected batch x n x n");
    TORCH_CHECK(A.scalar_type() == at::kFloat, "expected float32");
    A = A.contiguous();
    const int batch = (int)A.size(0);
    const int n = (int)A.size(1);
    auto L = at::empty_like(A);

    if (n == 32) {
        constexpr int TPB = 128;
        const int warps = TPB / 32;
        const int blocks = (batch + warps - 1) / warps;
        const size_t shmem = (size_t)warps * 32 * 33 * sizeof(float);
        potrf32_warp<<<blocks, TPB, shmem>>>(A.data_ptr<float>(),
                                             L.data_ptr<float>(), batch);
        return L;
    }
    bool ok;
    // Measured and rejected: n=32 through launch_blk<32,8,64,true,4> (18.1 ->
    // 36.1 us). At n=32 the blocked kernel is all diagonal block and TRSM with
    // a 24x24 trailing matrix -- there is no update to amortise the barriers.
    // TM = 8 was measured on B200 and lost: it halves the tile count, and
    // both shapes are already short of tiles (256x128 45.2 -> 73.3 us).
    // NBO = 64 at n=256 lost (87.2 -> 91.8): narrow-update work grows as NBO^2.
    // NBO = 16 wins where the trailing matrix is small (n=64 22.6 -> 21.2,
    // n=128 37.0 -> 34.9) but loses at n=256 (88.2 -> 90.0), where the halved
    // big-update depth costs more than the halved narrow update.
    switch (n) {
        case 64:  ok = launch_blk2<64, 8, 16, 128, 4>(A, L, batch);   break;
        case 128: ok = launch_blk2<128, 8, 16, 512, 4>(A, L, batch);  break;
        case 256: ok = launch_blk2<256, 8, 32, 1024, 4>(A, L, batch); break;
        default: TORCH_CHECK(false, "unsupported n");
    }
    // Empty result tells the caller to use the library path (n=256 needs 132 KB
    // of SMEM per CTA, which only Hopper/Blackwell can opt into).
    if (!ok) return at::empty({0}, A.options());
    return L;
}
"""

CPP_SRC = (
    "at::Tensor potrf(at::Tensor A);\n"
    "void gemm_nt_f16(at::Tensor C, at::Tensor A, at::Tensor B, double alpha);"
)

_ext = load_inline(
    name="cholesky_adapt",
    cpp_sources=CPP_SRC,
    cuda_sources=CUDA_SRC,
    functions=["potrf", "gemm_nt_f16"],
    extra_cuda_cflags=["-O3", "--expt-relaxed-constexpr", "-use_fast_math"],
    extra_ldflags=["-lcublas"],
    verbose=False,
)

_CUSTOM_N = (32, 64, 128, 256)




# Blocked right-looking Cholesky for large n. The diagonal blocks and panel
# solves stay in fp32; the trailing updates, which hold essentially all of the
# n^3/3 flops, run on TF32 tensor cores. cuSOLVER does not use TF32 for potrf,
# which is where the headroom comes from.
# Gate on total work, not n: 1x4096 loses to cuSOLVER but 2x4096 wins by 1.4x,
# so batch matters as much as size.
_TF32_MIN_WORK = 1.0e11
_NB = 2048
# Base block for _tri_inv. The fp32 base solve dominates, not the launch count,
# so smaller is better: B200 1x32768 measured 512 -> 37.9ms, 256 -> 37.6,
# 128 -> 29.5, 64 -> 28.8.
_TRI_BASE = int(os.environ.get("CHOL_TRI_BASE", 64))
# Same knob for the mid regime (n in 512/1024). Smaller wins there too:
# B200 640x512 measured 64 -> 1854, 32 -> 1854/1067, 16 -> 1812.
_MID_TRI_BASE = int(os.environ.get("CHOL_MID_TRI_BASE", 16))


def _h(x):
    # solve_triangular returns a column-major result and .half() preserves those
    # strides, which the cuBLAS wrapper rejects. Cast and transpose in one pass.
    return x.to(torch.float16, memory_format=torch.contiguous_format)


def _blk(T, off, cnt, s, step):
    # cnt sub-blocks of size s x s, the first at element offset `off` into T's
    # storage, successive ones `step` elements apart. Regular strides, so the
    # whole set is one batched op.
    lead = tuple(T.shape[:-2])
    lstr = tuple(T.stride()[:-2])
    return T.as_strided(
        lead + (cnt, s, s),
        lstr + (step, T.stride(-2), 1),
        T.storage_offset() + off,
    )


def _tri_inv(Lkk, base):
    # Inverse of a lower-triangular m x m block, bottom-up 2x2 recursion:
    #   inv([[A,0],[C,B]]) = [[Ai,0],[-Bi C Ai, Bi]]
    # Every level is ONE strided-batched op over all blocks of that level, so the
    # whole inverse costs ~9 launches regardless of `base`, and only base^2/m^2 of
    # the work is the non-tensor-core fp32 solve. A per-block recursion instead
    # emits ~25 small launches per step and measured slower on B200.
    # cuSOLVER hands back a column-major L; _blk assumes stride(-1) == 1, and
    # reading L^T as lower-triangular silently yields a diagonal-only "inverse"
    # (-> indefinite trailing matrix -> NaN), so normalise the layout first.
    Lkk = Lkk.contiguous()
    m = Lkk.shape[-1]
    ld = Lkk.stride(-2)
    X = torch.zeros_like(Lkk)
    s = min(base, m)
    cnt = m // s
    D = _blk(Lkk, 0, cnt, s, s * (ld + 1))
    eye = torch.eye(s, device=Lkk.device, dtype=Lkk.dtype).expand_as(D)
    _blk(X, 0, cnt, s, s * (ld + 1)).copy_(
        torch.linalg.solve_triangular(D, eye, upper=False, left=True)
    )
    while s < m:
        cnt = m // (2 * s)
        step = 2 * s * (ld + 1)
        Ai = _blk(X, 0, cnt, s, step)
        Bi = _blk(X, s * (ld + 1), cnt, s, step)
        C = _blk(Lkk, s * ld, cnt, s, step)
        _blk(X, s * ld, cnt, s, step).copy_(-(Bi @ C @ Ai))
        s *= 2
    return X


# Left-looking crossover. Right-looking rewrites the whole trailing matrix once
# per step (~8n^3/(3nb) bytes of fp32 read-modify-write, ~45 GB at n=32768) plus
# an n^2 clone; left-looking touches each block column exactly once and pulls its
# update out of an fp16 mirror (~n^3/(3nb) bytes, ~6 GB) with no clone. That only
# pays once there are enough block columns to amortise the mirror: measured on
# B200, n=8192 loses (3910 -> 4130, four columns), n=16384 is a wash (+0.7%),
# n=32768 wins (28000 -> 27100, reproduced twice).
_LEFT_MIN_N = int(os.environ.get("CHOL_LEFT_MIN_N", 32768))
# Must be _TRI_BASE * 2^k for _tri_inv. 1024 measured worse on B200 (1x32768
# 26700 -> 27300): halving nb doubles the number of cuSOLVER diagonal potrf
# calls and C(nb)/nb only improves slightly, while the left-looking traffic it
# would save is already the small term.
_LEFT_NB = int(os.environ.get("CHOL_LEFT_NB", 2048))


def _left_tf32(A, nb):
    # Left-looking blocked Cholesky, batch 1, fp16 trailing math.
    # Block column k is A[k:, k:e] - Lh[k:, :k] @ Lh[k:e, :k]^T, factored in
    # place; the fp32 working buffer is n x nb, not n x n.
    n = A.shape[-1]
    # Must be zeros, not empty: the panel gemm_nt_f16 accumulates (beta=1) into
    # L[e:, k:e], so that region has to start at zero.
    L = torch.zeros_like(A)
    # empty, not zeros: step k reads only Lh[k:, :k], and every one of those rows
    # was written by an earlier step's diagonal block or panel. Skips a 2 GB
    # memset at n=32768.
    Lh = torch.empty(n, n, device=A.device, dtype=torch.float16)
    buf = torch.empty(n, nb, device=A.device, dtype=torch.float32)
    for k in range(0, n, nb):
        e = min(k + nb, n)
        w = e - k
        W = buf[: n - k, :w]
        W.copy_(A[k:, k:e])
        if k:
            _ext.gemm_nt_f16(W, Lh[k:, :k], Lh[k:e, :k], -1.0)
        Lkk = torch.linalg.cholesky_ex(W[:w], check_errors=False).L
        L[k:e, k:e] = Lkk
        Lh[k:e, k:e] = Lkk
        if e >= n:
            break
        Linv = _tri_inv(Lkk, _TRI_BASE)
        L21 = L[e:, k:e]
        _ext.gemm_nt_f16(L21, _h(W[w:]), _h(Linv), 1.0)
        Lh[e:, k:e] = L21
    return L


def _blocked_tf32(A, nb=None):
    n = A.shape[-1]
    # Block size tuned per size: small n needs several steps to amortise
    # overhead, large n wants fewer, larger GEMMs. Once the trailing updates run
    # in fp16 the leftover critical path is the O(n*nb^2) fp32 panel work, so the
    # optimum moves down from 4096 to 2048 -- but only at n >= 16384; nb=2048
    # measurably regresses n=8192.
    if nb is None:
        nb = 1024 if n <= 4096 else (2048 if n >= 16384 else 4096)
    # fp16 in / fp32 out for the trailing update and the panel GEMM: same 10-bit
    # mantissa as tf32, ~2x the tensor-core rate. Below 8192 the cast traffic is
    # not amortised.
    fp16 = n >= 8192
    # A batch of 1 still routes cuSOLVER through its *batched* potrf/trsm, which
    # are built for many-small and degrade badly for one-large. Squeezing the
    # singleton dim puts the nb x nb diagonal block on the classical path; that
    # block is the largest remaining non-tensor-core item at n=32768.
    squeeze = A.dim() == 3 and A.shape[0] == 1
    if squeeze:
        A = A[0]
    prev = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        if squeeze and n >= _LEFT_MIN_N:
            return _left_tf32(A, _LEFT_NB).unsqueeze(0)
        W = A.clone()
        L = torch.zeros_like(A)
        for k in range(0, n, nb):
            e = min(k + nb, n)
            Lkk = torch.linalg.cholesky_ex(W[..., k:e, k:e], check_errors=False).L
            L[..., k:e, k:e] = Lkk
            if e >= n:
                break
            # Panel: X @ Lkk^T = A21. A direct TRSM here is fp32 and never
            # touches tensor cores, and its cost (nb^2 * (n-e)) dominates the
            # trailing tf32 GEMMs at large n. Invert the nb x nb block once
            # (nb^3/2, ~14x cheaper at the first step) and turn the panel into a
            # tf32 GEMM instead.
            # With _tri_inv (batched 2x2 recursion) the inverse is cheap enough to
            # pay off from n=8192 up: B200 1x8192 5010 -> 3950us.
            if n >= 8192:
                Linv = _tri_inv(Lkk, _TRI_BASE)
                if fp16:
                    # Accumulate straight into L's panel (still all zeros there,
                    # so beta=1 is an overwrite). The obvious form -- zeros_like,
                    # GEMM into it, then copy into L -- costs three extra full
                    # passes over an (n-e) x nb fp32 buffer per step, ~6.5 GB of
                    # traffic over the whole factorization at n=32768.
                    L21 = L[..., e:, k:e]
                    _ext.gemm_nt_f16(L21, _h(W[..., e:, k:e]), _h(Linv), 1.0)
                else:
                    L21 = W[..., e:, k:e] @ Linv.transpose(-1, -2)
                    L[..., e:, k:e] = L21
            else:
                L21 = torch.linalg.solve_triangular(
                    Lkk.transpose(-1, -2), W[..., e:, k:e], upper=True, left=False
                )
                L[..., e:, k:e] = L21
            # Trailing update by block column, touching only the lower triangle.
            # Updating the full square would double the flops.
            L21h = _h(L21) if fp16 else None
            for jb in range(e, n, nb):
                je = min(jb + nb, n)
                if fp16:
                    _ext.gemm_nt_f16(
                        W[..., jb:, jb:je],
                        L21h[..., jb - e:, :],
                        L21h[..., jb - e:je - e, :],
                        -1.0,
                    )
                else:
                    W[..., jb:, jb:je] -= (
                        L21[..., jb - e:, :]
                        @ L21[..., jb - e:je - e, :].transpose(-1, -2)
                    )
        return L.unsqueeze(0) if squeeze else L
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev


def _blocked_mid(A, nb, tri_base=None):
    # n in {512, 1024}, high batch. Same right-looking structure as
    # _blocked_tf32, but the panel solve is an inverted-diagonal tf32 GEMM
    # rather than an fp32 batched TRSM: at these sizes the TRSM, not the
    # trailing SYRK, is the bottleneck (nb=256+TRSM measured 4090us at 640x512
    # vs 3280us for this). Kept separate so the large-n path is untouched.
    n = A.shape[-1]
    prev = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        W = A.clone()
        L = torch.zeros_like(A)
        for k in range(0, n, nb):
            e = min(k + nb, n)
            Dkk = W[..., k:e, k:e]
            Lkk = _ext.potrf(Dkk.contiguous()) if (e - k) in _CUSTOM_N else None
            if Lkk is None or not Lkk.numel():
                Lkk = torch.linalg.cholesky_ex(Dkk, check_errors=False).L
            L[..., k:e, k:e] = Lkk
            if e >= n:
                break
            # Batched 2x2 recursion instead of solve_triangular(Lkk, eye): only
            # base^2/nb^2 of the panel-inverse work stays on the non-TC fp32
            # solve. B200 640x512 2390 -> 1854, 60x1024 1431 -> 1067.
            Linv = _tri_inv(Lkk, tri_base or _MID_TRI_BASE)
            L21 = W[..., e:, k:e] @ Linv.transpose(-1, -2)
            L[..., e:, k:e] = L21
            # Trailing update in fp16 in / fp32 accumulate: same 10-bit mantissa
            # as tf32 for ~2x the tensor-core rate, so residuals are unchanged.
            # Measured and rejected: doing the panel GEMM in fp16 too (the extra
            # W21 cast costs more than the 2x rate saves), and cutting the copy
            # traffic here (block-lower-triangle W copy, matmul out= into L) --
            # 1-3% worse on B200 even though it cut local 4070 time 17%.
            L21h = _h(L21)
            for jb in range(e, n, nb):
                je = min(jb + nb, n)
                _ext.gemm_nt_f16(
                    W[..., jb:, jb:je],
                    L21h[..., jb - e:, :],
                    L21h[..., jb - e:je - e, :],
                    -1.0,
                )
        return L
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    if data.dim() == 3:
        if n in _CUSTOM_N:
            out = _ext.potrf(data)
            if out.numel():
                return out
            # SMEM opt-in failed for this n on this device; fall through.
        # 60x1024 is below _TF32_MIN_WORK but still beats batched cuSOLVER on
        # the blocked tf32 path with a small nb (B200: 2890 -> 2520 us).
        # Explicit batch floors rather than a b*n^3 work gate: with the custom
        # potrf, fp16 trailing and _tri_inv panel, the low-batch pair now wins
        # here too (4x1024 1281 -> 683, 16x512 596 -> 580). The b>=4 floor at
        # n=1024 is load-bearing -- 2x1024 lowrank NaNs on this path.
        # Batch floors on top of the work gate: five B200 rounds (waves
        # 8/10/15/16/17/18) measure the low-batch pair winning here too
        # (4x1024 1281 -> 745, 16x512 596 -> 416). Settled; do not re-litigate.
        b = data.shape[0]
        if n in (512, 1024) and (
            b * float(n) ** 3 >= 5e10
            or (n == 1024 and b >= 4)
            or (n == 512 and b >= 16)
        ):
            nb = 128 if n == 512 else 256
            # Low-batch n=512: collapse _tri_inv to a single batched
            # solve_triangular (base == nb, no 2x2 recursion). 16x512 is
            # dispatch-bound so the ~20 launches saved beat the wider fp32 base
            # solve; at n=1024 the same change costs (base=16 -> 745,
            # 128 -> 1314, 256 -> 1337), so this is n==512-only by measurement.
            tb = nb if (n == 512 and b * float(n) ** 3 < 5e10) else None
            return _blocked_mid(data, nb=nb, tri_base=tb)
        # Few-large: cuSOLVER's batched potrf is built for many-small and loses
        # badly here -- 2x2048 gets 1.7 TF/s where 1x4096 gets 15 on the very same
        # library via the classical single-matrix routine. One call per matrix
        # keeps every one of them on the classical path.
        # b <= 4 only: measured on B200, 8x2048 loses to the batched call
        # (5060 -> 5400us) while 2x2048 wins 3450 -> 1349 and 2x4096 6930 -> 3210.
        # b == 1 is already on the classical path, so looping only adds an n^2
        # copy into the output buffer (1x4096 measured 1528 -> 1603us).
        # Measured and rejected on B200: routing 8x2048 and 2x4096 to
        # _blocked_batch(256, 1024) -- a blocked tf32 factorization with the
        # custom single-CTA kernel on the diagonal, i.e. no cuSOLVER anywhere.
        # 8x2048 4910 -> 5520, 2x4096 3210 -> 4660. The 4070 said the opposite
        # (1.3-1.4x wins); on the B200 cuSOLVER is fast enough at n>=2048 that a
        # Python-level blocked loop cannot pay for its launches. The crossover to
        # _blocked_tf32 winning is above 4096, not below it.
        # 8x2048 is the worst-throughput entry in the large regime (4.5 TF/s):
        # cuSOLVER's batched potrf at n=2048 is slow no matter the batch. Route it
        # through the mid blocked path with nb=256, whose diagonal block goes to
        # the custom single-CTA potrf, so cuSOLVER is out of the loop entirely.
        # b >= 8 only: measured on B200, at b == 2 this path is far worse than one
        # classical cuSOLVER call per matrix (2x2048 1348 -> 3730, 2x4096
        # 3210 -> 9300) -- nb=256 leaves too little parallelism per launch.
        if n == 2048 and data.shape[0] >= 8:
            return _blocked_mid(data, nb=256, tri_base=128)
        if 2048 <= n <= 4096 and 2 <= data.shape[0] <= 4:
            L = torch.empty_like(data)
            for i in range(data.shape[0]):
                L[i] = torch.linalg.cholesky_ex(data[i], check_errors=False).L
            return L
        if data.shape[0] * float(n) ** 3 >= _TF32_MIN_WORK:
            return _blocked_tf32(data)
    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 942 lines total

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

Best evidence level for this revision: reported

JSON