Skip to content
KernelIndex
Search⌘K

submission 835885

FlamingoPg · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-835885?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
130.5ms
#460 of 515
2026-06-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:90a371dbc9a8ebc2c0144e10456b41e9e7ecc2bdae72a739a6d443d9d386874d
license declaredunknown
license concludedunknown
authorsFlamingoPg
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float sh[];

Kernel source

submission.py247 lines
"""
Fully fused blocked Householder QR in a single CUDA kernel.
Panel QR + T construction + WY trailing update — all in one kernel launch per block.
No Python loops over columns. No torch.bmm overhead.
"""
import torch
from task import input_t, output_t

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

__device__ __forceinline__ float warp_sum(float v) {
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) v += __shfl_down_sync(0xffffffff, v, o);
    return v;
}

__device__ float block_sum(float v, float* sh) {
    int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
    v = warp_sum(v);
    if (lane == 0) sh[wid] = v;
    __syncthreads();
    int nw = blockDim.x >> 5;
    v = (threadIdx.x < nw) ? sh[threadIdx.x] : 0.f;
    if (wid == 0) v = warp_sum(v);
    return v;
}

/*
 * Fused blocked Householder QR kernel.
 * Each thread block handles one matrix.
 * Shared memory: v[n] + scratch[8] + VtV[nb*nb] + tau_panel[nb]
 *
 * For each block of nb columns:
 *   1. Panel QR (sequential columns, parallel rows)
 *   2. Build V and T (DLARFT)
 *   3. WY trailing update: trailing -= V @ T @ V^T @ trailing
 */
__global__ void __launch_bounds__(256)
fused_blocked_qr(
    float* __restrict__ A,   // (batch, n, n) in-place
    float* __restrict__ tau,  // (batch, n)
    const int n,
    const int nb
) {
    const int bid = blockIdx.x;
    const int tid = threadIdx.x;
    const int NT  = blockDim.x;

    float* a = A + (long long)bid * n * n;
    float* t = tau + (long long)bid * n;

    extern __shared__ float sh[];
    float* v_sh   = sh;                    // n floats
    float* r_sh   = sh + n;                // 8 floats
    float* vt     = sh + n + 8;            // nb floats (tau_panel)
    float* VtV_sh = sh + n + 8 + nb;       // nb*nb floats

    for (int jb_start = 0; jb_start < n; jb_start += nb) {
        int jb = (jb_start + nb <= n) ? nb : (n - jb_start);
        int panel_rows = n - jb_start;

        // Zero VtV
        for (int i = tid; i < jb * jb; i += NT) VtV_sh[i] = 0.f;
        __syncthreads();

        // ===== Panel QR =====
        for (int k = 0; k < jb; k++) {
            int col = jb_start + k;
            int rows = n - col;

            // Norm
            float ls = 0.f;
            for (int i = tid; i < rows; i += NT) {
                float val = a[(long long)(col + i) * n + col];
                ls += val * val;
            }
            ls = block_sum(ls, r_sh);
            float norm_x = sqrtf(ls);

            // Householder params
            float x0 = a[(long long)col * n + col];
            float sgn = (x0 >= 0.f) ? 1.f : -1.f;
            float alpha = -sgn * norm_x;
            float tau_k = (norm_x < 1e-30f) ? 0.f : (alpha - x0) / alpha;
            if (tid == 0) { t[col] = tau_k; vt[k] = tau_k; }

            // Compute v (v[0]=1, v[i>0]=a[col+i,col]/v0)
            float v0 = x0 - alpha;
            float inv = (fabsf(v0) > 1e-30f) ? (1.f / v0) : 0.f;
            for (int i = tid; i < rows; i += NT) {
                float vi = (i == 0) ? 1.f : a[(long long)(col + i) * n + col] * inv;
                v_sh[i] = vi;
                if (i > 0) a[(long long)(col + i) * n + col] = vi;
            }
            if (tid == 0) a[(long long)col * n + col] = alpha;
            __syncthreads();

            // Store VtV[:, k] = V^T @ v_k (dot products with previous vectors)
            if (k > 0) {
                for (int prev = tid; prev < k; prev += NT) {
                    int pcol = jb_start + prev;
                    int poffset = col - pcol;  // offset into v_prev
                    // v_prev[poffset..] overlaps with v_sh[0..]
                    // v_prev[poffset] = a[col * n + pcol] (the Householder vector component)
                    float dot = 0.f;
                    for (int j = 0; j < rows && (poffset + j) < (n - pcol); j++) {
                        float vp;
                        if (poffset + j == 0) vp = 1.f;  // v_prev[0] = 1
                        else vp = a[(long long)(pcol + poffset + j) * n + pcol];
                        dot += vp * v_sh[j];
                    }
                    VtV_sh[prev * jb + k] = dot;
                }
                __syncthreads();
            }

            // Build T[:, k] using DLARFT recurrence
            if (k > 0 && tau_k != 0.f) {
                // T[:k, k] = T[:k, :k] @ (-tau_k * VtV[:k, k])
                // First, scale VtV by -tau_k
                for (int i = tid; i < k; i += NT)
                    VtV_sh[i * jb + k] *= -tau_k;
                __syncthreads();
                // Then T[i, k] = sum_{j=i}^{k-1} T[i, j] * VtV_sh[j, k]
                // Note: T is stored in VtV_sh (reusing upper triangle)
                for (int row = tid; row < k; row += NT) {
                    float sum = 0.f;
                    for (int j = row; j < k; j++)
                        sum += VtV_sh[row * jb + j] * VtV_sh[j * jb + k];
                    VtV_sh[row * jb + k] = sum;
                }
                __syncthreads();
            }
            if (tid == 0) VtV_sh[k * jb + k] = tau_k;

            // Apply reflection to remaining PANEL columns
            if (tau_k != 0.f) {
                for (int j = col + 1 + tid; j < jb_start + jb; j += NT) {
                    float dot = v_sh[0] * a[(long long)col * n + j];
                    for (int i = 1; i < rows; i++)
                        dot += v_sh[i] * a[(long long)(col + i) * n + j];
                    dot *= tau_k;
                    a[(long long)col * n + j] -= dot;
                    for (int i = 1; i < rows; i++)
                        a[(long long)(col + i) * n + j] -= dot * v_sh[i];
                }
            }
            __syncthreads();
        }

        // ===== WY Trailing Update =====
        // trailing -= V @ T @ V^T @ trailing
        // trailing is A[jb_start:, jb_start+jb:] = a[row*n + col] for row>=jb_start, col>=jb_start+jb
        int trail_cols = n - jb_start - jb;
        if (trail_cols <= 0) continue;

        // We have T in VtV_sh (upper triangular, jb x jb)
        // V is implicit: V[row - jb_start, k] = 1 if row == jb_start+k, else a[row*n + jb_start+k]
        //
        // Step 1: Compute VtA = V^T @ trailing  (jb x trail_cols)
        // Step 2: Compute Z = T @ VtA           (jb x trail_cols)
        // Step 3: Compute V @ Z                 (panel_rows x trail_cols)
        // Step 4: trailing -= V @ Z
        //
        // We'll compute column by column of trailing to minimize shared memory.
        // For each trail column c:
        //   VtA[:, c] = V^T @ trailing[:, c]  (jb dot products)
        //   Z[:, c] = T @ VtA[:, c]           (jb x jb @ jb)
        //   trailing[:, c] -= V @ Z[:, c]      (panel_rows elements)

        // Process trail columns in tiles
        const int TILE = 1;  // Process one column at a time for simplicity
        for (int c_base = tid; c_base < trail_cols; c_base += NT) {
            int c = jb_start + jb + c_base;  // actual column index

            // Step 1: VtA[k, c_base] = V[:, k]^T @ trailing[:, c]
            // V[:, k] has 1 at row k and a[row*n + jb_start+k] for row > k
            float VtA_local[64];  // max nb=64
            for (int k = 0; k < jb; k++) {
                int vrow = jb_start + k;
                float dot = a[(long long)vrow * n + c];  // V[k, k]=1 * trailing[k, c]
                for (int r = k + 1; r < panel_rows; r++) {
                    float v_val = a[(long long)(jb_start + r) * n + vrow];
                    dot += v_val * a[(long long)(jb_start + r) * n + c];
                }
                VtA_local[k] = dot;
            }

            // Step 2: Z = T @ VtA (T is upper triangular in VtV_sh)
            float Z_local[64];
            for (int i = 0; i < jb; i++) {
                float sum = 0.f;
                for (int j = i; j < jb; j++)
                    sum += VtV_sh[i * jb + j] * VtA_local[j];
                Z_local[i] = sum;
            }

            // Step 3 & 4: trailing[:, c] -= V @ Z
            for (int r = 0; r < panel_rows; r++) {
                float vz = 0.f;
                for (int k = 0; k < jb; k++) {
                    float v_val;
                    if (r == k) v_val = 1.f;
                    else if (r > k) v_val = a[(long long)(jb_start + r) * n + (jb_start + k)];
                    else v_val = 0.f;
                    vz += v_val * Z_local[k];
                }
                a[(long long)(jb_start + r) * n + c] -= vz;
            }
        }
        __syncthreads();
    }
}

void run_fused_qr(torch::Tensor A, torch::Tensor tau, int nb) {
    int batch = A.size(0), n = A.size(1);
    int threads = 256;
    int smem = (n + 8 + nb + nb * nb) * sizeof(float);
    fused_blocked_qr<<<batch, threads, smem>>>(
        A.data_ptr<float>(), tau.data_ptr<float>(), n, nb
    );
}
"""

_CPP = "void run_fused_qr(torch::Tensor A, torch::Tensor tau, int nb);"

try:
    from torch.utils.cpp_extension import load_inline
    _ext = load_inline(name="fused_qr", cpp_sources=_CPP, cuda_sources=_CUDA, verbose=False, with_cuda=True)
    _OK = True
except Exception:
    _OK = False


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if not _OK or n <= 352:
        return torch.geqrf(data)

    H = data.clone().contiguous()
    tau = torch.zeros(batch, n, device=data.device, dtype=torch.float32)
    nb = 64  # fits in shared memory: 64*64*4 = 16KB + n*4 + overhead
    _ext.run_fused_qr(H, tau, nb)
    return H, tau
scrolls · 247 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