Skip to content
KernelIndex
Search⌘K

submission 901618

damngamerz · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-901618?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
674.4µs
#59 of 337
2026-07-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b8fed08771dcb872c1b559fe68c4bcc13a134edaad6b8024e746393255d3401c
license declaredunknown
license concludedunknown
authorsdamngamerz
imported2026-08-26

Techniques

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

cluster__global__ __cluster_dims__(BPC,1,1) __launch_bounds__(P::max_threads_per_block) void cga_chol(T*a,unsigned ld,int*info,unsigned batches){
fused-epilogueusing SyrkEpilogue=typename cutlass::epilogue::collective::CollectiveBuilder<cutlass::arch::Sm100,cutlass::arch::OpClassTensorOp,SyrkTile,SyrkCluster,cutlass::epilogue::collective:…
mmav -= tl.dot(prev, tl.trans(prev), input_precision=PREC)
num-warps = 1num_warps=1 if n == 512 else 2)
shared-memoryextern __shared__ __align__(16) cusolverdx::byte storage[]; float* a=(float*)storage;
vector-width = float4for(int z=0;z<BPB;z++){int bi=first+z;if(bi<batches)for(int i=threadIdx.x;i<N*N/4;i+=blockDim.x)((float4*)(a+z*N*N))[i]=((const float4*)(input+size_t(bi)*N*N))[i];}
warp-specialization…izeof(SyrkEpilogue::SharedStorage)>,cutlass::gemm::KernelTmaWarpSpecialized2SmSm100>::CollectiveOp;…

Kernel source

submission.py579 lines
import hashlib
import importlib.util
import os
import pathlib
import subprocess
import sysconfig


def _load_dx():
    root = pathlib.Path.home() / ".cache" / "chol_dx_sm100_tcTriInvV2"
    root.mkdir(parents=True, exist_ok=True)
    so = root / "chol_dx_sm100_tcTriInvV2.so"
    cpp = r'''
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/CUDABlas.h>
#include <c10/cuda/CUDAGuard.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
extern void chol_cga_launch(const float*,float*,int,int);
extern void chol_launch(const float*,float*,int,int);
extern void chol_blocked_launch(const float*,float*,float*,int,int,cublasHandle_t);
extern void chol_vendor_launch(const float*,float*,float*,int*,int,int,cusolverDnHandle_t);
extern void lower_syrk_launch(const void*,float*,void*,size_t,int,int,int,int);
extern void lower_bsyrk_launch(const void*,float*,void*,size_t,int,int,int,int);
extern void lower_copy_launch(const float*,float*,int,int);
torch::Tensor chol_dx(torch::Tensor x){c10::cuda::CUDAGuard g(x.device());int b=x.size(0),n=x.size(1);auto out=torch::empty_like(x);chol_launch(x.data_ptr<float>(),out.data_ptr<float>(),b,n);return out;}
torch::Tensor chol_cga(torch::Tensor x){c10::cuda::CUDAGuard g(x.device());auto out=torch::empty_like(x);chol_cga_launch(x.data_ptr<float>(),out.data_ptr<float>(),x.size(1),x.size(0));return out.transpose(1,2);}
torch::Tensor chol_blocked(torch::Tensor x){c10::cuda::CUDAGuard g(x.device());int64_t b=x.size(0),n=x.size(1),count=b*n*n;auto storage=torch::empty({count+b*4096},x.options());auto out=storage.narrow(0,0,count).view_as(x);float* inv=storage.data_ptr<float>()+count;chol_blocked_launch(x.data_ptr<float>(),out.data_ptr<float>(),inv,b,n,at::cuda::getCurrentCUDABlasHandle());return out.transpose(1,2);}
static cusolverDnHandle_t vendor_handle(){static cusolverDnHandle_t h=[](){cusolverDnHandle_t v;cusolverDnCreate(&v);return v;}();return h;}
torch::Tensor chol_vendor(torch::Tensor x){c10::cuda::CUDAGuard g(x.device());int b=x.size(0),n=x.size(1);auto out=torch::empty_like(x);auto info=torch::empty({b},x.options().dtype(torch::kInt32));auto work=torch::empty({n,n},x.options());chol_vendor_launch(x.data_ptr<float>(),out.data_ptr<float>(),work.data_ptr<float>(),info.data_ptr<int>(),b,n,vendor_handle());return out.transpose(1,2);}
torch::Tensor lower_copy(torch::Tensor x){c10::cuda::CUDAGuard g(x.device());auto out=torch::empty_like(x);lower_copy_launch(x.data_ptr<float>(),out.data_ptr<float>(),x.size(0),x.size(1));return out;}
void lower_bsyrk(torch::Tensor panel,torch::Tensor trail,int64_t band){c10::cuda::CUDAGuard g(panel.device());lower_bsyrk_launch(panel.data_ptr(),trail.data_ptr<float>(),nullptr,0,panel.size(0),panel.size(1),trail.stride(0),band);}
void lower_syrk(torch::Tensor panel,torch::Tensor trail,int64_t band){c10::cuda::CUDAGuard g(panel.device());lower_syrk_launch(panel.data_ptr(),trail.data_ptr<float>(),nullptr,0,panel.size(0),panel.size(1),trail.stride(0),band);}
PYBIND11_MODULE(TORCH_EXTENSION_NAME,m){m.def("chol_dx",&chol_dx);m.def("chol_cga",&chol_cga);m.def("chol_blocked",&chol_blocked);m.def("chol_vendor",&chol_vendor);m.def("lower_copy",&lower_copy);m.def("lower_syrk",&lower_syrk);m.def("lower_bsyrk",&lower_bsyrk);}
'''
    cu = r'''
#include <cuda_runtime.h>
#include <cooperative_groups.h>
#include <cublasdx.hpp>
#include <cusolverdx_io.hpp>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cutlass/cutlass.h>
#include <cutlass/half.h>
#include <cutlass/gemm/device/gemm_universal_adapter.h>
#include <cutlass/gemm/collective/collective_builder.hpp>
#include <cutlass/epilogue/collective/collective_builder.hpp>
#include <cutlass/gemm/kernel/gemm_universal.hpp>
#undef __CUDA_NO_HALF_OPERATORS__
#undef __CUDA_NO_HALF_CONVERSIONS__
#undef __CUDA_NO_BFLOAT16_CONVERSIONS__
#undef __CUDA_NO_HALF2_OPERATORS__
#include <cusolverdx.hpp>

template<int N,int BPB,class Solver>
__global__ __launch_bounds__(Solver::max_threads_per_block)
void kernel(const float* input,float* output,int batches){
 CUSOLVERDX_SKIP_IF_NOT_APPLICABLE_SM(Solver); int first=blockIdx.x*BPB;if(first>=batches)return;
 extern __shared__ __align__(16) cusolverdx::byte storage[]; float* a=(float*)storage;
 for(int z=0;z<BPB;z++){int bi=first+z;if(bi<batches)for(int i=threadIdx.x;i<N*N/4;i+=blockDim.x)((float4*)(a+z*N*N))[i]=((const float4*)(input+size_t(bi)*N*N))[i];}
 // The status words are dead after POTRF, so borrow output storage that the
 // fused lower-triangle write below immediately overwrites.
 __syncthreads(); Solver().execute(a,reinterpret_cast<int*>(output+size_t(first)*N*N)); __syncthreads();
 for(int z=0;z<BPB;z++){int bi=first+z;if(bi<batches)for(int i=threadIdx.x;i<N*N;i+=blockDim.x){int r=i/N,c=i-r*N;output[size_t(bi)*N*N+i]=r>=c?a[z*N*N+i]:0.0f;}}
}
template<int N,int BPB,int NT>void go(const float*in,float*out,int b ){
 using namespace cusolverdx;using Solver=decltype(Size<N>()+Precision<float>()+Type<type::real>()+Function<function::potrf>()+FillMode<fill_mode::lower>()+Arrangement<row_major>()+SM<1000>()+Block()+BatchesPerBlock<BPB>()+BlockDim<NT>());
 auto fn=kernel<N,BPB,Solver>;static bool ok=[](auto f){cudaFuncSetAttribute(f,cudaFuncAttributeMaxDynamicSharedMemorySize,Solver::shared_memory_size);return true;}(fn);(void)ok;
 fn<<<(b+BPB-1)/BPB,Solver::block_dim,Solver::shared_memory_size>>>(in,out,b);
}

template<class Solver>
__global__ __launch_bounds__(Solver::max_threads_per_block)
void diag64(float* out,int* info,int n,int batches,int k){
 CUSOLVERDX_SKIP_IF_NOT_APPLICABLE_SM(Solver);int bi=blockIdx.x;if(bi>=batches)return;
 extern __shared__ __align__(16) cusolverdx::byte storage[];float* a=(float*)storage;
 float* src=out+size_t(bi)*n*n+k*n+k;
 for(int q=threadIdx.x;q<1024;q+=blockDim.x){int r=q>>4,c=(q&15)<<2;float4 v=*((const float4*)(src+r*n+c));float*d=a+r*65+c;d[0]=v.x;d[1]=v.y;d[2]=v.z;d[3]=v.w;}
 __syncthreads();Solver().execute(a,info+bi);__syncthreads();
 for(int i=threadIdx.x;i<4096;i+=blockDim.x){int r=i/64,c=i%64;if(r>=c)src[r*n+c]=a[r*65+c];}
}
using DiagSolver64=decltype(cusolverdx::Size<64>()+cusolverdx::LeadingDimension<65>()+cusolverdx::Precision<float>()+cusolverdx::Type<cusolverdx::type::real>()+cusolverdx::Function<cusolverdx::function::potrf>()+cusolverdx::FillMode<cusolverdx::fill_mode::lower>()+cusolverdx::Arrangement<cusolverdx::row_major>()+cusolverdx::SM<1000>()+cusolverdx::Block()+cusolverdx::BatchesPerBlock<1>()+cusolverdx::BlockDim<64>());
using DiagSolver128=decltype(cusolverdx::Size<64>()+cusolverdx::LeadingDimension<65>()+cusolverdx::Precision<float>()+cusolverdx::Type<cusolverdx::type::real>()+cusolverdx::Function<cusolverdx::function::potrf>()+cusolverdx::FillMode<cusolverdx::fill_mode::lower>()+cusolverdx::Arrangement<cusolverdx::row_major>()+cusolverdx::SM<1000>()+cusolverdx::Block()+cusolverdx::BatchesPerBlock<1>()+cusolverdx::BlockDim<128>());
using DiagSolver256=decltype(cusolverdx::Size<64>()+cusolverdx::LeadingDimension<65>()+cusolverdx::Precision<float>()+cusolverdx::Type<cusolverdx::type::real>()+cusolverdx::Function<cusolverdx::function::potrf>()+cusolverdx::FillMode<cusolverdx::fill_mode::lower>()+cusolverdx::Arrangement<cusolverdx::row_major>()+cusolverdx::SM<1000>()+cusolverdx::Block()+cusolverdx::BatchesPerBlock<1>()+cusolverdx::BlockDim<256>());
template<class Solver> __global__ __launch_bounds__(256) void diag_inv64(float*out,float*dst,int*info,int n,int batches,int k){int bi=blockIdx.x;if(bi>=batches)return;extern __shared__ __align__(16) cusolverdx::byte sm[];float*a=(float*)sm;float*inv=(float*)(sm+Solver::shared_memory_size);float*tmp=inv+4096;float*src=out+size_t(bi)*n*n+k*n+k;for(int q=threadIdx.x;q<1024;q+=blockDim.x){int r=q>>4,c=(q&15)<<2;float4 v=*((const float4*)(src+r*n+c));float*d=a+r*65+c;d[0]=v.x;d[1]=v.y;d[2]=v.z;d[3]=v.w;}for(int q=threadIdx.x;q<4096;q+=blockDim.x){inv[q]=0;tmp[q]=0;}__syncthreads();Solver().execute(a,info+bi);__syncthreads();for(int q=threadIdx.x;q<4096;q+=blockDim.x){int r=q/64,c=q%64;if(r>=c)src[r*n+c]=a[r*65+c];}if(threadIdx.x<64){int leaf=threadIdx.x/16,j=threadIdx.x%16,c=leaf*16+j;float x[16];for(int i=0;i<16;i++){int r=leaf*16+i;float v=(i==j);if(j<=i){for(int q=0;q<i;q++)v-=a[r*65+leaf*16+q]*x[q];v/=a[r*65+r];}x[i]=v;}for(int i=0;i<16;i++)if(j<=i)inv[(leaf*16+i)*64+c]=x[i];}__syncthreads();for(int width=16;width<=32;width*=2){int pairs=64/(2*width),count=pairs*width*width;for(int q=threadIdx.x;q<count;q+=blockDim.x){int pair=q/(width*width),z=q%(width*width),i=z/width,j=z%width,base=pair*2*width;float v=0;for(int p=0;p<width;p++)v+=a[(base+width+i)*65+base+p]*inv[(base+p)*64+base+j];tmp[(base+width+i)*64+base+j]=v;}__syncthreads();for(int q=threadIdx.x;q<count;q+=blockDim.x){int pair=q/(width*width),z=q%(width*width),i=z/width,j=z%width,base=pair*2*width;float v=0;for(int p=0;p<width;p++)v+=inv[(base+width+i)*64+base+width+p]*tmp[(base+width+p)*64+base+j];inv[(base+width+i)*64+base+j]=-v;}__syncthreads();}for(int q=threadIdx.x;q<4096;q+=blockDim.x){int r=q/64,c=q%64;a[c*65+r]=inv[q];}__syncthreads();float*d=dst+size_t(bi)*4096;for(int q=threadIdx.x;q<4096;q+=blockDim.x){int r=q/64,c=q%64;d[q]=a[r*65+c];}}
void launch_diag_inv(float*out,float*dst,int*info,int n,int b,int k){auto fn=diag_inv64<DiagSolver256>;constexpr int bytes=DiagSolver256::shared_memory_size+2*4096*sizeof(float);static bool ok=[](auto f){cudaFuncSetAttribute(f,cudaFuncAttributeMaxDynamicSharedMemorySize,bytes);return true;}(fn);(void)ok;fn<<<b,256,bytes>>>(out,dst,info,n,b,k);}
void launch_diag(float*out,int*info,int n,int b,int k ){
 if(n==512&&b<=16){auto fn=diag64<DiagSolver128>;static bool ok=[](auto f){cudaFuncSetAttribute(f,cudaFuncAttributeMaxDynamicSharedMemorySize,DiagSolver128::shared_memory_size);return true;}(fn);(void)ok;fn<<<b,DiagSolver128::block_dim,DiagSolver128::shared_memory_size>>>(out,info,n,b,k);}
 else {auto fn=diag64<DiagSolver64>;static bool ok=[](auto f){cudaFuncSetAttribute(f,cudaFuncAttributeMaxDynamicSharedMemorySize,DiagSolver64::shared_memory_size);return true;}(fn);(void)ok;fn<<<b,DiagSolver64::block_dim,DiagSolver64::shared_memory_size>>>(out,info,n,b,k);}
}

// Invert four 16x16 diagonal leaves, then form the 32x32 and 64x64
// lower-left inverse quadrants: [A 0; C D]^-1 has -D^-1 C A^-1 below A^-1.
__global__ void inverse64(const float* out,float* dst,int n,int batches,int k){
 int bi=blockIdx.x;if(bi>=batches)return;__shared__ float l[4096],inv[4096],tmp[4096];
 const float* src=out+size_t(bi)*n*n+k*n+k;
 for(int q=threadIdx.x;q<2048;q+=blockDim.x){int r=q/32,c=2*(q%32);float2 v=((const float2*)(src+r*n))[c/2];((float2*)l)[q]=r>=c+1?v:make_float2(r>=c?v.x:0.0f,0.0f);((float2*)inv)[q]=make_float2(0,0);((float2*)tmp)[q]=make_float2(0,0);}__syncthreads();
 if(threadIdx.x<64){int leaf=threadIdx.x/16,j=threadIdx.x%16,c=leaf*16+j;float x[16];
   for(int i=0;i<16;i++){int r=leaf*16+i;float v=(i==j);if(j<=i){for(int q=0;q<i;q++)v-=l[r*64+leaf*16+q]*x[q];v/=l[r*64+r];}x[i]=v;}
   for(int i=0;i<16;i++)if(j<=i)inv[(leaf*16+i)*64+c]=x[i];
 }
 __syncthreads();
 for(int width=16;width<=32;width*=2){int pairs=64/(2*width),count=pairs*width*width;
   for(int q=threadIdx.x;q<count;q+=blockDim.x){int pair=q/(width*width),z=q%(width*width),i=z/width,j=z%width,base=pair*2*width;float v=0.0f;
     for(int p=0;p<width;p++)v+=l[(base+width+i)*64+base+p]*inv[(base+p)*64+base+j];tmp[(base+width+i)*64+base+j]=v;}
   __syncthreads();
   for(int q=threadIdx.x;q<count;q+=blockDim.x){int pair=q/(width*width),z=q%(width*width),i=z/width,j=z%width,base=pair*2*width;float v=0.0f;
     for(int p=0;p<width;p++)v+=inv[(base+width+i)*64+base+width+p]*tmp[(base+width+p)*64+base+j];inv[(base+width+i)*64+base+j]=-v;}
   __syncthreads();
 }
 float* d=dst+size_t(bi)*4096;for(int q=threadIdx.x;q<4096;q+=blockDim.x){int r=q/64,c=q%64;d[c*64+r]=inv[q];}
}
// Initialize huge factorizations in one pass: retain the lower input and make
// the unused upper triangle deterministically zero.
__global__ void lower_copy_kernel(const float* input,float* out,size_t total4,int n){
 size_t q=size_t(blockIdx.x)*blockDim.x+threadIdx.x;if(q>=total4)return;size_t z=(q*4)%(size_t(n)*n);int r=z/n,c=z-r*n;
 float4 v=*((const float4*)input+q);v.x=c<=r?v.x:0.0f;v.y=c+1<=r?v.y:0.0f;v.z=c+2<=r?v.z:0.0f;v.w=c+3<=r?v.w:0.0f;*((float4*)out+q)=v;
}
void lower_copy_launch(const float* input,float* out,int b,int n){size_t total4=size_t(b)*n*n/4;lower_copy_kernel<<<(total4+255)/256,256>>>(input,out,total4,n);}

// Panels are written transposed into the otherwise-unused upper triangle.
// The final pass moves them to the lower triangle and clears all upper entries.
__global__ void finish_blocked(float* out,int n,int tiles,int batches){int q=blockIdx.x%tiles,bi=blockIdx.x/tiles;int br=(int)((sqrtf(8.0f*q+1.0f)-1.0f)*0.5f);while((br+1)*(br+2)/2<=q)++br;while(br*(br+1)/2>q)--br;int bc=q-br*(br+1)/2;float*a=out+size_t(bi)*n*n;for(int ii=threadIdx.y;ii<64;ii+=16)for(int jj=threadIdx.x;jj<64;jj+=16){int r=br*64+ii,c=bc*64+jj;if(br==bc){if(ii>jj){a[c*n+r]=a[r*n+c];a[r*n+c]=0.0f;}}else a[r*n+c]=0.0f;}}
void native_lower_update(float*,float*,int,int,int);
void chol_blocked_launch(const float* input,float* out,float* inverses,int b,int n,cublasHandle_t h){
 int* info=(int*)inverses;cudaMemcpy(out,input,size_t(b)*n*n*sizeof(float),cudaMemcpyDeviceToDevice);
 cublasMath_t oldmath;cublasGetMathMode(h,&oldmath);cublasSetMathMode(h,(n==256&&b<64||n==1024&&b<4)?CUBLAS_PEDANTIC_MATH:CUBLAS_TF32_TENSOR_OP_MATH);
 const float one=1.0f,minus=-1.0f,zero=0.0f;long long matrix_stride=(long long)n*n,inv_stride=4096;
 for(int k=0;k<n;k+=64){int e=k+64;if(e==n){launch_diag(out,info,n,b,k);break;}int rem=n-e;launch_diag_inv(out,inverses,info,n,b,k);
   float* panel_t=out+k*n+e;float* schur_panel=out+e*n+k;
   cublasSgemmStridedBatched(h,CUBLAS_OP_T,CUBLAS_OP_T,rem,64,64,&one,schur_panel,n,matrix_stride,inverses,64,inv_stride,&zero,panel_t,n,matrix_stride,b);
   float* trail=out+e*n+e;
   if((n==1024&&b>8)||(n==512&&b==16))native_lower_update(panel_t,trail,rem,n,b);
   else cublasSgemmStridedBatched(h,CUBLAS_OP_N,CUBLAS_OP_T,rem,rem,64,&minus,panel_t,n,matrix_stride,panel_t,n,matrix_stride,&one,trail,n,matrix_stride,b);
 }
 cublasSetMathMode(h,oldmath);int nt=n/64,tiles=nt*(nt+1)/2;finish_blocked<<<b*tiles,dim3(16,16)>>>(out,n,tiles,b);
}
void chol_launch(const float*in,float*out,int b,int n ){if(n==32)go<32,2,64>(in,out,b);else if(n==64)go<64,2,96>(in,out,b);else if(n==128)go<128,1,128>(in,out,b);}


#define CDIV(a,b) (((a)+(b)-1)/(b))
template<unsigned NB,cusolverdx::arrangement AR,class T> __device__ T* ctile(T* a,unsigned ld,unsigned i,unsigned j){if constexpr(AR==cusolverdx::col_major)return a+i*NB+j*NB*ld;else return a+i*NB*ld+j*NB;}
template<unsigned NB,cusolverdx::arrangement AR,unsigned NT,class T> __device__ void cdiag_load(const T*a,int ld,T*s,int ls){for(int q=threadIdx.x;q<NB*NB;q+=NT){int i=q%NB,j=q/NB;bool take=AR==cusolverdx::col_major?i<=j:i>=j;if(take)s[i+j*ls]=__ldcg(a+i+j*ld);}__syncthreads();}
template<unsigned NB,cusolverdx::arrangement AR,unsigned NT,class T> __device__ void cdiag_store(const T*s,int ls,T*a,int ld){__syncthreads();for(int q=threadIdx.x;q<NB*NB;q+=NT){int i=q%NB,j=q/NB;bool take=AR==cusolverdx::col_major?i<=j:i>=j;if(take)__stcg(a+i+j*ld,s[i+j*ls]);}}
template<class P,class R,class GD,unsigned N,unsigned BPC,class T=typename P::a_data_type>
__global__ __cluster_dims__(BPC,1,1) __launch_bounds__(P::max_threads_per_block) void cga_chol(T*a,unsigned ld,int*info,unsigned batches){
 CUSOLVERDX_SKIP_IF_NOT_APPLICABLE_SM(P);namespace cg=cooperative_groups;auto cl=cg::this_cluster();unsigned bi=blockIdx.x/BPC;if(bi>=batches)return;a+=size_t(bi)*N*ld;
 constexpr unsigned NB=P::m_size,LS=P::lda,NT=P::max_threads_per_block,NTILES=N/NB,MAXP=NTILES-1,PPT=CDIV(MAXP,BPC);constexpr auto AR=P::a_arrangement;
 extern __shared__ __align__(16) cusolverdx::byte sm[];auto [diag,pan,work,si]=cusolverdx::shared_memory::slice<T,T,T,int>(sm,alignof(T),NB*LS,alignof(T),PPT*NB*LS,alignof(T),NB*LS,alignof(int),1);
 cusolverdx::byte* ds[BPC];for(int i=0;i<BPC;i++)ds[i]=cl.map_shared_rank(sm,i);auto doff=(cusolverdx::byte*)diag-sm,poff=(cusolverdx::byte*)pan-sm;constexpr unsigned tb=NB*LS*sizeof(T);
 auto dptr=[&](unsigned owner){return (T*)(ds[owner]+doff);};auto pptr=[&](unsigned k,unsigned j){unsigned x=j-k-1,owner=x%BPC,local=x/BPC;return (T*)(ds[owner]+poff+local*tb);};int ri=0;
 for(unsigned k=0;k<NTILES;k++){
  T*akk=ctile<NB,AR>(a,ld,k,k);if(k==0){if(cl.block_rank()==0){cdiag_load<NB,AR,NT>(akk,ld,diag,LS);P().execute(diag,LS,si);cdiag_store<NB,AR,NT>(diag,LS,akk,ld);if(threadIdx.x==0&&*si)ri=*si;}cl.sync();}
  int psz=NTILES-k-1;if(!psz)continue;if(cl.block_rank()!=0){T*rd=dptr(0);for(int x=threadIdx.x;x<NB*LS;x+=NT)diag[x]=rd[x];}
  for(int x=cl.block_rank();x<psz;x+=BPC){unsigned j=k+1+x;cusolverdx::copy_2d<NT,NB,NB,AR,1>(ctile<NB,AR>(a,ld,k,j),ld,pptr(k,j),LS);}__syncthreads();
  for(int x=cl.block_rank();x<psz;x+=BPC){unsigned j=k+1+x;R().execute(diag,LS,pptr(k,j),LS);}cl.sync();
  for(int x=0;x<psz;x++){unsigned j=k+1+x;if(x%BPC==cl.block_rank()){T*akj=pptr(k,j);for(int i=k+1;i<=int(j);i++){T*aij=ctile<NB,AR>(a,ld,i,j);if(i==int(j)){cdiag_load<NB,AR,NT>(aij,ld,work,LS);GD().execute(T(-1),akj,akj,T(1),work);if(x==0){__syncthreads();P().execute(work,LS,si);if(threadIdx.x==0&&ri==0&&*si)ri=*si+(k+1)*NB;cdiag_store<NB,AR,NT>(work,LS,aij,ld);for(int z=threadIdx.x;z<NB*LS;z+=NT)diag[z]=work[z];}else{cdiag_store<NB,AR,NT>(work,LS,aij,ld);__syncthreads();}}else{T*aki=pptr(k,i);cusolverdx::copy_2d<NT,NB,NB,AR,1>(aij,ld,work,LS);__syncthreads();GD().execute(T(-1),aki,akj,T(1),work);__syncthreads();cusolverdx::copy_2d<NT,NB,NB,AR,1>(work,LS,aij,ld);__syncthreads();}}cusolverdx::copy_2d<NT,NB,NB,AR,1>(akj,LS,ctile<NB,AR>(a,ld,k,j),ld);}}
  cl.sync();
 }
 (void)ri;
}template<class P,class R,class GD,unsigned N,class T=typename P::a_data_type>
__global__ __launch_bounds__(P::max_threads_per_block) void cga_chol_one(T*a,unsigned ld,int*info,unsigned batches){
 CUSOLVERDX_SKIP_IF_NOT_APPLICABLE_SM(P);namespace cg=cooperative_groups;unsigned bi=blockIdx.x;if(bi>=batches)return;a+=size_t(bi)*N*ld;
 constexpr unsigned BPC=1,NB=P::m_size,LS=P::lda,NT=P::max_threads_per_block,NTILES=N/NB,MAXP=NTILES-1,PPT=MAXP;constexpr auto AR=P::a_arrangement;
 extern __shared__ __align__(16) cusolverdx::byte sm[];auto [diag,pan,work,si]=cusolverdx::shared_memory::slice<T,T,T,int>(sm,alignof(T),NB*LS,alignof(T),PPT*NB*LS,alignof(T),NB*LS,alignof(int),1);
 cusolverdx::byte* ds[BPC];for(int i=0;i<BPC;i++)ds[i]=sm;auto doff=(cusolverdx::byte*)diag-sm,poff=(cusolverdx::byte*)pan-sm;constexpr unsigned tb=NB*LS*sizeof(T);
 auto dptr=[&](unsigned owner){return (T*)(ds[owner]+doff);};auto pptr=[&](unsigned k,unsigned j){unsigned x=j-k-1,owner=x%BPC,local=x/BPC;return (T*)(ds[owner]+poff+local*tb);};int ri=0;
 for(unsigned k=0;k<NTILES;k++){
  T*akk=ctile<NB,AR>(a,ld,k,k);if(k==0){if(0==0){for(int q=threadIdx.x;q<NB*NB/2;q+=NT)((float2*)diag)[q]=((const float2*)akk)[(q/(NB/2))*(ld/2)+q%(NB/2)];__syncthreads();P().execute(diag,LS,si);__syncthreads();for(int q=threadIdx.x;q<NB*NB/2;q+=NT)((float2*)akk)[(q/(NB/2))*(ld/2)+q%(NB/2)]=((float2*)diag)[q];if(threadIdx.x==0&&*si)ri=*si;}__syncthreads();}
  int psz=NTILES-k-1;if(!psz)continue;if(0!=0){T*rd=dptr(0);for(int x=threadIdx.x;x<NB*LS;x+=NT)diag[x]=rd[x];}
  for(int x=0;x<psz;x+=BPC){unsigned j=k+1+x;cusolverdx::copy_2d<NT,NB,NB,AR,1>(ctile<NB,AR>(a,ld,k,j),ld,pptr(k,j),LS);}__syncthreads();
  for(int x=0;x<psz;x+=BPC){unsigned j=k+1+x;R().execute(diag,LS,pptr(k,j),LS);}__syncthreads();
  for(int x=0;x<psz;x++){unsigned j=k+1+x;if(x%BPC==0){T*akj=pptr(k,j);for(int i=k+1;i<=int(j);i++){T*aij=ctile<NB,AR>(a,ld,i,j);if(i==int(j)){for(int q=threadIdx.x;q<NB*NB/2;q+=NT)((float2*)work)[q]=((const float2*)aij)[(q/(NB/2))*(ld/2)+q%(NB/2)];__syncthreads();GD().execute(T(-1),akj,akj,T(1),work);if(x==0){__syncthreads();P().execute(work,LS,si);if(threadIdx.x==0&&ri==0&&*si)ri=*si+(k+1)*NB;__syncthreads();for(int q=threadIdx.x;q<NB*NB/2;q+=NT)((float2*)aij)[(q/(NB/2))*(ld/2)+q%(NB/2)]=((float2*)work)[q];for(int z=threadIdx.x;z<NB*LS;z+=NT)diag[z]=work[z];}else{__syncthreads();for(int q=threadIdx.x;q<NB*NB/2;q+=NT)((float2*)aij)[(q/(NB/2))*(ld/2)+q%(NB/2)]=((float2*)work)[q];__syncthreads();}}else{T*aki=pptr(k,i);cusolverdx::copy_2d<NT,NB,NB,AR,1>(aij,ld,work,LS);__syncthreads();GD().execute(T(-1),aki,akj,T(1),work);__syncthreads();cusolverdx::copy_2d<NT,NB,NB,AR,1>(work,LS,aij,ld);__syncthreads();}}cusolverdx::copy_2d<NT,NB,NB,AR,1>(akj,LS,ctile<NB,AR>(a,ld,k,j),ld);}}
  __syncthreads();
 }
 __syncthreads();for(int q=threadIdx.x;q<NTILES*NB*NB;q+=NT){int t=q/(NB*NB),z=q%(NB*NB),i=z/NB,j=z%NB;if(i>j)a[(t*NB+i)*ld+t*NB+j]=0.0f;}__syncthreads();
 (void)ri;
}
constexpr unsigned CGNB=64,CGNT=256;constexpr auto CGAR=cusolverdx::row_major;
using CGP=decltype(cusolverdx::Function<cusolverdx::function::potrf>()+cusolverdx::FillMode<cusolverdx::fill_mode::upper>()+cusolverdx::Size<CGNB>()+cusolverdx::LeadingDimension<CGNB>()+cusolverdx::Precision<float>()+cusolverdx::Type<cusolverdx::type::real>()+cusolverdx::Arrangement<CGAR>()+cusolverdx::Block()+cusolverdx::BlockDim<CGNT>()+cusolverdx::SM<1000>());
using CGR=decltype(cusolverdx::Function<cusolverdx::function::trsm>()+cusolverdx::Size<CGNB,CGNB,CGNB>()+cusolverdx::LeadingDimension<CGNB,CGNB>()+cusolverdx::Precision<float>()+cusolverdx::Type<cusolverdx::type::real>()+cusolverdx::Side<cusolverdx::side::left>()+cusolverdx::Diag<cusolverdx::diag::non_unit>()+cusolverdx::TransposeMode<cusolverdx::transpose::transposed>()+cusolverdx::Arrangement<CGAR,CGAR>()+cusolverdx::FillMode<cusolverdx::fill_mode::upper>()+cusolverdx::Block()+cusolverdx::BlockDim<CGNT>()+cusolverdx::SM<1000>());
using CGGA=cublasdx::Arrangement<cublasdx::col_major,cublasdx::row_major,cublasdx::row_major>;using CGGD=decltype(cublasdx::Size<CGNB,CGNB,CGNB>()+CGGA()+cublasdx::Alignment<16,16,16>()+cublasdx::Precision<float>()+cublasdx::Type<cublasdx::type::real>()+cublasdx::Function<cublasdx::function::MM>()+cublasdx::LeadingDimension<CGNB,CGNB,CGNB>()+cublasdx::Block()+cublasdx::BlockDim<CGNT>()+cublasdx::SM<1000>());
template<unsigned NN> __global__ void cga_init_upper(const float*in,float*out,size_t total4){size_t q=size_t(blockIdx.x)*blockDim.x+threadIdx.x;if(q<total4){size_t z4=q%(size_t(NN)*NN/4);int r=z4/(NN/4),c=4*(z4%(NN/4));float4 v=((const float4*)in)[q];v.x=c>=r?v.x:0.0f;v.y=c+1>=r?v.y:0.0f;v.z=c+2>=r?v.z:0.0f;v.w=c+3>=r?v.w:0.0f;((float4*)out)[q]=v;}}
template<unsigned NN> __global__ void cga_finish_diag(float*a){int bi=blockIdx.x/(NN/64),t=blockIdx.x%(NN/64);a+=size_t(bi)*NN*NN;for(int i=threadIdx.y;i<64;i+=16)for(int j=threadIdx.x;j<64;j+=16)if(i>j)a[(t*64+i)*NN+t*64+j]=0.0f;}
template<unsigned NN,unsigned BPC> void cga_go(const float*in,float*out,int b){size_t total4=size_t(b)*NN*NN/4;cga_init_upper<NN><<<(total4+255)/256,256>>>(in,out,total4);auto fn=[](){if constexpr(BPC==1)return cga_chol_one<CGP,CGR,CGGD,NN,float>;else return cga_chol<CGP,CGR,CGGD,NN,BPC,float>;}();constexpr int sm=(1+CDIV(NN/CGNB-1,BPC)+1)*CGNB*CGNB*sizeof(float)+sizeof(int);static bool ok=[](auto f){cudaFuncSetAttribute(f,cudaFuncAttributeMaxDynamicSharedMemorySize,sm);return true;}(fn);(void)ok;fn<<<b*BPC,CGNT,sm>>>(out,NN,nullptr,b);}
void chol_cga_launch(const float*in,float*out,int n,int b){if(n==128)cga_go<128,1>(in,out,b);else cga_go<256,1>(in,out,b);}
using SyrkC=float;using SyrkLayoutA=cutlass::layout::RowMajor;using SyrkLayoutB=cutlass::layout::ColumnMajor;using SyrkLayoutC=cutlass::layout::RowMajor;using SyrkTile=cute::Shape<cute::_256,cute::_256,cute::_64>;using SyrkCluster=cute::Shape<cute::_2,cute::_1,cute::_1>;using SyrkProblem=cute::Shape<int,int,int,int>;
using SyrkEpilogue=typename cutlass::epilogue::collective::CollectiveBuilder<cutlass::arch::Sm100,cutlass::arch::OpClassTensorOp,SyrkTile,SyrkCluster,cutlass::epilogue::collective::EpilogueTileAuto,float,float,SyrkC,SyrkLayoutC,4,SyrkC,SyrkLayoutC,4,cutlass::epilogue::collective::EpilogueScheduleAuto>::CollectiveOp;
template<class E>using SyrkMainloopT=typename cutlass::gemm::collective::CollectiveBuilder<cutlass::arch::Sm100,cutlass::arch::OpClassTensorOp,E,SyrkLayoutA,8,E,SyrkLayoutB,8,float,SyrkTile,SyrkCluster,cutlass::gemm::collective::StageCountAutoCarveout<sizeof(SyrkEpilogue::SharedStorage)>,cutlass::gemm::KernelTmaWarpSpecialized2SmSm100>::CollectiveOp;
template<class E>using SyrkKernelT=cutlass::gemm::kernel::GemmUniversal<SyrkProblem,SyrkMainloopT<E>,SyrkEpilogue>;
template<class E>using SyrkGemmT=cutlass::gemm::device::GemmUniversalAdapter<SyrkKernelT<E>>;
using NativeLayoutA=cutlass::layout::ColumnMajor; using NativeLayoutB=cutlass::layout::RowMajor;
using NativeTile=cute::Shape<cute::_64,cute::_128,cute::_64>; using NativeCluster=cute::Shape<cute::_1,cute::_1,cute::_1>;
using NativeEpilogue=typename cutlass::epilogue::collective::CollectiveBuilder<cutlass::arch::Sm100,cutlass::arch::OpClassTensorOp,NativeTile,NativeCluster,cutlass::epilogue::collective::EpilogueTileAuto,float,float,float,SyrkLayoutC,4,float,SyrkLayoutC,4,cutlass::epilogue::collective::EpilogueScheduleAuto>::CollectiveOp;
using NativeMainloop=typename cutlass::gemm::collective::CollectiveBuilder<cutlass::arch::Sm100,cutlass::arch::OpClassTensorOp,float,NativeLayoutA,4,float,NativeLayoutB,4,float,NativeTile,NativeCluster,cutlass::gemm::collective::StageCount<3>,cutlass::gemm::collective::KernelScheduleAuto>::CollectiveOp;
using NativeKernel=cutlass::gemm::kernel::GemmUniversal<SyrkProblem,NativeMainloop,NativeEpilogue>; using NativeGemm=cutlass::gemm::device::GemmUniversalAdapter<NativeKernel>;
void native_lower_update(float* panel,float* trail,int rem,int ld,int batches){
 using StrideA=typename NativeKernel::StrideA;using StrideB=typename NativeKernel::StrideB;using StrideC=typename NativeKernel::StrideC;long long matrix_stride=(long long)ld*ld;
 auto launch=[&](int r,int m){int nn=r+m;auto shape=cute::make_shape(m,nn,64,batches);
   StrideA sa{cute::Int<1>{},int64_t(ld),int64_t(matrix_stride)};StrideB sb{cute::Int<1>{},int64_t(ld),int64_t(matrix_stride)};StrideC sc{int64_t(ld),cute::Int<1>{},int64_t(matrix_stride)};
   typename NativeGemm::Arguments args{cutlass::gemm::GemmUniversalMode::kGemm,shape,{panel+r,sa,panel,sb},{{-1.0f,1.0f},trail+r*ld,sc,trail+r*ld,sc}};
   NativeGemm op;if(op.can_implement(args)!=cutlass::Status::kSuccess)return false;if(op.initialize(args)!=cutlass::Status::kSuccess)return false;return op.run()==cutlass::Status::kSuccess;
 };
 int band=batches==16?512:256,r=0;for(;r+band<=rem;r+=band)if(!launch(r,band))return;if(r<rem)launch(r,rem-r);
}
__global__ void clear_syrk_upper(float* trail,size_t total,int rem,int ld,int band){
 size_t q=size_t(blockIdx.x)*blockDim.x+threadIdx.x;if(q>=total)return;int row=q/band,col=(row/band)*band+q%band;if(col<rem&&col>row)trail[size_t(row)*ld+col]=0.0f;
}
template<class SyrkA> void lower_syrk_typed(const void* panel_ptr,float* trail,void* workspace,size_t workspace_size,int rem,int k,int ld,int band){
 const SyrkA* panel=(const SyrkA*)panel_ptr;using Kernel=SyrkKernelT<SyrkA>;using Gemm=SyrkGemmT<SyrkA>;using StrideA=typename Kernel::StrideA;using StrideB=typename Kernel::StrideB;using StrideC=typename Kernel::StrideC;
 for(int r=0;r<rem;r+=band){int m=min(band,rem-r),nn=r+m;auto shape=cute::make_shape(m,nn,k,1);
   StrideA sa{k,cute::Int<1>{},int64_t(m)*k};
   StrideB sb{int64_t(k),cute::Int<1>{},int64_t(k)*nn};
   StrideC sc{ld,cute::Int<1>{},int64_t(ld)*rem};
   typename Gemm::Arguments args{cutlass::gemm::GemmUniversalMode::kGemm,shape,{(SyrkA*)panel+r*k,sa,(SyrkA*)panel,sb},{{-1.0f,1.0f},trail+r*ld,sc,trail+r*ld,sc}};
   size_t need=Gemm::get_workspace_size(args);if(need>workspace_size)return;Gemm op;if(op.can_implement(args)!=cutlass::Status::kSuccess)return;if(op.initialize(args,workspace)!=cutlass::Status::kSuccess)return;if(op.run()!=cutlass::Status::kSuccess)return;
 }
 size_t total=size_t(rem)*band;clear_syrk_upper<<<(total+255)/256,256>>>(trail,total,rem,ld,band);
}
void lower_syrk_launch(const void*p,float*t,void*w,size_t z,int r,int k,int l,int b){lower_syrk_typed<cutlass::half_t>(p,t,w,z,r,k,l,b);}
void lower_bsyrk_launch(const void*p,float*t,void*w,size_t z,int r,int k,int l,int b){lower_syrk_typed<cutlass::bfloat16_t>(p,t,w,z,r,k,l,b);}

// cuSOLVER writes a column-major lower factor. Since the symmetric row-major
// input has the same storage, POTRF can operate in place; this kernel converts
// its column-major result to row-major lower and clears the opposite triangle.
__global__ void vendor_finish(float* out,int n,int pairs){
 int q=blockIdx.x%pairs,bi=blockIdx.x/pairs;
 int tr=(int)((sqrtf(8.0f*q+1.0f)-1.0f)*0.5f);
 while((tr+1)*(tr+2)/2<=q)++tr;while(tr*(tr+1)/2>q)--tr;
 int tc=q-tr*(tr+1)/2,i=threadIdx.y,j=threadIdx.x;
 int r=tr*16+i,c=tc*16+j;if(r>=n||c>=n)return;
 float* a=out+size_t(bi)*n*n;
 if(tr==tc){if(i>j)a[r*n+c]=0.0f;}
 else a[r*n+c]=0.0f;
}
__global__ void vendor_init(const float* in,float* out,float** ptrs,int b,int n,size_t total4){
 size_t q=size_t(blockIdx.x)*blockDim.x+threadIdx.x;if(q<total4){size_t z4=q%(size_t(n)*n/4),bi=q/(size_t(n)*n/4);int r=z4/(n/4),c=4*(z4%(n/4));float4 v=((const float4*)(in+bi*size_t(n)*n))[z4];v.x=c>=r?v.x:0.0f;v.y=c+1>=r?v.y:0.0f;v.z=c+2>=r?v.z:0.0f;v.w=c+3>=r?v.w:0.0f;((float4*)(out+bi*size_t(n)*n))[z4]=v;}if(q<size_t(b)&&ptrs)ptrs[q]=out+q*size_t(n)*n;
}
void chol_vendor_launch(const float*in,float*out,float*work,int*info,int b,int n,cusolverDnHandle_t h){
 float** ptrs=n==2048&&b==8?(float**)work:nullptr;size_t total4=size_t(b)*n*n/4;vendor_init<<<(max(total4,size_t(b))+255)/256,256>>>(in,out,ptrs,b,n,total4);
 if(n==2048&&b==8){
   cusolverDnSpotrfBatched(h,CUBLAS_FILL_MODE_LOWER,n,ptrs,n,info,b);
 }else{
   int lwork=0;cusolverDnSpotrf_bufferSize(h,CUBLAS_FILL_MODE_LOWER,n,out,n,&lwork);
   for(int i=0;i<b;i++)cusolverDnSpotrf(h,CUBLAS_FILL_MODE_LOWER,n,out+size_t(i)*n*n,n,work,lwork,info+i);
 }

}
'''
    if not so.exists():
        import torch
        (root / "binding.cpp").write_text(cpp)
        (root / "kernel.cu").write_text(cu)
        ti = pathlib.Path(torch.__file__).parent / "include"
        py = sysconfig.get_paths()["include"]
        lib = pathlib.Path(torch.__file__).parent / "lib"
        defs = "-DTORCH_EXTENSION_NAME=chol_dx_sm100_tcTriInvV2 -DTORCH_API_INCLUDE_EXTENSION_H"
        inc = f"-I/opt/mathdx/include -I/opt/mathdx/external/cutlass/include -isystem {ti} -isystem {ti}/torch/csrc/api/include -isystem /usr/local/cuda/include -isystem {py}"
        un = "-U__CUDA_NO_HALF_OPERATORS__ -U__CUDA_NO_HALF_CONVERSIONS__ -U__CUDA_NO_BFLOAT16_CONVERSIONS__ -U__CUDA_NO_HALF2_OPERATORS__"
        cmds = [
            f"g++-13 -O3 -fPIC -std=c++17 {defs} {inc} -c {root/'binding.cpp'} -o {root/'binding.o'}",
            f"nvcc -O3 --expt-relaxed-constexpr -std=c++17 -arch=sm_100a -rdc=true -dlto -ccbin g++-13 -Xcompiler -fPIC {defs} {inc} {un} -c {root/'kernel.cu'} -o {root/'kernel.o'}",
            f"nvcc -dlink -arch=sm_100a -dlto -Xcompiler -fPIC {root/'kernel.o'} /opt/mathdx/lib/libcusolverdx.a -o {root/'dlink.o'}",
            f"g++-13 -shared {root/'binding.o'} {root/'kernel.o'} {root/'dlink.o'} /opt/mathdx/lib/libcusolverdx.a -L{lib} -lc10 -lc10_cuda -ltorch_cpu -ltorch_cuda -ltorch -ltorch_python -L/usr/local/cuda/lib64 -lcusolver -lcublas -lcudart -o {so}",
        ]
        for cmd in cmds: subprocess.check_call(cmd, shell=True)
    spec = importlib.util.spec_from_file_location("chol_dx_sm100_tcTriInvV2", so)
    module = importlib.util.module_from_spec(spec); spec.loader.exec_module(module)
    return module

_dx = _load_dx()

import torch
import triton
import triton.language as tl

from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = True
try:
    torch.set_float32_matmul_precision("medium")
except AttributeError:
    pass


@triton.jit
def _chol32(input_ptr, output_ptr, stride: tl.constexpr):
    m = tl.program_id(0)
    ids = tl.arange(0, 32)
    r, c = ids[:, None], ids[None, :]
    off = m * stride + r * 32 + c
    v = tl.where(r >= c, tl.load(input_ptr + off), 0.0)
    for k in range(32):
        row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
        d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
        d -= tl.sum(tl.where(ids < k, row * row, 0.0), axis=0)
        d = tl.sqrt(tl.maximum(d, 0.0))
        col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
        col = (col - tl.sum(tl.where(c < k, v * row[None, :], 0.0), axis=1)) / d
        v = tl.where((r == k) & (c == k), d, v)
        v = tl.where((r > k) & (c == k), col[:, None], v)
    tl.store(output_ptr + off, v)


@triton.jit
def _block_diag(inp, out, stride: tl.constexpr, N: tl.constexpr, K: tl.constexpr,
                PREC: tl.constexpr, STAGED: tl.constexpr, SPLIT16: tl.constexpr,
                INVERSE: tl.constexpr):
    m = tl.program_id(0); ids = tl.arange(0, 32); r, c = ids[:, None], ids[None, :]
    base = m * stride; off = base + (K + r) * N + K + c
    if INVERSE and K > 0:
        scratch = base + (K - 32 + r) * N + K - 32 + c
        tl.store(out + scratch, 0.0, mask=c > r)
    v = tl.load(inp + off)
    for p in range(0, K, 32):
        prev = tl.load(out + base + (K + r) * N + p + c)
        v -= tl.dot(prev, tl.trans(prev), input_precision=PREC)
    v = tl.where(r >= c, v, 0.0)
    if INVERSE:
        inv = tl.where(r == c, 1.0, 0.0)
    if SPLIT16:
        for k in range(0, 16):
            row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
            d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
            d -= tl.sum(tl.where(ids < k, row * row, 0.0), axis=0)
            d = tl.sqrt(tl.maximum(d, 0.0))
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            col = (col - tl.sum(tl.where(c < k, v * row[None, :], 0.0), axis=1)) / d
            v = tl.where((r == k) & (c == k), d, v)
            v = tl.where((r > k) & (c == k), col[:, None], v)
            if INVERSE:
                rhs = tl.where(ids == k, 1.0, 0.0)
                rhs -= tl.sum(tl.where(r < k, row[:, None] * inv, 0.0), axis=0)
                inv = tl.where(r == k, (rhs / d)[None, :], inv)
        panel = tl.where(c < 16, v, 0.0)
        update = tl.dot(panel, tl.trans(panel), input_precision=PREC)
        v = tl.where((r >= 16) & (c >= 16) & (r >= c), v - update, v)
        for k in range(16, 32):
            row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
            d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
            d -= tl.sum(tl.where((ids >= 16) & (ids < k), row * row, 0.0), axis=0)
            d = tl.sqrt(tl.maximum(d, 0.0))
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            prod = tl.where((c >= 16) & (c < k), v * row[None, :], 0.0)
            col = (col - tl.sum(prod, axis=1)) / d
            v = tl.where((r == k) & (c == k), d, v)
            v = tl.where((r > k) & (c == k), col[:, None], v)
            if INVERSE:
                rhs = tl.where(ids == k, 1.0, 0.0)
                rhs -= tl.sum(tl.where(r < k, row[:, None] * inv, 0.0), axis=0)
                inv = tl.where(r == k, (rhs / d)[None, :], inv)
    elif STAGED:
        # Factor the 32-wide diagonal tile recursively as 8 + 8 + 16.
        # Materializing each completed sub-panel as a dot update keeps the
        # following scalar recurrence short and reduces its live state.
        for k in range(0, 8):
            row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
            d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
            d -= tl.sum(tl.where(ids < k, row * row, 0.0), axis=0)
            d = tl.sqrt(tl.maximum(d, 0.0))
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            col = (col - tl.sum(tl.where(c < k, v * row[None, :], 0.0), axis=1)) / d
            v = tl.where((r == k) & (c == k), d, v)
            v = tl.where((r > k) & (c == k), col[:, None], v)
        panel = tl.where(c < 8, v, 0.0)
        update = tl.dot(panel, tl.trans(panel), input_precision=PREC)
        v = tl.where((r >= 8) & (c >= 8) & (r >= c), v - update, v)

        for k in range(8, 16):
            row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
            d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
            d -= tl.sum(tl.where((ids >= 8) & (ids < k), row * row, 0.0), axis=0)
            d = tl.sqrt(tl.maximum(d, 0.0))
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            prod = tl.where((c >= 8) & (c < k), v * row[None, :], 0.0)
            col = (col - tl.sum(prod, axis=1)) / d
            v = tl.where((r == k) & (c == k), d, v)
            v = tl.where((r > k) & (c == k), col[:, None], v)
        panel = tl.where((c >= 8) & (c < 16), v, 0.0)
        update = tl.dot(panel, tl.trans(panel), input_precision=PREC)
        v = tl.where((r >= 16) & (c >= 16) & (r >= c), v - update, v)

        for k in range(16, 32):
            row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
            d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
            d -= tl.sum(tl.where((ids >= 16) & (ids < k), row * row, 0.0), axis=0)
            d = tl.sqrt(tl.maximum(d, 0.0))
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            prod = tl.where((c >= 16) & (c < k), v * row[None, :], 0.0)
            col = (col - tl.sum(prod, axis=1)) / d
            v = tl.where((r == k) & (c == k), d, v)
            v = tl.where((r > k) & (c == k), col[:, None], v)
    else:
        for k in range(32):
            row = tl.sum(tl.where(r == k, v, 0.0), axis=0)
            d = tl.sum(tl.where(ids == k, row, 0.0), axis=0)
            d -= tl.sum(tl.where(ids < k, row * row, 0.0), axis=0)
            d = tl.sqrt(tl.maximum(d, 0.0))
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            col = (col - tl.sum(tl.where(c < k, v * row[None, :], 0.0), axis=1)) / d
            v = tl.where((r == k) & (c == k), d, v)
            v = tl.where((r > k) & (c == k), col[:, None], v)
    if INVERSE and K + 32 < N:
        # Reuse the otherwise-zero upper triangle to pass L^-T to the panel.
        # Build the inverse alongside factorization to avoid a second pass.
        stored = tl.where(c > r, tl.trans(inv), v)
        tl.store(out + off, stored)
    else:
        tl.store(out + off, v)


@triton.jit
def _block_panel(inp, out, stride: tl.constexpr, N: tl.constexpr, K: tl.constexpr,
                 PREC: tl.constexpr, B: tl.constexpr, STAGED: tl.constexpr,
                 GROUP: tl.constexpr, INVERSE: tl.constexpr):
    pid = tl.program_id(0)
    if GROUP:
        group = pid // (B * GROUP)
        first_br = group * GROUP
        group_size = tl.minimum(GROUP, (N - K - 32) // 32 - first_br)
        in_group = pid % (B * GROUP)
        m, br = in_group // group_size, first_br + in_group % group_size
    else:
        m, br = pid % B, pid // B
    ids = tl.arange(0, 32); r, c = ids[:, None], ids[None, :]
    rs = K + 32 + br * 32; base = m * stride
    off = base + (rs + r) * N + K + c
    v = tl.load(inp + off)
    for p in range(0, K, 32):
        left = tl.load(out + base + (rs + r) * N + p + c)
        pivot = tl.load(out + base + (K + r) * N + p + c)
        v -= tl.dot(left, tl.trans(pivot), input_precision=PREC)
    dblock = tl.load(out + base + (K + r) * N + K + c)
    if INVERSE:
        invt = tl.where(c > r, dblock, tl.where(c == r, 1.0 / dblock, 0.0))
        v = tl.dot(v, invt, input_precision="tf32")
    elif STAGED:
        # Complete 8, then 8, then 16 columns and use tensor-core updates
        # between stages. This bounds the scalar recurrence and live ranges.
        for k in range(8):
            prow = tl.sum(tl.where(r == k, dblock, 0.0), axis=0)
            piv = tl.sum(tl.where(ids == k, prow, 0.0), axis=0)
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            prod = tl.where(c < k, v * prow[None, :], 0.0)
            col = (col - tl.sum(prod, axis=1)) / piv
            v = tl.where(c == k, col[:, None], v)
        solved = tl.where(c < 8, v, 0.0)
        lower = tl.where((r >= 8) & (c < 8), dblock, 0.0)
        update = tl.dot(solved, tl.trans(lower), input_precision=PREC)
        v = tl.where(c >= 8, v - update, v)
        for k in range(8, 16):
            prow = tl.sum(tl.where(r == k, dblock, 0.0), axis=0)
            piv = tl.sum(tl.where(ids == k, prow, 0.0), axis=0)
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            prod = tl.where((c >= 8) & (c < k), v * prow[None, :], 0.0)
            col = (col - tl.sum(prod, axis=1)) / piv
            v = tl.where(c == k, col[:, None], v)
        solved = tl.where((c >= 8) & (c < 16), v, 0.0)
        lower = tl.where((r >= 16) & (c >= 8) & (c < 16), dblock, 0.0)
        update = tl.dot(solved, tl.trans(lower), input_precision=PREC)
        v = tl.where(c >= 16, v - update, v)
        for k in range(16, 32):
            prow = tl.sum(tl.where(r == k, dblock, 0.0), axis=0)
            piv = tl.sum(tl.where(ids == k, prow, 0.0), axis=0)
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            prod = tl.where((c >= 16) & (c < k), v * prow[None, :], 0.0)
            col = (col - tl.sum(prod, axis=1)) / piv
            v = tl.where(c == k, col[:, None], v)
    else:
        for k in range(32):
            prow = tl.sum(tl.where(r == k, dblock, 0.0), axis=0)
            piv = tl.sum(tl.where(ids == k, prow, 0.0), axis=0)
            col = tl.sum(tl.where(c == k, v, 0.0), axis=1)
            col = (col - tl.sum(tl.where(c < k, v * prow[None, :], 0.0), axis=1)) / piv
            v = tl.where(c == k, col[:, None], v)
    tl.store(out + off, v)
    tl.store(out + base + (K + r) * N + rs + c, 0.0)


def _small(data):
    b, n, _ = data.shape; out = torch.empty_like(data); stride = n * n
    precision = "tf32x3"
    panel_warps = 1 if n == 64 else 2
    diag_warps = 4 if n == 256 else panel_warps
    for k in range(0, n, 32):
        _block_diag[(b,)](data, out, stride, N=n, K=k, PREC=precision,
                          STAGED=n == 512, SPLIT16=False, INVERSE=False,
                          num_warps=diag_warps)
        rem = (n - k - 32) // 32
        if rem:
            _block_panel[(b * rem,)](data, out, stride, N=n, K=k, PREC=precision,
                                      B=b, STAGED=False, GROUP=0, INVERSE=False,
                                      num_warps=panel_warps)
    return out


def _small_throughput(data):
    b, n, _ = data.shape; out = torch.empty_like(data); stride = n * n
    for k in range(0, n, 32):
        _block_diag[(b,)](data, out, stride, N=n, K=k, PREC="tf32",
                          STAGED=True, SPLIT16=n == 512, INVERSE=n == 512,
                          num_warps=1 if n == 512 else 2)
        rem = (n - k - 32) // 32
        if rem:
            _block_panel[(b * rem,)](data, out, stride, N=n, K=k, PREC="tf32",
                                      B=b, STAGED=n == 512, GROUP=16 if n == 512 else 0,
                                      INVERSE=n == 512, num_warps=2 if n == 512 else 1)
    return out


@torch.compile(mode="reduce-overhead", fullgraph=False)
def _tri_inverse(diag, leaf=256):
    b, n, _ = diag.shape
    inv = torch.zeros_like(diag)
    eye = torch.eye(leaf, device=diag.device, dtype=diag.dtype)
    blocks = torch.stack([diag[:, k:k+leaf, k:k+leaf] for k in range(0, n, leaf)], dim=1)
    solved = torch.linalg.solve_triangular(blocks.flatten(0, 1), eye.expand(b*(n//leaf),-1,-1), upper=False)
    solved = solved.unflatten(0, (b, n//leaf))
    for j,k in enumerate(range(0,n,leaf)): inv[:,k:k+leaf,k:k+leaf]=solved[:,j]
    width=leaf
    while width<n:
        for k in range(0,n,2*width):
            m,e=min(k+width,n),min(k+2*width,n)
            if m==e: continue
            left=torch.bmm(diag[:,m:e,k:m].half(),inv[:,k:m,k:m].half(),out_dtype=torch.float32)
            off=torch.bmm(inv[:,m:e,m:e].half(),left.half(),out_dtype=torch.float32)
            inv[:,m:e,k:m]=-off
        width*=2
    return inv


@torch.compile(mode="reduce-overhead", fullgraph=False)
def _blocked(data, block, inverse=False):
    out = _dx.lower_copy(data) if data.shape[-1] % 4 == 0 else data.clone().tril_(); n = out.shape[-1]
    for k in range(0, n, block):
        e = min(k + block, n)
        diag = torch.linalg.cholesky_ex(out[:, k:e, k:e], check_errors=False).L
        out[:, k:e, k:e] = diag
        if e == n: break
        if inverse:
            inv = _tri_inverse(diag)
            if out.shape[0] == 1:
                panel = torch.mm(out[0,e:,k:e].bfloat16(), inv[0].T.bfloat16()).unsqueeze(0) if n==32768 else torch.mm(out[0,e:,k:e].half(), inv[0].T.half()).unsqueeze(0)
            else:
                panel = torch.bmm(out[:,e:,k:e].half(), inv.transpose(-1,-2).half(), out_dtype=torch.float32)
        else:
            panel = torch.linalg.solve_triangular(diag, out[:,e:,k:e].transpose(-1,-2), upper=False).transpose(-1,-2)
        out[:,e:,k:e] = panel
        if out.shape[0] == 1:
            (_dx.lower_bsyrk if n==32768 else _dx.lower_syrk)(panel[0], out[0,e:,e:], 3072)
        else:
            ph = panel.half()
            out[:,e:,e:].sub_(torch.bmm(ph, ph.transpose(-1,-2), out_dtype=torch.float32))
    return out if out.shape[0] == 1 else out.tril_()


def _single_matrix_calls(data):
    return torch.cat([
        torch.linalg.cholesky_ex(data[i:i+1], check_errors=False).L
        for i in range(data.shape[0])
    ], 0)


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if n in (32, 64): return _dx.chol_dx(data)
    if n == 128: return _dx.chol_cga(data)
    if n == 256: return _dx.chol_cga(data)
    if n == 512 and batch > 16: return _small_throughput(data)
    if n == 512: return _dx.chol_blocked(data)

    if n == 1024 and batch > 8: return _dx.chol_blocked(data)
    if n == 1024: return _dx.chol_blocked(data)
    if n in (2048, 4096): return _dx.chol_blocked(data)
    if n >= 16384: return _blocked(data, 4096, inverse=True)
    if n == 8192: return _blocked(data, 4096, inverse=True)
    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 579 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