submission 844577
sergeylitvinov_ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 566 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-844577?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:8fd93a3988ec5c8abdc18b2bd5be8ed964588e5494c5b3bc02cce7bee65ccede
license declaredunknown
license concludedunknown
authorssergeylitvinov_
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
nvcuda::wmma::fragment<nvcuda::wmma::accumulator,16,16,16,float> acc[2]; for(int k=0;k<nk;k++) nvcuda::wmma::fill_fragment(acc[k],0.f);shared-memory
extern __shared__ float s[];Kernel source
submission.py566 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r"""
/* qrlib.cuh -- Recursive Householder QR (compact-WY) library, batched, GPU.
*
* Paper: Leng et al., "High Performance Householder QR ... Using Tensor Cores", IEEE TPDS 2025.
* Algorithm 1 (recursive QR) + eq (4) (recursive WY combination). Compact-WY form A = Q R,
* Q = I - V T V^T. The recursion turns the trailing-apply / WY-combine into SQUARE GEMMs (good
* for tensor cores). Optional Ootomo fp16 split (3 fp16 GEMMs = ~fp32 accuracy) for those GEMMs.
*
* Public entry: void qr(long Hp, long taup, int b, int n)
* Hp: device ptr to b*n*n fp32, ROW-MAJOR. Input matrix; overwritten with packed reflectors+R (geqrf layout).
* taup: device ptr to b*n fp32. Output Householder coefficients.
* Env: QR_FP16=1 enables Ootomo fp16 tensor-core GEMMs; QR_TF32=1 enables raw TF32 (coarser).
*
* Clients (test_qr.cu, bench_qr.cu) #include this header and add their own main().
*/
#ifndef QRLIB_CUH
#define QRLIB_CUH
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <mma.h>
#include <cublas_v2.h>
#include <cstdlib>
#include <cstdio>
#ifndef QR_NB
#define QR_NB 64 /* MAX base-case panel width (compile-time array size); actual width g_nb<=QR_NB
is chosen at runtime to fit the device shared memory. */
#endif
#define QR_PT 256 /* panel threads */
/* ===== criteria-based method registry (X-macro). Each row = a specialized blocked-QR method,
* selected by the FIRST row whose N_MAX >= n. Columns:
* N_MAX : method applies to all n <= N_MAX (ascending)
* NB : panel block width = warp-per-column count = threads/32 (panel launches NB*32 threads)
* TR : register rows per lane in the panel (col[TR], TR*32 >= N_MAX)
* Register budget: the panel is __launch_bounds__(NB*32), so reg cap = 65536/(NB*32). Pick NB so
* col[TR] (+overhead) fits without spilling: NB=32 -> 64 regs (TR<=16 ok); NB=16 -> 128 regs (TR<=32 ok).
* Add a row to specialize a new size; n outside the table routes to cuSOLVER (torch.geqrf). */
#ifndef QR_N1024_NB
#define QR_N1024_NB 32 /* n=1024 panel width. B200/sm_120 measured: NB=32 (32 panels, 16B spill absorbed by fast
* L1/L2) BEATS NB=16 (64 panels, no spill): 8.0ms vs 9.1ms on eu-g7. A100 sm_80 is the
* OPPOSITE (slow mem -> spill catastrophic, NB16 2x faster) -> A100 is correctness-only. */
#endif
#define QR_SIZE_TABLE(X) \
/* N_MAX NB TR */ \
X( 512, 32, 16 ) \
X( 1024, QR_N1024_NB, 32 )
#define QR_MAX_N 1024 /* largest n the blocked path handles (= last QR_SIZE_TABLE row) */
namespace qrlib {
static inline void qck(cudaError_t e){ if(e){ printf("qrlib cuda: %s\n",cudaGetErrorString(e)); } }
/* ---- state (single-threaded host driver) ---- */
static cublasHandle_t g_H = 0;
static float *g_V=0, *g_T=0, *g_W=0; /* Vg b*(m x n), Tg b*(n x n), scratch b*(n x n) */
static __half *g_Ah=0,*g_Al=0,*g_Bh=0,*g_Bl=0;/* Ootomo fp16 split scratch, b*(n x n) each */
static __half *g_Af=0,*g_Vh=0,*g_Vl=0,*g_Th=0,*g_Wh=0,*g_Yh=0,*g_Yl=0; /* blocked-fp16: trailing, panel V_hi/V_lo/T, apply scratch (Y split) */
static float *g_Vf=0,*g_G=0; /* wpc path: fp32 clean V (for Gram) + Gram V^T V */
static __half *g_V3=0,*g_Y3=0; /* stacked apply: V3=[Vh|Vl|Vh] (m x 3w), Y3=[Yh;Yh;Yl] (3w x rest) */
static int g_M=0, g_N=0, g_B=0, g_fp16=0, g_nb=QR_NB, g_bnb=0;
static long g_sA=0, g_sV=0, g_sT=0, g_sW=0; /* per-matrix strides */
/* NOTE: all kernels/cuBLAS use the default execution context (leaderboard rule forbids others). */
static void ensure_state(int b,int n); /* fwd decl (defined with the drivers below) */
/* ================= kernels ================= */
/* Panel (one block per matrix): col-major A[bi][p:m, p:p+w] -> packed A; clean V; compact-WY T. */
__global__ void k_panel(float* A, int m, int p, int w, int lda, float* Vg, int ldv, float* Tg, int ldt,
long sA, long sV, long sT) {
extern __shared__ float s[];
const int mm=m-p, tid=threadIdx.x, bi=blockIdx.x;
float* Ab=A+bi*sA; float* Vb=Vg+bi*sV; float* Tb=Tg+bi*sT;
float* P=s; float* Tsh=P+(long)mm*w;
__shared__ float red[QR_PT];
for(int idx=tid; idx<mm*w; idx+=QR_PT){ int i=idx%mm,j=idx/mm; P[i+(long)j*mm]=Ab[(p+i)+(long)(p+j)*lda]; }
__syncthreads();
for(int j=0;j<w;++j){
float ss=0.f; for(int i=j+tid;i<mm;i+=QR_PT){ float v=P[i+(long)j*mm]; ss+=v*v; }
for(int o=16;o>0;o>>=1) ss+=__shfl_down_sync(0xffffffffu,ss,o);
if((tid&31)==0) red[tid>>5]=ss; __syncthreads();
float nrm2=0.f; for(int q=0;q<QR_PT/32;q++) nrm2+=red[q];
float alpha=P[j+(long)j*mm], nrm=sqrtf(nrm2), beta=(alpha>=0?-nrm:nrm);
float tau=(nrm==0.f)?0.f:(beta-alpha)/beta, scal=(nrm==0.f)?0.f:1.f/(alpha-beta);
for(int i=j+1+tid;i<mm;i+=QR_PT) P[i+(long)j*mm]*=scal;
__syncthreads();
if(tid==0){ P[j+(long)j*mm]=beta; Tsh[j+(long)j*w]=tau; }
__syncthreads();
for(int c=j+1;c<w;++c){
float dot=0.f; for(int i=j+tid;i<mm;i+=QR_PT){ float vi=(i==j)?1.f:P[i+(long)j*mm]; dot+=vi*P[i+(long)c*mm]; }
for(int o=16;o>0;o>>=1) dot+=__shfl_down_sync(0xffffffffu,dot,o);
if((tid&31)==0) red[tid>>5]=dot; __syncthreads();
float d=0.f; for(int q=0;q<QR_PT/32;q++) d+=red[q]; d*=tau;
for(int i=j+tid;i<mm;i+=QR_PT){ float vi=(i==j)?1.f:P[i+(long)j*mm]; P[i+(long)c*mm]-=vi*d; }
__syncthreads();
}
}
for(int j=1;j<w;++j){
float tau=Tsh[j+(long)j*w];
for(int i=tid;i<j;i+=QR_PT){ float zz=P[j+(long)i*mm]; for(int k=j+1;k<mm;++k) zz+=P[k+(long)i*mm]*P[k+(long)j*mm]; red[i]=zz; } /* z[i] in shared red[] */
__syncthreads();
for(int i=tid;i<j;i+=QR_PT){ float acc=0; for(int k=i;k<j;++k) acc+=Tsh[i+(long)k*w]*red[k]; Tsh[i+(long)j*w]=-tau*acc; } /* parallel matvec over i */
__syncthreads();
}
for(int idx=tid;idx<mm*w;idx+=QR_PT){ int i=idx%mm,j=idx/mm; float v=P[i+(long)j*mm];
Ab[(p+i)+(long)(p+j)*lda]=v; Vb[(p+i)+(long)(p+j)*ldv]=(i<j)?0.f:(i==j)?1.f:v; }
for(int idx=tid;idx<w*w;idx+=QR_PT){ int i=idx%w,j=idx/w; Tb[(p+i)+(long)(p+j)*ldt]=(i<=j)?Tsh[i+(long)j*w]:0.f; }
}
__global__ void k_zero(float* X,long n){ for(long i=(long)blockIdx.x*blockDim.x+threadIdx.x;i<n;i+=(long)gridDim.x*blockDim.x) X[i]=0.f; }
/* split fp32 (rows x cols, ld, stride s) -> fp16 hi + lo (packed col-major ld=rows, stride rows*cols) */
__global__ void k_split(const float* src,int ld,long s,__half* hi,__half* lo,int rows,int cols,long hs){
long bi=blockIdx.y; const float* S=src+bi*s; __half* Hh=hi+bi*hs; __half* Ll=lo+bi*hs;
for(long t=(long)blockIdx.x*blockDim.x+threadIdx.x;t<(long)rows*cols;t+=(long)gridDim.x*blockDim.x){
int i=t%rows,j=t/rows; float x=S[i+(long)j*ld]; __half h=__float2half(x);
Hh[i+(long)j*rows]=h; Ll[i+(long)j*rows]=__float2half(x-__half2float(h)); }
}
__global__ void k_r2c(const float* in,float* out,int n,long s){ long bi=blockIdx.y; const float* I=in+bi*s; float* O=out+bi*s;
for(long t=(long)blockIdx.x*blockDim.x+threadIdx.x;t<(long)n*n;t+=(long)gridDim.x*blockDim.x){ int i=t/n,j=t%n; O[i+(long)j*n]=I[i*(long)n+j]; } }
__global__ void k_c2r(const float* in,float* out,int n,long s){ long bi=blockIdx.y; const float* I=in+bi*s; float* O=out+bi*s;
for(long t=(long)blockIdx.x*blockDim.x+threadIdx.x;t<(long)n*n;t+=(long)gridDim.x*blockDim.x){ int i=t/n,j=t%n; O[i*(long)n+j]=I[i+(long)j*n]; } }
/* Coalesced 32x32 tiled transpose (replaces naive k_r2c/k_c2r; profile: naive = 1840us @ 10.6% DRAM uncoalesced).
* Canonical: computes out[i,j]=in[j,i] (both row-major), which IS r2c (Acm_colmajor=transpose(Hrm)) AND c2r (its inverse).
* Launch dim3 block(32,8), grid((n+31)/32,(n+31)/32,b). +1 pad kills shared-bank conflicts. */
#define QR_TT 32
/* CHK=false drops per-element bounds checks (valid when n%32==0: our 512/1024/352) -> instruction-bound transpose
* (B200: 87% SM / 35% DRAM) gets leaner. CHK=true is the safe path for n not a multiple of 32 (e.g. 176). */
template<bool CHK>
__global__ void k_tpose(const float* in,float* out,int n,long s){
__shared__ float tile[QR_TT][QR_TT+1];
long bi=blockIdx.z; const float* I=in+bi*s; float* O=out+bi*s;
int x=blockIdx.x*QR_TT+threadIdx.x, y=blockIdx.y*QR_TT+threadIdx.y; /* read source coalesced (x=fast=col) */
for(int r=0;r<QR_TT;r+=blockDim.y) if(!CHK || (x<n && y+r<n)) tile[threadIdx.y+r][threadIdx.x]=I[(long)(y+r)*n+x];
__syncthreads();
x=blockIdx.y*QR_TT+threadIdx.x; y=blockIdx.x*QR_TT+threadIdx.y; /* write dest coalesced (transposed block) */
for(int r=0;r<QR_TT;r+=blockDim.y) if(!CHK || (x<n && y+r<n)) O[(long)(y+r)*n+x]=tile[threadIdx.x][threadIdx.y+r];
}
__global__ void k_extau(const float* T,float* tau,int n){ long bi=blockIdx.y; const float* Tb=T+bi*(long)n*n; float* tb=tau+bi*(long)n;
for(int j=blockIdx.x*blockDim.x+threadIdx.x;j<n;j+=gridDim.x*blockDim.x) tb[j]=Tb[j+(long)j*n]; }
/* ===== blocked fp16-trailing path: cuBLAS fp16 tensor-core apply, fp32 reflectors+R output ===== */
__global__ void k_f2h(const float* a,__half* h,long tot){ for(long i=(long)blockIdx.x*blockDim.x+threadIdx.x;i<tot;i+=(long)gridDim.x*blockDim.x) h[i]=__float2half(a[i]); }
/* input Hrm (fp32 row-major) -> Af (fp16 COLUMN-major): Af[i+j*n]=fp16(Hrm[i*n+j]). FUSES r2c+f2h in ONE
* coalesced tiled pass (the fp32 col-major Acm input is never read — panels read g_Af, write Acm as output).
* Launch dim3 block(32,8), grid((n+31)/32,(n+31)/32,b). */
template<bool CHK>
__global__ void k_r2ch(const float* Hrm,__half* Af,int n,long s){
__shared__ float tile[QR_TT][QR_TT+1];
long bi=blockIdx.z; const float* H=Hrm+bi*s; __half* A=Af+bi*s;
int x=blockIdx.x*QR_TT+threadIdx.x, y=blockIdx.y*QR_TT+threadIdx.y;
for(int r=0;r<QR_TT;r+=blockDim.y) if(!CHK || (x<n && y+r<n)) tile[threadIdx.y+r][threadIdx.x]=H[(long)(y+r)*n+x];
__syncthreads();
x=blockIdx.y*QR_TT+threadIdx.x; y=blockIdx.x*QR_TT+threadIdx.y;
for(int r=0;r<QR_TT;r+=blockDim.y) if(!CHK || (x<n && y+r<n)) A[(long)(y+r)*n+x]=__float2half(tile[threadIdx.x][threadIdx.y+r]);
}
/* FAST warp-per-column panel (NB=32 cols=32 warps, register-resident, m<=512). COLUMN-MAJOR, reads fp16
* trailing Af; writes fp32 packed Acm, fp32 clean V (Vf for Gram/T), fp16 V-split (Vh/Vl), tau. */
template<int TR,int NW> /* TR=register rows/lane (col[TR]); NW=warps=columns=threads/32. launch_bounds=NW*32 sets reg cap. */
__global__ void __launch_bounds__(NW*32) k_panel_wpc(const __half* Af,float* Acm,float* Vf,__half* Vh,__half* Vl,__half* V3,float* tauout,
int n,int p,int w,long sAf,long sAc,long sVc,long sV3){
const int bi=blockIdx.x, ww=threadIdx.x>>5, lane=threadIdx.x&31, m=n-p;
const __half* Ab=Af+bi*sAf; float* Cb=Acm+bi*sAc; float* Vfb=Vf+bi*sVc; __half* Vhb=Vh+bi*sVc; __half* Vlb=Vl+bi*sVc; __half* V3b=V3+bi*sV3;
__shared__ float vsh[TR*32]; float col[TR];
for(int i=0;i<TR;i++){ int r=lane+i*32; col[i]=(r<m&&ww<w)?__half2float(Ab[(long)(p+r)+(long)(p+ww)*n]):0.f; }
for(int j=0;j<w;j++){
if(ww==j){
float s=0.f; for(int i=0;i<TR;i++){ int r=lane+i*32; if(r>=j&&r<m) s+=col[i]*col[i]; }
for(int o=16;o>0;o>>=1) s+=__shfl_xor_sync(0xffffffffu,s,o);
float alpha=__shfl_sync(0xffffffffu,col[j>>5],j&31);
float beta=0,rden=1,t=0;
if(s>0.f){ float nrm=sqrtf(s),sg=(alpha>=0)?1.f:-1.f; beta=-sg*nrm; rden=1.f/(alpha-beta); t=(beta-alpha)/beta; }
for(int i=0;i<TR;i++){ int r=lane+i*32; if(r>j&&r<m) col[i]*=rden; if(r==j) col[i]=beta; }
for(int i=0;i<TR;i++){ int r=lane+i*32; if(r<m) vsh[r]=(r>j)?col[i]:(r==j?1.f:0.f); }
if(lane==0) tauout[(long)bi*n+(p+j)]=t;
}
__syncthreads();
float t=tauout[(long)bi*n+(p+j)];
if(ww>j&&ww<w&&t!=0.f){
float dot=0.f; for(int i=0;i<TR;i++){ int r=lane+i*32; if(r<m) dot+=vsh[r]*col[i]; }
for(int o=16;o>0;o>>=1) dot+=__shfl_xor_sync(0xffffffffu,dot,o);
float wc=t*dot; for(int i=0;i<TR;i++){ int r=lane+i*32; if(r<m) col[i]-=vsh[r]*wc; }
}
__syncthreads();
}
if(ww<w) for(int i=0;i<TR;i++){ int r=lane+i*32; if(r<m){
Cb[(long)(p+r)+(long)(p+ww)*n]=col[i]; /* fp32 packed reflectors+R (col-major Acm) */
float vc=(r<ww)?0.f:(r==ww)?1.f:col[i]; __half vhh=__float2half(vc); __half vll=__float2half(vc-__half2float(vhh));
Vfb[r+(long)ww*n]=vc; Vhb[r+(long)ww*n]=vhh; Vlb[r+(long)ww*n]=vll;
V3b[r+(long)ww*n]=vhh; V3b[r+(long)(w+ww)*n]=vll; V3b[r+(long)(2*w+ww)*n]=vhh; } } /* V3=[Vh|Vl|Vh] for stacked apply */
}
/* ===== fused WMMA apply kernels (fp16-storage trailing, V-split -> ~fp32 accuracy, fewer C passes) ===== */
/* W = V^T C (fp32). V split (Vh,Vl), C fp16. col-major, ld=ldn for V&C, ldw for W. mm,w,rest mult of 16. */
__global__ void qr_k_VtC(const __half* Vh,const __half* Vl,const __half* C,float* W,int mm,int w,int rest,int ldn,int ldw,long sV,long sC,long sW){
long bi=blockIdx.y; int ctile=blockIdx.x*16, nk=w/16; /* one block per ctile -> read C once, all w-rows */
const __half* Vhb=Vh+bi*sV;const __half* Vlb=Vl+bi*sV;const __half* Cb=C+bi*sC;float* Wb=W+bi*sW;
nvcuda::wmma::fragment<nvcuda::wmma::accumulator,16,16,16,float> acc[2]; for(int k=0;k<nk;k++) nvcuda::wmma::fill_fragment(acc[k],0.f);
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a,16,16,16,__half,nvcuda::wmma::row_major> ah,al; /* V^T: a[k,i]=V[i,k], ld=ldn */
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b,16,16,16,__half,nvcuda::wmma::col_major> b; /* C[i,c], ld=ldn */
for(int i0=0;i0<mm;i0+=16){
nvcuda::wmma::load_matrix_sync(b,Cb+i0+(long)ctile*ldn,ldn);
for(int kt=0;kt<nk;kt++){
nvcuda::wmma::load_matrix_sync(ah,Vhb+i0+(long)(kt*16)*ldn,ldn); nvcuda::wmma::mma_sync(acc[kt],ah,b,acc[kt]);
nvcuda::wmma::load_matrix_sync(al,Vlb+i0+(long)(kt*16)*ldn,ldn); nvcuda::wmma::mma_sync(acc[kt],al,b,acc[kt]);
}
}
for(int kt=0;kt<nk;kt++) nvcuda::wmma::store_matrix_sync(Wb+(kt*16)+(long)ctile*ldw,acc[kt],ldw,nvcuda::wmma::mem_col_major);
}
/* C -= V Y. 64x64 tile, 4 warps, V pre-split (Vh,Vl) + Y (fp32) staged/split in shared. Bounds-checked. */
__global__ void __launch_bounds__(128) qr_k_CmVY(__half* C,const __half* Vh,const __half* Vl,const float* Yf,int mm,int w,int rest,int ldn,int ldw,long sV,long sC,long sY){
long bi=blockIdx.z; int ri=blockIdx.y*64, ci=blockIdx.x*64;
const __half* Vhb=Vh+bi*sV;const __half* Vlb=Vl+bi*sV;const float* Yfb=Yf+bi*sY;__half* Cb=C+bi*sC;
__shared__ __half sVh[64*64],sVl[64*64],sYh[64*64],sYl[64*64];
int tid=threadIdx.x;
for(int t=tid;t<64*w;t+=128){ int i=t%64,k=t/64; int gi=ri+i; sVh[i+k*64]=(gi<mm)?Vhb[gi+(long)k*ldn]:__float2half(0.f); sVl[i+k*64]=(gi<mm)?Vlb[gi+(long)k*ldn]:__float2half(0.f); }
for(int t=tid;t<w*64;t+=128){ int k=t%w,c=t/w; int gc=ci+c; float y=(gc<rest)?Yfb[k+(long)gc*ldw]:0.f; __half h=__float2half(y); sYh[k+c*w]=h; sYl[k+c*w]=__float2half(y-__half2float(h)); }
__syncthreads();
int warp=tid/32, wr=warp/2, wc=warp%2, nk=w/16;
nvcuda::wmma::fragment<nvcuda::wmma::matrix_a,16,16,16,__half,nvcuda::wmma::col_major> avh,avl;
nvcuda::wmma::fragment<nvcuda::wmma::matrix_b,16,16,16,__half,nvcuda::wmma::col_major> byh,byl;
nvcuda::wmma::fragment<nvcuda::wmma::accumulator,16,16,16,float> acc[2][2];
for(int a=0;a<2;a++)for(int bb=0;bb<2;bb++) nvcuda::wmma::fill_fragment(acc[a][bb],0.f);
for(int kk=0;kk<nk;kk++){
for(int a=0;a<2;a++){ int rf=wr*32+a*16;
nvcuda::wmma::load_matrix_sync(avh,sVh+rf+(long)(kk*16)*64,64); nvcuda::wmma::load_matrix_sync(avl,sVl+rf+(long)(kk*16)*64,64);
for(int bb=0;bb<2;bb++){ int cf=wc*32+bb*16;
nvcuda::wmma::load_matrix_sync(byh,sYh+kk*16+(long)cf*w,w); nvcuda::wmma::load_matrix_sync(byl,sYl+kk*16+(long)cf*w,w);
nvcuda::wmma::mma_sync(acc[a][bb],avh,byh,acc[a][bb]); nvcuda::wmma::mma_sync(acc[a][bb],avh,byl,acc[a][bb]); nvcuda::wmma::mma_sync(acc[a][bb],avl,byh,acc[a][bb]);
}
}
}
__shared__ float tl[128*16]; float* myt=tl+warp*256;
for(int a=0;a<2;a++)for(int bb=0;bb<2;bb++){
nvcuda::wmma::store_matrix_sync(myt,acc[a][bb],16,nvcuda::wmma::mem_col_major);
int r0=ri+wr*32+a*16, c0=ci+wc*32+bb*16;
for(int t=tid%32;t<256;t+=32){ int i=t%16,c=t/16; int gr=r0+i,gc=c0+c; if(gr<mm&&gc<rest){ long idx=(long)gr+(long)gc*ldn; Cb[idx]=__float2half(__half2float(Cb[idx])-myt[t]); } }
}
}
/* W = Wstack[0:w] + Wstack[w:2w] (sum the two halves of [Vh|Vl]^T C). col-major, Ws ld=2w, Wf ld=w. */
__global__ void k_wadd(const float* Ws,float* Wf,int w,int rest,long sWs,long sWf){
long bi=blockIdx.y; const float* W=Ws+bi*sWs; float* O=Wf+bi*sWf;
for(long t=(long)blockIdx.x*blockDim.x+threadIdx.x;t<(long)w*rest;t+=(long)gridDim.x*blockDim.x){
int k=t%w,c=t/w; O[k+(long)c*w]=W[k+(long)c*2*w]+W[(k+w)+(long)c*2*w]; }
}
/* Y3 = [Yh ; Yh ; Yl] (3w x rest) directly from fp32 Yf (w x rest, ld=w) -- splits Yf in-kernel (fuses k_split). */
__global__ void k_y3build(const float* Yf,__half* Y3,int w,int rest,long sY,long sY3){
long bi=blockIdx.y; const float* yf=Yf+bi*sY; __half* o=Y3+bi*sY3; int w3=3*w;
for(long t=(long)blockIdx.x*blockDim.x+threadIdx.x;t<(long)w3*rest;t+=(long)gridDim.x*blockDim.x){
int k=t%w3,c=t/w3, kk=(k<w)?k:(k<2*w?k-w:k-2*w); float y=yf[kk+(long)c*w]; __half h=__float2half(y);
o[k+(long)c*w3]=(k<2*w)?h:__float2half(y-__half2float(h)); } /* [0,2w)=Yh, [2w,3w)=Yl */
}
/* compact-WY T (w x w upper, fp32) from Gram G=V^T V (col-major) + tau -> g_T[p:p+w,p:p+w] (ld=n). */
__global__ void k_formT(const float* G,const float* tau,float* T,int n,int p,int w,long sG,long sT){
extern __shared__ float sm[]; float* Gs=sm; float* Ts=Gs+w*w; float* ts=Ts+w*w;
int bi=blockIdx.x, tid=threadIdx.x;
const float* Gb=G+bi*sG; const float* tb=tau+bi*n; float* Tb=T+bi*sT;
for(int idx=tid;idx<w*w;idx+=blockDim.x){ Gs[idx]=Gb[idx]; Ts[idx]=0.f; }
for(int j=tid;j<w;j+=blockDim.x) ts[j]=tb[p+j];
__syncthreads();
if(tid==0) Ts[0]=ts[0];
__syncthreads();
for(int j=1;j<w;++j){ float tj=ts[j];
for(int i=tid;i<j;i+=blockDim.x){ float acc=0; for(int k=i;k<j;++k) acc+=Ts[i+(long)k*w]*Gs[k+(long)j*w]; Ts[i+(long)j*w]=-tj*acc; }
if(tid==0) Ts[j+(long)j*w]=tj; __syncthreads(); }
for(int idx=tid;idx<w*w;idx+=blockDim.x){ int i=idx%w,j=idx/w; Tb[(p+i)+(long)(p+j)*n]=(i<=j)?Ts[i+(long)j*w]:0.f; }
}
/* Panel (one block per matrix): read fp16 trailing Af[p:m,p:p+w] (col-major ld=n), factor in fp32;
* write packed reflectors+R -> fp32 Acm (col-major), clean V -> fp16 Vh (mm x w, ld=n, local rows),
* compact-WY T -> fp16 Th (w x w, ld=nb), tau diag stays in Th[j,j]. */
__global__ void k_panel_h(const __half* Af,float* Acm,int m,int p,int w,int n,
__half* Vh,__half* Vl,int ldv,__half* Th,int ldt,float* Tf,float* tauout,long sAf,long sAc,long sVh,long sTh,long sTf){
extern __shared__ float s[];
const int mm=m-p, tid=threadIdx.x, bi=blockIdx.x;
const __half* Ab=Af+bi*sAf; float* Cb=Acm+bi*sAc; __half* Vb=Vh+bi*sVh; __half* Vbl=Vl+bi*sVh; __half* Tb=Th+bi*sTh; float* Tfb=Tf+bi*sTf;
float* P=s; float* Tsh=P+(long)mm*w; __shared__ float red[QR_PT];
for(int idx=tid; idx<mm*w; idx+=QR_PT){ int i=idx%mm,j=idx/mm; P[i+(long)j*mm]=__half2float(Ab[(p+i)+(long)(p+j)*n]); }
__syncthreads();
for(int j=0;j<w;++j){
float ss=0.f; for(int i=j+tid;i<mm;i+=QR_PT){ float v=P[i+(long)j*mm]; ss+=v*v; }
for(int o=16;o>0;o>>=1) ss+=__shfl_down_sync(0xffffffffu,ss,o);
if((tid&31)==0) red[tid>>5]=ss; __syncthreads();
float nrm2=0.f; for(int q=0;q<QR_PT/32;q++) nrm2+=red[q];
float alpha=P[j+(long)j*mm], nrm=sqrtf(nrm2), beta=(alpha>=0?-nrm:nrm);
float tau=(nrm==0.f)?0.f:(beta-alpha)/beta, scal=(nrm==0.f)?0.f:1.f/(alpha-beta);
for(int i=j+1+tid;i<mm;i+=QR_PT) P[i+(long)j*mm]*=scal;
__syncthreads();
if(tid==0){ P[j+(long)j*mm]=beta; Tsh[j+(long)j*w]=tau; tauout[(long)bi*n+(p+j)]=tau; }
__syncthreads();
for(int c=j+1;c<w;++c){
float dot=0.f; for(int i=j+tid;i<mm;i+=QR_PT){ float vi=(i==j)?1.f:P[i+(long)j*mm]; dot+=vi*P[i+(long)c*mm]; }
for(int o=16;o>0;o>>=1) dot+=__shfl_down_sync(0xffffffffu,dot,o);
if((tid&31)==0) red[tid>>5]=dot; __syncthreads();
float d=0.f; for(int q=0;q<QR_PT/32;q++) d+=red[q]; d*=tau;
for(int i=j+tid;i<mm;i+=QR_PT){ float vi=(i==j)?1.f:P[i+(long)j*mm]; P[i+(long)c*mm]-=vi*d; }
__syncthreads();
}
}
for(int j=1;j<w;++j){
float tau=Tsh[j+(long)j*w];
for(int i=tid;i<j;i+=QR_PT){ float zz=P[j+(long)i*mm]; for(int k=j+1;k<mm;++k) zz+=P[k+(long)i*mm]*P[k+(long)j*mm]; red[i]=zz; }
__syncthreads();
for(int i=tid;i<j;i+=QR_PT){ float acc=0; for(int k=i;k<j;++k) acc+=Tsh[i+(long)k*w]*red[k]; Tsh[i+(long)j*w]=-tau*acc; }
__syncthreads();
}
for(int idx=tid;idx<mm*w;idx+=QR_PT){ int i=idx%mm,j=idx/mm; float v=P[i+(long)j*mm];
Cb[(p+i)+(long)(p+j)*n]=v; /* fp32 packed reflectors+R */
float vc=(i<j)?0.f:(i==j)?1.f:v; __half vh=__float2half(vc);
Vb[i+(long)j*ldv]=vh; Vbl[i+(long)j*ldv]=__float2half(vc-__half2float(vh)); } /* V = V_hi + V_lo (~fp32) */
for(int idx=tid;idx<w*w;idx+=QR_PT){ int i=idx%w,j=idx/w; float t=(i<=j)?Tsh[i+(long)j*w]:0.f;
Tb[i+(long)j*ldt]=__float2half(t); Tfb[(p+i)+(long)(p+j)*n]=t; } /* fp16 + fp32 panel T (fp32 for Y=T^T W) */
}
/* copy R-rows produced by the apply: Af[p:p+w, p+w:n] (fp16) -> Acm (fp32). */
__global__ void k_rrows(const __half* Af,float* Acm,int p,int w,int rest,int n,long sAf,long sAc){
long bi=blockIdx.y; const __half* Ab=Af+bi*sAf; float* Cb=Acm+bi*sAc;
for(long t=(long)blockIdx.x*blockDim.x+threadIdx.x;t<(long)w*rest;t+=(long)gridDim.x*blockDim.x){
int i=t%w, c=t/w; int gr=p+i, gc=p+w+c; Cb[gr+(long)gc*n]=__half2float(Ab[gr+(long)gc*n]); } /* R-rows -> col-major Acm */
}
/* ================= host GEMM helpers ================= */
/* Ootomo 3-GEMM: C = alpha op(A) op(B) + beta C in ~fp32 accuracy via fp16 tensor cores. */
static void ogemm(cublasOperation_t opA,cublasOperation_t opB,int m,int n,int k,
float alpha,const float* A,int lda,long sA,const float* B,int ldb,long sB,
float beta,float* C,int ldc,long sC){
int ar=(opA==CUBLAS_OP_N)?m:k, ac=(opA==CUBLAS_OP_N)?k:m;
int br=(opB==CUBLAS_OP_N)?k:n, bc=(opB==CUBLAS_OP_N)?n:k;
long has=(long)ar*ac, hbs=(long)br*bc;
dim3 ga((ar*ac+255)/256,g_B), gb((br*bc+255)/256,g_B);
k_split<<<ga,256>>>(A,lda,sA,g_Ah,g_Al,ar,ac,has);
k_split<<<gb,256>>>(B,ldb,sB,g_Bh,g_Bl,br,bc,hbs);
float one=1.f;
cublasGemmStridedBatchedEx(g_H,opA,opB,m,n,k,&alpha,g_Ah,CUDA_R_16F,ar,has,g_Bh,CUDA_R_16F,br,hbs,&beta,C,CUDA_R_32F,ldc,sC,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
cublasGemmStridedBatchedEx(g_H,opA,opB,m,n,k,&alpha,g_Ah,CUDA_R_16F,ar,has,g_Bl,CUDA_R_16F,br,hbs,&one,C,CUDA_R_32F,ldc,sC,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
cublasGemmStridedBatchedEx(g_H,opA,opB,m,n,k,&alpha,g_Al,CUDA_R_16F,ar,has,g_Bh,CUDA_R_16F,br,hbs,&one,C,CUDA_R_32F,ldc,sC,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
}
/* dispatch: fp32 cuBLAS (TF32 if math mode set) or Ootomo fp16 for the larger GEMMs. */
static void gemm(cublasOperation_t opA,cublasOperation_t opB,int m,int n,int k,
float alpha,const float* A,int lda,long sA,const float* B,int ldb,long sB,
float beta,float* C,int ldc,long sC){
if(g_fp16 && k>=48) ogemm(opA,opB,m,n,k,alpha,A,lda,sA,B,ldb,sB,beta,C,ldc,sC);
else cublasSgemmStridedBatched(g_H,opA,opB,m,n,k,&alpha,A,lda,sA,B,ldb,sB,&beta,C,ldc,sC,g_B);
}
/* B[bi] = A[bi][off:m, c0:c0+ncol] := Q^T B (Q = panel cols [off,off+w], rows off..m). */
static void apply_QtB(float* A,int lda,int off,int w,int c0,int ncol){
int mm=g_M-off;
float* Vp=g_V+off+(long)off*g_M; float* Tp=g_T+off+(long)off*g_N; float* B=A+off+(long)c0*lda;
float* Wt=g_W; float* Yt=g_W+(long)w*ncol;
gemm(CUBLAS_OP_T,CUBLAS_OP_N,w,ncol,mm,1.f,Vp,g_M,g_sV,B,lda,g_sA,0.f,Wt,w,g_sW);
gemm(CUBLAS_OP_T,CUBLAS_OP_N,w,ncol,w, 1.f,Tp,g_N,g_sT,Wt,w,g_sW,0.f,Yt,w,g_sW);
gemm(CUBLAS_OP_N,CUBLAS_OP_N,mm,ncol,w,-1.f,Vp,g_M,g_sV,Yt,w,g_sW,1.f,B,lda,g_sA);
}
/* combine: Tg[off:off+n1, off+n1:off+n1+n2] = -T1 (V1^T V2) T2 */
static void combine(int off,int n1,int n2){
int mm=g_M-off;
float* V1=g_V+off+(long)off*g_M; float* V2=g_V+off+(long)(off+n1)*g_M;
float* T1=g_T+off+(long)off*g_N; float* T2=g_T+(off+n1)+(long)(off+n1)*g_N;
float* T12=g_T+off+(long)(off+n1)*g_N;
float* G=g_W; float* M=g_W+(long)n1*n2;
gemm(CUBLAS_OP_T,CUBLAS_OP_N,n1,n2,mm,1.f,V1,g_M,g_sV,V2,g_M,g_sV,0.f,G,n1,g_sW);
gemm(CUBLAS_OP_N,CUBLAS_OP_N,n1,n2,n1,1.f,T1,g_N,g_sT,G,n1,g_sW,0.f,M,n1,g_sW);
gemm(CUBLAS_OP_N,CUBLAS_OP_N,n1,n2,n2,-1.f,M,n1,g_sW,T2,g_N,g_sT,0.f,T12,g_N,g_sT);
}
/* recursive QR (Algorithm 1) on cols [off,off+ncol], rows off..m. */
static void recurse(float* A,int lda,int off,int ncol){
if(ncol<=g_nb){
size_t shb=((size_t)(g_M-off)*ncol+(size_t)ncol*ncol)*sizeof(float);
k_panel<<<g_B,QR_PT,shb>>>(A,g_M,off,ncol,lda,g_V,g_M,g_T,g_N,g_sA,g_sV,g_sT);
return;
}
int n1=ncol/2,n2=ncol-n1;
recurse(A,lda,off,n1);
apply_QtB(A,lda,off,n1,off+n1,n2);
recurse(A,lda,off+n1,n2);
combine(off,n1,n2);
}
/* ===== blocked fp16-trailing apply: C := Q^T C via cuBLAS fp16 tensor cores (all fp16, fp32 accum) ===== */
/* fp32-accurate V-split apply: W=(Vh+Vl)^T C in fp32, Y=T^T W in fp32 (fp32 T), then
* C -= (Vh+Vl)(Yh+Yl) via 3 fp16-TC terms -> ~fp32 accuracy with fp16-stored C. */
static void apply_h(int n,int p,int w,int rest){
int mm=n-p;
long sAf=(long)n*n, sVh=(long)n*QR_NB, sWh=(long)QR_NB*n, sTn=(long)n*n;
const __half* Vh=g_Vh; const __half* Vl=g_Vl; /* mm x w, ld=n */
__half* C=g_Af + p + (long)(p+w)*n; /* mm x rest, ld=n */
float* Wf=g_W; /* fp32 W (w x rest, ld=w) */
float* Yf=g_W + (long)g_B*QR_NB*n; /* fp32 Y after W */
const float* Tp=g_T + p + (long)p*n; /* fp32 panel T at [p,p], ld=n */
__half* Yh=g_Yh; __half* Yl=g_Yl; /* fp16 Y split */
float one=1.f, neg=-1.f, zero=0.f;
static int s_mode = getenv("QR_APPLY") ? atoi(getenv("QR_APPLY")) : 3; /* 3=STACKED(default) 0=hybrid 1=cuBLAS6 2=WMMA */
if(s_mode==3){ /* STACKED: 3 GEMMs, C read 2x in fp16 (half baseline fp32 traffic). Exact 3-term precision. */
long sV3=(long)n*3*QR_NB, sY3=(long)3*QR_NB*n, sWs=(long)2*QR_NB*n, sWf=(long)QR_NB*n;
/* g_V3 holds THIS panel's [Vh|Vl|Vh] at local rows 0..mm-1 (aligned with C sub-block local rows). */
float* Ws=g_W; Wf=g_W+(long)g_B*2*QR_NB*n; Yf=g_W+(long)g_B*3*QR_NB*n;
/* Wstack = [Vh|Vl]^T C = V3[:,0:2w]^T C (2w x rest, fp32) */
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_T,CUBLAS_OP_N,2*w,rest,mm,&one,g_V3,CUDA_R_16F,n,sV3,C,CUDA_R_16F,n,sAf,&zero,Ws,CUDA_R_32F,2*w,sWs,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
k_wadd<<<dim3((w*rest+255)/256,g_B),256>>>(Ws,Wf,w,rest,sWs,sWf); /* W = Ws[0:w]+Ws[w:2w] */
cublasSgemmStridedBatched(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,w,&one,Tp,n,sTn,Wf,w,sWf,&zero,Yf,w,sWf,g_B); /* Y=T^T W */
k_y3build<<<dim3((3*w*rest+255)/256,g_B),256>>>(Yf,g_Y3,w,rest,sWf,sY3); /* Y3=[Yh;Yh;Yl], split fused */
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_N,CUBLAS_OP_N,mm,rest,3*w,&neg,g_V3,CUDA_R_16F,n,sV3,g_Y3,CUDA_R_16F,3*w,sY3,&one,C,CUDA_R_16F,n,sAf,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP); /* C -= V3 Y3 */
return;
}
if(s_mode==0){ /* HYBRID: cuBLAS for precise W,Y; FUSED WMMA for C-=V Y (reads C once, full 3-term) */
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,mm,&one,Vh,CUDA_R_16F,n,sVh,C,CUDA_R_16F,n,sAf,&zero,Wf,CUDA_R_32F,w,sWh,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,mm,&one,Vl,CUDA_R_16F,n,sVh,C,CUDA_R_16F,n,sAf,&one, Wf,CUDA_R_32F,w,sWh,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
cublasSgemmStridedBatched(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,w,&one,Tp,n,sTn,Wf,w,sWh,&zero,Yf,w,sWh,g_B);
qr_k_CmVY<<<dim3((rest+63)/64,(mm+63)/64,g_B),128>>>(C,Vh,Vl,Yf,mm,w,rest,n,w,sVh,sAf,sWh);
return;
}
if(s_mode==2){ /* full WMMA */
qr_k_VtC<<<dim3(rest/16,g_B),32>>>(Vh,Vl,C,Wf,mm,w,rest,n,w,sVh,sAf,sWh);
cublasSgemmStridedBatched(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,w,&one,Tp,n,sTn,Wf,w,sWh,&zero,Yf,w,sWh,g_B);
qr_k_CmVY<<<dim3((rest+63)/64,(mm+63)/64,g_B),128>>>(C,Vh,Vl,Yf,mm,w,rest,n,w,sVh,sAf,sWh);
return;
}
/* full cuBLAS V-split apply (6 GEMMs) */
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,mm,&one,Vh,CUDA_R_16F,n,sVh,C,CUDA_R_16F,n,sAf,&zero,Wf,CUDA_R_32F,w,sWh,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,mm,&one,Vl,CUDA_R_16F,n,sVh,C,CUDA_R_16F,n,sAf,&one, Wf,CUDA_R_32F,w,sWh,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP);
cublasSgemmStridedBatched(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,rest,w,&one,Tp,n,sTn,Wf,w,sWh,&zero,Yf,w,sWh,g_B);
k_split<<<dim3((w*rest+255)/256,g_B),256>>>(Yf,w,sWh,Yh,Yl,w,rest,sWh);
auto GT=[&](float a,const __half*A,const __half*B,float bt){
cublasGemmStridedBatchedEx(g_H,CUBLAS_OP_N,CUBLAS_OP_N,mm,rest,w,&a,A,CUDA_R_16F,n,sVh,B,CUDA_R_16F,w,sWh,&bt,C,CUDA_R_16F,n,sAf,g_B,CUBLAS_COMPUTE_32F,CUBLAS_GEMM_DEFAULT_TENSOR_OP); };
GT(neg,Vh,Yh,one); GT(neg,Vh,Yl,one); GT(neg,Vl,Yh,one); /* full 3-term C -= (Vh+Vl)(Yh+Yl) approx (safe precision) */
}
/* Blocked Householder QR, fp16 trailing storage + cuBLAS fp16 apply, fp32 reflectors+R into Acm. */
static void factor_blocked(const float* Hrm,float* Acm,float* tau,int b,int n){
ensure_state(b,n);
g_bnb=g_nb; /* panel width = max that fits shared mem */
size_t shmax=((size_t)n*g_bnb+(size_t)g_bnb*g_bnb)*sizeof(float);
qck(cudaFuncSetAttribute(k_panel_h,cudaFuncAttributeMaxDynamicSharedMemorySize,(int)shmax));
long sAf=(long)n*n, sVh=(long)n*QR_NB, sTh=(long)QR_NB*QR_NB, sG=(long)QR_NB*QR_NB;
int usewpc = (n<=QR_MAX_N); /* warp-per-column panel handles these sizes */
if(usewpc){ /* pick panel width NB from the method registry (first row n<=N_MAX) */
#define X(NMAX,NB,TR) else if(n<=(NMAX)) g_bnb=(NB);
if(0){} QR_SIZE_TABLE(X)
#undef X
}
float one=1.f, zero=0.f;
long sV3=(long)n*3*QR_NB;
{ dim3 tb(QR_TT,8), tg((n+QR_TT-1)/QR_TT,(n+QR_TT-1)/QR_TT,b);
if(n%QR_TT==0) k_r2ch<false><<<tg,tb>>>(Hrm,g_Af,n,(long)n*n);
else k_r2ch<true ><<<tg,tb>>>(Hrm,g_Af,n,(long)n*n); } /* fused coalesced r2c+f2h -> g_Af fp16 col-major */
for(int p=0;p<n;p+=g_bnb){
int w=(g_bnb<n-p)?g_bnb:(n-p), mm=n-p;
if(usewpc && w==g_bnb){
/* size-specialized panel from the method registry (first row with n<=N_MAX): NB warps, TR register rows */
#define X(NMAX,NB,TR) else if(n<=(NMAX)) k_panel_wpc<TR,NB><<<b,(NB)*32>>>(g_Af,Acm,g_Vf,g_Vh,g_Vl,g_V3,tau,n,p,w,sAf,sAf,sVh,sV3);
if(0){} QR_SIZE_TABLE(X)
#undef X
cublasSgemmStridedBatched(g_H,CUBLAS_OP_T,CUBLAS_OP_N,w,w,mm,&one,g_Vf,n,sVh,g_Vf,n,sVh,&zero,g_G,w,sG,b); /* Gram=V^T V */
k_formT<<<b,256,(size_t)(2*w*w+w)*sizeof(float)>>>(g_G,tau,g_T,n,p,w,sG,(long)n*n); /* fp32 compact-WY T */
} else {
size_t shb=((size_t)mm*w+(size_t)w*w)*sizeof(float);
qck(cudaFuncSetAttribute(k_panel_h,cudaFuncAttributeMaxDynamicSharedMemorySize,(int)shb));
k_panel_h<<<b,QR_PT,shb>>>(g_Af,Acm,n,p,w,n,g_Vh,g_Vl,n,g_Th,g_bnb,g_T,tau,sAf,sAf,sVh,sTh,(long)n*n);
}
int rest=n-(p+w);
if(rest>0){ apply_h(n,p,w,rest);
dim3 gr((w*rest+255)/256,b); k_rrows<<<gr,256>>>(g_Af,Acm,p,w,rest,n,sAf,sAf); }
}
}
/* ================= public column-major driver (for direct testing) ================= */
/* Factor Acm (b * n*n, COLUMN-MAJOR, in place packed) -> g_V, g_T filled. Sets up state. */
static void ensure_state(int b,int n){
if(!g_H){ cublasCreate(&g_H); /* default execution context */
if(getenv("QR_TF32")&&atoi(getenv("QR_TF32"))) cublasSetMathMode(g_H,CUBLAS_TF32_TENSOR_OP_MATH);
g_fp16 = getenv("QR_FP16")?atoi(getenv("QR_FP16")):0;
if(getenv("QR_VERBOSE")) fprintf(stderr,"[qrlib] g_fp16=%d\n",g_fp16);
}
g_M=n; g_N=n; g_B=b; g_sA=(long)n*n; g_sV=(long)n*n; g_sT=(long)n*n; g_sW=(long)n*n;
/* Choose g_nb = largest base-case width (<=QR_NB) whose largest panel (mm=n, w=g_nb) fits
* the device opt-in shared memory, with margin. Then raise the opt-in limit for k_panel. */
int dev=0; cudaGetDevice(&dev); int smax=0;
cudaDeviceGetAttribute(&smax,cudaDevAttrMaxSharedMemoryPerBlockOptin,dev);
size_t budget=(size_t)smax - 1024; /* leave a little room */
g_nb=QR_NB;
while(g_nb>1 && ((size_t)n*g_nb+(size_t)g_nb*g_nb)*sizeof(float) > budget) g_nb-=8;
if(g_nb<8) g_nb=8;
size_t shmax=((size_t)n*g_nb+(size_t)g_nb*g_nb)*sizeof(float);
qck(cudaFuncSetAttribute(k_panel,cudaFuncAttributeMaxDynamicSharedMemorySize,(int)shmax));
if(getenv("QR_VERBOSE")) fprintf(stderr,"[qrlib] n=%d g_nb=%d shmax=%zuKB (devmax=%dKB)\n",n,g_nb,shmax/1024,smax/1024);
static long cap=0; long need=(long)b*n*n;
if(need>cap){
if(g_V){ cudaFree(g_V);cudaFree(g_T);cudaFree(g_W);cudaFree(g_Ah);cudaFree(g_Al);cudaFree(g_Bh);cudaFree(g_Bl);
cudaFree(g_Af);cudaFree(g_Vh);cudaFree(g_Vl);cudaFree(g_Th);cudaFree(g_Wh);cudaFree(g_Yh);cudaFree(g_Yl);
cudaFree(g_Vf);cudaFree(g_G);cudaFree(g_V3);cudaFree(g_Y3); }
cudaMalloc(&g_V,need*4); cudaMalloc(&g_T,need*4); cudaMalloc(&g_W,need*4);
cudaMalloc(&g_Ah,need*2); cudaMalloc(&g_Al,need*2); cudaMalloc(&g_Bh,need*2); cudaMalloc(&g_Bl,need*2);
cudaMalloc(&g_Af,need*2); /* fp16 trailing */
cudaMalloc(&g_Vh,(long)b*n*QR_NB*2); cudaMalloc(&g_Vl,(long)b*n*QR_NB*2); cudaMalloc(&g_Th,(long)b*QR_NB*QR_NB*2);
cudaMalloc(&g_Wh,(long)b*QR_NB*n*2); cudaMalloc(&g_Yh,(long)b*QR_NB*n*2); cudaMalloc(&g_Yl,(long)b*QR_NB*n*2);
cudaMalloc(&g_Vf,(long)b*n*QR_NB*4); cudaMalloc(&g_G,(long)b*QR_NB*QR_NB*4);
cudaMalloc(&g_V3,(long)b*n*3*QR_NB*2); cudaMalloc(&g_Y3,(long)b*3*QR_NB*n*2);
cap=need;
}
}
/* Factor Acm (b * n*n, COLUMN-MAJOR) in place -> packed A, g_V, g_T. */
static void factor_colmajor(float* Acm,int b,int n){
ensure_state(b,n);
long need=(long)b*n*n;
k_zero<<<256,256>>>(g_V,need); k_zero<<<256,256>>>(g_T,need);
recurse(Acm,n,0,n);
}
} /* namespace qrlib */
/* ================= public entry (row-major, geqrf layout) ================= */
extern "C" void qr(long Hp,long taup,int b,int n){
using namespace qrlib;
float* Hrm=(float*)Hp; float* tau=(float*)taup;
ensure_state(b,n);
static float* Acm=0; static long acap=0; long need=(long)b*n*n;
if(need>acap){ if(Acm)cudaFree(Acm); cudaMalloc(&Acm,need*4); acap=need; }
dim3 g((n*n+255)/256,b);
dim3 tb(QR_TT,8), tg((n+QR_TT-1)/QR_TT,(n+QR_TT-1)/QR_TT,b); /* coalesced tiled transpose launch */
#define TPOSE(SRC,DST) do{ if(n%QR_TT==0) k_tpose<false><<<tg,tb>>>(SRC,DST,n,(long)n*n); \
else k_tpose<true ><<<tg,tb>>>(SRC,DST,n,(long)n*n); }while(0)
int recursive = getenv("QR_RECURSIVE") && atoi(getenv("QR_RECURSIVE"));
if(recursive){
TPOSE(Hrm,Acm);
factor_colmajor(Acm,b,n);
dim3 gt((n+255)/256,b); k_extau<<<gt,256>>>(g_T,tau,n);
TPOSE(Acm,Hrm);
} else {
factor_blocked(Hrm,Acm,tau,b,n); /* fused input r2c+f2h inside; blocked fp16-trailing (writes tau) */
TPOSE(Acm,Hrm); /* packed col-major -> row-major (geqrf layout, coalesced) */
}
#undef TPOSE
}
#endif /* QRLIB_CUH */
"""
CPP_SRC = 'extern "C" void qr(long H, long tau, int b, int n);'
_mod = load_inline(name="qrmod", cpp_sources=[CPP_SRC], cuda_sources=[CUDA_SRC],
functions=["qr"], extra_cuda_cflags=["-O3"], extra_ldflags=["-lcublas"],
verbose=False)
def custom_kernel(data: input_t) -> output_t:
b, n, _ = data.shape # assume contiguous, on GPU, float32
if n < 64 or n > 1024: # qrlib panel fits shared mem for n<=1024; else cuSOLVER
return torch.geqrf(data)
H = data.clone()
tau = torch.empty(b, n, device=data.device, dtype=data.dtype)
_mod.qr(H.data_ptr(), tau.data_ptr(), b, n)
return H, tau
scrolls · 566 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