Skip to content
KernelIndex
Search⌘K

submission 798785

Diablo! · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_fused.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-798785?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
54.3ms
#397 of 515
2026-06-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c2a6a409155c0899c6417da69f2d6e987853f9a0f648fae6327c10cabf585fb3
license declaredunknown
license concludedunknown
authorsDiablo!
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float sm[];

Kernel source

submission_fused.py304 lines
import torch
from task import input_t, output_t

torch.backends.cuda.matmul.allow_tf32 = False

_NB = 16
_eye_cache: dict = {}

_CPP = "void geqr2_panel(at::Tensor H, at::Tensor tau, at::Tensor T_out, int64_t col, int64_t w, int64_t nt);"

_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>

__global__ void geqr2_panel_kernel(float* __restrict__ H, float* __restrict__ tau, float* __restrict__ T_out,
                                   int B, int m, int n, int col, int w, int ldp) {
    const int b = blockIdx.x;
    if (b >= B) return;
    const int p = m - col;
    extern __shared__ float sm[];
    float* P   = sm;                      // panel [p][ldp], ldp = w+1
    float* red = sm + (size_t)p * ldp;    // [nt]
    
    __shared__ double warp_red[32];
    __shared__ double T_sm[32 * 32];
    __shared__ double y_sm[32];

    const int tid = threadIdx.x;
    const int nt = blockDim.x;
    const long base = (long)b * m * n;
    const int ldt = (m < n ? m : n);
    const int nwarps = nt >> 5;

    for (int idx = tid; idx < p * w; idx += nt) {
        const int i = idx / w;
        const int jj = idx - i * w;
        P[i * ldp + jj] = H[base + (long)(col + i) * n + (col + jj)];
    }
    __syncthreads();

    for (int j = 0; j < w; ++j) {
        float partial = 0.f;
        for (int i = j + 1 + tid; i < p; i += nt) {
            float v = P[i * ldp + j];
            partial += v * v;
        }
        partial += __shfl_xor_sync(0xffffffff, partial, 16);
        partial += __shfl_xor_sync(0xffffffff, partial, 8);
        partial += __shfl_xor_sync(0xffffffff, partial, 4);
        partial += __shfl_xor_sync(0xffffffff, partial, 2);
        partial += __shfl_xor_sync(0xffffffff, partial, 1);
        int lane = tid & 31;
        int wid = tid >> 5;
        if (lane == 0) red[wid] = partial;
        __syncthreads();
        for (int s = nwarps >> 1; s > 0; s >>= 1) {
            if (tid < s) red[tid] += red[tid + s];
            __syncthreads();
        }
        const float tailsq = red[0];

        const float alpha = P[j * ldp + j];
        float beta, tauj, invden;
        if (tailsq > 0.f) {
            const float normx = sqrtf(alpha * alpha + tailsq);
            beta = (alpha >= 0.f) ? -normx : normx;
            tauj = (beta - alpha) / beta;
            invden = 1.f / (alpha - beta);
        } else {
            beta = alpha; tauj = 0.f; invden = 0.f;
        }
        for (int i = j + 1 + tid; i < p; i += nt) P[i * ldp + j] *= invden;
        if (tid == 0) {
            P[j * ldp + j] = beta;
            tau[(long)b * ldt + (col + j)] = tauj;
            T_sm[j * w + j] = (double)tauj;
        }
        __syncthreads();

        for (int c = j + 1; c < w; ++c) {
            float pd = 0.f;
            for (int i = j + tid; i < p; i += nt) {
                float vi = (i == j) ? 1.f : P[i * ldp + j];
                pd += vi * P[i * ldp + c];
            }
            pd += __shfl_xor_sync(0xffffffff, pd, 16);
            pd += __shfl_xor_sync(0xffffffff, pd, 8);
            pd += __shfl_xor_sync(0xffffffff, pd, 4);
            pd += __shfl_xor_sync(0xffffffff, pd, 2);
            pd += __shfl_xor_sync(0xffffffff, pd, 1);
            if (lane == 0) red[wid] = pd;
            __syncthreads();
            for (int s = nwarps >> 1; s > 0; s >>= 1) {
                if (tid < s) red[tid] += red[tid + s];
                __syncthreads();
            }
            const float coef = tauj * red[0];
            for (int i = j + tid; i < p; i += nt) {
                float vi = (i == j) ? 1.f : P[i * ldp + j];
                P[i * ldp + c] -= coef * vi;
            }
            __syncthreads();
        }
        
        if (j > 0) {
            for (int i = 0; i < j; ++i) {
                double pd = 0.0;
                for (int row = j + tid; row < p; row += nt) {
                    double vi = (row == i) ? 1.0 : (double)P[row * ldp + i];
                    double vj = (row == j) ? 1.0 : (double)P[row * ldp + j];
                    pd += vi * vj;
                }
                pd += __shfl_xor_sync(0xffffffff, pd, 16);
                pd += __shfl_xor_sync(0xffffffff, pd, 8);
                pd += __shfl_xor_sync(0xffffffff, pd, 4);
                pd += __shfl_xor_sync(0xffffffff, pd, 2);
                pd += __shfl_xor_sync(0xffffffff, pd, 1);
                if (lane == 0) warp_red[wid] = pd;
                __syncthreads();
                for (int s = nwarps >> 1; s > 0; s >>= 1) {
                    if (tid < s) warp_red[tid] += warp_red[tid + s];
                    __syncthreads();
                }
                if (tid == 0) {
                    y_sm[i] = warp_red[0];
                }
                __syncthreads();
            }
            if (tid < j) {
                double sum = 0.0;
                for (int k = tid; k < j; ++k) {
                    sum += T_sm[tid * w + k] * y_sm[k];
                }
                T_sm[tid * w + j] = -(double)tauj * sum;
            }
            __syncthreads();
        }
    }

    for (int idx = tid; idx < p * w; idx += nt) {
        const int i = idx / w;
        const int jj = idx - i * w;
        H[base + (long)(col + i) * n + (col + jj)] = P[i * ldp + jj];
    }
    
    // Write out T
    for (int idx = tid; idx < w * w; idx += nt) {
        const int i = idx / w;
        const int jj = idx - i * w;
        if (i <= jj) {
            T_out[(long)b * w * w + i * w + jj] = T_sm[i * w + jj];
        } else {
            T_out[(long)b * w * w + i * w + jj] = 0.f;
        }
    }
}

void geqr2_panel(at::Tensor H, at::Tensor tau, at::Tensor T_out, int64_t col, int64_t w, int64_t nt) {
    const int B = H.size(0), m = H.size(1), n = H.size(2);
    const int p = m - (int)col;
    const int ldp = (int)w + 1;
    const size_t shmem = ((size_t)p * ldp + nt + w * w + w) * sizeof(float);
    cudaFuncSetAttribute(geqr2_panel_kernel,
                         cudaFuncAttributeMaxDynamicSharedMemorySize, 226 * 1024);
    geqr2_panel_kernel<<<B, (int)nt, shmem>>>(
        H.data_ptr<float>(), tau.data_ptr<float>(), T_out.data_ptr<float>(), B, m, n, (int)col, (int)w, ldp);
    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "geqr2_panel launch failed: ", cudaGetErrorString(err));
}
"""

_ext = None
try:
    from torch.utils.cpp_extension import load_inline
    _ext = load_inline(
        name="qr_panel_v2",
        cpp_sources=_CPP,
        cuda_sources=_CUDA,
        functions=["geqr2_panel"],
        extra_cuda_cflags=["-O3"],
        verbose=False,
    )
except Exception:
    _ext = None


def _eye(w: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
    key = (w, device, dtype)
    e = _eye_cache.get(key)
    if e is None:
        e = torch.eye(w, device=device, dtype=dtype)
        _eye_cache[key] = e
    return e


def _factor_panel_eager(H: torch.Tensor, tau: torch.Tensor, col: int, w: int, m: int) -> None:
    for jj in range(w):
        j = col + jj
        x = H[:, j:m, j]
        alpha = x[:, 0]
        tail = x[:, 1:]
        tailsq = (tail * tail).sum(1)
        normx = torch.sqrt(alpha * alpha + tailsq)
        beta = -torch.copysign(normx, alpha)
        reflect = tailsq > 0
        ones = torch.ones_like(beta)
        zero = torch.zeros_like(beta)
        beta_safe = torch.where(reflect, beta, ones)
        tau_j = torch.where(reflect, (beta - alpha) / beta_safe, zero)
        denom = alpha - beta
        denom_safe = torch.where(reflect, denom, ones)
        invden = torch.where(reflect, 1.0 / denom_safe, zero)
        vtail = tail * invden[:, None]
        H[:, j, j] = torch.where(reflect, beta, alpha)
        H[:, j + 1:m, j] = vtail
        tau[:, j] = tau_j
        if jj + 1 < w:
            sub = H[:, j:m, j + 1:col + w]
            wv = sub[:, 0, :] + torch.einsum('bp,bpc->bc', vtail, sub[:, 1:, :])
            coef = tau_j[:, None] * wv
            sub[:, 0, :] -= coef
            sub[:, 1:, :] -= vtail[:, :, None] * coef[:, None, :]


def _apply_block(H: torch.Tensor, tau: torch.Tensor, T: torch.Tensor, col: int, w: int, m: int, n: int) -> None:
    cstart = col + w
    p = m - col
    V = H[:, col:m, col:cstart]
    Vtop = torch.tril(V[:, :w, :], -1) + _eye(w, H.device, H.dtype)
    Vmat = torch.cat([Vtop, V[:, w:, :]], dim=1) if p > w else Vtop

    C = H[:, col:m, cstart:n]
    Wm = Vmat.transpose(1, 2) @ C
    if T is not None:
        X = T.transpose(1, 2) @ Wm
    else:
        tcol = tau[:, col:cstart]
        G = Vmat.transpose(1, 2) @ Vmat
        safe = torch.where(tcol != 0, tcol, torch.ones_like(tcol))
        invtau = torch.where(tcol != 0, 1.0 / safe, torch.full_like(tcol, 1e20))
        Tinv = torch.triu(G, 1) + torch.diag_embed(invtau)
        X = torch.linalg.solve_triangular(Tinv.transpose(1, 2), Wm, upper=False)
        
    C.baddbmm_(Vmat, X, alpha=-1, beta=1)


def _qr_core(A: torch.Tensor, use_cuda: bool) -> output_t:
    B, m, n = A.shape
    H = A.clone()
    k = min(m, n)
    tau = torch.zeros(B, k, device=A.device, dtype=A.dtype)
    col = 0
    while col < k:
        p = m - col
        nb = _nb_for_panel(p, n) if p <= min(128, n // 2) else _nb_for(n)
        w = min(nb, k - col)
        
        T_out = None
        if use_cuda:
            T_out = torch.zeros(B, w, w, device=A.device, dtype=A.dtype)
            w32 = (p + 31) // 32
            wp2 = 1 if w32 <= 1 else 2 if w32 <= 2 else 4 if w32 <= 4 else 8
            nt = 32 * wp2
            _ext.geqr2_panel(H, tau, T_out, col, w, nt)
        else:
            _factor_panel_eager(H, tau, col, w, m)
            
        if col + w < n:
            _apply_block(H, tau, T_out, col, w, m, n)
        col += w
    return H, tau


def _nb_for(n: int) -> int:
    fit = 53000 // n - 1
    if fit < 1:
        fit = 1
    base = 16 if n <= 192 else 24 if n <= 384 else 32
    return min(base, fit)


def _nb_for_panel(p: int, n: int) -> int:
    fit = 53000 // p - 1
    if fit < 1:
        fit = 1
    base = 16 if n <= 192 else 24 if n <= 384 else 32
    return min(base, fit)


@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
    A = data
    B, n, _ = A.shape
    if 64 < n <= 2048:
        use_cuda = _ext is not None and A.is_cuda and A.dtype == torch.float32
        try:
            return _qr_core(A, use_cuda)
        except Exception:
            try:
                return _qr_core(A, False)
            except Exception:
                return torch.geqrf(A)
    return torch.geqrf(A)
scrolls · 304 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