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
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.
mma
wmma::fragment<wmma::accumulator,16,16,8,float> acc; wmma::fill_fragment(acc,0.f);shared-memory
extern __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