Skip to content
KernelIndex
Search⌘K

submission 804359

afrenkai · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_cool.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-804359?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
6.55ms
#224 of 515
2026-06-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e0f9d47cc00b2268ff7b9f0e977d2b2cbdb9e3c1b239284bc0be87e529fdfe58
license declaredunknown
license concludedunknown
authorsafrenkai
imported2026-08-26

Techniques

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

mmawmma::fragment<wmma::accumulator,16,16,8,float> acc; wmma::fill_fragment(acc,0.f);
shared-memoryextern __shared__ float W[];

Kernel source

submission_cool.py405 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
from torch.utils.cpp_extension import load_inline

torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")

_SINGLE_NMAX = 224
_PB = 64
_SMEM_BYTES = 200 * 1024
_TWO_LEVEL_NMIN = 896
_TWO_LEVEL_NMAX = 1536
_PB_OUT_TARGET = 128

_CUDA_SRC = r"""
#include <cuda_runtime.h>
#include <math.h>

extern __shared__ float W[];

template <int LANES, int BLK>
__global__ void qr_panel_kernel(const float* __restrict__ Ain, float* __restrict__ H, float* __restrict__ tau,
                                int n, int row0, int col0, int R, int C, int use_shfl) {
    const int b = blockIdx.x, t = threadIdx.x;
    const int rg = BLK / LANES, tx = t & (LANES - 1), ty = t / LANES;
    const size_t base = (size_t)b * n * n + (size_t)row0 * n + col0;
    float* taub = tau + (size_t)b * n + col0;

    for (int i = t; i < R; i += BLK)
        for (int c = 0; c < C; ++c) W[i * C + c] = Ain[base + (size_t)i * n + c];
    __syncthreads();

    __shared__ float sred[BLK], wpart[BLK], wsh[64];
    __shared__ float stau, sdenom;

    const int warp = t >> 5, lane = t & 31, NW = BLK / 32;
    for (int j = 0; j < C; ++j) {
        float partial = 0.f;
        for (int i = j + t; i < R; i += BLK) { float v = W[i * C + j]; partial += v * v; }
        if (use_shfl) {
            #pragma unroll
            for (int o = 16; o > 0; o >>= 1) partial += __shfl_down_sync(0xffffffffu, partial, o);
            if (lane == 0) sred[warp] = partial;
            __syncthreads();
        } else {
            sred[t] = partial; __syncthreads();
            for (int s = BLK / 2; s > 0; s >>= 1) { if (t < s) sred[t] += sred[t + s]; __syncthreads(); }
        }
        if (t == 0) {
            float xnorm2 = 0.f;
            if (use_shfl) { for (int w = 0; w < NW; ++w) xnorm2 += sred[w]; }
            else xnorm2 = sred[0];
            float x0 = W[j * C + j], xn = sqrtf(xnorm2), beta, tj, denom;
            if (xn == 0.f) { beta = x0; tj = 0.f; denom = 1.f; }
            else { beta = (x0 >= 0.f) ? -xn : xn; denom = x0 - beta; tj = (beta - x0) / beta; }
            stau = tj; sdenom = denom; taub[j] = tj; W[j * C + j] = beta;
        }
        __syncthreads();
        const float tj = stau, denom = sdenom;
        for (int i = j + 1 + t; i < R; i += BLK) W[i * C + j] /= denom;
        __syncthreads();

        if (tj != 0.f) {
            for (int cb = 0; cb < C; cb += LANES) {
                int c = cb + tx;
                float part = 0.f;
                if (c > j && c < C)
                    for (int i = j + 1 + ty; i < R; i += rg) part += W[i * C + j] * W[i * C + c];
                wpart[t] = part; __syncthreads();
                if (ty == 0 && c > j && c < C) {
                    float w = W[j * C + c];
                    for (int g = 0; g < rg; ++g) w += wpart[g * LANES + tx];
                    wsh[tx] = w * tj;
                }
                __syncthreads();
                if (c > j && c < C) {
                    float w = wsh[tx];
                    if (ty == 0) W[j * C + c] -= w;
                    for (int i = j + 1 + ty; i < R; i += rg) W[i * C + c] -= W[i * C + j] * w;
                }
                if (cb + LANES >= C) __syncthreads();
            }
        }
    }

    for (int i = t; i < R; i += BLK)
        for (int c = 0; c < C; ++c) H[base + (size_t)i * n + c] = W[i * C + c];
}

#define INST(L,B) cudaFuncSetAttribute(qr_panel_kernel<L,B>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_BYTES_PLACEHOLDER);
#define LAUNCH(B) do { \
    if (lanes == 64)      qr_panel_kernel<64,B><<<batch, B, shmem>>>(ap, hp, tp, n, row0, col0, R, C, use_shfl); \
    else if (lanes == 16) qr_panel_kernel<16,B><<<batch, B, shmem>>>(ap, hp, tp, n, row0, col0, R, C, use_shfl); \
    else                  qr_panel_kernel<32,B><<<batch, B, shmem>>>(ap, hp, tp, n, row0, col0, R, C, use_shfl); \
} while (0)

void qr_panel(torch::Tensor A, torch::Tensor H, torch::Tensor tau, int row0, int col0, int R, int C, int lanes, int threads, int use_shfl) {
    const int batch = H.size(0), n = H.size(1);
    const size_t shmem = (size_t)R * C * sizeof(float);
    static int configured = 0;
    if (!configured) {
        INST(16,256) INST(32,256) INST(64,256)
        INST(16,512) INST(32,512) INST(64,512)
        INST(16,1024) INST(32,1024) INST(64,1024)
        configured = 1;
    }
    const float* ap = A.data_ptr<float>();
    float* hp = H.data_ptr<float>();
    float* tp = tau.data_ptr<float>();
    if (threads == 1024)     LAUNCH(1024);
    else if (threads == 512) LAUNCH(512);
    else                     LAUNCH(256);
    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) throw std::runtime_error(cudaGetErrorString(err));
}
""".replace(
    "SMEM_BYTES_PLACEHOLDER", str(_SMEM_BYTES)
)

_CPP_SRC = "void qr_panel(torch::Tensor A, torch::Tensor H, torch::Tensor tau, int row0, int col0, int R, int C, int lanes, int threads, int use_shfl);"

_mod = load_inline(
    name="qr_blk_mod",
    cpp_sources=[_CPP_SRC],
    cuda_sources=[_CUDA_SRC],
    functions=["qr_panel"],
    verbose=True,
)


_SRC_ST = r"""
#include <cuda_runtime.h>
#include <mma.h>
#include <math.h>
using namespace nvcuda;

template<int N>
__device__ void gemm_AB(const float* A,const float* B,float* C,int M,int K,int warp,int nw,float alpha){
    const int TM=M/16,TN=N/16;
    for(int tile=warp;tile<TM*TN;tile+=nw){int tm=tile/TN,tn=tile%TN;
        wmma::fragment<wmma::accumulator,16,16,8,float> acc; wmma::fill_fragment(acc,0.f);
        for(int k0=0;k0<K;k0+=8){
            wmma::fragment<wmma::matrix_a,16,16,8,wmma::precision::tf32,wmma::row_major> ah,al;
            wmma::fragment<wmma::matrix_b,16,16,8,wmma::precision::tf32,wmma::row_major> bh,bl;
            wmma::load_matrix_sync(ah,A+(size_t)(tm*16)*K+k0,K);
            wmma::load_matrix_sync(bh,B+(size_t)k0*N+tn*16,N);
            for(int i=0;i<ah.num_elements;i++){float o=ah.x[i],h=wmma::__float_to_tf32(o);al.x[i]=wmma::__float_to_tf32(o-h);ah.x[i]=h;}
            for(int i=0;i<bh.num_elements;i++){float o=bh.x[i],h=wmma::__float_to_tf32(o);bl.x[i]=wmma::__float_to_tf32(o-h);bh.x[i]=h;}
            wmma::mma_sync(acc,ah,bh,acc);wmma::mma_sync(acc,ah,bl,acc);wmma::mma_sync(acc,al,bh,acc);
        }
        float* Cp=C+(size_t)(tm*16)*N+tn*16;
        if(alpha<0.f){wmma::fragment<wmma::accumulator,16,16,8,float> cf;wmma::load_matrix_sync(cf,Cp,N,wmma::mem_row_major);
            for(int i=0;i<cf.num_elements;i++)cf.x[i]-=acc.x[i];wmma::store_matrix_sync(Cp,cf,N,wmma::mem_row_major);}
        else wmma::store_matrix_sync(Cp,acc,N,wmma::mem_row_major);
    }
}
template<int N>
__device__ void gemm_AtB(const float* A,const float* B,float* C,int M,int K,int warp,int nw){
    const int TM=M/16,TN=N/16;
    for(int tile=warp;tile<TM*TN;tile+=nw){int tm=tile/TN,tn=tile%TN;
        wmma::fragment<wmma::accumulator,16,16,8,float> acc; wmma::fill_fragment(acc,0.f);
        for(int k0=0;k0<K;k0+=8){
            wmma::fragment<wmma::matrix_a,16,16,8,wmma::precision::tf32,wmma::col_major> ah,al;
            wmma::fragment<wmma::matrix_b,16,16,8,wmma::precision::tf32,wmma::row_major> bh,bl;
            wmma::load_matrix_sync(ah,A+(size_t)k0*M+tm*16,M);
            wmma::load_matrix_sync(bh,B+(size_t)k0*N+tn*16,N);
            for(int i=0;i<ah.num_elements;i++){float o=ah.x[i],h=wmma::__float_to_tf32(o);al.x[i]=wmma::__float_to_tf32(o-h);ah.x[i]=h;}
            for(int i=0;i<bh.num_elements;i++){float o=bh.x[i],h=wmma::__float_to_tf32(o);bl.x[i]=wmma::__float_to_tf32(o-h);bh.x[i]=h;}
            wmma::mma_sync(acc,ah,bh,acc);wmma::mma_sync(acc,ah,bl,acc);wmma::mma_sync(acc,al,bh,acc);
        }
        wmma::store_matrix_sync(C+(size_t)(tm*16)*N+tn*16,acc,N,wmma::mem_row_major);
    }
}
__device__ void chol_up(float* G,int pe,float shift,int t,int nt){
    if(t<pe) G[t*pe+t]+=shift; __syncthreads();
    for(int k=0;k<pe;++k){
        float d=sqrtf(fmaxf(G[k*pe+k],1e-30f));
        for(int j=k+t;j<pe;j+=nt){ if(j==k)G[k*pe+k]=d; else if(j>k)G[k*pe+j]/=d; }
        __syncthreads();
        for(int idx=t;idx<(pe-k-1)*(pe-k-1);idx+=nt){int i=k+1+idx/(pe-k-1),j=k+1+idx%(pe-k-1);
            if(j>=i)G[i*pe+j]-=G[k*pe+i]*G[k*pe+j];}
        __syncthreads();
    }
    for(int idx=t;idx<pe*pe;idx+=nt){int i=idx/pe,j=idx%pe; if(j<i)G[idx]=0.f;} __syncthreads();
}
__device__ void trisolve_R(float* B,const float* R,int m,int pe,int t,int nt){
    for(int j=0;j<pe;++j){float rjj=R[j*pe+j];
        for(int i=t;i<m;i+=nt){float s=B[i*pe+j];for(int l=0;l<j;++l)s-=B[i*pe+l]*R[l*pe+j];B[i*pe+j]=s/rjj;}
        __syncthreads();
    }
}
__device__ void lu_np(float* M,int pe,int t,int nt){
    for(int k=0;k<pe;++k){float d=M[k*pe+k];
        for(int i=k+1+t;i<pe;i+=nt)M[i*pe+k]/=d; __syncthreads();
        for(int idx=t;idx<(pe-k-1)*(pe-k-1);idx+=nt){int i=k+1+idx/(pe-k-1),j=k+1+idx%(pe-k-1);
            M[i*pe+j]-=M[i*pe+k]*M[k*pe+j];}
        __syncthreads();
    }
}

extern __shared__ float SM[];
template<int BLK,int PB,int TILE>
__global__ void qr_stiefel(float* __restrict__ A,float* __restrict__ tau,int n){
    const int bb=blockIdx.x,t=threadIdx.x,warp=t/32,nw=BLK/32;
    float* base=A+(size_t)bb*n*n; float* taub=tau+(size_t)bb*n;
    for(int k=0;k<n;k+=PB){
        const int pe=(PB<n-k)?PB:(n-k); const int R=n-k;
        float* Pp=SM; float* Gm=Pp+(size_t)R*pe; float* R1=Gm+(size_t)pe*pe;
        float* R2=R1+(size_t)pe*pe; float* Ct=R2+(size_t)pe*pe; float* Wm=Ct+(size_t)R*TILE;
        for(int i=t;i<R*pe;i+=BLK) Pp[i]=base[(size_t)(k+i/pe)*n+(k+i%pe)];
        __syncthreads();
        if (R < 2*pe) {
            __shared__ float sred[BLK]; __shared__ float bstau,bsden;
            for(int j=0;j<pe;++j){
                float part=0.f; for(int i=j+t;i<R;i+=BLK){float v=Pp[i*pe+j];part+=v*v;}
                sred[t]=part; __syncthreads();
                for(int s=BLK/2;s>0;s>>=1){if(t<s)sred[t]+=sred[t+s];__syncthreads();}
                if(t==0){float x0=Pp[j*pe+j],xn=sqrtf(sred[0]),beta,tj,den;
                    if(xn==0.f){beta=x0;tj=0.f;den=1.f;}else{beta=(x0>=0.f)?-xn:xn;den=x0-beta;tj=(beta-x0)/beta;}
                    bstau=tj;bsden=den;taub[k+j]=tj;Pp[j*pe+j]=beta;}
                __syncthreads();
                float tj=bstau,den=bsden;
                for(int i=j+1+t;i<R;i+=BLK)Pp[i*pe+j]/=den; __syncthreads();
                if(tj!=0.f){for(int c=j+1+t;c<pe;c+=BLK){float w=Pp[j*pe+c];
                    for(int i=j+1;i<R;++i)w+=Pp[i*pe+j]*Pp[i*pe+c]; w*=tj; Pp[j*pe+c]-=w;
                    for(int i=j+1;i<R;++i)Pp[i*pe+c]-=Pp[i*pe+j]*w;}}
                __syncthreads();
            }
            for(int i=t;i<R*pe;i+=BLK) base[(size_t)(k+i/pe)*n+(k+i%pe)]=Pp[i];
            __syncthreads();
            const int ncolb=n-k-pe; if(ncolb<=0) continue;
            for(int idx=t;idx<pe*pe;idx+=BLK){int i=idx/pe,j=idx%pe; if(i<j)Pp[i*pe+j]=0.f; else if(i==j)Pp[i*pe+j]=1.0f;}
            __syncthreads();
            gemm_AtB<PB>(Pp,Pp,R1,pe,R,warp,nw); __syncthreads();
            for(int i=0;i<pe;++i){float ti=taub[k+i];
                for(int c=t;c<pe;c+=BLK){float acc=(i==c)?1.0f:0.0f;for(int j=0;j<i;++j)acc-=R1[j*pe+i]*R2[j*pe+c];R2[i*pe+c]=acc*ti;}
                __syncthreads();}
            for(int jc=0;jc<ncolb;jc+=TILE){const int w=(TILE<ncolb-jc)?TILE:(ncolb-jc);const int col0=k+pe+jc;
                for(int i=t;i<R*TILE;i+=BLK){int ii=i/TILE,cc=i%TILE;Ct[i]=(cc<w)?base[(size_t)(k+ii)*n+(col0+cc)]:0.f;}
                __syncthreads();
                gemm_AtB<TILE>(Pp,Ct,Wm,pe,R,warp,nw);__syncthreads();
                gemm_AB<TILE>(R2,Wm,Gm,pe,pe,warp,nw,1.0f);__syncthreads();
                gemm_AB<TILE>(Pp,Gm,Ct,R,pe,warp,nw,-1.0f);__syncthreads();
                for(int i=t;i<R*TILE;i+=BLK){int ii=i/TILE,cc=i%TILE;if(cc<w)base[(size_t)(k+ii)*n+(col0+cc)]=Ct[i];}
                __syncthreads();}
            continue;
        }
        gemm_AtB<PB>(Pp,Pp,Gm,pe,R,warp,nw); __syncthreads();
        __shared__ float smax; if(t==0){float mx=0.f;for(int j=0;j<pe;++j)mx=fmaxf(mx,Gm[j*pe+j]);smax=mx;} __syncthreads();
        float shift=11.0f*((float)R*pe + (float)pe*(pe+1))*1.1920929e-7f*smax;  // Fukaya et al. shifted-CholeskyQR shift (pass 1 only)
        for(int i=t;i<pe*pe;i+=BLK) R1[i]=Gm[i]; __syncthreads();
        chol_up(R1,pe,shift,t,BLK);                 // R1 = chol(G+shift)
        trisolve_R(Pp,R1,R,pe,t,BLK);               // Pp <- Q = Pp R1^-1
        gemm_AtB<PB>(Pp,Pp,Gm,pe,R,warp,nw); __syncthreads();   // G2 = Q^T Q
        for(int i=t;i<pe*pe;i+=BLK) R2[i]=Gm[i]; __syncthreads();
        chol_up(R2,pe,0.0f,t,BLK);   // pass 2 UNSHIFTED (refine); pivot floor handles zero cols
        trisolve_R(Pp,R2,R,pe,t,BLK);               // Pp <- Q (pass2)
        gemm_AB<PB>(R2,R1,Gm,pe,pe,warp,nw,1.0f); __syncthreads();   // Gm = R2 R1
        gemm_AtB<PB>(Pp,Pp,R1,pe,R,warp,nw); __syncthreads();   // G3 = Q^T Q -> R1
        chol_up(R1,pe,0.0f,t,BLK);                 // pass 3 UNSHIFTED (refine); R1 = R3
        trisolve_R(Pp,R1,R,pe,t,BLK);                // Pp <- Q (orthonormal, pass3)
        gemm_AB<PB>(R1,Gm,R2,pe,pe,warp,nw,1.0f); __syncthreads();   // R2 = R3 (R2 R1) = R
        for(int i=t;i<pe*pe;i+=BLK) Gm[i]=R2[i]; __syncthreads();    // Gm = R (final)
        __shared__ float sgn[PB];
        if(t<pe){float q=Pp[t*pe+t];sgn[t]=(q>=0.f)?-1.f:1.f;} __syncthreads();
        for(int i=t;i<R*pe;i+=BLK){int ii=i/pe,jj=i%pe; float m=-Pp[i]*sgn[jj]; if(ii==jj)m+=1.0f; Pp[i]=m;}
        __syncthreads();
        lu_np(Pp,pe,t,BLK);                         
        for(int j=0;j<pe;++j){float ujj=Pp[j*pe+j];
            for(int i=pe+t;i<R;i+=BLK){float s=Pp[i*pe+j];for(int l=0;l<j;++l)s-=Pp[i*pe+l]*Pp[l*pe+j];Pp[i*pe+j]=s/ujj;}
            __syncthreads();
        }
        for(int i=t;i<R*pe;i+=BLK){int ii=i/pe,jj=i%pe; float val;
            if(ii<jj) val = sgn[ii]*Gm[ii*pe+jj];      // R_signed upper = s_i * R[i,j] (ROW sign)
            else if(ii==jj) val = sgn[jj]*Gm[jj*pe+jj];
            else val = Pp[i];                          // strict-lower reflector
            base[(size_t)(k+ii)*n+(k+jj)] = val;
        }
        for(int idx=t;idx<pe*pe;idx+=BLK){int i=idx/pe,j=idx%pe; if(i<j)Pp[i*pe+j]=0.f; else if(i==j)Pp[i*pe+j]=1.0f;}
        __syncthreads();
        const int ncol=n-k-pe; if(ncol<=0) continue;
        gemm_AtB<PB>(Pp,Pp,R1,pe,R,warp,nw); __syncthreads();
        for(int i=0;i<pe;++i){float ti=taub[k+i];
            for(int c=t;c<pe;c+=BLK){float acc=(i==c)?1.0f:0.0f;for(int j=0;j<i;++j)acc-=R1[j*pe+i]*R2[j*pe+c];R2[i*pe+c]=acc*ti;}
            __syncthreads();
        }
        for(int jc=0;jc<ncol;jc+=TILE){const int w=(TILE<ncol-jc)?TILE:(ncol-jc);const int col0=k+pe+jc;
            for(int i=t;i<R*TILE;i+=BLK){int ii=i/TILE,cc=i%TILE; Ct[i]=(cc<w)?base[(size_t)(k+ii)*n+(col0+cc)]:0.f;}
            __syncthreads();
            gemm_AtB<TILE>(Pp,Ct,Wm,pe,R,warp,nw); __syncthreads();   // W=V^T C (pe x TILE)
            gemm_AB<TILE>(R2,Wm,Gm,pe,pe,warp,nw,1.0f); __syncthreads(); // Z=T W -> Gm
            gemm_AB<TILE>(Pp,Gm,Ct,R,pe,warp,nw,-1.0f); __syncthreads(); // C -= V Z
            for(int i=t;i<R*TILE;i+=BLK){int ii=i/TILE,cc=i%TILE; if(cc<w)base[(size_t)(k+ii)*n+(col0+cc)]=Ct[i];}
            __syncthreads();
        }
    }
}
#define CFG(B,PB,TILE) cudaFuncSetAttribute(qr_stiefel<B,PB,TILE>,cudaFuncAttributeMaxDynamicSharedMemorySize,SMEM_PLACEHOLDER);
void qr_stiefel_launch(torch::Tensor A,torch::Tensor tau,int threads){
    const int batch=A.size(0),n=A.size(1); static int cfg=0; if(!cfg){CFG(256,48,16) CFG(512,48,16) cfg=1;}
    size_t sh=((size_t)n*48 + 3*48*48 + (size_t)n*16 + 48*16)*4;
    float* ap=A.data_ptr<float>();float* tp=tau.data_ptr<float>();
    if(threads==256) qr_stiefel<256,48,16><<<batch,256,sh>>>(ap,tp,n);
    else             qr_stiefel<512,48,16><<<batch,512,sh>>>(ap,tp,n);
    cudaError_t e=cudaGetLastError(); if(e)throw std::runtime_error(cudaGetErrorString(e));
}
""".replace(
    "SMEM_PLACEHOLDER", str(200 * 1024)
)
_CPP_ST = "void qr_stiefel_launch(torch::Tensor A, torch::Tensor tau, int threads);"
_mod_st = None
if False:
    _mod_st = load_inline(
        name="qr_stiefel_mod",
        cpp_sources=[_CPP_ST],
        cuda_sources=[_SRC_ST],
        functions=["qr_stiefel_launch"],
        extra_cuda_cflags=["-arch=sm_100a"],
        verbose=True,
    )


def _panel_block_size(n: int) -> int:
    return max(8, min(_PB, _SMEM_BYTES // (4 * n)))


def _lanes(r: int) -> int:
    if r <= 256:
        return 64
    if r <= 1024:
        return 32
    return 16


def _threads(n: int) -> int:
    if n <= 64:
        return 256
    return 512


def _wy(h, tau, k, pe, c0, c1, inverse, fp32):
    if c1 <= c0:
        return
    dev, dt, batch = h.device, h.dtype, h.shape[0]
    vr = h[:, k:, k : k + pe].clone()
    vr[:, :pe, :].tril_(-1)
    idx = torch.arange(pe, device=dev)
    vr[:, idx, idx] = 1.0
    taus = tau[:, k : k + pe]
    inv_tau = torch.where(taus != 0, 1.0 / taus, taus.new_full((), 1e30))
    tinv = vr.transpose(1, 2) @ vr
    tinv.triu_(1)
    tinv.diagonal(dim1=-2, dim2=-1).copy_(inv_tau)
    c = h[:, k:, c0:c1]
    y = vr.transpose(1, 2) @ c
    if inverse:
        eye = torch.eye(pe, device=dev, dtype=dt).expand(batch, pe, pe)
        minv = torch.linalg.solve_triangular(tinv.transpose(1, 2), eye, upper=False)
        if fp32:
            prev = torch.backends.cuda.matmul.allow_tf32
            torch.backends.cuda.matmul.allow_tf32 = False
            z = minv @ y
            torch.backends.cuda.matmul.allow_tf32 = prev
        else:
            z = minv @ y
    else:
        z = torch.linalg.solve_triangular(tinv.transpose(1, 2), y, upper=False)
    c.baddbmm_(vr, z, beta=1.0, alpha=-1.0)


def custom_kernel(data: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    a = data
    batch, n, t = a.shape
    tau = torch.empty(batch, n, device=a.device, dtype=a.dtype)
    thr = _threads(n)

    if n <= _SINGLE_NMAX:
        h = torch.empty_like(a)
        _mod.qr_panel(a, h, tau, 0, 0, n, n, _lanes(n), thr, 1)
        return h, tau
    h = a.clone()
    pb = _panel_block_size(n)

    if _TWO_LEVEL_NMIN <= n <= _TWO_LEVEL_NMAX:
        pb_out = min(max(pb, (_PB_OUT_TARGET // pb) * pb), n)
        for kb in range(0, n, pb_out):
            peb = min(pb_out, n - kb)
            bend = kb + peb
            for ki in range(kb, bend, pb):
                pei = min(pb, bend - ki)
                _mod.qr_panel(h, h, tau, ki, ki, n - ki, pei, _lanes(n - ki), thr, 0)
                _wy(h, tau, ki, pei, ki + pei, bend, inverse=False, fp32=False)
            _wy(h, tau, kb, peb, bend, n, inverse=True, fp32=False)
        return h, tau

    inv = 448 <= n <= 768
    for k in range(0, n, pb):
        pe = min(pb, n - k)
        r = n - k
        _mod.qr_panel(h, h, tau, k, k, r, pe, _lanes(r), thr, 0)
        _wy(h, tau, k, pe, k + pe, n, inverse=inv, fp32=True)
    return h, tau
scrolls · 405 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