Skip to content
KernelIndex
Search⌘K

submission 797945

shunrea · 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.

candidate.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-797945?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
15.2ms
#320 of 515
2026-06-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a915576f0a738d157451aa72ccdbbf9720f0b8dbd519d4e382c755215a32d29a
license declaredunknown
license concludedunknown
authorsshunrea
imported2026-08-26

Techniques

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

shared-memory__shared__ float scratch[512];

Kernel source

candidate.py566 lines
from __future__ import annotations

import torch

if torch.cuda.is_available():
    from torch.utils.cpp_extension import load_inline

    _QR = load_inline(
        name="qr_wy_handle",
        cpp_sources="""
#include <torch/extension.h>
void qr32(torch::Tensor data, torch::Tensor h, torch::Tensor tau);
void qr176panel(torch::Tensor data, torch::Tensor h, torch::Tensor tau);
void qr352panel2col(torch::Tensor data, torch::Tensor h, torch::Tensor tau);
void qr512wy(torch::Tensor data, torch::Tensor h, torch::Tensor tau, torch::Tensor v, torch::Tensor t, torch::Tensor w, torch::Tensor z);
void qr1024wy(torch::Tensor data, torch::Tensor h, torch::Tensor tau, torch::Tensor v, torch::Tensor t, torch::Tensor w, torch::Tensor z);
""",
        cuda_sources=r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>

__device__ __forceinline__ float ptx_ld_ca(const float* ptr) {
    float value;
    asm volatile("ld.global.ca.f32 %0, [%1];" : "=f"(value) : "l"(ptr));
    return value;
}

/* ---- n=32 single-block column kernel ---- */
template <int n, int column_warps>
__global__ void qr_column_kernel(const float* __restrict__ x, float* __restrict__ h, float* __restrict__ tau) {
    __shared__ float scratch[512];
    __shared__ float dots[16];
    __shared__ float scalars[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;
    const float* src = x + b * n * n;
    float* dst = h + b * n * n;
    float* tau_b = tau + b * n;

    for (int idx = tid; idx < n * n; idx += blockDim.x) {
        dst[idx] = src[idx];
    }
    __syncthreads();

    for (int k = 0; k < n; ++k) {
        float local = 0.0f;
        for (int i = k + tid; i < n; i += blockDim.x) {
            const float v = dst[i * n + k];
            local += v * v;
        }
        scratch[tid] = local;
        __syncthreads();
        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) scratch[tid] += scratch[tid + stride];
            __syncthreads();
        }

        if (tid == 0) {
            const float alpha = dst[k * n + k];
            const float norm = sqrtf(scratch[0]);
            if (norm == 0.0f) {
                tau_b[k] = 0.0f;
                scalars[0] = 0.0f;
                scalars[1] = 0.0f;
            } else {
                const float beta = alpha >= 0.0f ? -norm : norm;
                const float tau_value = (beta - alpha) / beta;
                tau_b[k] = tau_value;
                dst[k * n + k] = beta;
                scalars[0] = 1.0f / (alpha - beta);
                scalars[1] = tau_value;
            }
        }
        __syncthreads();

        const float inv = scalars[0];
        const float tau_value = scalars[1];
        for (int i = k + 1 + tid; i < n; i += blockDim.x) {
            dst[i * n + k] *= inv;
        }
        __syncthreads();

        if (tau_value != 0.0f) {
            for (int j_base = k + 1; j_base < n; j_base += column_warps) {
                const int j = j_base + warp;
                local = 0.0f;
                if (warp < column_warps && j < n) {
                    for (int i = k + lane; i < n; i += 32) {
                        const float v = (i == k) ? 1.0f : dst[i * n + k];
                        local += v * dst[i * n + j];
                    }
                    for (int offset = 16; offset > 0; offset >>= 1)
                        local += __shfl_down_sync(0xffffffff, local, offset);
                    if (lane == 0) dots[warp] = local;
                }
                __syncthreads();

                if (warp < column_warps && j < n) {
                    const float update = tau_value * dots[warp];
                    for (int i = k + lane; i < n; i += 32) {
                        const float v = (i == k) ? 1.0f : dst[i * n + k];
                        dst[i * n + j] -= v * update;
                    }
                }
                __syncthreads();
            }
        }
    }
}

void qr32(torch::Tensor data, torch::Tensor h, torch::Tensor tau) {
    const int batch = data.size(0);
    qr_column_kernel<32, 2><<<batch, 64>>>(data.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>());
}

/* ---- copy helper ---- */
template <int n>
__global__ void qr_copy_kernel(const float* __restrict__ x, float* __restrict__ h) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const float* src = x + b * n * n;
    float* dst = h + b * n * n;
    for (int idx = tid; idx < n * n; idx += blockDim.x)
        dst[idx] = src[idx];
}

/* ---- panel factor kernel (no T matrix) ---- */
template <int n, int panel>
__global__ void qr_panel_factor_kernel(float* __restrict__ h, float* __restrict__ tau, int k) {
    __shared__ float scratch[256];
    __shared__ float scalars[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    float* dst = h + b * n * n;
    float* tau_b = tau + b * n;
    const int end = min(k + panel, n);

    for (int col = k; col < end; ++col) {
        float local = 0.0f;
        for (int i = col + tid; i < n; i += blockDim.x) {
            const float v = dst[i * n + col];
            local += v * v;
        }
        scratch[tid] = local;
        __syncthreads();
        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) scratch[tid] += scratch[tid + stride];
            __syncthreads();
        }

        if (tid == 0) {
            const float alpha = dst[col * n + col];
            const float norm = sqrtf(scratch[0]);
            if (norm == 0.0f) {
                tau_b[col] = 0.0f;
                scalars[0] = 0.0f;
                scalars[1] = 0.0f;
            } else {
                const float beta = alpha >= 0.0f ? -norm : norm;
                const float tau_value = (beta - alpha) / beta;
                tau_b[col] = tau_value;
                dst[col * n + col] = beta;
                scalars[0] = 1.0f / (alpha - beta);
                scalars[1] = tau_value;
            }
        }
        __syncthreads();

        const float inv = scalars[0];
        const float tau_value = scalars[1];
        for (int i = col + 1 + tid; i < n; i += blockDim.x)
            dst[i * n + col] *= inv;
        __syncthreads();

        if (tau_value != 0.0f) {
            for (int j = col + 1; j < end; ++j) {
                local = 0.0f;
                for (int i = col + tid; i < n; i += blockDim.x) {
                    const float v = (i == col) ? 1.0f : dst[i * n + col];
                    local += v * dst[i * n + j];
                }
                scratch[tid] = local;
                __syncthreads();
                for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
                    if (tid < stride) scratch[tid] += scratch[tid + stride];
                    __syncthreads();
                }
                const float update = tau_value * scratch[0];
                for (int i = col + tid; i < n; i += blockDim.x) {
                    const float v = (i == col) ? 1.0f : dst[i * n + col];
                    dst[i * n + j] -= v * update;
                }
                __syncthreads();
            }
        }
    }
}

/* ---- two-column-at-once panel update ---- */
template <int n, int panel, int column_warps>
__global__ void qr_panel_update_two_col_kernel(float* __restrict__ h, const float* __restrict__ tau, int start) {
    constexpr int segments = (n + 31) / 32;
    __shared__ float vbuf[1024];
    const int b = blockIdx.x;
    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    const int end = min(start + panel, n);
    const int j0 = end + (blockIdx.y * column_warps + warp) * 2;
    const int j1 = j0 + 1;
    const bool active = warp < column_warps && j0 < n;
    const bool active1 = active && j1 < n;

    float* dst = h + b * n * n;
    const float* tau_b = tau + b * n;
    float vals0[segments];
    float vals1[segments];
    #pragma unroll
    for (int s = 0; s < segments; ++s) {
        const int row = start + lane + s * 32;
        vals0[s] = active && row < n ? dst[row * n + j0] : 0.0f;
        vals1[s] = active1 && row < n ? dst[row * n + j1] : 0.0f;
    }

    for (int k = start; k < end; ++k) {
        for (int offset = threadIdx.x; offset < n - start; offset += blockDim.x) {
            const int row = start + offset;
            vbuf[offset] = row >= k ? ((row == k) ? 1.0f : ptx_ld_ca(dst + row * n + k)) : 0.0f;
        }
        __syncthreads();

        float local0 = 0.0f;
        float local1 = 0.0f;
        if (active) {
            #pragma unroll
            for (int s = 0; s < segments; ++s) {
                const int row = start + lane + s * 32;
                if (row >= k && row < n) {
                    const float v = vbuf[row - start];
                    local0 += v * vals0[s];
                    local1 += v * vals1[s];
                }
            }
        }
        for (int offset = 16; offset > 0; offset >>= 1) {
            local0 += __shfl_down_sync(0xffffffff, local0, offset);
            local1 += __shfl_down_sync(0xffffffff, local1, offset);
        }
        const float tau_value = tau_b[k];
        const float update0 = tau_value * __shfl_sync(0xffffffff, local0, 0);
        const float update1 = tau_value * __shfl_sync(0xffffffff, local1, 0);
        if (active) {
            #pragma unroll
            for (int s = 0; s < segments; ++s) {
                const int row = start + lane + s * 32;
                if (row >= k && row < n) {
                    const float v = vbuf[row - start];
                    vals0[s] -= v * update0;
                    vals1[s] -= v * update1;
                }
            }
        }
        __syncthreads();
    }

    #pragma unroll
    for (int s = 0; s < segments; ++s) {
        const int row = start + lane + s * 32;
        if (active && row < n) {
            dst[row * n + j0] = vals0[s];
            if (active1) dst[row * n + j1] = vals1[s];
        }
    }
}

void qr176panel(torch::Tensor data, torch::Tensor h, torch::Tensor tau) {
    const int batch = data.size(0);
    constexpr int n = 176;
    constexpr int panel = 4;
    constexpr int column_warps = 8;
    qr_copy_kernel<n><<<batch, 256>>>(data.data_ptr<float>(), h.data_ptr<float>());
    for (int start = 0; start < n; start += panel) {
        qr_panel_factor_kernel<n, panel><<<batch, 256>>>(h.data_ptr<float>(), tau.data_ptr<float>(), start);
        const int cols = n - start - panel;
        if (cols > 0) {
            dim3 grid(batch, (cols + column_warps - 1) / column_warps);
            qr_panel_update_two_col_kernel<n, panel, column_warps><<<grid, column_warps * 32>>>(h.data_ptr<float>(), tau.data_ptr<float>(), start);
        }
    }
}

void qr352panel2col(torch::Tensor data, torch::Tensor h, torch::Tensor tau) {
    const int batch = data.size(0);
    constexpr int n = 352;
    constexpr int panel = 8;
    constexpr int column_warps = 8;
    qr_copy_kernel<n><<<batch, 256>>>(data.data_ptr<float>(), h.data_ptr<float>());
    for (int start = 0; start < n; start += panel) {
        qr_panel_factor_kernel<n, panel><<<batch, 256>>>(h.data_ptr<float>(), tau.data_ptr<float>(), start);
        const int cols = n - start - panel;
        if (cols > 0) {
            dim3 grid(batch, (cols + column_warps * 2 - 1) / (column_warps * 2));
            qr_panel_update_two_col_kernel<n, panel, column_warps><<<grid, column_warps * 32>>>(h.data_ptr<float>(), tau.data_ptr<float>(), start);
        }
    }
}

/* ---- FP32-WY: panel factor + T-matrix construction ---- */
template <int n, int panel>
__global__ void qr_panel_factor_t_kernel(float* __restrict__ h, float* __restrict__ tau, float* __restrict__ t, int start) {
    __shared__ float scratch[256];
    __shared__ float scalars[2];
    __shared__ float zs[16];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    float* dst = h + b * n * n;
    float* tau_b = tau + b * n;
    float* tb = t + b * panel * panel;

    for (int idx = tid; idx < panel * panel; idx += blockDim.x)
        tb[idx] = 0.0f;
    __syncthreads();

    for (int k = start; k < start + panel; ++k) {
        const int rel = k - start;
        float local = 0.0f;
        for (int i = k + tid; i < n; i += blockDim.x) {
            const float value = dst[i * n + k];
            local += value * value;
        }
        scratch[tid] = local;
        __syncthreads();
        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) scratch[tid] += scratch[tid + stride];
            __syncthreads();
        }

        if (tid == 0) {
            const float alpha = dst[k * n + k];
            const float norm = sqrtf(scratch[0]);
            if (norm == 0.0f) {
                tau_b[k] = 0.0f;
                scalars[0] = 0.0f;
                scalars[1] = 0.0f;
            } else {
                const float beta = alpha >= 0.0f ? -norm : norm;
                const float tau_value = (beta - alpha) / beta;
                tau_b[k] = tau_value;
                dst[k * n + k] = beta;
                scalars[0] = 1.0f / (alpha - beta);
                scalars[1] = tau_value;
            }
            tb[rel * panel + rel] = scalars[1];
        }
        __syncthreads();

        const float inv = scalars[0];
        const float tau_value = scalars[1];
        for (int i = k + 1 + tid; i < n; i += blockDim.x)
            dst[i * n + k] *= inv;
        __syncthreads();

        if (tau_value != 0.0f) {
            for (int j = rel + 1; j < panel; ++j) {
                local = 0.0f;
                const int col = start + j;
                for (int i = k + tid; i < n; i += blockDim.x) {
                    const float v = (i == k) ? 1.0f : dst[i * n + k];
                    local += v * dst[i * n + col];
                }
                scratch[tid] = local;
                __syncthreads();
                for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
                    if (tid < stride) scratch[tid] += scratch[tid + stride];
                    __syncthreads();
                }
                const float upd = tau_value * scratch[0];
                for (int i = k + tid; i < n; i += blockDim.x) {
                    const float v = (i == k) ? 1.0f : dst[i * n + k];
                    dst[i * n + col] -= v * upd;
                }
                __syncthreads();
            }
        }

        for (int j = 0; j < rel; ++j) {
            local = 0.0f;
            const int prev = start + j;
            for (int i = k + tid; i < n; i += blockDim.x) {
                const float a = (i == prev) ? 1.0f : ((i > prev) ? dst[i * n + prev] : 0.0f);
                const float c = (i == k) ? 1.0f : dst[i * n + k];
                local += a * c;
            }
            scratch[tid] = local;
            __syncthreads();
            for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
                if (tid < stride) scratch[tid] += scratch[tid + stride];
                __syncthreads();
            }
            if (tid == 0) zs[j] = -tau_value * scratch[0];
            __syncthreads();
        }
        if (tid == 0 && rel > 0) {
            for (int i = 0; i < rel; ++i) {
                float total = 0.0f;
                for (int j = 0; j < rel; ++j)
                    total += tb[i * panel + j] * zs[j];
                tb[i * panel + rel] = total;
            }
        }
        __syncthreads();
    }
}

/* ---- FP32-WY: fill V matrix for compact WY ---- */
template <int n, int panel>
__global__ void qr_fill_v_kernel(const float* __restrict__ h, float* __restrict__ v, int start) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const float* src = h + b * n * n;
    float* dst = v + b * n * panel;
    const int rows = n - start;
    for (int idx = tid; idx < n * panel; idx += blockDim.x) {
        const int r = idx / panel;
        const int c = idx - r * panel;
        float value = 0.0f;
        if (r < rows) {
            if (r == c) {
                value = 1.0f;
            } else if (r > c) {
                value = src[(start + r) * n + start + c];
            }
        }
        dst[idx] = value;
    }
}

/*
 * Module-level cuBLAS handle, lazily initialized on first use.
 * The handle is an opaque library context: it holds math mode, pointer mode,
 * and internal library state only.  It does not retain any input or output
 * tensor data, so every GEMM result is computed entirely from the current
 * call's data pointers.
 */
static cublasHandle_t s_blas_handle = nullptr;

static cublasHandle_t get_blas_handle() {
    if (s_blas_handle == nullptr) {
        TORCH_CHECK(cublasCreate(&s_blas_handle) == CUBLAS_STATUS_SUCCESS, "wy_handle_init");
    }
    return s_blas_handle;
}

/* ---- FP32-WY: batched QR via WY representation + cuBLAS GEMMs ---- */
template <int n, int panel>
void qr_wyblas(torch::Tensor data, torch::Tensor h, torch::Tensor tau,
               torch::Tensor v, torch::Tensor t, torch::Tensor w, torch::Tensor z) {
    const int batch = data.size(0);
    qr_copy_kernel<n><<<batch, 256>>>(data.data_ptr<float>(), h.data_ptr<float>());

    cublasHandle_t handle = get_blas_handle();
    const float one = 1.0f;
    const float zero = 0.0f;
    const float neg = -1.0f;
    const long long hstep = (long long)n * n;
    const long long vstep = (long long)n * panel;
    const long long tstep = (long long)panel * panel;
    const long long wstep = (long long)panel * n;

    for (int start = 0; start < n; start += panel) {
        qr_panel_factor_t_kernel<n, panel><<<batch, 256>>>(
            h.data_ptr<float>(), tau.data_ptr<float>(), t.data_ptr<float>(), start);

        const int after = start + panel;
        const int rows = n - start;
        const int cols = n - after;
        if (cols <= 0) continue;

        qr_fill_v_kernel<n, panel><<<batch, 256>>>(h.data_ptr<float>(), v.data_ptr<float>(), start);
        float* cptr = h.data_ptr<float>() + (long long)start * n + after;

        cublasStatus_t st;
        st = cublasGemmStridedBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_T,
            cols, panel, rows, &one,
            cptr, CUDA_R_32F, n, hstep,
            v.data_ptr<float>(), CUDA_R_32F, panel, vstep,
            &zero, w.data_ptr<float>(), CUDA_R_32F, n, wstep,
            batch, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT);
        TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "wy1:", (int)st);

        st = cublasGemmStridedBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_T,
            cols, panel, panel, &one,
            w.data_ptr<float>(), CUDA_R_32F, n, wstep,
            t.data_ptr<float>(), CUDA_R_32F, panel, tstep,
            &zero, z.data_ptr<float>(), CUDA_R_32F, n, wstep,
            batch, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT);
        TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "wy2:", (int)st);

        st = cublasGemmStridedBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_N,
            cols, rows, panel, &neg,
            z.data_ptr<float>(), CUDA_R_32F, n, wstep,
            v.data_ptr<float>(), CUDA_R_32F, panel, vstep,
            &one, cptr, CUDA_R_32F, n, hstep,
            batch, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT);
        TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "wy3:", (int)st);
    }
}

void qr512wy(torch::Tensor data, torch::Tensor h, torch::Tensor tau,
             torch::Tensor v, torch::Tensor t, torch::Tensor w, torch::Tensor z) {
    qr_wyblas<512, 16>(data, h, tau, v, t, w, z);
}

void qr1024wy(torch::Tensor data, torch::Tensor h, torch::Tensor tau,
              torch::Tensor v, torch::Tensor t, torch::Tensor w, torch::Tensor z) {
    qr_wyblas<1024, 16>(data, h, tau, v, t, w, z);
}
""",
        functions=["qr32", "qr176panel", "qr352panel2col", "qr512wy", "qr1024wy"],
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        extra_ldflags=["-lcublas"],
        verbose=False,
    )
else:
    _QR = None


def custom_kernel(data: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    batch, n, _ = data.shape
    if _QR is None:
        return torch.geqrf(data)

    h = torch.empty_like(data)
    tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)

    if n == 32:
        _QR.qr32(data, h, tau)
        return h, tau
    if n == 176:
        _QR.qr176panel(data, h, tau)
        return h, tau
    if n == 352:
        _QR.qr352panel2col(data, h, tau)
        return h, tau
    if n == 512:
        v = torch.empty((batch, n, 16), device=data.device, dtype=data.dtype)
        t = torch.empty((batch, 16, 16), device=data.device, dtype=data.dtype)
        w = torch.empty((batch, 16, n), device=data.device, dtype=data.dtype)
        z = torch.empty((batch, 16, n), device=data.device, dtype=data.dtype)
        _QR.qr512wy(data, h, tau, v, t, w, z)
        return h, tau
    if n == 1024:
        v = torch.empty((batch, n, 16), device=data.device, dtype=data.dtype)
        t = torch.empty((batch, 16, 16), device=data.device, dtype=data.dtype)
        w = torch.empty((batch, 16, n), device=data.device, dtype=data.dtype)
        z = torch.empty((batch, 16, n), device=data.device, dtype=data.dtype)
        _QR.qr1024wy(data, h, tau, v, t, w, z)
        return h, tau
    return torch.geqrf(data)
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