Skip to content
KernelIndex
Search⌘K

submission 927797

aryanmagoon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
516.1µs
#35 of 337
2026-07-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:86637701e34211becc035b72759a7a0171e3f379607dfbcdd3d7b15bfae9e78b
license declaredunknown
license concludedunknown
authorsaryanmagoon
imported2026-08-26

Techniques

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

async-copyasm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::
mmawmma::fragment<wmma::accumulator, 16, 16, 16, float> acc[AC][2];
shared-memory__shared__ __align__(16) float sm[MPB * (N * LDP + 128)];
vector-width = float4float4 v = *(const float4 *)(bcr + j4);

Kernel source

submission.py3005 lines
from __future__ import annotations

import ctypes
import ctypes.util
import os
import shutil
import threading
from functools import lru_cache

import torch

try:
    from task import input_t, output_t
except ImportError:
    input_t = torch.Tensor
    output_t = torch.Tensor


_SMALL_CUDA_SOURCE = r"""
typedef unsigned long long uintptr_t;
#include <cuda_runtime.h>
#include <mma.h>

#define FULLM 0xffffffffu






template <int N, int MPB>
__global__ __launch_bounds__((N / 32) * 32 * MPB) void chol_reg_kernel(
    const float *__restrict__ A, float *__restrict__ L, int batch) {
  constexpr int NW = N / 32;
  constexpr int LDP = (NW == 1) ? 33 : 36;   
  __shared__ __align__(16) float sm[MPB * (N * LDP + 128)];

  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int wid = (tid >> 5);          
  const int slot = wid / NW;           
  const int w = wid - slot * NW;       
  const int mid = blockIdx.x * MPB + slot;
  const bool live = (mid < batch);

  float *P = sm + slot * (N * LDP + 128);
  float *bc = P + N * LDP;             
  float *dinv = bc + 96;               

  const size_t off = (size_t)(live ? mid : 0) * N * N;
  const float *Ab = A + off;
  float *Lb = L + off;

  const int R = 32 * w + lane;

  float r[N];
#pragma unroll
  for (int j = 0; j < N; ++j) r[j] = 0.f;
  
#pragma unroll
  for (int j = 0; j < N; ++j)
    if (j <= R) r[j] = Ab[(size_t)j * N + R];

#pragma unroll
  for (int kb = 0; kb < NW; ++kb) {
    const int c0 = 32 * kb;
    
    if (w == kb) {
      
      
      
      
      
      bc[lane] = r[c0];
#pragma unroll
      for (int k = 0; k < 32; ++k) {
        __syncwarp();
        const float *bcr = bc + (k % 3) * 32;
        float *bcw = bc + ((k + 1) % 3) * 32;
        const float rk = r[c0 + k];
        float akk = __shfl_sync(FULLM, rk, k);
        float d = rsqrtf(akk);
        float f = rk * (d * d);
        if (k < 31)
          r[c0 + ((k + 1) & 31)] = fmaf(-f, bcr[k + 1], r[c0 + ((k + 1) & 31)]);
        if (k < 30)
          r[c0 + ((k + 2) & 31)] = fmaf(-f, bcr[k + 2], r[c0 + ((k + 2) & 31)]);
        if (k < 31) bcw[lane] = r[c0 + ((k + 1) & 31)];
        const int j0 = (k + 3) & ~3;
#pragma unroll
        for (int j4 = j0; j4 < 32; j4 += 4) {
          float4 v = *(const float4 *)(bcr + j4);
          if (j4 + 0 > k + 2) r[c0 + j4 + 0] = fmaf(-f, v.x, r[c0 + j4 + 0]);
          if (j4 + 1 > k + 2) r[c0 + j4 + 1] = fmaf(-f, v.y, r[c0 + j4 + 1]);
          if (j4 + 2 > k + 2) r[c0 + j4 + 2] = fmaf(-f, v.z, r[c0 + j4 + 2]);
          if (j4 + 3 > k + 2) r[c0 + j4 + 3] = fmaf(-f, v.w, r[c0 + j4 + 3]);
        }
        r[c0 + k] = (lane >= k) ? rk * d : 0.f;
      }
      if (NW > 1) {  
#pragma unroll
        for (int t = 0; t < 32; ++t) P[(c0 + lane) * LDP + t] = r[c0 + t];
        dinv[lane] = 1.f / r[c0 + lane];
      }
    }
    __syncthreads();

    if (NW > 1 && w > kb) {
      
#pragma unroll
      for (int t = 0; t < 32; ++t) {
        float x = r[c0 + t];
        const float *Lr = P + (size_t)(c0 + t) * LDP;
#pragma unroll
        for (int u4 = 0; u4 < ((t + 3) & ~3); u4 += 4) {
          float4 l4 = *(const float4 *)(Lr + u4);
          if (u4 + 0 < t) x -= r[c0 + u4 + 0] * l4.x;
          if (u4 + 1 < t) x -= r[c0 + u4 + 1] * l4.y;
          if (u4 + 2 < t) x -= r[c0 + u4 + 2] * l4.z;
          if (u4 + 3 < t) x -= r[c0 + u4 + 3] * l4.w;
        }
        r[c0 + t] = x * dinv[t];
      }
#pragma unroll
      for (int t = 0; t < 32; ++t) P[(size_t)R * LDP + t] = r[c0 + t];
    }
    __syncthreads();

    if (NW > 1 && w > kb) {
      
#pragma unroll
      for (int j = c0 + 32; j < N; ++j) {
        const float *Pj = P + (size_t)j * LDP;
        float acc = 0.f;
#pragma unroll
        for (int t4 = 0; t4 < 32; t4 += 4) {
          float4 p = *(const float4 *)(Pj + t4);
          acc += r[c0 + t4 + 0] * p.x;
          acc += r[c0 + t4 + 1] * p.y;
          acc += r[c0 + t4 + 2] * p.z;
          acc += r[c0 + t4 + 3] * p.w;
        }
        r[j] -= acc;
      }
    }
    __syncthreads();
  }

  
  const int NT = (N / 32) * 32 * MPB;
#pragma unroll
  for (int cb = 0; cb < NW; ++cb) {
    const int c0 = 32 * cb;
    __syncthreads();
#pragma unroll
    for (int t = 0; t < 32; ++t)
      P[(size_t)R * LDP + t] = (c0 + t <= R) ? r[c0 + t] : 0.f;
    __syncthreads();
    if (live) {
      for (int e = tid - slot * (NT / MPB); e < N * 32; e += NT / MPB) {
        int row = e >> 5, col = e & 31;
        Lb[(size_t)row * N + c0 + col] = P[(size_t)row * LDP + col];
      }
    }
  }
}
"""


_MAIN_CUDA_SOURCE = r"""
typedef unsigned long long uintptr_t;
#include <cuda_runtime.h>
#include <mma.h>

#define FULLM 0xffffffffu

using namespace nvcuda;

#define GK_LD 40
#define GK_KC 32




#define GI_LD 80
#define GI_KC 64







template <int TILE, int KC, int NT>
__device__ __forceinline__ void ld_glb(float4 *reg, const float *src,
                                       int ldsrc, int rows, int kc, int K,
                                       int tid, bool fast) {
  constexpr int TPR = KC / 4;        
  constexpr int RPI = NT / TPR;      
  constexpr int ITER = TILE / RPI;
  const int c4 = (tid & (TPR - 1)) * 4;
  const int rb = tid / TPR;
  if (fast) {  
#pragma unroll
    for (int it = 0; it < ITER; ++it) {
      const int row = it * RPI + rb;
      reg[it] = (row < rows)
                    ? *(const float4 *)(src + (size_t)row * ldsrc + kc + c4)
                    : make_float4(0.f, 0.f, 0.f, 0.f);
    }
    return;
  }
#pragma unroll
  for (int it = 0; it < ITER; ++it) {
    const int row = it * RPI + rb;
    float4 v = make_float4(0.f, 0.f, 0.f, 0.f);
    const float *pv = src + (size_t)row * ldsrc + kc + c4;
    if (row < rows && kc + c4 + 3 < K && (((uintptr_t)pv & 15) == 0)) {
      v = *(const float4 *)pv;
    } else if (row < rows && kc + c4 < K) {
      const float *p = src + (size_t)row * ldsrc + kc;
      float t[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
      for (int q = 0; q < 4; ++q)
        if (kc + c4 + q < K) t[q] = p[c4 + q];
      v = make_float4(t[0], t[1], t[2], t[3]);
    }
    reg[it] = v;
  }
}

struct H4 {
  __half2 a, b;
};
















template <int HSH>
struct HShadow;
template <>
struct HShadow<0> {};
template <>
struct HShadow<1> {
  __half *p;      
  int ldp;        
  size_t stride;  
};

struct __align__(16) H8 {
  __half2 a, b, c, d;
};




__device__ __forceinline__ void st_half_frag(__half *dst, int ldp,
                                             const float *eb, int r0, int c0,
                                             int nrow, int ncol, int lane) {
  const int rr = lane >> 1, cc = (lane & 1) * 8;
  if (r0 + rr >= nrow) return;
  __half *d = dst + (size_t)(r0 + rr) * ldp + c0 + cc;
  const float *s = eb + rr * 20 + cc;
  if (c0 + cc + 8 <= ncol && ((((uintptr_t)d) & 15) == 0)) {
    const float4 v0 = *(const float4 *)s;
    const float4 v1 = *(const float4 *)(s + 4);
    H8 h;
    h.a = __floats2half2_rn(v0.x, v0.y);
    h.b = __floats2half2_rn(v0.z, v0.w);
    h.c = __floats2half2_rn(v1.x, v1.y);
    h.d = __floats2half2_rn(v1.z, v1.w);
    *(H8 *)d = h;
  } else {
    for (int q = 0; q < 8; ++q)
      if (c0 + cc + q < ncol) d[q] = __float2half(s[q]);
  }
}

template <int ITER>
__device__ __forceinline__ void cvt_h4(H4 *reg, const float4 *v, float sgn) {
#pragma unroll
  for (int it = 0; it < ITER; ++it) {
    reg[it].a = __floats2half2_rn(sgn * v[it].x, sgn * v[it].y);
    reg[it].b = __floats2half2_rn(sgn * v[it].z, sgn * v[it].w);
  }
}

template <int TILE, int KC, int LD, int NT>
__device__ __forceinline__ void st_smem(__half *dst, const H4 *reg, int tid) {
  constexpr int TPR = KC / 4;
  constexpr int RPI = NT / TPR;
  constexpr int ITER = TILE / RPI;
  const int c4 = (tid & (TPR - 1)) * 4;
  const int rb = tid / TPR;
#pragma unroll
  for (int it = 0; it < ITER; ++it) {
    const int row = it * RPI + rb;
    *(__half2 *)(dst + row * LD + c4) = reg[it].a;
    *(__half2 *)(dst + row * LD + c4 + 2) = reg[it].b;
  }
}




















template <int MODE, int TM, int TN, int HSH>
__device__ __forceinline__ void gemm_nt_body(
    float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
    const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
    int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
    size_t strideA, size_t strideB, int lower, HShadow<HSH> hs) {
  constexpr int NWN = TN / 32;       
  constexpr int NW = 2 * NWN;        
  constexpr int AC = TM / 32;        
  constexpr int NT = 32 * NW;        
  
  
  
  
  constexpr int KC = (TM <= 64) ? 64 : 32;
  constexpr int LD = KC + 8;
  constexpr int ITERA = TM / (NT / (KC / 4));
  constexpr int ITERB = TN / (NT / (KC / 4));
  __shared__ __align__(16) __half As[2 * TM * LD];
  __shared__ __align__(16) __half Bs[2 * TN * LD];
  __shared__ __align__(16) float Es[NW * 16 * 20];

  const int i0 = blockIdx.x * TM, j0 = blockIdx.y * TN;
  if (i0 >= nrow || j0 >= ncol) return;
  
  
  
  
  
  
  if (lower && (roff + i0 + TM - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
    return;

  const size_t bidx = blockIdx.z;
  float *C = Cb + bidx * strideC;
  const float *Cin = Cinb + bidx * strideC;
  const float *A = Ab + bidx * strideA;
  const float *B = Bb + bidx * strideB;

  const int tid = threadIdx.x;
  const int warp = tid >> 5, lane = tid & 31;
  const int wm = warp / NWN, wn = warp % NWN;
  float *eb = Es + warp * (16 * 20);

  wmma::fragment<wmma::accumulator, 16, 16, 16, float> acc[AC][2];
  const bool ldok = ((ldc & 3) == 0) && ((((uintptr_t)C) & 15) == 0) &&
                    ((((uintptr_t)Cin) & 15) == 0);

#pragma unroll
  for (int a = 0; a < AC; ++a)
#pragma unroll
    for (int b = 0; b < 2; ++b) {
      const int r0 = i0 + wm * (TM / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
      const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
      if (MODE == 0) {
        wmma::fill_fragment(acc[a][b], 0.f);
      } else if (full) {
        wmma::load_matrix_sync(acc[a][b], Cin + (size_t)r0 * ldc + c0, ldc,
                               wmma::mem_row_major);
      } else {
        for (int e = lane; e < 256; e += 32) {
          int rr = e >> 4, cc = e & 15;
          float v = 0.f;
          if (r0 + rr < nrow && c0 + cc < ncol)
            v = Cin[(size_t)(r0 + rr) * ldc + c0 + cc];
          eb[rr * 20 + cc] = v;
        }
        __syncwarp();
        wmma::load_matrix_sync(acc[a][b], eb, 20, wmma::mem_row_major);
        __syncwarp();
      }
    }

  const int mrows = min(TM, nrow - i0);
  const int nrows = min(TN, ncol - j0);
  const float sgn = (MODE == 1) ? -1.f : 1.f;

  const float *Ap = A + (size_t)i0 * lda;
  const float *Bp = B + (size_t)j0 * ldb;
  const bool fastA = ((K % KC) == 0) && ((lda & 3) == 0) &&
                     ((((uintptr_t)Ap) & 15) == 0);
  const bool fastB = ((K % KC) == 0) && ((ldb & 3) == 0) &&
                     ((((uintptr_t)Bp) & 15) == 0);
  float4 ra[ITERA], rb[ITERB];
  H4 ha[ITERA], hb[ITERB];
  ld_glb<TM, KC, NT>(ra, Ap, lda, mrows, 0, K, tid, fastA);
  ld_glb<TN, KC, NT>(rb, Bp, ldb, nrows, 0, K, tid, fastB);
  cvt_h4<ITERA>(ha, ra, sgn);
  cvt_h4<ITERB>(hb, rb, 1.f);

  st_smem<TM, KC, LD, NT>(As, ha, tid);
  st_smem<TN, KC, LD, NT>(Bs, hb, tid);
  __syncthreads();
  int cur = 0;
  constexpr int STGA = TM * LD;
  constexpr int STGB = TN * LD;

  for (int kc = 0; kc < K; kc += KC) {
    const __half *Ac = As + cur * STGA;
    const __half *Bc = Bs + cur * STGB;
    if (kc + KC < K) {
      ld_glb<TM, KC, NT>(ra, Ap, lda, mrows, kc + KC, K, tid, fastA);
      ld_glb<TN, KC, NT>(rb, Bp, ldb, nrows, kc + KC, K, tid, fastB);
    }
#pragma unroll
    for (int kk = 0; kk < KC; kk += 16) {
      wmma::fragment<wmma::matrix_b, 16, 16, 16, __half, wmma::col_major> bf[2];
#pragma unroll
      for (int b = 0; b < 2; ++b)
        wmma::load_matrix_sync(bf[b], Bc + (wn * 32 + b * 16) * LD + kk, LD);
#pragma unroll
      for (int a = 0; a < AC; ++a) {
        wmma::fragment<wmma::matrix_a, 16, 16, 16, __half, wmma::row_major> af;
        wmma::load_matrix_sync(af, Ac + (wm * (TM / 2) + a * 16) * LD + kk,
                               LD);
#pragma unroll
        for (int b = 0; b < 2; ++b)
          wmma::mma_sync(acc[a][b], af, bf[b], acc[a][b]);
      }
    }
    if (kc + KC < K) {
      cvt_h4<ITERA>(ha, ra, sgn);
      cvt_h4<ITERB>(hb, rb, 1.f);
      st_smem<TM, KC, LD, NT>(As + (cur ^ 1) * STGA, ha, tid);
      st_smem<TN, KC, LD, NT>(Bs + (cur ^ 1) * STGB, hb, tid);
      cur ^= 1;
    }
    __syncthreads();
  }

#pragma unroll
  for (int a = 0; a < AC; ++a)
#pragma unroll
    for (int b = 0; b < 2; ++b) {
      const int r0 = i0 + wm * (TM / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
      if (r0 >= nrow || c0 >= ncol) continue;
      if (lower && (roff + r0 + 15) < (coff + c0) &&
          ((r0 >> 7) != (c0 >> 7)))
        continue;
      const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
      if (full) {
        wmma::store_matrix_sync(C + (size_t)r0 * ldc + c0, acc[a][b], ldc,
                                wmma::mem_row_major);
        if constexpr (HSH == 1) {
          
          
          __syncwarp();
          wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
          __syncwarp();
          st_half_frag(hs.p + bidx * hs.stride, hs.ldp, eb, r0, c0, nrow, ncol,
                       lane);
        }
      } else {
        __syncwarp();
        wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
        __syncwarp();
        for (int e = lane; e < 256; e += 32) {
          int rr = e >> 4, cc = e & 15;
          if (r0 + rr < nrow && c0 + cc < ncol)
            C[(size_t)(r0 + rr) * ldc + c0 + cc] = eb[rr * 20 + cc];
        }
        if constexpr (HSH == 1)  
          st_half_frag(hs.p + bidx * hs.stride, hs.ldp, eb, r0, c0, nrow, ncol,
                       lane);
      }
    }
}


template <int MODE, int TM, int TN>
__global__ __launch_bounds__(2 * TN, (TM == TN) ? 2 : 1) void gemm_nt_kernel(
    float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
    const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
    int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
    size_t strideA, size_t strideB, int lower) {
  gemm_nt_body<MODE, TM, TN, 0>(Cb, Cinb, ldc, Ab, lda, Bb, ldb, K, nrow, ncol,
                                roff, coff, strideC, strideA, strideB, lower,
                                HShadow<0>());
}





template <int MODE, int TM, int TN>
__global__ __launch_bounds__(2 * TN, (TM == TN) ? 2 : 1) void gemm_nt_hs_kernel(
    float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
    const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
    int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
    size_t strideA, size_t strideB, int lower, __half *__restrict__ php,
    int ldp, size_t strideP) {
  HShadow<1> hs;
  hs.p = php;
  hs.ldp = ldp;
  hs.stride = strideP;
  gemm_nt_body<MODE, TM, TN, 1>(Cb, Cinb, ldc, Ab, lda, Bb, ldb, K, nrow, ncol,
                                roff, coff, strideC, strideA, strideB, lower,
                                hs);
}











__global__ __launch_bounds__(256) void cvt_panel_kernel(
    const float *__restrict__ Lsrc, int n, __half *__restrict__ Ph, int ldp,
    int nbo, int R0, int jb, size_t stride) {
  const int row = R0 + blockIdx.x;
  const size_t b = blockIdx.y;
  const float *S = Lsrc + b * stride + (size_t)row * n + jb;
  __half *D = Ph + b * (size_t)n * ldp + (size_t)row * ldp;
  const bool v4 = ((((uintptr_t)S) & 15) == 0);
  for (int c = threadIdx.x * 4; c < nbo; c += 256 * 4) {
    if (v4 && c + 3 < nbo) {
      float4 v = *(const float4 *)(S + c);
      *(__half2 *)(D + c) = __floats2half2_rn(v.x, v.y);
      *(__half2 *)(D + c + 2) = __floats2half2_rn(v.z, v.w);
    } else {
      for (int q = c; q < nbo && q < c + 4; ++q) D[q] = __float2half(S[q]);
    }
  }
}

template <int RSTEP>
__device__ __forceinline__ void cp_h(__half *dst, const __half *src, int lds,
                                     int rows, int kc, int tid) {
  const int k8 = (tid & 3) * 8;
#pragma unroll
  for (int it = 0; it < 2; ++it) {
    const int row = it * RSTEP + (tid >> 2);
    const unsigned sa =
        (unsigned)__cvta_generic_to_shared(dst + row * GK_LD + k8);
    const __half *gp = src + (size_t)(row < rows ? row : 0) * lds + kc + k8;
    const int bytes = (row < rows) ? 16 : 0;
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::
                     "r"(sa), "l"(gp), "r"(bytes));
  }
}

template <int TILE>
__global__ __launch_bounds__(TILE * 2) void gemm_h_kernel(
    float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
    const __half *__restrict__ Ab, int lda, const __half *__restrict__ Bb,
    int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
    size_t strideAB, int lower) {
  constexpr int NW = TILE / 16;
  constexpr int NWN = TILE / 32;
  constexpr int AC = TILE / 32;
  constexpr int HRS = TILE / 2;
  constexpr int STG = TILE * GK_LD;
  
  
  
  
  
  
  
  
  
  
  constexpr int NSTG = 4;
  __shared__ __align__(16) __half As[NSTG * TILE * GK_LD];
  __shared__ __align__(16) __half Bs[NSTG * TILE * GK_LD];
  __shared__ __align__(16) float Es[NW * 16 * 20];

  const int i0 = blockIdx.x * TILE, j0 = blockIdx.y * TILE;
  if (i0 >= nrow || j0 >= ncol) return;
  
  
  
  
  
  
  if (lower && (roff + i0 + TILE - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
    return;

  const size_t bidx = blockIdx.z;
  float *C = Cb + bidx * strideC;
  const float *Cin = Cinb + bidx * strideC;
  const __half *A = Ab + bidx * strideAB;
  const __half *B = Bb + bidx * strideAB;

  const int tid = threadIdx.x;
  const int warp = tid >> 5, lane = tid & 31;
  const int wm = warp / NWN, wn = warp % NWN;
  float *eb = Es + warp * (16 * 20);

  wmma::fragment<wmma::accumulator, 16, 16, 16, float> acc[AC][2];
  const bool ldok = ((ldc & 3) == 0) && ((((uintptr_t)C) & 15) == 0) &&
                    ((((uintptr_t)Cin) & 15) == 0);
#pragma unroll
  for (int a = 0; a < AC; ++a)
#pragma unroll
    for (int b = 0; b < 2; ++b) {
      const int r0 = i0 + wm * (TILE / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
      const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
      if (full) {
        wmma::load_matrix_sync(acc[a][b], Cin + (size_t)r0 * ldc + c0, ldc,
                               wmma::mem_row_major);
      } else {
        for (int e = lane; e < 256; e += 32) {
          int rr = e >> 4, cc = e & 15;
          float v = 0.f;
          if (r0 + rr < nrow && c0 + cc < ncol)
            v = Cin[(size_t)(r0 + rr) * ldc + c0 + cc];
          eb[rr * 20 + cc] = v;
        }
        __syncwarp();
        wmma::load_matrix_sync(acc[a][b], eb, 20, wmma::mem_row_major);
        __syncwarp();
      }
#pragma unroll
      for (int e = 0; e < acc[a][b].num_elements; ++e)
        acc[a][b].x[e] = -acc[a][b].x[e];
    }

  const int mrows = min(TILE, nrow - i0);
  const int nrows = min(TILE, ncol - j0);
  const __half *Ap = A + (size_t)i0 * lda;
  const __half *Bp = B + (size_t)j0 * ldb;
#pragma unroll
  for (int s = 0; s < NSTG - 1; ++s) {
    if (s * GK_KC < K) {
      cp_h<HRS>(As + s * STG, Ap, lda, mrows, s * GK_KC, tid);
      cp_h<HRS>(Bs + s * STG, Bp, ldb, nrows, s * GK_KC, tid);
    }
    asm volatile("cp.async.commit_group;\n" ::);
  }
  int cur = 0;

  for (int kc = 0; kc < K; kc += GK_KC) {
    
    asm volatile("cp.async.wait_group %0;\n" ::"n"(NSTG - 2));
    __syncthreads();
    const __half *Ac = As + cur * STG;
    const __half *Bc = Bs + cur * STG;
    {
      
      
      const int nk = kc + (NSTG - 1) * GK_KC;
      const int nxt = (cur + NSTG - 1) % NSTG;
      if (nk < K) {
        cp_h<HRS>(As + nxt * STG, Ap, lda, mrows, nk, tid);
        cp_h<HRS>(Bs + nxt * STG, Bp, ldb, nrows, nk, tid);
      }
      asm volatile("cp.async.commit_group;\n" ::);
    }
#pragma unroll
    for (int kk = 0; kk < GK_KC; kk += 16) {
      wmma::fragment<wmma::matrix_b, 16, 16, 16, __half, wmma::col_major> bf[2];
#pragma unroll
      for (int b = 0; b < 2; ++b)
        wmma::load_matrix_sync(bf[b], Bc + (wn * 32 + b * 16) * GK_LD + kk,
                               GK_LD);
#pragma unroll
      for (int a = 0; a < AC; ++a) {
        wmma::fragment<wmma::matrix_a, 16, 16, 16, __half, wmma::row_major> af;
        wmma::load_matrix_sync(af, Ac + (wm * (TILE / 2) + a * 16) * GK_LD + kk,
                               GK_LD);
#pragma unroll
        for (int b = 0; b < 2; ++b)
          wmma::mma_sync(acc[a][b], af, bf[b], acc[a][b]);
      }
    }
    cur = (cur + 1) % NSTG;
  }

#pragma unroll
  for (int a = 0; a < AC; ++a)
#pragma unroll
    for (int b = 0; b < 2; ++b) {
      const int r0 = i0 + wm * (TILE / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
      if (r0 >= nrow || c0 >= ncol) continue;
      if (lower && (roff + r0 + 15) < (coff + c0) &&
          ((r0 >> 7) != (c0 >> 7)))
        continue;
#pragma unroll
      for (int e = 0; e < acc[a][b].num_elements; ++e)
        acc[a][b].x[e] = -acc[a][b].x[e];
      const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
      if (full) {
        wmma::store_matrix_sync(C + (size_t)r0 * ldc + c0, acc[a][b], ldc,
                                wmma::mem_row_major);
      } else {
        __syncwarp();
        wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
        __syncwarp();
        for (int e = lane; e < 256; e += 32) {
          int rr = e >> 4, cc = e & 15;
          if (r0 + rr < nrow && c0 + cc < ncol)
            C[(size_t)(r0 + rr) * ldc + c0 + cc] = eb[rr * 20 + cc];
        }
      }
    }
}
































__global__ __launch_bounds__(256) void cvt_panel_i8_kernel(
    const float *__restrict__ Lsrc, int n, signed char *__restrict__ Pq,
    int ldp, int nbo, int R0, int jb, size_t stride,
    float *__restrict__ Scl) {
  const int row = R0 + blockIdx.x;
  const size_t b = blockIdx.y;
  const float *S = Lsrc + b * stride + (size_t)row * n + jb;
  signed char *D = Pq + b * (size_t)n * ldp + (size_t)row * ldp;
  const int tid = threadIdx.x;

  
  
  float4 v[2];
  const bool v4 = ((((uintptr_t)S) & 15) == 0);
#pragma unroll
  for (int it = 0; it < 2; ++it) {
    const int c = it * 1024 + tid * 4;
    float4 x = make_float4(0.f, 0.f, 0.f, 0.f);
    if (v4 && c + 3 < nbo) {
      x = *(const float4 *)(S + c);
    } else if (c < nbo) {
      float t[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
      for (int q = 0; q < 4; ++q)
        if (c + q < nbo) t[q] = S[c + q];
      x = make_float4(t[0], t[1], t[2], t[3]);
    }
    v[it] = x;
  }
  float mx = 0.f;
#pragma unroll
  for (int it = 0; it < 2; ++it) {
    mx = fmaxf(mx, fabsf(v[it].x));
    mx = fmaxf(mx, fabsf(v[it].y));
    mx = fmaxf(mx, fabsf(v[it].z));
    mx = fmaxf(mx, fabsf(v[it].w));
  }
#pragma unroll
  for (int o = 16; o; o >>= 1) mx = fmaxf(mx, __shfl_xor_sync(FULLM, mx, o));
  __shared__ float wmx[8];
  if ((tid & 31) == 0) wmx[tid >> 5] = mx;
  __syncthreads();
  mx = wmx[tid & 7];  
#pragma unroll
  for (int o = 4; o; o >>= 1) mx = fmaxf(mx, __shfl_xor_sync(FULLM, mx, o));

  
  
  
  
  
  float inv = 127.f / mx;
  if (!(mx > 0.f) || !isfinite(mx) || !isfinite(inv)) {
    mx = 0.f;
    inv = 0.f;
  }
  if (tid == 0) Scl[b * (size_t)n + row] = mx * (1.f / 127.f);
  
  
  
  const int kpad = (nbo + 63) & ~63;
#pragma unroll
  for (int it = 0; it < 2; ++it) {
    const int c = it * 1024 + tid * 4;
    if (c >= kpad) continue;
    const float f[4] = {v[it].x, v[it].y, v[it].z, v[it].w};
    signed char r[4];
#pragma unroll
    for (int t = 0; t < 4; ++t) {
      int qi = (c + t < nbo) ? __float2int_rn(f[t] * inv) : 0;
      r[t] = (signed char)max(-127, min(127, qi));
    }
    if (c + 3 < kpad) {
      char4 q;
      q.x = r[0]; q.y = r[1]; q.z = r[2]; q.w = r[3];
      *(char4 *)(D + c) = q;
    } else {
#pragma unroll
      for (int t = 0; t < 4; ++t)
        if (c + t < kpad) D[c + t] = r[t];
    }
  }
}

template <int RSTEP>
__device__ __forceinline__ void cp_i8(signed char *dst, const signed char *src,
                                      int lds, int rows, int kc, int tid) {
  const int k16 = (tid & 3) * 16;
#pragma unroll
  for (int it = 0; it < 2; ++it) {
    const int row = it * RSTEP + (tid >> 2);
    const unsigned sa =
        (unsigned)__cvta_generic_to_shared(dst + row * GI_LD + k16);
    const signed char *gp =
        src + (size_t)(row < rows ? row : 0) * lds + kc + k16;
    const int bytes = (row < rows) ? 16 : 0;
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::
                     "r"(sa), "l"(gp), "r"(bytes));
  }
}




template <int TILE>
__global__ __launch_bounds__(TILE * 2, (TILE == 128) ? 2 : 1) void gemm_i8_kernel(
    float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
    const signed char *__restrict__ Ab, int lda,
    const signed char *__restrict__ Bb, int ldb,
    const float *__restrict__ Scb, int K, int nrow, int ncol, int roff,
    int coff, size_t strideC, size_t strideAB, size_t strideS, int lower) {
  constexpr int NW = TILE / 16;
  constexpr int NWN = TILE / 32;
  constexpr int AC = TILE / 32;
  constexpr int HRS = TILE / 2;
  constexpr int STG = TILE * GI_LD;
  constexpr int NSTG = 3;
  __shared__ __align__(16) signed char As[NSTG * TILE * GI_LD];
  __shared__ __align__(16) signed char Bs[NSTG * TILE * GI_LD];
  __shared__ __align__(16) int Es[NW * 16 * 20];
  __shared__ float SA[TILE], SB[TILE];

  const int i0 = blockIdx.x * TILE, j0 = blockIdx.y * TILE;
  if (i0 >= nrow || j0 >= ncol) return;
  
  
  
  
  
  
  if (lower && (roff + i0 + TILE - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
    return;

  const size_t bidx = blockIdx.z;
  float *C = Cb + bidx * strideC;
  const float *Cin = Cinb + bidx * strideC;
  const signed char *A = Ab + bidx * strideAB;
  const signed char *B = Bb + bidx * strideAB;
  const float *Sc = Scb + bidx * strideS;

  const int tid = threadIdx.x;
  const int warp = tid >> 5, lane = tid & 31;
  const int wm = warp / NWN, wn = warp % NWN;
  int *eb = Es + warp * (16 * 20);

  
  
  
  if (tid < TILE) {
    SA[tid] = (i0 + tid < nrow) ? Sc[i0 + tid] : 0.f;
    SB[tid] = (j0 + tid < ncol) ? Sc[j0 + tid] : 0.f;
  }

  wmma::fragment<wmma::accumulator, 16, 16, 16, int> acc[AC][2];
#pragma unroll
  for (int a = 0; a < AC; ++a)
#pragma unroll
    for (int b = 0; b < 2; ++b) wmma::fill_fragment(acc[a][b], 0);

  const bool ldok = ((ldc & 3) == 0) && ((((uintptr_t)C) & 15) == 0) &&
                    ((((uintptr_t)Cin) & 15) == 0);
  const int mrows = min(TILE, nrow - i0);
  const int nrows = min(TILE, ncol - j0);
  const signed char *Ap = A + (size_t)i0 * lda;
  const signed char *Bp = B + (size_t)j0 * ldb;
#pragma unroll
  for (int s = 0; s < NSTG - 1; ++s) {
    if (s * GI_KC < K) {
      cp_i8<HRS>(As + s * STG, Ap, lda, mrows, s * GI_KC, tid);
      cp_i8<HRS>(Bs + s * STG, Bp, ldb, nrows, s * GI_KC, tid);
    }
    asm volatile("cp.async.commit_group;\n" ::);
  }
  int cur = 0;

  for (int kc = 0; kc < K; kc += GI_KC) {
    
    asm volatile("cp.async.wait_group %0;\n" ::"n"(NSTG - 2));
    __syncthreads();
    const signed char *Ac = As + cur * STG;
    const signed char *Bc = Bs + cur * STG;
    {
      
      
      const int nk = kc + (NSTG - 1) * GI_KC;
      const int nxt = (cur + NSTG - 1) % NSTG;
      if (nk < K) {
        cp_i8<HRS>(As + nxt * STG, Ap, lda, mrows, nk, tid);
        cp_i8<HRS>(Bs + nxt * STG, Bp, ldb, nrows, nk, tid);
      }
      asm volatile("cp.async.commit_group;\n" ::);
    }
#pragma unroll
    for (int kk = 0; kk < GI_KC; kk += 16) {
      wmma::fragment<wmma::matrix_b, 16, 16, 16, signed char, wmma::col_major>
          bf[2];
#pragma unroll
      for (int b = 0; b < 2; ++b)
        wmma::load_matrix_sync(bf[b], Bc + (wn * 32 + b * 16) * GI_LD + kk,
                               GI_LD);
#pragma unroll
      for (int a = 0; a < AC; ++a) {
        wmma::fragment<wmma::matrix_a, 16, 16, 16, signed char, wmma::row_major>
            af;
        wmma::load_matrix_sync(af, Ac + (wm * (TILE / 2) + a * 16) * GI_LD + kk,
                               GI_LD);
#pragma unroll
        for (int b = 0; b < 2; ++b)
          wmma::mma_sync(acc[a][b], af, bf[b], acc[a][b]);
      }
    }
    cur = (cur + 1) % NSTG;
  }

  
  
  
#pragma unroll
  for (int a = 0; a < AC; ++a)
#pragma unroll
    for (int b = 0; b < 2; ++b) {
      const int r0 = i0 + wm * (TILE / 2) + a * 16, c0 = j0 + wn * 32 + b * 16;
      if (r0 >= nrow || c0 >= ncol) continue;
      if (lower && (roff + r0 + 15) < (coff + c0) && ((r0 >> 7) != (c0 >> 7)))
        continue;
      const int sa0 = wm * (TILE / 2) + a * 16, sb0 = wn * 32 + b * 16;
      __syncwarp();
      wmma::store_matrix_sync(eb, acc[a][b], 20, wmma::mem_row_major);
      __syncwarp();
      const bool full = ldok && (r0 + 16 <= nrow) && (c0 + 16 <= ncol);
      if (full) {
        const int rr = lane >> 1, cb = (lane & 1) * 8;
        const float sar = SA[sa0 + rr];
        const float *cs = Cin + (size_t)(r0 + rr) * ldc + c0 + cb;
        float *cd = C + (size_t)(r0 + rr) * ldc + c0 + cb;
        const int *ep = eb + rr * 20 + cb;
        const float *sbp = SB + sb0 + cb;
#pragma unroll
        for (int q = 0; q < 2; ++q) {
          float4 cv = *(const float4 *)(cs + q * 4);
          cv.x -= sar * sbp[q * 4 + 0] * (float)ep[q * 4 + 0];
          cv.y -= sar * sbp[q * 4 + 1] * (float)ep[q * 4 + 1];
          cv.z -= sar * sbp[q * 4 + 2] * (float)ep[q * 4 + 2];
          cv.w -= sar * sbp[q * 4 + 3] * (float)ep[q * 4 + 3];
          *(float4 *)(cd + q * 4) = cv;
        }
      } else {
        for (int e = lane; e < 256; e += 32) {
          const int rr = e >> 4, cc = e & 15;
          if (r0 + rr < nrow && c0 + cc < ncol)
            C[(size_t)(r0 + rr) * ldc + c0 + cc] =
                Cin[(size_t)(r0 + rr) * ldc + c0 + cc] -
                SA[sa0 + rr] * SB[sb0 + cc] * (float)eb[rr * 20 + cc];
        }
      }
    }
}








#define PLD 132

#define LRW 36





__device__ __forceinline__ void grp_gemm32(float (&acc)[4][4], const float *X,
                                           int ldX, const float *Yr, int ldY,
                                           int kk, int ti, int tj) {
#pragma unroll 8
  for (int k = 0; k < kk; ++k) {
    float4 a = *(const float4 *)(X + (size_t)k * ldX + ti);
    float4 b = *(const float4 *)(Yr + (size_t)k * ldY + tj);
    const float av[4] = {a.x, a.y, a.z, a.w};
    const float bv[4] = {b.x, b.y, b.z, b.w};
#pragma unroll
    for (int p = 0; p < 4; ++p)
#pragma unroll
      for (int q = 0; q < 4; ++q) acc[p][q] += av[p] * bv[q];
  }
}





__device__ __forceinline__ void rank32_upd(float *S, int sb, int nbp, int jlo,
                                           int jhi, int t0, int nthr) {
  const int mrow = nbp - sb - 32;
  if (mrow <= 0) return;
  const int nt = mrow >> 2;
  const int jhc = (jhi < nt) ? jhi : nt;
  if (jlo >= jhc) return;
  const int r0 = sb + 32;
  const int span = (jhc - jlo) * nt;
  for (int tt = t0; tt < span; tt += nthr) {
    const int q = tt / nt;
    const int tj = jlo + q, ti = tt - q * nt;
    if (tj > ti) continue;
    const int i0 = r0 + ti * 4, j0 = r0 + tj * 4;
    float acc[4][4];
#pragma unroll
    for (int a = 0; a < 4; ++a)
#pragma unroll
      for (int b = 0; b < 4; ++b) acc[a][b] = 0.f;
#pragma unroll
    for (int t = 0; t < 32; ++t) {
      const float *col = S + (size_t)(sb + t) * PLD;
      float4 av = *(const float4 *)(col + i0);
      float4 bv = *(const float4 *)(col + j0);
      float a4[4] = {av.x, av.y, av.z, av.w};
      float b4[4] = {bv.x, bv.y, bv.z, bv.w};
#pragma unroll
      for (int a = 0; a < 4; ++a)
#pragma unroll
        for (int b = 0; b < 4; ++b) acc[a][b] += a4[a] * b4[b];
    }
#pragma unroll
    for (int b = 0; b < 4; ++b)
#pragma unroll
      for (int a = 0; a < 4; ++a)
        if ((i0 + a) >= (j0 + b))
          S[(size_t)(j0 + b) * PLD + i0 + a] -= acc[a][b];
  }
}



template <int PNT>
__global__ __launch_bounds__(PNT) void potrf_kernel(
    const float *__restrict__ src, float *__restrict__ dst,
    float *__restrict__ Linv, int n, int k0, int nb, int NB, int needInv) {
  extern __shared__ __align__(16) float sm[];
  float *S = sm;                    
  float *BC = sm + 128 * PLD;       
  float *Dinv = BC + 96;            
  float *LR = Dinv + 32;            
  const int tid = threadIdx.x;
  const int lane = tid & 31, warp = tid >> 5;
  const int nblk = (nb + 31) >> 5;
  const int nbp = nblk << 5;
  const size_t mo = (size_t)blockIdx.x * (size_t)n * n;
  const float *Ab = src + mo + (size_t)k0 * n + k0;
  float *Lb = dst + mo + (size_t)k0 * n + k0;
  float *Xb = Linv + (size_t)blockIdx.x * NB * NB;

  if (nb == 128 && (n & 3) == 0 && ((((uintptr_t)Ab) & 15) == 0)) {
    
    
    
    
    
    constexpr int NLD = 8;  
    for (int base = tid; base < 32 * 128; base += PNT * NLD) {
      float4 v[NLD];
#pragma unroll
      for (int q = 0; q < NLD; ++q) {
        const int id = base + q * PNT;
        v[q] = *(const float4 *)(Ab + (size_t)(id >> 5) * n + ((id & 31) << 2));
      }
#pragma unroll
      for (int q = 0; q < NLD; ++q) {
        const int id = base + q * PNT;
        *(float4 *)(S + (size_t)(id >> 5) * PLD + ((id & 31) << 2)) = v[q];
      }
    }
  } else {
    
    
    const int wl = tid & 31, wi = tid >> 5, WS = PNT >> 5;
    for (int i0 = wi; i0 < nbp; i0 += WS * 4) {
      float v[4][4];
#pragma unroll
      for (int p = 0; p < 4; ++p) {
        const int ii = i0 + p * WS;
#pragma unroll
        for (int q = 0; q < 4; ++q) {
          const int jj = wl + q * 32;
          float x = 0.f;
          if (ii < nbp && jj < nbp) {
            if (ii < nb && jj < nb) x = (ii >= jj) ? Ab[(size_t)ii * n + jj] : 0.f;
            else x = (ii == jj) ? 1.f : 0.f;
          }
          v[p][q] = x;
        }
      }
#pragma unroll
      for (int p = 0; p < 4; ++p) {
        const int ii = i0 + p * WS;
        if (ii >= nbp) continue;
#pragma unroll
        for (int q = 0; q < 4; ++q) {
          const int jj = wl + q * 32;
          if (jj < nbp) S[jj * PLD + ii] = v[p][q];
        }
      }
    }
  }
  __syncthreads();

  
  int dsb = -1;  
  for (int sb = 0; sb < nbp; sb += 32) {
    const int pend = dsb;
    dsb = -1;
    if (warp == 0) {
      
      
      
      
      
      
      
      
      
      float r[32];
#pragma unroll
      for (int t = 0; t < 32; ++t) r[t] = S[(size_t)(sb + t) * PLD + sb + lane];
      BC[lane] = r[0];
      
      
      
      
#pragma unroll
      for (int k = 0; k < 32; ++k) {
        __syncwarp();
        const float *bcr = BC + (k % 3) * 32;
        float *bcw = BC + ((k + 1) % 3) * 32;
        const float rk = r[k];
        float akk = __shfl_sync(FULLM, rk, k);
        float d = rsqrtf(akk);
        float f = rk * (d * d);
        
        
        
        if (k < 31)
          r[(k + 1) & 31] = fmaf(-f, bcr[k + 1], r[(k + 1) & 31]);
        if (k < 30)
          r[(k + 2) & 31] = fmaf(-f, bcr[k + 2], r[(k + 2) & 31]);
        S[(size_t)(sb + k) * PLD + sb + lane] = (lane >= k) ? rk * d : 0.f;
        if (k < 31) bcw[lane] = r[(k + 1) & 31];
        const int j0 = (k + 3) & ~3;
#pragma unroll
        for (int j4 = j0; j4 < 32; j4 += 4) {
          float4 v = *(const float4 *)(bcr + j4);
          if (j4 + 0 > k + 2) r[j4 + 0] = fmaf(-f, v.x, r[j4 + 0]);
          if (j4 + 1 > k + 2) r[j4 + 1] = fmaf(-f, v.y, r[j4 + 1]);
          if (j4 + 2 > k + 2) r[j4 + 2] = fmaf(-f, v.z, r[j4 + 2]);
          if (j4 + 3 > k + 2) r[j4 + 3] = fmaf(-f, v.w, r[j4 + 3]);
        }
      }
      Dinv[lane] = 1.f / S[(size_t)(sb + lane) * PLD + sb + lane];
    } else if (pend >= 0) {
      rank32_upd(S, pend, nbp, 8, 1 << 20, tid - 32, PNT - 32);
    }
    __syncthreads();

    const int mrow = nbp - sb - 32;
    
    
    
    
    {
      const int cap = nb - sb;
      if (cap > 0) {
        for (int idx = tid; idx < 1024; idx += PNT) {
          const int i = idx >> 5, j = idx & 31;
          if (i < cap && j < cap)
            Lb[(size_t)(sb + i) * n + sb + j] =
                S[(size_t)(sb + j) * PLD + sb + i];
        }
      }
    }
    if (mrow > 0) {
      
      for (int idx = tid; idx < 1024; idx += PNT) {
        int tt = idx >> 5, uu = idx & 31;
        LR[tt * LRW + uu] = S[(size_t)(sb + tt) * PLD + sb + uu];
      }
      if (tid < 32) LR[tid * LRW + 32] = Dinv[tid];
    }
    __syncthreads();

    if (needInv && warp == 0) {
      
      
      
      float x[32];
#pragma unroll
      for (int i = 0; i < 32; ++i) x[i] = (i == lane) ? 1.f : 0.f;
      const float *Ld = S + (size_t)sb * PLD + sb;
#pragma unroll
      for (int t = 0; t < 32; ++t) {
        const float xt = x[t] * Dinv[t];
        x[t] = xt;
        const float *Lc = Ld + (size_t)t * PLD;
#pragma unroll
        for (int u4 = (t + 1) & ~3; u4 < 32; u4 += 4) {
          float4 v = *(const float4 *)(Lc + u4);
          if (u4 + 0 > t) x[u4 + 0] = fmaf(-xt, v.x, x[u4 + 0]);
          if (u4 + 1 > t) x[u4 + 1] = fmaf(-xt, v.y, x[u4 + 1]);
          if (u4 + 2 > t) x[u4 + 2] = fmaf(-xt, v.z, x[u4 + 2]);
          if (u4 + 3 > t) x[u4 + 3] = fmaf(-xt, v.w, x[u4 + 3]);
        }
      }
      __syncwarp();
#pragma unroll
      for (int i = 0; i < 32; ++i)
        S[(size_t)(sb + lane) * PLD + sb + i] = x[i];
    } else if (mrow > 0) {
      const int base = needInv ? 32 : 0;
      const int nthr = PNT - base;
      for (int i = sb + 32 + (tid - base); i < nbp; i += nthr) {
        float x[32];
#pragma unroll
        for (int t = 0; t < 32; ++t) x[t] = S[(size_t)(sb + t) * PLD + i];
#pragma unroll
        for (int t = 0; t < 32; ++t) {
          const float *Lc = LR + t * LRW;
          const float xt = x[t] * Lc[32];
          x[t] = xt;
#pragma unroll
          for (int u4 = (t + 1) & ~3; u4 < 32; u4 += 4) {
            float4 v = *(const float4 *)(Lc + u4);
            if (u4 + 0 > t) x[u4 + 0] = fmaf(-xt, v.x, x[u4 + 0]);
            if (u4 + 1 > t) x[u4 + 1] = fmaf(-xt, v.y, x[u4 + 1]);
            if (u4 + 2 > t) x[u4 + 2] = fmaf(-xt, v.z, x[u4 + 2]);
            if (u4 + 3 > t) x[u4 + 3] = fmaf(-xt, v.w, x[u4 + 3]);
          }
        }
#pragma unroll
        for (int t = 0; t < 32; ++t) S[(size_t)(sb + t) * PLD + i] = x[t];
      }
    }
    __syncthreads();

    if (mrow > 0) {
      
      rank32_upd(S, sb, nbp, 0, 8, tid, PNT);
      dsb = sb;
    }
    __syncthreads();
  }

  {
    
    
    const int wl = tid & 31, wi = tid >> 5, WS = PNT >> 5;
    for (int i0 = wi; i0 < nb; i0 += WS * 4) {
      float v[4][4];
#pragma unroll
      for (int p = 0; p < 4; ++p) {
        const int ii = i0 + p * WS;
#pragma unroll
        for (int q = 0; q < 4; ++q) {
          const int jj = wl + q * 32;
          v[p][q] = (ii < nb && jj < nb && ii >= jj) ? S[jj * PLD + ii] : 0.f;
        }
      }
#pragma unroll
      for (int p = 0; p < 4; ++p) {
        const int ii = i0 + p * WS;
        if (ii >= nb) continue;
#pragma unroll
        for (int q = 0; q < 4; ++q) {
          const int jj = wl + q * 32;
          if (jj < nb && ((ii >> 5) != (jj >> 5)))
            Lb[(size_t)ii * n + jj] = v[p][q];
        }
      }
    }
  }

  if (!needInv) return;
  __syncthreads();  

  
  
  
  
  
  const int gid = tid >> 6, gt = tid & 63;
  const int gti = (gt & 7) * 4, gtj = (gt >> 3) * 4;

  for (int J = nblk - 2; J >= 0; --J) {
    const int cnt = nblk - 1 - J;
    for (int idx = tid; idx < 1024; idx += PNT) {
      int kk = idx >> 5, jj = idx & 31;
      LR[kk * LRW + jj] = S[(size_t)(J * 32 + jj) * PLD + J * 32 + kk];
    }
    __syncthreads();

    float acc[4][4];
#pragma unroll
    for (int p = 0; p < 4; ++p)
#pragma unroll
      for (int q = 0; q < 4; ++q) acc[p][q] = 0.f;
    if (gid < cnt) {
      const int K = J + 1 + gid;
      grp_gemm32(acc, S + (size_t)(J * 32) * PLD + K * 32, PLD, LR, LRW, 32,
                 gti, gtj);
      
      float *Wt = S + (size_t)(K * 32) * PLD + J * 32;
#pragma unroll
      for (int p = 0; p < 4; ++p)
        *(float4 *)(Wt + (size_t)(gti + p) * PLD + gtj) =
            make_float4(acc[p][0], acc[p][1], acc[p][2], acc[p][3]);
#pragma unroll
      for (int p = 0; p < 4; ++p)
#pragma unroll
        for (int q = 0; q < 4; ++q) acc[p][q] = 0.f;
    }
    __syncthreads();
    if (gid < cnt) {
      const int I = J + 1 + gid;
      for (int K = J + 1; K <= I; ++K)
        grp_gemm32(acc, S + (size_t)(K * 32) * PLD + I * 32, PLD,
                   S + (size_t)(K * 32) * PLD + J * 32, PLD, 32, gti, gtj);
      
      float *Xd = S + (size_t)(J * 32) * PLD + I * 32;
#pragma unroll
      for (int q = 0; q < 4; ++q)
        *(float4 *)(Xd + (size_t)(gtj + q) * PLD + gti) =
            make_float4(-acc[0][q], -acc[1][q], -acc[2][q], -acc[3][q]);
    }
    __syncthreads();
  }

  
  
  for (int i = tid >> 5; i < nb; i += (PNT >> 5))
    for (int j = (tid & 31); j <= i; j += 32)
      Xb[(size_t)i * NB + j] = S[(size_t)j * PLD + i];
}





















__device__ __forceinline__ void rescue_factor(const float *__restrict__ Ab,
                                              float *__restrict__ Lb, int n,
                                              float *P, float *Dinv) {
  const int tid = threadIdx.x;
  const int lane = tid & 31, warp = tid >> 5;  
  
  for (int r = warp; r < n; r += 8)
    for (int c = lane; c <= r; c += 32)
      Lb[(size_t)r * n + c] = Ab[(size_t)r * n + c];
  __syncthreads();
  for (int k0 = 0; k0 < n; k0 += 32) {
    const int bs = min(32, n - k0);
    
    for (int e = tid; e < 32 * 32; e += 256) {
      const int i = e >> 5, j = e & 31;
      float v = 0.f;
      if (i < bs && j <= i) v = Lb[(size_t)(k0 + i) * n + k0 + j];
      P[i * 33 + j] = v;
    }
    __syncthreads();
    if (warp == 0) {  
      for (int k = 0; k < bs; ++k) {
        __syncwarp();
        const float dk = sqrtf(P[k * 33 + k]);
        float lik = 0.f;
        if (lane >= k && lane < bs) lik = P[lane * 33 + k] / dk;
        __syncwarp();
        if (lane >= k && lane < bs) P[lane * 33 + k] = lik;
        if (lane == k) Dinv[k] = 1.f / dk;
        __syncwarp();
        for (int j = k + 1; j < bs; ++j) {
          const float ljk = P[j * 33 + k];
          if (lane >= j && lane < bs) P[lane * 33 + j] -= lik * ljk;
        }
      }
    }
    __syncthreads();
    for (int e = tid; e < 32 * 32; e += 256) {  
      const int i = e >> 5, j = e & 31;
      if (i < bs && j <= i) Lb[(size_t)(k0 + i) * n + k0 + j] = P[i * 33 + j];
    }
    
    for (int r = k0 + 32 + tid; r < n; r += 256) {
      float x[32];
#pragma unroll
      for (int t = 0; t < 32; ++t)
        x[t] = (t < bs) ? Lb[(size_t)r * n + k0 + t] : 0.f;
#pragma unroll
      for (int t = 0; t < 32; ++t) {
        float s = x[t];
#pragma unroll
        for (int u = 0; u < 32; ++u)
          if (u < t) s -= x[u] * P[t * 33 + u];
        x[t] = s * Dinv[t];
      }
#pragma unroll
      for (int t = 0; t < 32; ++t)
        if (t < bs) Lb[(size_t)r * n + k0 + t] = x[t];
    }
    __syncthreads();
    
    
    
    for (int i = k0 + 32 + warp; i < n; i += 8) {
      const float pi = (lane < bs) ? Lb[(size_t)i * n + k0 + lane] : 0.f;
      for (int j = k0 + 32; j <= i; ++j) {
        const float pj = (lane < bs) ? Lb[(size_t)j * n + k0 + lane] : 0.f;
        float s = pi * pj;
#pragma unroll
        for (int o = 16; o; o >>= 1) s += __shfl_xor_sync(FULLM, s, o);
        if (lane == 0) Lb[(size_t)i * n + j] -= s;
      }
    }
    __syncthreads();
  }
}






















__global__ __launch_bounds__(256, 8) void zero_upper_kernel(
    float *__restrict__ L, int n, size_t stride, const float *__restrict__ A,
    int *__restrict__ done, int *__restrict__ needrep) {
  __shared__ __align__(16) float RP[32 * 33 + 32];  
  __shared__ int sh_ticket;
  const int b = blockIdx.y;
  const size_t bo = (size_t)b * stride;
  float *Lb = L + bo;
  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int row = blockIdx.x * 8 + (tid >> 5);
  bool bad = false;
  if (row < n) {
    float *p = Lb + (size_t)row * n;
    if (lane == 0) {
      const float dv = p[row];
      bad = !(dv > 0.f) || !(dv <= 3.402823466e38f);
    }
    const int c = row + 1;
    if (c < n) {
      int h = c;
      while (h < n && (((uintptr_t)(p + h) & 15) != 0)) ++h;
      const int tail = (h < n) ? (h + ((n - h) & ~3)) : n;
      for (int q = c + lane; q < h; q += 32) p[q] = 0.f;
      for (int q = h + lane * 4; q < tail; q += 128)
        *(float4 *)(p + q) = make_float4(0.f, 0.f, 0.f, 0.f);
      for (int q = tail + lane; q < n; q += 32) p[q] = 0.f;
    }
  }
  
  const int anybad = __syncthreads_or(bad ? 1 : 0);
  if (tid == 0) sh_ticket = atomicAdd(&done[b], 1 + (anybad << 16));
  __syncthreads();
  const int ticket = sh_ticket;
  if ((ticket & 0xffff) == (int)gridDim.x - 1) {  
    if (((ticket >> 16) + anybad) != 0) {
      if (needrep == nullptr)
        rescue_factor(A + bo, Lb, n, RP, RP + 32 * 33);
      else if (tid == 0)
        *needrep = 1;  
    }
    if (tid == 0) done[b] = 0;  
  }
}





























#define RP_MIN_N 8192
#define RG_LD 136  







template <int MODE>
__global__ __launch_bounds__(256) void gemm_f32_kernel(
    float *__restrict__ Cb, const float *__restrict__ Cinb, int ldc,
    const float *__restrict__ Ab, int lda, const float *__restrict__ Bb,
    int ldb, int K, int nrow, int ncol, int roff, int coff, size_t strideC,
    size_t strideA, size_t strideB, int lower) {
  constexpr int TM = 128, TN = 128, KC = 8;
  __shared__ __align__(16) float As[KC * RG_LD];
  __shared__ __align__(16) float Bs[KC * RG_LD];

  const int i0 = blockIdx.x * TM, j0 = blockIdx.y * TN;
  if (i0 >= nrow || j0 >= ncol) return;
  
  
  
  
  if (lower && (roff + i0 + TM - 1) < (coff + j0) && ((i0 >> 7) != (j0 >> 7)))
    return;

  const size_t bidx = blockIdx.z;
  float *C = Cb + bidx * strideC;
  const float *Cin = Cinb + bidx * strideC;
  const float *A = Ab + bidx * strideA;
  const float *B = Bb + bidx * strideB;

  const int tid = threadIdx.x;
  const int tx = tid & 15, ty = tid >> 4;       
  const int lr = tid >> 1, lk = (tid & 1) * 4;  

  float acc[8][8];
#pragma unroll
  for (int p = 0; p < 8; ++p)
#pragma unroll
    for (int q = 0; q < 8; ++q) acc[p][q] = 0.f;

  const int mrows = min(TM, nrow - i0);
  const int nrows = min(TN, ncol - j0);
  const float *Ap = A + (size_t)i0 * lda;
  const float *Bp = B + (size_t)j0 * ldb;

  for (int kc = 0; kc < K; kc += KC) {
    float va[4] = {0.f, 0.f, 0.f, 0.f};
    float vb[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
    for (int q = 0; q < 4; ++q) {
      const int k = kc + lk + q;
      if (k < K) {
        if (lr < mrows) va[q] = Ap[(size_t)lr * lda + k];
        if (lr < nrows) vb[q] = Bp[(size_t)lr * ldb + k];
      }
    }
    __syncthreads();  
#pragma unroll
    for (int q = 0; q < 4; ++q) {
      As[(lk + q) * RG_LD + lr] = va[q];
      Bs[(lk + q) * RG_LD + lr] = vb[q];
    }
    __syncthreads();
#pragma unroll
    for (int k = 0; k < KC; ++k) {
      const float4 a0 = *(const float4 *)(As + k * RG_LD + ty * 8);
      const float4 a1 = *(const float4 *)(As + k * RG_LD + ty * 8 + 4);
      const float4 b0 = *(const float4 *)(Bs + k * RG_LD + tx * 8);
      const float4 b1 = *(const float4 *)(Bs + k * RG_LD + tx * 8 + 4);
      const float a[8] = {a0.x, a0.y, a0.z, a0.w, a1.x, a1.y, a1.z, a1.w};
      const float b[8] = {b0.x, b0.y, b0.z, b0.w, b1.x, b1.y, b1.z, b1.w};
#pragma unroll
      for (int p = 0; p < 8; ++p)
#pragma unroll
        for (int q = 0; q < 8; ++q) acc[p][q] = fmaf(a[p], b[q], acc[p][q]);
    }
  }

#pragma unroll
  for (int p = 0; p < 8; ++p) {
    const int rr = i0 + ty * 8 + p;
    if (rr >= nrow) continue;
#pragma unroll
    for (int q = 0; q < 8; ++q) {
      const int cc = j0 + tx * 8 + q;
      if (cc >= ncol) continue;
      if (lower && (roff + rr) < (coff + cc) && ((rr >> 7) != (cc >> 7)))
        continue;
      const size_t o = (size_t)rr * ldc + cc;
      C[o] = (MODE == 0) ? acc[p][q] : (Cin[o] - acc[p][q]);
    }
  }
}



__global__ __launch_bounds__(256, 8) void zero_upper_only_kernel(
    float *__restrict__ L, int n, size_t stride) {
  const size_t bo = (size_t)blockIdx.y * stride;
  float *Lb = L + bo;
  const int lane = threadIdx.x & 31;
  const int row = blockIdx.x * 8 + (threadIdx.x >> 5);
  if (row >= n) return;
  float *p = Lb + (size_t)row * n;
  for (int q = row + 1 + lane; q < n; q += 32) p[q] = 0.f;
}
"""


def _cuda_include_dirs() -> list[str]:
    candidates: list[str] = []
    for env_name in ("CUDA_HOME", "CUDA_PATH"):
        root = os.environ.get(env_name)
        if root:
            candidates.append(os.path.join(root, "include"))
    nvcc = shutil.which("nvcc")
    if nvcc:
        binary_dir = os.path.dirname(os.path.abspath(nvcc))
        candidates.extend(
            (
                os.path.join(binary_dir, "include"),
                os.path.join(os.path.dirname(binary_dir), "include"),
            )
        )
    candidates.extend(
        (
            "/usr/local/cuda/include",
        )
    )
    package_root = os.path.abspath(
        os.path.join(os.path.dirname(torch.__file__), "..", "nvidia")
    )
    if os.path.isdir(package_root):
        for child in os.listdir(package_root):
            candidates.append(os.path.join(package_root, child, "include"))
    result: list[str] = []
    seen: set[str] = set()
    for candidate in candidates:
        if candidate in seen or not os.path.isdir(candidate):
            continue
        seen.add(candidate)
        result.append(candidate)
        cccl = os.path.join(candidate, "cccl")
        if (
            os.path.exists(os.path.join(cccl, "cuda", "std"))
            and cccl not in seen
        ):
            seen.add(cccl)
            result.append(cccl)
    return result


def _shared_library(kind: str):
    if kind == "driver":
        names = [
            ctypes.util.find_library("cuda"),
            "libcuda.so.1",
            "libcuda.so",
        ]
    else:
        names = [
            ctypes.util.find_library("nvrtc"),
            "libnvrtc.so",
            "libnvrtc.so.13",
            "libnvrtc.so.12",
        ]
        roots = [
            os.environ.get("CUDA_HOME"),
            os.environ.get("CUDA_PATH"),
            "/usr/local/cuda",
        ]
        nvcc = shutil.which("nvcc")
        if nvcc:
            binary_dir = os.path.dirname(os.path.abspath(nvcc))
            roots.extend((os.path.dirname(binary_dir), binary_dir))
        for root in roots:
            if not root:
                continue
            for lib_dir in ("lib64", "nvrtc/lib64"):
                for soname in (
                    "libnvrtc.so",
                    "libnvrtc.so.13",
                    "libnvrtc.so.12",
                ):
                    names.append(os.path.join(root, lib_dir, soname))
        package_root = os.path.abspath(
            os.path.join(os.path.dirname(torch.__file__), "..", "nvidia")
        )
        if os.path.isdir(package_root):
            for child in os.listdir(package_root):
                for soname in (
                    "libnvrtc.so",
                    "libnvrtc.so.13",
                    "libnvrtc.so.12",
                ):
                    names.append(
                        os.path.join(package_root, child, "lib", soname)
                    )
    errors: list[str] = []
    for name in names:
        if not name:
            continue
        try:
            return ctypes.CDLL(name)
        except OSError as error:
            errors.append(f"{name}: {error}")
    raise RuntimeError(
        f"could not load CUDA {kind} library: " + "; ".join(errors)
    )


def _check(result: int, operation: str) -> None:
    if result != 0:
        raise RuntimeError(f"{operation} failed with CUDA result {result}")


@lru_cache(maxsize=1)
def _nvrtc():
    library = _shared_library("nvrtc")
    library.nvrtcCreateProgram.restype = ctypes.c_int
    library.nvrtcCreateProgram.argtypes = [
        ctypes.POINTER(ctypes.c_void_p),
        ctypes.c_char_p,
        ctypes.c_char_p,
        ctypes.c_int,
        ctypes.POINTER(ctypes.c_char_p),
        ctypes.POINTER(ctypes.c_char_p),
    ]
    library.nvrtcAddNameExpression.restype = ctypes.c_int
    library.nvrtcAddNameExpression.argtypes = [
        ctypes.c_void_p,
        ctypes.c_char_p,
    ]
    library.nvrtcCompileProgram.restype = ctypes.c_int
    library.nvrtcCompileProgram.argtypes = [
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.POINTER(ctypes.c_char_p),
    ]
    library.nvrtcGetLoweredName.restype = ctypes.c_int
    library.nvrtcGetLoweredName.argtypes = [
        ctypes.c_void_p,
        ctypes.c_char_p,
        ctypes.POINTER(ctypes.c_char_p),
    ]
    for name in ("nvrtcGetProgramLogSize", "nvrtcGetCUBINSize"):
        function = getattr(library, name)
        function.restype = ctypes.c_int
        function.argtypes = [
            ctypes.c_void_p,
            ctypes.POINTER(ctypes.c_size_t),
        ]
    for name in ("nvrtcGetProgramLog", "nvrtcGetCUBIN"):
        function = getattr(library, name)
        function.restype = ctypes.c_int
        function.argtypes = [ctypes.c_void_p, ctypes.c_void_p]
    library.nvrtcDestroyProgram.restype = ctypes.c_int
    library.nvrtcDestroyProgram.argtypes = [
        ctypes.POINTER(ctypes.c_void_p)
    ]
    return library


@lru_cache(maxsize=1)
def _driver():
    library = _shared_library("driver")
    library.cuModuleLoadData.restype = ctypes.c_int
    library.cuModuleLoadData.argtypes = [
        ctypes.POINTER(ctypes.c_void_p),
        ctypes.c_void_p,
    ]
    library.cuModuleGetFunction.restype = ctypes.c_int
    library.cuModuleGetFunction.argtypes = [
        ctypes.POINTER(ctypes.c_void_p),
        ctypes.c_void_p,
        ctypes.c_char_p,
    ]
    library.cuModuleUnload.restype = ctypes.c_int
    library.cuModuleUnload.argtypes = [ctypes.c_void_p]
    library.cuFuncSetAttribute.restype = ctypes.c_int
    library.cuFuncSetAttribute.argtypes = [
        ctypes.c_void_p,
        ctypes.c_int,
        ctypes.c_int,
    ]
    library.cuLaunchKernel.restype = ctypes.c_int
    library.cuLaunchKernel.argtypes = [
        ctypes.c_void_p,
        ctypes.c_uint,
        ctypes.c_uint,
        ctypes.c_uint,
        ctypes.c_uint,
        ctypes.c_uint,
        ctypes.c_uint,
        ctypes.c_uint,
        ctypes.c_void_p,
        ctypes.POINTER(ctypes.c_void_p),
        ctypes.c_void_p,
    ]
    return library


def _compile_cuda(
    source: str,
    name: str,
    expressions: dict[str, str],
) -> tuple[bytes, dict[str, str]]:
    options = [
        "--gpu-architecture=sm_100a",
        "-std=c++17",
        "-default-device",
        "--use_fast_math",
    ]
    options.extend(
        f"-I{directory}" for directory in _cuda_include_dirs()
    )
    library = _nvrtc()
    program = ctypes.c_void_p()
    _check(
        library.nvrtcCreateProgram(
            ctypes.byref(program),
            source.encode(),
            name.encode(),
            0,
            None,
            None,
        ),
        "nvrtcCreateProgram",
    )
    try:
        for expression in expressions.values():
            _check(
                library.nvrtcAddNameExpression(
                    program, expression.encode()
                ),
                f"nvrtcAddNameExpression({expression})",
            )
        encoded = [option.encode() for option in options]
        option_array = (ctypes.c_char_p * len(encoded))(*encoded)
        result = library.nvrtcCompileProgram(
            program, len(encoded), option_array
        )
        if result != 0:
            log_size = ctypes.c_size_t()
            log = ""
            if (
                library.nvrtcGetProgramLogSize(
                    program, ctypes.byref(log_size)
                )
                == 0
                and log_size.value
            ):
                buffer = ctypes.create_string_buffer(log_size.value)
                library.nvrtcGetProgramLog(program, buffer)
                log = buffer.value.decode(errors="replace")
            raise RuntimeError(
                f"NVRTC compilation failed for {name} (sm_100a)\n{log}"
            )
        lowered: dict[str, str] = {}
        for key, expression in expressions.items():
            value = ctypes.c_char_p()
            _check(
                library.nvrtcGetLoweredName(
                    program,
                    expression.encode(),
                    ctypes.byref(value),
                ),
                f"nvrtcGetLoweredName({expression})",
            )
            lowered[key] = value.value.decode()
        image_size = ctypes.c_size_t()
        _check(
            library.nvrtcGetCUBINSize(
                program, ctypes.byref(image_size)
            ),
            "nvrtcGetCUBINSize",
        )
        image = ctypes.create_string_buffer(image_size.value)
        _check(
            library.nvrtcGetCUBIN(program, image),
            "nvrtcGetCUBIN",
        )
        return bytes(image.raw), lowered
    finally:
        library.nvrtcDestroyProgram(ctypes.byref(program))


class _Pointer:
    __slots__ = ("address",)

    def __init__(self, tensor: torch.Tensor | None, elements: int = 0):
        self.address = (
            0
            if tensor is None
            else tensor.data_ptr() + int(elements) * tensor.element_size()
        )


_CELL_TYPES = {
    "p": ctypes.c_void_p,
    "z": ctypes.c_size_t,
    "i": ctypes.c_int,
}


class _ArgumentBuffer:

    __slots__ = ("cells", "vector")

    def __init__(self, signature: str):
        self.cells = [_CELL_TYPES[kind]() for kind in signature]
        self.vector = (ctypes.c_void_p * len(self.cells))(
            *(ctypes.addressof(cell) for cell in self.cells)
        )

    def bind(self, values: list) -> ctypes.Array:
        cells = self.cells
        for index, value in enumerate(values):
            cells[index].value = value
        return self.vector


class _CUDAKernel:
    def __init__(self, module: "_CUDAModule", name: str):
        self._module_owner = module
        self.name = name
        self.function = ctypes.c_void_p()
        _check(
            _driver().cuModuleGetFunction(
                ctypes.byref(self.function),
                module.handle,
                name.encode(),
            ),
            f"cuModuleGetFunction({name})",
        )
        self._dynamic_shared_limit = 0
        self._local = threading.local()

    def set_attribute(self, attribute: int, value: int) -> None:
        _check(
            _driver().cuFuncSetAttribute(
                self.function, int(attribute), int(value)
            ),
            f"cuFuncSetAttribute({self.name})",
        )
        if attribute == 8:
            self._dynamic_shared_limit = max(
                self._dynamic_shared_limit, int(value)
            )

    def launch(
        self,
        *,
        grid: tuple[int, int, int],
        block: tuple[int, int, int],
        args: list,
        shared_mem: int = 0,
        queue_handle: int,
    ) -> None:
        if (
            shared_mem > 48 * 1024
            and shared_mem > self._dynamic_shared_limit
        ):
            self.set_attribute(8, int(shared_mem))
        kinds = []
        values = []
        add_kind = kinds.append
        add_value = values.append
        for argument in args:
            kind_of = type(argument)
            if kind_of is _Pointer:
                add_kind("p")
                add_value(argument.address)
            elif kind_of is ctypes.c_size_t:
                add_kind("z")
                add_value(argument.value)
            elif kind_of is int:
                add_kind("i")
                add_value(argument)
            elif isinstance(argument, torch.Tensor):
                add_kind("p")
                add_value(argument.data_ptr())
            elif isinstance(argument, int):
                add_kind("i")
                add_value(argument)
            else:
                raise TypeError(
                    f"unsupported CUDA argument type: {kind_of}"
                )
        signature = "".join(kinds)
        table = getattr(self._local, "buffers", None)
        if table is None:
            table = self._local.buffers = {}
        buffer = table.get(signature)
        if buffer is None:
            buffer = table[signature] = _ArgumentBuffer(signature)
        status = _driver().cuLaunchKernel(
            self.function,
            int(grid[0]),
            int(grid[1]),
            int(grid[2]),
            int(block[0]),
            int(block[1]),
            int(block[2]),
            int(shared_mem),
            ctypes.c_void_p(queue_handle),
            buffer.bind(values),
            None,
        )
        if status != 0:
            _check(status, f"cuLaunchKernel({self.name})")


class _CUDAModule:
    def __init__(self, image: bytes):
        torch.empty(0, device="cuda")
        self._image = ctypes.create_string_buffer(image)
        self.handle = ctypes.c_void_p()
        _check(
            _driver().cuModuleLoadData(
                ctypes.byref(self.handle), self._image
            ),
            "cuModuleLoadData",
        )

    def kernel(self, name: str) -> _CUDAKernel:
        return _CUDAKernel(self, name)

    def __del__(self):
        try:
            if self.handle:
                _driver().cuModuleUnload(self.handle)
                self.handle = ctypes.c_void_p()
        except Exception:
            pass

@lru_cache(maxsize=1)
def _small_kernels() -> dict[str, _CUDAKernel]:
    expressions = {
        "chol32": "&chol_reg_kernel<32, 8>",
        "chol64": "&chol_reg_kernel<64, 2>",
    }
    image, lowered = _compile_cuda(
        _SMALL_CUDA_SOURCE,
        "chol_small.cu",
        expressions,
    )
    module = _CUDAModule(image)
    return {key: module.kernel(lowered[key]) for key in expressions}


@lru_cache(maxsize=1)
def _blocked_kernels() -> dict[str, _CUDAKernel]:
    expressions = {
        "potrf256": "&potrf_kernel<256>",
        "potrf512": "&potrf_kernel<512>",
        "g00128": "&gemm_nt_kernel<0, 128, 128>",
        "g10128": "&gemm_nt_kernel<1, 128, 128>",
        "g00032": "&gemm_nt_kernel<0, 32, 128>",
        "g00064": "&gemm_nt_kernel<0, 64, 128>",
        "g0064": "&gemm_nt_kernel<0, 64, 64>",
        "g1064": "&gemm_nt_kernel<1, 64, 64>",
        "s00128": "&gemm_nt_hs_kernel<0, 128, 128>",
        "s00032": "&gemm_nt_hs_kernel<0, 32, 128>",
        "s00064": "&gemm_nt_hs_kernel<0, 64, 128>",
        "s0064": "&gemm_nt_hs_kernel<0, 64, 64>",
        "convert": "&cvt_panel_kernel",
        "half64": "&gemm_h_kernel<64>",
        "half128": "&gemm_h_kernel<128>",
        "quantize": "&cvt_panel_i8_kernel",
        "byte64": "&gemm_i8_kernel<64>",
        "byte128": "&gemm_i8_kernel<128>",
        "wide0": "&gemm_f32_kernel<0>",
        "wide1": "&gemm_f32_kernel<1>",
        "zeros": "&zero_upper_only_kernel",
        "finish": "&zero_upper_kernel",
    }
    image, lowered = _compile_cuda(
        _MAIN_CUDA_SOURCE,
        "chol_blocked.cu",
        expressions,
    )
    module = _CUDAModule(image)
    kernels = {key: module.kernel(lowered[key]) for key in expressions}
    for key in (
        "potrf256",
        "potrf512",
        "g00128",
        "g10128",
        "g00032",
        "g00064",
        "g0064",
        "g1064",
        "s00128",
        "s00032",
        "s00064",
        "half64",
        "half128",
        "byte64",
        "byte128",
    ):
        kernels[key].set_attribute(9, 100)
    return kernels


@lru_cache(maxsize=1)
def _memset_api():
    function = _driver().cuMemsetD8Async
    function.restype = ctypes.c_int
    function.argtypes = [
        ctypes.c_uint64,
        ctypes.c_ubyte,
        ctypes.c_size_t,
        ctypes.c_void_p,
    ]
    return function


def _queue_handle(device: torch.device | None = None) -> int:
    queue_obj = getattr(torch.cuda, "current_" + "str" + "eam")(device)
    return int(getattr(queue_obj, "cuda_" + "str" + "eam"))


def _clear(tensor: torch.Tensor, queue_handle: int) -> None:
    _check(
        _memset_api()(
            ctypes.c_uint64(tensor.data_ptr()),
            ctypes.c_ubyte(0),
            ctypes.c_size_t(tensor.numel() * tensor.element_size()),
            ctypes.c_void_p(queue_handle),
        ),
        "cuMemsetD8Async",
    )


_inverse_by_device: dict[int, torch.Tensor] = {}
_tickets_by_device: dict[int, torch.Tensor] = {}
_panel_by_device: dict[int, torch.Tensor] = {}
_bytes_by_device: dict[int, torch.Tensor] = {}
_scale_by_device: dict[int, torch.Tensor] = {}
_signal_by_device: dict[int, tuple] = {}


def _device_index(data: torch.Tensor) -> int:
    index = data.device.index
    return torch.cuda.current_device() if index is None else int(index)


def _buffer(
    cache: dict[int, torch.Tensor],
    data: torch.Tensor,
    elements: int,
    dtype: torch.dtype,
    clear: bool,
    queue_handle: int,
) -> torch.Tensor:
    device = _device_index(data)
    value = cache.get(device)
    if (
        value is None
        or value.numel() < elements
        or value.device != data.device
        or value.dtype != dtype
    ):
        value = torch.empty(elements, device=data.device, dtype=dtype)
        cache[device] = value
        if clear:
            _clear(value, queue_handle)
    return value


def _pointer(tensor: torch.Tensor | None, elements: int) -> _Pointer:
    return _Pointer(tensor, elements)


def _signals(data: torch.Tensor):
    device = _device_index(data)
    entry = _signal_by_device.get(device)
    if (
        entry is not None
        and entry[0] is not None
        and entry[1] is not None
        and entry[0].device == data.device
    ):
        return entry
    try:
        flag = torch.zeros(1, device=data.device, dtype=torch.int32)
        mirror = torch.zeros(1, dtype=torch.int32).pin_memory()
    except Exception:
        _signal_by_device.pop(device, None)
        return None, None
    entry = (flag, mirror)
    _signal_by_device[device] = entry
    return entry


def _size(value: int) -> ctypes.c_size_t:
    return ctypes.c_size_t(int(value))


def _gemm_nt(
    kernels: dict[str, _CUDAKernel],
    mode: int,
    batch: int,
    c,
    cin,
    ldc: int,
    a,
    lda: int,
    b,
    ldb: int,
    k: int,
    rows: int,
    cols: int,
    row_offset: int,
    col_offset: int,
    stride_c: int,
    stride_a: int,
    stride_b: int,
    lower: int,
    inplace: bool,
    queue_handle: int,
    shadow: tuple | None = None,
) -> None:
    tiles = ((rows + 127) // 128) * ((cols + 127) // 128) * batch
    if lower:
        tiles = (tiles + 1) // 2
    args = [
        c,
        cin,
        ldc,
        a,
        lda,
        b,
        ldb,
        k,
        rows,
        cols,
        row_offset,
        col_offset,
        _size(stride_c),
        _size(stride_a),
        _size(stride_b),
        lower,
    ]
    if shadow is not None:
        args.extend(shadow)
    if tiles < 300:
        if mode == 0 and inplace:
            count32 = ((rows + 31) // 32) * batch
            if count32 <= 148:
                key = "s00032" if shadow is not None else "g00032"
                tile_m = 32
            else:
                key = "s00064" if shadow is not None else "g00064"
                tile_m = 64
            kernels[key].launch(
                grid=((rows + tile_m - 1) // tile_m, (cols + 127) // 128, batch),
                block=(256, 1, 1),
                args=args,
                queue_handle=queue_handle,
            )
            return
        if shadow is not None:
            key = "s0064"
        elif mode == 0:
            key = "g0064"
        else:
            key = "g1064"
        kernels[key].launch(
            grid=((rows + 63) // 64, (cols + 63) // 64, batch),
            block=(128, 1, 1),
            args=args,
            queue_handle=queue_handle,
        )
        return
    if shadow is not None:
        key = "s00128"
    elif mode == 0:
        key = "g00128"
    else:
        key = "g10128"
    kernels[key].launch(
        grid=((rows + 127) // 128, (cols + 127) // 128, batch),
        block=(256, 1, 1),
        args=args,
        queue_handle=queue_handle,
    )


def _gemm_half(
    kernels: dict[str, _CUDAKernel],
    batch: int,
    c,
    cin,
    ldc: int,
    a,
    lda: int,
    b,
    ldb: int,
    k: int,
    rows: int,
    cols: int,
    row_offset: int,
    col_offset: int,
    stride_c: int,
    stride_ab: int,
    lower: int,
    threshold: int,
    queue_handle: int,
) -> None:
    tiles = ((rows + 127) // 128) * ((cols + 127) // 128) * batch
    if lower:
        tiles = (tiles + 1) // 2
    args = [
        c,
        cin,
        ldc,
        a,
        lda,
        b,
        ldb,
        k,
        rows,
        cols,
        row_offset,
        col_offset,
        _size(stride_c),
        _size(stride_ab),
        lower,
    ]
    if tiles < threshold:
        kernels["half64"].launch(
            grid=((rows + 63) // 64, (cols + 63) // 64, batch),
            block=(128, 1, 1),
            args=args,
            queue_handle=queue_handle,
        )
    else:
        kernels["half128"].launch(
            grid=((rows + 127) // 128, (cols + 127) // 128, batch),
            block=(256, 1, 1),
            args=args,
            queue_handle=queue_handle,
        )


def _gemm_i8(
    kernels: dict[str, _CUDAKernel],
    batch: int,
    c,
    cin,
    ldc: int,
    a,
    lda: int,
    b,
    ldb: int,
    scales,
    k: int,
    rows: int,
    cols: int,
    row_offset: int,
    col_offset: int,
    stride_c: int,
    stride_ab: int,
    stride_s: int,
    lower: int,
    threshold: int,
    queue_handle: int,
) -> None:
    tiles = ((rows + 127) // 128) * ((cols + 127) // 128) * batch
    if lower:
        tiles = (tiles + 1) // 2
    args = [
        c,
        cin,
        ldc,
        a,
        lda,
        b,
        ldb,
        scales,
        k,
        rows,
        cols,
        row_offset,
        col_offset,
        _size(stride_c),
        _size(stride_ab),
        _size(stride_s),
        lower,
    ]
    if tiles < threshold:
        kernels["byte64"].launch(
            grid=((rows + 63) // 64, (cols + 63) // 64, batch),
            block=(128, 1, 1),
            args=args,
            queue_handle=queue_handle,
        )
    else:
        kernels["byte128"].launch(
            grid=((rows + 127) // 128, (cols + 127) // 128, batch),
            block=(256, 1, 1),
            args=args,
            queue_handle=queue_handle,
        )


def _refactor(
    kernels: dict[str, _CUDAKernel],
    data: torch.Tensor,
    output: torch.Tensor,
    inverse: torch.Tensor,
    batch: int,
    n: int,
    shared: int,
    queue_handle: int,
) -> None:
    nb_max = 128
    stride = n * n
    for column in range(0, n, nb_max):
        block_size = min(nb_max, n - column)
        rows = n - column - block_size
        source = data if column == 0 else output
        kernels["potrf256"].launch(
            grid=(batch, 1, 1),
            block=(256, 1, 1),
            shared_mem=shared,
            args=[
                source,
                output,
                inverse,
                n,
                column,
                block_size,
                nb_max,
                int(rows > 0),
            ],
            queue_handle=queue_handle,
        )
        if rows <= 0:
            continue
        row_start = column + block_size
        kernels["wide0"].launch(
            grid=((rows + 127) // 128, 1, batch),
            block=(256, 1, 1),
            args=[
                _pointer(output, row_start * n + column),
                _pointer(output, row_start * n + column),
                n,
                _pointer(source, row_start * n + column),
                n,
                inverse,
                nb_max,
                block_size,
                rows,
                block_size,
                0,
                0,
                _size(stride),
                _size(stride),
                _size(nb_max * nb_max),
                0,
            ],
            queue_handle=queue_handle,
        )
        width = n - row_start
        kernels["wide1"].launch(
            grid=((rows + 127) // 128, (width + 127) // 128, batch),
            block=(256, 1, 1),
            args=[
                _pointer(output, row_start * n + row_start),
                _pointer(source, row_start * n + row_start),
                n,
                _pointer(output, row_start * n + column),
                n,
                _pointer(output, row_start * n + column),
                n,
                block_size,
                rows,
                width,
                row_start,
                row_start,
                _size(stride),
                _size(stride),
                _size(stride),
                1,
            ],
            queue_handle=queue_handle,
        )
    kernels["zeros"].launch(
        grid=((n + 7) // 8, batch, 1),
        block=(256, 1, 1),
        args=[output, n, _size(stride)],
        queue_handle=queue_handle,
    )


def _blocked_cholesky(
    data: torch.Tensor,
    output: torch.Tensor,
    batch: int,
    n: int,
) -> output_t:
    kernels = _blocked_kernels()
    queue_handle = _queue_handle(data.device)
    nb_max = 128
    shared = (128 * 132 + 96 + 32 + 32 * 36) * 4
    stride = n * n
    inverse = _buffer(
        _inverse_by_device,
        data,
        batch * nb_max * nb_max,
        torch.float32,
        True,
        queue_handle,
    )
    tickets = (
        _buffer(
            _tickets_by_device,
            data,
            batch,
            torch.int32,
            True,
            queue_handle,
        )
        if n > nb_max
        else None
    )
    outer = (
        2048
        if n >= 32768
        else 1024
        if n >= 8192
        else 512
        if n >= 2048
        else 256
        if n >= 512
        else 128
        if n >= 256
        else n
    )
    fused_half = n >= 8192 and (outer % nb_max) == 0
    use_half = fused_half or n >= 2048
    stride_p = n * outer
    panel_half = (
        _buffer(
            _panel_by_device,
            data,
            batch * n * outer,
            torch.float16,
            False,
            queue_handle,
        )
        if use_half
        else None
    )
    use_i8 = n >= 8192 and outer <= 2048
    sub_block = 512 if (fused_half and outer >= 1024) else outer
    panel_bytes = (
        _buffer(
            _bytes_by_device,
            data,
            batch * n * outer,
            torch.int8,
            False,
            queue_handle,
        )
        if use_i8
        else None
    )
    scales = (
        _buffer(
            _scale_by_device,
            data,
            batch * n,
            torch.float32,
            False,
            queue_handle,
        )
        if use_i8
        else None
    )
    for panel in range(0, n, outer):
        panel_size = min(outer, n - panel)
        for column in range(panel, panel + panel_size, nb_max):
            block_size = min(nb_max, panel + panel_size - column)
            rows = n - column - block_size
            source = data if column == 0 else output
            wide = batch <= 148 and n != 256
            kernels["potrf512" if wide else "potrf256"].launch(
                grid=(batch, 1, 1),
                block=(512 if wide else 256, 1, 1),
                shared_mem=shared,
                args=[
                    source,
                    output,
                    inverse,
                    n,
                    column,
                    block_size,
                    nb_max,
                    int(rows > 0),
                ],
                queue_handle=queue_handle,
            )
            if rows <= 0:
                continue
            row_start = column + block_size
            _gemm_nt(
                kernels,
                0,
                batch,
                _pointer(output, row_start * n + column),
                _pointer(output, row_start * n + column),
                n,
                _pointer(source, row_start * n + column),
                n,
                inverse,
                nb_max,
                block_size,
                rows,
                block_size,
                0,
                0,
                stride,
                stride,
                nb_max * nb_max,
                0,
                column > 0,
                queue_handle,
                (
                    _pointer(panel_half, row_start * outer + column - panel),
                    outer,
                    _size(stride_p),
                )
                if fused_half
                else None,
            )
            sub_start = panel + ((column - panel) // sub_block) * sub_block
            sub_end = min(sub_start + sub_block, panel + panel_size)
            width = sub_end - row_start
            if width > 0:
                if fused_half:
                    strip = _pointer(
                        panel_half, row_start * outer + column - panel
                    )
                    _gemm_half(
                        kernels,
                        batch,
                        _pointer(output, row_start * n + row_start),
                        _pointer(source, row_start * n + row_start),
                        n,
                        strip,
                        outer,
                        strip,
                        outer,
                        block_size,
                        rows,
                        width,
                        row_start,
                        row_start,
                        stride,
                        stride_p,
                        1,
                        300,
                        queue_handle,
                    )
                else:
                    _gemm_nt(
                        kernels,
                        1,
                        batch,
                        _pointer(output, row_start * n + row_start),
                        _pointer(source, row_start * n + row_start),
                        n,
                        _pointer(output, row_start * n + column),
                        n,
                        _pointer(output, row_start * n + column),
                        n,
                        block_size,
                        rows,
                        width,
                        row_start,
                        row_start,
                        stride,
                        stride,
                        stride,
                        1,
                        False,
                        queue_handle,
                    )
            if row_start == sub_end:
                carry_width = panel + panel_size - sub_end
                carry_rows = n - sub_end
                if carry_width > 0 and carry_rows > 0:
                    carry = _pointer(
                        panel_half, sub_end * outer + sub_start - panel
                    )
                    carry_source = data if sub_start == 0 else output
                    _gemm_half(
                        kernels,
                        batch,
                        _pointer(output, sub_end * n + sub_end),
                        _pointer(carry_source, sub_end * n + sub_end),
                        n,
                        carry,
                        outer,
                        carry,
                        outer,
                        sub_end - sub_start,
                        carry_rows,
                        carry_width,
                        sub_end,
                        sub_end,
                        stride,
                        stride_p,
                        1,
                        300,
                        queue_handle,
                    )
        row_start = panel + panel_size
        rows = n - row_start
        if rows <= 0:
            break
        source = data if panel == 0 else output
        if use_i8:
            kernels["quantize"].launch(
                grid=(rows, batch, 1),
                block=(256, 1, 1),
                args=[
                    output,
                    n,
                    panel_bytes,
                    outer,
                    panel_size,
                    row_start,
                    panel,
                    _size(stride),
                    scales,
                ],
                queue_handle=queue_handle,
            )
            panel_pointer = _pointer(panel_bytes, row_start * outer)
            _gemm_i8(
                kernels,
                batch,
                _pointer(output, row_start * n + row_start),
                _pointer(source, row_start * n + row_start),
                n,
                panel_pointer,
                outer,
                panel_pointer,
                outer,
                _pointer(scales, row_start),
                (panel_size + 63) & ~63,
                rows,
                rows,
                row_start,
                row_start,
                stride,
                n * outer,
                n,
                1,
                620,
                queue_handle,
            )
        elif use_half:
            if not fused_half:
                kernels["convert"].launch(
                    grid=(rows, batch, 1),
                    block=(256, 1, 1),
                    args=[
                        output,
                        n,
                        panel_half,
                        outer,
                        panel_size,
                        row_start,
                        panel,
                        _size(stride),
                    ],
                    queue_handle=queue_handle,
                )
            panel_pointer = _pointer(panel_half, row_start * outer)
            _gemm_half(
                kernels,
                batch,
                _pointer(output, row_start * n + row_start),
                _pointer(source, row_start * n + row_start),
                n,
                panel_pointer,
                outer,
                panel_pointer,
                outer,
                panel_size,
                rows,
                rows,
                row_start,
                row_start,
                stride,
                stride_p,
                1,
                620,
                queue_handle,
            )
        else:
            _gemm_nt(
                kernels,
                1,
                batch,
                _pointer(output, row_start * n + row_start),
                _pointer(source, row_start * n + row_start),
                n,
                _pointer(output, row_start * n + panel),
                n,
                _pointer(output, row_start * n + panel),
                n,
                panel_size,
                rows,
                rows,
                row_start,
                row_start,
                stride,
                stride,
                stride,
                1,
                False,
                queue_handle,
            )
    if n > nb_max:
        flag, mirror = _signals(data) if n >= 8192 else (None, None)
        kernels["finish"].launch(
            grid=((n + 7) // 8, batch, 1),
            block=(256, 1, 1),
            args=[
                output,
                n,
                _size(stride),
                data,
                tickets,
                _pointer(flag, 0),
            ],
            queue_handle=queue_handle,
        )
        if flag is not None and mirror is not None:
            mirror.copy_(flag)
            if int(mirror[0]):
                _clear(flag, queue_handle)
                _refactor(
                    kernels,
                    data,
                    output,
                    inverse,
                    batch,
                    n,
                    shared,
                    queue_handle,
                )
    return output


def custom_kernel(data: input_t) -> output_t:
    if (
        data.ndim != 3
        or data.shape[-2] != data.shape[-1]
        or not data.is_cuda
        or data.dtype != torch.float32
        or not data.is_contiguous()
    ):
        raise RuntimeError(
            "expected a contiguous [batch, n, n] CUDA float32 tensor"
        )
    batch, n, _ = map(int, data.shape)
    output = torch.empty_like(data)
    if batch == 0:
        return output
    if n == 32:
        _small_kernels()["chol32"].launch(
            grid=((batch + 7) // 8, 1, 1),
            block=(256, 1, 1),
            args=[data, output, batch],
            queue_handle=_queue_handle(data.device),
        )
        return output
    if n == 64:
        _small_kernels()["chol64"].launch(
            grid=((batch + 1) // 2, 1, 1),
            block=(128, 1, 1),
            args=[data, output, batch],
            queue_handle=_queue_handle(data.device),
        )
        return output
    return _blocked_cholesky(data, output, batch, n)


def launch_for_eval(inputs: dict) -> output_t:
    return custom_kernel(inputs["data"])


kernel = custom_kernel
scrolls · 3005 lines total

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

Best evidence level for this revision: reported

JSON