Skip to content
KernelIndex
Search⌘K

submission 834652

shiva_92757 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

cand_km49.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-834652?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
8.79ms
#269 of 515
2026-06-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6eb77479aab004b5de22359a1c861ddfa5cdc73fcaff8a2ca280e3554f5eb55e
license declaredunknown
license concludedunknown
authorsshiva_92757
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float Ts[];

Kernel source

cand_km49.py226 lines
import torch
# km49: panel warp-count doubled (BM=512: NW 4->8, BM>=1024: NW 8->16) — tests if the 63%% panel reduction is warp-parallelism-limited. Identical math to km41 (correctness unchanged).
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# km19: trailing via cublas COMPUTE_32F_FAST_16F -> fp32 in/out, fp16 TC compute, fp32 accumulate,
# ZERO .half() copies (the profile #1). Same precision as the B200-passing torch-fp16 path (rel_err
# 4.5e-4), no copies. Stride-aware on the trailing slice (no .contiguous either).
_TB = r'''
#include <torch/extension.h>
__global__ void tbuild(const float* __restrict__ W, const float* __restrict__ tau,
                       float* __restrict__ Tout, int NB) {
    extern __shared__ float Ts[];
    int mat=blockIdx.x, tid=threadIdx.x, nt=blockDim.x;
    const float* Wm=W+(long)mat*NB*NB; const float* tm=tau+(long)mat*NB; float* To=Tout+(long)mat*NB*NB;
    for(int i=tid;i<NB*NB;i+=nt) Ts[i]=0.f;
    __syncthreads();
    for(int k=0;k<NB;k++){
        if(tid==0) Ts[k*NB+k]=tm[k];
        if(k>0){
            if(tid<k){ float s=0.f; for(int j=0;j<k;j++) s+=Ts[tid*NB+j]*Wm[j*NB+k]; Ts[tid*NB+k]=-tm[k]*s; }
            __syncthreads();
        }
    }
    for(int i=tid;i<NB*NB;i+=nt) To[i]=Ts[i];
}
void tbuild_launch(torch::Tensor W, torch::Tensor tau, torch::Tensor Tout){
    int b=W.size(0), NB=W.size(1);
    tbuild<<<b, 64, (size_t)NB*NB*sizeof(float)>>>(W.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), NB);
    TORCH_CHECK(cudaGetLastError()==cudaSuccess,"tbuild");
}
'''
_F16 = r'''
#include <torch/extension.h>
#include <cublas_v2.h>
static cublasHandle_t HBL=nullptr;
static void gx(int tA,int tB,int m,int n,int k,float al,const float*A,long sA,int lda,
               const float*B,long sB,int ldb,float be,float*C,long sC,int ldc,int b,cublasComputeType_t ct=CUBLAS_COMPUTE_32F_FAST_16F){
    if(!HBL) cublasCreate(&HBL);
    cublasOperation_t oa= tA?CUBLAS_OP_T:CUBLAS_OP_N, ob= tB?CUBLAS_OP_T:CUBLAS_OP_N;
    cublasGemmStridedBatchedEx(HBL, ob, oa, n, m, k, &al, B,CUDA_R_32F,ldb,sB, A,CUDA_R_32F,lda,sA,
        &be, C,CUDA_R_32F,ldc,sC, b, ct, CUBLAS_GEMM_DEFAULT);
}
void vtv(torch::Tensor V, torch::Tensor W){
    int b=V.size(0),M=V.size(1),jb=V.size(2);
    gx(1,0,jb,jb,M, 1.f, V.data_ptr<float>(),(long)M*jb,jb, V.data_ptr<float>(),(long)M*jb,jb, 0.f, W.data_ptr<float>(),(long)jb*jb,jb, b);
    TORCH_CHECK(cudaGetLastError()==cudaSuccess,"vtv");
}
void trail(torch::Tensor V, torch::Tensor T, torch::Tensor C, torch::Tensor W1, torch::Tensor W){
    int b=V.size(0),M=V.size(1),jb=V.size(2),Tw=C.size(2);
    long cs0=C.stride(0), cs1=C.stride(1);
    gx(1,0,jb,Tw,M, 1.f, V.data_ptr<float>(),(long)M*jb,jb, C.data_ptr<float>(),cs0,cs1, 0.f, W1.data_ptr<float>(),(long)jb*Tw,Tw, b);
    gx(1,0,jb,Tw,jb,1.f, T.data_ptr<float>(),(long)jb*jb,jb, W1.data_ptr<float>(),(long)jb*Tw,Tw, 0.f, W.data_ptr<float>(),(long)jb*Tw,Tw, b);
    gx(0,0,M,Tw,jb,-1.f, V.data_ptr<float>(),(long)M*jb,jb, W.data_ptr<float>(),(long)jb*Tw,Tw, 1.f, C.data_ptr<float>(),cs0,cs1, b);
    TORCH_CHECK(cudaGetLastError()==cudaSuccess,"trail");
}
void trail_t(torch::Tensor V, torch::Tensor T, torch::Tensor C, torch::Tensor W1, torch::Tensor W){
    int b=V.size(0),M=V.size(1),jb=V.size(2),Tw=C.size(2);
    long cs0=C.stride(0), cs1=C.stride(1);
    auto CT=CUBLAS_COMPUTE_32F_FAST_TF32;
    gx(1,0,jb,Tw,M, 1.f, V.data_ptr<float>(),(long)M*jb,jb, C.data_ptr<float>(),cs0,cs1, 0.f, W1.data_ptr<float>(),(long)jb*Tw,Tw, b, CT);
    gx(1,0,jb,Tw,jb,1.f, T.data_ptr<float>(),(long)jb*jb,jb, W1.data_ptr<float>(),(long)jb*Tw,Tw, 0.f, W.data_ptr<float>(),(long)jb*Tw,Tw, b, CT);
    gx(0,0,M,Tw,jb,-1.f, V.data_ptr<float>(),(long)M*jb,jb, W.data_ptr<float>(),(long)jb*Tw,Tw, 1.f, C.data_ptr<float>(),cs0,cs1, b, CT);
    TORCH_CHECK(cudaGetLastError()==cudaSuccess,"trail_t");
}
// ---- fused n=176 ----


__global__ void fused_qr(float* __restrict__ H, float* __restrict__ TAU,
                         int n, int NB, long sb, long sr, long sc, long tsb) {
    extern __shared__ float s[];
    int mat = blockIdx.x, tid = threadIdx.x, nt = blockDim.x;
    float* base = H + (long)mat*sb;
    float* TAUm = TAU + (long)mat*tsb;
    for (int j = 0; j < n; j += NB) {
        int M = n - j, jb = (NB < n - j) ? NB : (n - j);
        float* P = s; float* red = s + (long)M*jb;
        for (int idx = tid; idx < M*jb; idx += nt) { int r=idx/jb,c=idx%jb; P[idx]=base[(long)(j+r)*sr+(long)(j+c)*sc]; }
        __syncthreads();
        for (int k = 0; k < jb; k++) {
            float ps = 0.f;
            for (int r = k + tid; r < M; r += nt) { float x = P[r*jb + k]; ps += x*x; }
            red[tid] = ps; __syncthreads();
            for (int off = nt>>1; off>0; off>>=1) { if (tid<off) red[tid]+=red[tid+off]; __syncthreads(); }
            float nrm = sqrtf(red[0]); __syncthreads();
            float alpha = P[k*jb + k]; __syncthreads();
            int active = nrm > 1e-30f;
            float beta = (alpha >= 0.f) ? -nrm : nrm;
            float tau_k = (active && beta != 0.f) ? (beta - alpha)/beta : 0.f;
            float denom = active ? (alpha - beta) : 1.f;
            if (tid == 0) TAUm[j + k] = tau_k;
            for (int r = k + tid; r < M; r += nt) {
                if (r == k)      P[r*jb + k] = active ? beta : alpha;
                else if (active) P[r*jb + k] = P[r*jb + k] / denom;
            }
            __syncthreads();
            for (int c = k + 1; c < jb; c++) {
                float w = 0.f;
                for (int r = k + tid; r < M; r += nt) { float v = (r==k)?1.f:P[r*jb+k]; w += v*P[r*jb+c]; }
                red[tid] = w; __syncthreads();
                for (int off = nt>>1; off>0; off>>=1) { if (tid<off) red[tid]+=red[tid+off]; __syncthreads(); }
                float wsum = red[0]; __syncthreads();
                for (int r = k + tid; r < M; r += nt) { float v=(r==k)?1.f:P[r*jb+k]; P[r*jb+c]-=tau_k*v*wsum; }
                __syncthreads();
            }
        }
        for (int idx = tid; idx < M*jb; idx += nt) { int r=idx/jb,c=idx%jb; base[(long)(j+r)*sr+(long)(j+c)*sc]=P[idx]; }
        __syncthreads();
        for (int cc = j + jb + tid; cc < n; cc += nt) {
            for (int k = 0; k < jb; k++) {
                float tau_k = TAUm[j + k];
                if (tau_k == 0.f) continue;
                float w = base[(long)(j+k)*sr + (long)cc*sc];
                for (int r = k + 1; r < M; r++) w += P[r*jb + k] * base[(long)(j+r)*sr + (long)cc*sc];
                w *= tau_k;
                base[(long)(j+k)*sr + (long)cc*sc] -= w;
                for (int r = k + 1; r < M; r++) base[(long)(j+r)*sr + (long)cc*sc] -= P[r*jb + k] * w;
            }
        }
        __syncthreads();
    }
}

void fused_launch(torch::Tensor H, torch::Tensor TAU, int NB, int threads) {
    int b = H.size(0), n = H.size(1);
    size_t smem = ((size_t)n*NB + threads) * sizeof(float);
    cudaFuncSetAttribute(fused_qr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
    fused_qr<<<b, threads, smem>>>(H.data_ptr<float>(), TAU.data_ptr<float>(), n, NB,
                                   H.stride(0), H.stride(1), H.stride(2), TAU.stride(0));
    TORCH_CHECK(cudaGetLastError()==cudaSuccess, "launch");
}

'''
_tb = load_inline(name="km21_tbuild", cpp_sources="void tbuild_launch(torch::Tensor,torch::Tensor,torch::Tensor);", cuda_sources=_TB, functions=["tbuild_launch"], verbose=False, extra_cuda_cflags=["-O3"])
_f16 = load_inline(name="km21_f16", cpp_sources="void trail(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor);void trail_t(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor);void vtv(torch::Tensor,torch::Tensor);void fused_launch(torch::Tensor,torch::Tensor,int,int);", cuda_sources=_F16, functions=["trail","trail_t","vtv","fused_launch"], verbose=False, extra_cuda_cflags=["-O3"], extra_ldflags=["-lcublas"])

@triton.jit
def _panel_kernel(H, TAU, M, j, niter,
                  sb, sr, sc, tsb, tsn,
                  NB: tl.constexpr, BM: tl.constexpr):
    pid = tl.program_id(0)
    rows = tl.arange(0, BM)
    cols = tl.arange(0, NB)
    rmask = rows < M
    ptr = H + pid * sb + (j + rows)[:, None] * sr + (j + cols)[None, :] * sc
    pmask = rmask[:, None] & (cols[None, :] < niter)
    P = tl.load(ptr, mask=pmask, other=0.0)
    tau_acc = tl.zeros((NB,), dtype=tl.float32)
    for k in range(niter):
        is_k = rows == k
        ge_k = (rows >= k) & rmask
        gt_k = (rows > k) & rmask
        colk = tl.sum(tl.where(cols[None, :] == k, P, 0.0), axis=1)
        x = tl.where(ge_k, colk, 0.0)
        alpha = tl.sum(tl.where(is_k, colk, 0.0))
        nrm = tl.sqrt(tl.sum(x * x))
        active = nrm > 1e-30
        sign = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = -sign * nrm
        denom = tl.where(active, alpha - beta, 1.0)
        tau_k = tl.where(active & (beta != 0.0), (beta - alpha) / beta, 0.0)
        af = tl.where(active, 1.0, 0.0)
        v_tail = (colk / denom) * af
        v = tl.where(is_k, 1.0, tl.where(gt_k, v_tail, 0.0))
        w = tl.sum(v[:, None] * P, axis=0)
        P = tl.where(cols[None, :] > k, P - tau_k * v[:, None] * w[None, :], P)
        beta_eff = tl.where(active, beta, alpha)
        colk_act = tl.where(rows < k, colk, tl.where(is_k, beta_eff, v_tail))
        colk_new = tl.where(rmask, colk_act, 0.0)
        P = tl.where(cols[None, :] == k, colk_new[:, None], P)
        tau_acc = tl.where(cols == k, tau_k, tau_acc)
    tl.store(ptr, P, mask=pmask)
    tl.store(TAU + pid * tsb + (j + cols) * tsn, tau_acc, mask=cols < niter)





def _build_VT(H, tau, j, jb, b, n, dev):
    panel = H[:, j:, j:j+jb]; V = torch.tril(panel,-1).contiguous(); ar=torch.arange(jb,device=dev); V[:,ar,ar]=1.0
    W = torch.empty(b,jb,jb,device=dev); _f16.vtv(V, W); T=torch.zeros(b,jb,jb,device=dev)
    _tb.tbuild_launch(W, tau[:,j:j+jb].contiguous(), T); return V,T


def _fused_qr(A, NB=32, threads=256):
    b,n,_=A.shape; H=A.to(torch.float32).contiguous().clone(); tau=torch.zeros(b,n,dtype=torch.float32,device=A.device).contiguous()
    _f16.fused_launch(H,tau,NB,threads); return H,tau

def _triton_qr(A):
    b,n,_=A.shape; dev=A.device
    H=A.to(torch.float32).contiguous().clone(); tau=torch.zeros(b,n,dtype=torch.float32,device=dev)
    cn=H.norm(dim=1); rn=H.norm(dim=2); sp=lambda x:x.amax(1)/x.clamp_min(1e-30).amin(1)
    zf=(H.abs()<1e-12).float().mean(dim=(1,2))
    spc=sp(cn); spr=sp(rn)
    ill_any=bool(((spc>1e3)|(spr>1e3)|(zf>0.2)).any())
    # 3-tier: well-cond->FAST_16F; TF32-safe ill->TF32 (1.7x); precision-critical ill->fp32
    need_fp32=bool((((zf>0.5)|(spc<10)|((spc>1e10)&(spr>1e3)))).any())
    BM=triton.next_power_of_2(n); NB=32; NW=16 if BM>=1024 else 8
    for j in range(0,n,NB):
        jb=min(NB,n-j); M=n-j
        _panel_kernel[(b,)](H,tau,M,j,jb,H.stride(0),H.stride(1),H.stride(2),tau.stride(0),tau.stride(1),NB=NB,BM=BM,num_warps=NW)
        if j+jb<n:
            V,T=_build_VT(H,tau,j,jb,b,n,dev); Tw=n-j-jb; C=H[:,j:,j+jb:]
            if ill_any and need_fp32:
                # precision-critical (band/rowscale/mixed): full fp32
                W=torch.bmm(V.transpose(1,2),C); W=torch.bmm(T.transpose(1,2),W)
                H[:,j:,j+jb:]=C-torch.bmm(V,W)
            elif ill_any:
                # TF32-safe ill (rankdef/clustered/nearcollinear): TF32 = 1.7x over fp32
                W1=torch.empty(b,jb,Tw,device=dev); W=torch.empty(b,jb,Tw,device=dev)
                _f16.trail_t(V, T, C, W1, W)
            else:
                # well-conditioned: FAST_16F zero-copy (the breakthrough)
                W1=torch.empty(b,jb,Tw,device=dev); W=torch.empty(b,jb,Tw,device=dev)
                _f16.trail(V, T, C, W1, W)
    return H,tau


def custom_kernel(data: input_t) -> output_t:
    b,n,_=data.shape
    if n>=2048: return torch.geqrf(data)
    if 128<=n<=256: return _fused_qr(data)
    return _triton_qr(data)
scrolls · 226 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