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
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-memory
extern __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