submission 837855
gct · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 423 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-837855?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:749ab52f4856ca86df86337ea271129a4c9efeb1a137951cc5f7003e7514d086
license declaredunknown
license concludedunknown
authorsgct
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
VtC = tl.dot(tl.trans(V), C, input_precision="tf32x3")num-warps = 4
num_warps=4,persistent-kernel
Panel factorization + trailing update ALL inside one persistent kernel per batch element.shared-memory
extern __shared__ float sh[];Kernel source
submission.py423 lines
"""Shape-routed hybrid QR — Fully-fused Triton: single kernel launch for entire QR.
Panel factorization + trailing update ALL inside one persistent kernel per batch element.
Eliminates Python loop overhead (31 launches → 1 launch for n=512).
Panel is computed column-by-column in registers (BM×BN tile).
Trailing is updated tile-by-tile within the same kernel.
For n>BM_MAX (panel tile limit), falls back to multi-launch or geqrf.
"""
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
NB = 32
import os
# Co-residence guardrail: B*SPLIT_M must fit resident CTAs. B200 ~148 SMs → SPLIT_M<=16 at B=8.
# 3060 has 28 SMs → use SPLIT_M=2 locally (B*2 co-resident at small B). Override via env.
_SPLIT_M_N2048 = int(os.environ.get("SPLIT_M_N2048", "16"))
# ---- CUDA whole-matrix Householder QR (one CTA per matrix) ----
_CUDA = r"""
#include <torch/extension.h>
#include <vector>
__global__ void hh_qr(float* __restrict__ A, float* __restrict__ tau, int n) {
extern __shared__ float sh[];
float* As = sh; float* v = sh + n * n;
__shared__ float red[256]; __shared__ float s_tau, s_scale, s_diag;
int b = blockIdx.x, t = threadIdx.x, T = blockDim.x;
float* Ab = A + (size_t)b * n * n;
for (int idx = t; idx < n * n; idx += T) As[idx] = Ab[idx];
__syncthreads();
for (int k = 0; k < n; k++) {
float loc = 0.f;
for (int i = k + 1 + t; i < n; i += T) { float x = As[i * n + k]; loc += x * x; }
red[t] = loc; __syncthreads();
for (int s = T / 2; s > 0; s >>= 1) { if (t < s) red[t] += red[t + s]; __syncthreads(); }
if (t == 0) {
float a = As[k * n + k], xn = red[0];
float sg = (a >= 0.f) ? 1.f : -1.f, be = -sg * sqrtf(a * a + xn);
bool ac = xn > 0.f;
s_tau = ac ? (be - a) / be : 0.f; s_scale = ac ? 1.f / (a - be) : 0.f; s_diag = ac ? be : a;
}
__syncthreads();
float tk = s_tau, sc = s_scale;
for (int i = t; i < n; i += T) v[i] = (i < k) ? 0.f : (i == k ? 1.f : As[i * n + k] * sc);
__syncthreads();
if (t == 0) { As[k * n + k] = s_diag; tau[(size_t)b * n + k] = tk; }
for (int i = k + 1 + t; i < n; i += T) As[i * n + k] = v[i];
__syncthreads();
for (int j = k + 1 + t; j < n; j += T) {
float w = 0.f;
for (int i = k; i < n; i++) w += v[i] * As[i * n + j];
w *= tk;
for (int i = k; i < n; i++) As[i * n + j] -= v[i] * w;
}
__syncthreads();
}
for (int idx = t; idx < n * n; idx += T) Ab[idx] = As[idx];
}
std::vector<torch::Tensor> cuda_qr(torch::Tensor A) {
int B = A.size(0), n = A.size(1);
auto H = A.clone().contiguous();
auto tau = torch::zeros({B, n}, A.options());
size_t shmem = ((size_t)n * n + n) * sizeof(float);
cudaFuncSetAttribute(hh_qr, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)shmem);
hh_qr<<<B, 256, shmem>>>(H.data_ptr<float>(), tau.data_ptr<float>(), n);
return {H, tau};
}
"""
_cu = load_inline(name="cuda_hh_qr_fused", cpp_sources="std::vector<torch::Tensor> cuda_qr(torch::Tensor A);",
cuda_sources=_CUDA, functions=["cuda_qr"], verbose=False)
_lim = torch.cuda.get_device_properties(0).shared_memory_per_block_optin // 4
_NMAX = int(_lim ** 0.5)
while (_NMAX * _NMAX + _NMAX) > _lim:
_NMAX -= 1
@triton.jit
def _fused_qr_kernel(A, TAU, N: tl.constexpr, NB_val: tl.constexpr,
sab, sai, saj, stb, stj,
BM: tl.constexpr, BN: tl.constexpr):
"""Fully-fused QR: one CTA per batch element, loops over all panels internally.
For each panel k0=0,NB,2*NB,...:
1. Load panel columns [k0:N, k0:k0+NB] into registers (BM×BN tile)
2. Factor NB columns (Householder reflections) → V, T, tau
3. For each trailing tile [k0:N, hi:hi+BN]:
- Load C tile
- VtC = V^T @ C
- TtVtC = T^T @ VtC
- C -= V @ TtVtC
- Store C tile
4. Store factored panel back
"""
b = tl.program_id(0)
row = tl.arange(0, BM)
col = tl.arange(0, BN)
num_panels = (N + NB_val - 1) // NB_val
for panel_idx in range(num_panels):
k0 = panel_idx * NB_val
nb = tl.minimum(NB_val, N - k0)
M = N - k0
grow = k0 + row
rmask = grow < N
cmask = col < nb
# Load panel
m2 = rmask[:, None] & cmask[None, :]
ptr = A + b * sab + grow[:, None] * sai + (k0 + col)[None, :] * saj
P = tl.load(ptr, mask=m2, other=0.0)
tauv = tl.zeros((BN,), dtype=tl.float32)
Tm = tl.zeros((BN, BN), dtype=tl.float32)
# Panel factorization (same as _panel_qr)
for j in tl.static_range(BN):
cj = tl.sum(tl.where(col[None, :] == j, P, 0.0), axis=1)
diag = row == j; below = row > j
alpha = tl.sum(tl.where(diag, cj, 0.0))
xnsq = tl.sum(tl.where(below & rmask, cj * cj, 0.0))
s = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = -s * tl.sqrt(alpha * alpha + xnsq)
active = xnsq > 0.0
tau_j = tl.where(active, (beta - alpha) / beta, 0.0)
scale = tl.where(active, 1.0 / (alpha - beta), 0.0)
v = tl.where(diag, 1.0, tl.where(below, cj * scale, 0.0))
w = tl.sum(v[:, None] * P, axis=0)
tcp = tl.where(col < j, -tau_j * w, 0.0)
mvec = tl.sum(Tm * tcp[None, :], axis=1)
newT = tl.where(col == j, tau_j, tl.where(col < j, mvec, 0.0))
Tm = tl.where(col[None, :] == j, newT[:, None], Tm)
upd = (col[None, :] > j) & (row[:, None] >= j)
P = P - tl.where(upd, tau_j * v[:, None] * w[None, :], 0.0)
dval = tl.where(active, beta, alpha)
newcol = tl.where(diag, dval, tl.where(below, cj * scale, 0.0))
P = tl.where(col[None, :] == j, tl.where(row[:, None] < j, P, newcol[:, None]), P)
tauv = tl.where(col == j, tau_j, tauv)
# Store factored panel
tl.store(ptr, P, mask=m2)
tl.store(TAU + b * stb + (k0 + col) * stj, tauv, mask=cmask)
# Build V from factored panel (V[i,j] = 1 if i==j, P[i,j] if i>j, 0 if i<j)
V = tl.where(row[:, None] > col[None, :], P,
tl.where(row[:, None] == col[None, :], 1.0, 0.0))
# T transpose for trailing update
Tt = tl.trans(Tm)
# Trailing update: for each BN-wide tile of columns after the panel
hi = k0 + NB_val
Nc = N - hi
num_tiles = (Nc + BN - 1) // BN if Nc > 0 else 0
for tile_idx in range(num_tiles):
jcol = tile_idx * BN + tl.arange(0, BN)
jmask = jcol < Nc
# Load C tile
c_ptr = A + b * sab + grow[:, None] * sai + (hi + jcol)[None, :] * saj
C = tl.load(c_ptr, mask=rmask[:, None] & jmask[None, :], other=0.0)
# VtC = V^T @ C (BN × BN)
VtC = tl.dot(tl.trans(V), C, input_precision="tf32x3")
# TtVtC = T^T @ VtC
TtVtC = tl.dot(Tt, VtC, input_precision="tf32x3")
# C -= V @ TtVtC
update = tl.dot(V, TtVtC, input_precision="tf32x3")
C = C - update
tl.store(c_ptr, C, mask=rmask[:, None] & jmask[None, :])
def _fused_triton_qr(A):
B, N, _ = A.shape
H = A.clone()
tau = torch.zeros((B, N), device=A.device, dtype=A.dtype)
BM = triton.next_power_of_2(N)
nw = 16 if BM >= 1024 else 8
_fused_qr_kernel[(B,)](
H, tau, N, NB,
H.stride(0), H.stride(1), H.stride(2),
tau.stride(0), tau.stride(1),
BM=BM, BN=NB, num_warps=nw,
)
return H, tau
# Multi-launch Triton path for sizes where BM would be too large for fused
@triton.jit
def _panel_qr(A, TAU, TT, N, K0, NBw, sab, sai, saj, stb, stj, ttb, tti, ttj,
BM: tl.constexpr, BN: tl.constexpr):
b = tl.program_id(0)
row = tl.arange(0, BM); col = tl.arange(0, BN)
grow = K0 + row; rmask = grow < N; cmask = col < NBw
m2 = rmask[:, None] & cmask[None, :]
ptr = A + b * sab + grow[:, None] * sai + (K0 + col)[None, :] * saj
P = tl.load(ptr, mask=m2, other=0.0)
tauv = tl.zeros((BN,), dtype=tl.float32)
Tm = tl.zeros((BN, BN), dtype=tl.float32)
for j in tl.static_range(BN):
cj = tl.sum(tl.where(col[None, :] == j, P, 0.0), axis=1)
diag = row == j; below = row > j
alpha = tl.sum(tl.where(diag, cj, 0.0))
xnsq = tl.sum(tl.where(below & rmask, cj * cj, 0.0))
s = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = -s * tl.sqrt(alpha * alpha + xnsq)
active = xnsq > 0.0
tau_j = tl.where(active, (beta - alpha) / beta, 0.0)
scale = tl.where(active, 1.0 / (alpha - beta), 0.0)
v = tl.where(diag, 1.0, tl.where(below, cj * scale, 0.0))
w = tl.sum(v[:, None] * P, axis=0)
tcp = tl.where(col < j, -tau_j * w, 0.0)
mvec = tl.sum(Tm * tcp[None, :], axis=1)
newT = tl.where(col == j, tau_j, tl.where(col < j, mvec, 0.0))
Tm = tl.where(col[None, :] == j, newT[:, None], Tm)
upd = (col[None, :] > j) & (row[:, None] >= j)
P = P - tl.where(upd, tau_j * v[:, None] * w[None, :], 0.0)
dval = tl.where(active, beta, alpha)
newcol = tl.where(diag, dval, tl.where(below, cj * scale, 0.0))
P = tl.where(col[None, :] == j, tl.where(row[:, None] < j, P, newcol[:, None]), P)
tauv = tl.where(col == j, tau_j, tauv)
tl.store(ptr, P, mask=m2)
tl.store(TAU + b * stb + (K0 + col) * stj, tauv, mask=cmask)
tptr = TT + b * ttb + col[:, None] * tti + col[None, :] * ttj
tl.store(tptr, Tm, mask=cmask[:, None] & cmask[None, :])
@triton.jit
def _fused_trailing(
H, TT, N, K0, NB_actual, Nc,
sab, sai, saj, ttb, tti, ttj,
VBUF, vb_batch, vb_row, vb_col,
BM_TILE: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
):
b = tl.program_id(0)
tile_j = tl.program_id(1)
hi = K0 + NB_actual
M = N - K0
kcol = tl.arange(0, BK)
kmask = kcol < NB_actual
trow = tl.arange(0, BK)
tcol_t = tl.arange(0, BK)
t_ptr = TT + b * ttb + trow[:, None] * tti + tcol_t[None, :] * ttj
T = tl.load(t_ptr, mask=kmask[:, None] & kmask[None, :], other=0.0)
Tt = tl.trans(T)
jcol = tile_j * BN + tl.arange(0, BN)
jmask = jcol < Nc
VtC = tl.zeros((BK, BN), dtype=tl.float32)
for m_start in range(0, M, BM_TILE):
row = m_start + tl.arange(0, BM_TILE)
grow = K0 + row
rmask = (row < M)
h_ptr = H + b * sab + grow[:, None] * sai + (K0 + kcol)[None, :] * saj
Hpanel = tl.load(h_ptr, mask=rmask[:, None] & kmask[None, :], other=0.0, eviction_policy="evict_last")
V_chunk = tl.where(row[:, None] > kcol[None, :], Hpanel,
tl.where(row[:, None] == kcol[None, :], 1.0, 0.0))
c_ptr = H + b * sab + grow[:, None] * sai + (hi + jcol)[None, :] * saj
C_chunk = tl.load(c_ptr, mask=rmask[:, None] & jmask[None, :], other=0.0)
VtC += tl.dot(tl.trans(V_chunk), C_chunk, input_precision="tf32x3")
v_ptr = VBUF + b * vb_batch + row[:, None] * vb_row + kcol[None, :] * vb_col
tl.store(v_ptr, V_chunk, mask=rmask[:, None] & kmask[None, :])
TtVtC = tl.dot(Tt, VtC, input_precision="tf32x3")
for m_start in range(0, M, BM_TILE):
row = m_start + tl.arange(0, BM_TILE)
grow = K0 + row
rmask = (row < M)
v_ptr = VBUF + b * vb_batch + row[:, None] * vb_row + kcol[None, :] * vb_col
V_chunk = tl.load(v_ptr, mask=rmask[:, None] & kmask[None, :], other=0.0)
c_ptr = H + b * sab + grow[:, None] * sai + (hi + jcol)[None, :] * saj
C_chunk = tl.load(c_ptr, mask=rmask[:, None] & jmask[None, :], other=0.0)
update = tl.dot(V_chunk, TtVtC, input_precision="tf32x3")
C_chunk = C_chunk - update
tl.store(c_ptr, C_chunk, mask=rmask[:, None] & jmask[None, :])
@triton.jit
def _splitm_panel_qr(A, TAU, TT, VAL, CNT, N, K0, NBw, SPLIT_M: tl.constexpr,
sab, sai, saj, stb, stj, ttb, tti, ttj,
M, ROWS_PER, BM: tl.constexpr, BN: tl.constexpr, W: tl.constexpr):
# Cross-CTA split-M Householder panel factorization (std Triton, atomic spin-barrier).
# grid = (B, SPLIT_M). program_id(1) owns a row-slice of the panel; per reflector j
# two cross-CTA reductions (phase0: [alpha,xnsq], phase1: w = v^T P) via per-(b,j,phase)
# scratch slots used exactly once (release/acquire arrival counter, no reset).
b = tl.program_id(0)
pm = tl.program_id(1)
base = pm * ROWS_PER # panel-local start row for this CTA
hi_local = tl.minimum(base + ROWS_PER, M) # exclusive upper bound for this slice
prow = base + tl.arange(0, BM) # panel-local row indices
rmask = prow < hi_local # own-slice guard (no double counting)
col = tl.arange(0, BN); cmask = col < NBw
grow = K0 + prow
m2 = rmask[:, None] & cmask[None, :]
ptr = A + b * sab + grow[:, None] * sai + (K0 + col)[None, :] * saj
P = tl.load(ptr, mask=m2, other=0.0)
wcol = tl.arange(0, W)
tauv = tl.zeros((BN,), dtype=tl.float32)
Tm = tl.zeros((BN, BN), dtype=tl.float32)
for j in range(BN):
cj = tl.sum(tl.where(col[None, :] == j, P, 0.0), axis=1)
diag = prow == j; below = prow > j
alpha_p = tl.sum(tl.where(diag & rmask, cj, 0.0))
xnsq_p = tl.sum(tl.where(below & rmask, cj * cj, 0.0))
# ---- reduction phase 0 : [alpha, xnsq] ----
slot = b * (BN * 2) + j * 2 + 0
vptr = VAL + slot * W + wcol
pvals = tl.where(wcol == 0, alpha_p, tl.where(wcol == 1, xnsq_p, 0.0))
tl.atomic_add(vptr, pvals, sem="release")
cptr = CNT + slot
tl.atomic_add(cptr, 1, sem="acq_rel")
done = tl.atomic_add(cptr, 0, sem="acquire")
while done < SPLIT_M:
done = tl.atomic_add(cptr, 0, sem="acquire")
red = tl.load(vptr)
alpha = tl.sum(tl.where(wcol == 0, red, 0.0))
xnsq = tl.sum(tl.where(wcol == 1, red, 0.0))
# ---- reflector (identical on all CTAs) ----
s = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = -s * tl.sqrt(alpha * alpha + xnsq)
active = xnsq > 0.0
tau_j = tl.where(active, (beta - alpha) / beta, 0.0)
scale = tl.where(active, 1.0 / (alpha - beta), 0.0)
v = tl.where(diag, 1.0, tl.where(below, cj * scale, 0.0))
v = tl.where(rmask, v, 0.0)
# ---- reduction phase 1 : w = v^T P (length NB) ----
w_p = tl.sum(v[:, None] * P, axis=0)
slot1 = b * (BN * 2) + j * 2 + 1
vptr1 = VAL + slot1 * W + wcol
wp_full = tl.where(wcol < BN,
tl.sum(tl.where(col[None, :] == wcol[:, None], w_p[None, :], 0.0), axis=1),
0.0)
tl.atomic_add(vptr1, wp_full, sem="release")
cptr1 = CNT + slot1
tl.atomic_add(cptr1, 1, sem="acq_rel")
done1 = tl.atomic_add(cptr1, 0, sem="acquire")
while done1 < SPLIT_M:
done1 = tl.atomic_add(cptr1, 0, sem="acquire")
wred = tl.load(vptr1)
w = tl.sum(tl.where(wcol[:, None] == col[None, :], wred[:, None], 0.0), axis=0) # (NB,)
# ---- T matrix (compact-WY), identical on all CTAs ----
tcp = tl.where(col < j, -tau_j * w, 0.0)
mvec = tl.sum(Tm * tcp[None, :], axis=1)
newT = tl.where(col == j, tau_j, tl.where(col < j, mvec, 0.0))
Tm = tl.where(col[None, :] == j, newT[:, None], Tm)
tauv = tl.where(col == j, tau_j, tauv)
# ---- local P update ----
upd = (col[None, :] > j) & (prow[:, None] >= j)
P = P - tl.where(upd & rmask[:, None], tau_j * v[:, None] * w[None, :], 0.0)
dval = tl.where(active, beta, alpha)
newcol = tl.where(diag, dval, tl.where(below, cj * scale, 0.0))
P = tl.where(col[None, :] == j, tl.where(prow[:, None] < j, P, newcol[:, None]), P)
tl.store(ptr, P, mask=m2)
if pm == 0:
tl.store(TAU + b * stb + (K0 + col) * stj, tauv, mask=cmask)
tptr = TT + b * ttb + col[:, None] * tti + col[None, :] * ttj
tl.store(tptr, Tm, mask=cmask[:, None] & cmask[None, :])
def _triton_qr(A, SPLIT_M=1):
B, N, _ = A.shape; dev, dt = A.device, A.dtype
H = A.clone(); tau = torch.zeros((B, N), device=dev, dtype=dt)
Tt = torch.empty((B, NB, NB), device=dev, dtype=dt)
Vbuf = torch.empty((B, N, NB), device=dev, dtype=dt)
W = max(triton.next_power_of_2(NB), 2)
if SPLIT_M > 1:
VAL = torch.empty((B * NB * 2 * W,), device=dev, dtype=torch.float32)
CNT = torch.empty((B * NB * 2,), device=dev, dtype=torch.int32)
for k0 in range(0, N, NB):
nb = min(NB, N - k0); M = N - k0; BM = triton.next_power_of_2(M)
nw = 16 if BM >= 256 else 8
if SPLIT_M > 1:
rows_per = (M + SPLIT_M - 1) // SPLIT_M
BMs = triton.next_power_of_2(rows_per)
nws = 16 if BMs >= 256 else 8
VAL.zero_(); CNT.zero_()
_splitm_panel_qr[(B, SPLIT_M)](
H, tau, Tt, VAL, CNT, N, k0, nb, SPLIT_M,
H.stride(0), H.stride(1), H.stride(2),
tau.stride(0), tau.stride(1), Tt.stride(0), Tt.stride(1), Tt.stride(2),
M, rows_per, BM=BMs, BN=NB, W=W, num_warps=nws)
else:
_panel_qr[(B,)](H, tau, Tt, N, k0, nb, H.stride(0), H.stride(1), H.stride(2),
tau.stride(0), tau.stride(1), Tt.stride(0), Tt.stride(1), Tt.stride(2),
BM=BM, BN=NB, num_warps=nw)
hi = k0 + nb
if hi < N:
Nc = N - hi
BN_tile = min(triton.next_power_of_2(Nc), 128)
grid = (B, triton.cdiv(Nc, BN_tile))
_fused_trailing[grid](
H, Tt, N, k0, nb, Nc,
H.stride(0), H.stride(1), H.stride(2),
Tt.stride(0), Tt.stride(1), Tt.stride(2),
Vbuf, Vbuf.stride(0), Vbuf.stride(1), Vbuf.stride(2),
BM_TILE=32, BN=BN_tile, BK=NB,
num_warps=4,
)
return H, tau
def custom_kernel(data: input_t) -> output_t:
B, n, _ = data.shape
if 16 <= n <= _NMAX:
H, tau = _cu.cuda_qr(data)
return H, tau
# Fused single-kernel path: single launch for entire QR
# BM=next_pow2(n), shmem ~= BM*32*4*several. B200 has 228KB, fits n<=512.
_shmem_limit = torch.cuda.get_device_properties(0).shared_memory_per_block_optin
_fused_max = 256 if _shmem_limit >= 228 * 1024 else 128
if n <= _fused_max:
return _fused_triton_qr(data)
if 128 <= n <= 2048:
# split-M panel ONLY at n=2048 (B=8 starves SMs); all other n keep single-CTA panel.
split_m = _SPLIT_M_N2048 if n == 2048 else 1
return _triton_qr(data, SPLIT_M=split_m)
return torch.geqrf(data)
scrolls · 423 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