submission 824867
obito092430 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 197 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-824867?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:0ffcbeb339bb907f8cc012a443a2c7d3d59c107cde1fc2a0ebff3a60f83ba631
license declaredunknown
license concludedunknown
authorsobito092430
imported2026-08-26
Kernel source
submission.py197 lines
"""Batched square compact-Householder QR factorization — Triton implementation.
Column-by-column batched Householder QR with three Triton kernels per step:
1. _hh_compute — sigma, alpha, tau, store v[1:] below diagonal
2. _compute_w — w = Trailing^T @ v
3. _apply_hh — trailing -= tau * outer(v, w)
Returns (H, tau) matching torch.geqrf convention.
"""
import torch
import triton
import triton.language as tl
# ---------------------------------------------------------------------------
# Kernel 1 — compute Householder reflector for column k
# ---------------------------------------------------------------------------
@triton.jit
def _hh_compute_kernel(
H_ptr, tau_ptr, k,
stride_hb, stride_hi, stride_hj,
stride_tb,
n: tl.constexpr,
BLOCK: tl.constexpr,
):
pid_b = tl.program_id(0)
m = n - k
m_is_1 = m == 1
# Load diagonal x0 = H[b, k, k]
x0 = tl.load(H_ptr + pid_b * stride_hb + k * stride_hi + k * stride_hj).to(tl.float32)
# sigma = ||x||^2
sigma = x0 * x0
offs = tl.arange(0, BLOCK)
for i0 in range(0, n, BLOCK):
i = i0 + offs
mask = (i >= 1) & (i < m)
xi = tl.load(H_ptr + pid_b * stride_hb + (k + i) * stride_hi + k * stride_hj,
mask=mask, other=0.0).to(tl.float32)
sigma += tl.sum(tl.where(mask, xi * xi, 0.0))
is_zero = sigma == 0.0
sqrt_sigma = tl.sqrt(sigma)
sign = tl.where(x0 >= 0, 1.0, -1.0)
alpha = -sign * sqrt_sigma
# tau = (alpha - x0) / alpha, overridden to 0 for last column or zero column
normal_tau = (alpha - x0) / alpha
tau_val = tl.where(m_is_1, 0.0, tl.where(is_zero, 0.0, normal_tau))
tl.store(tau_ptr + pid_b * stride_tb + k, tau_val)
# Store R diagonal (leave unchanged for last column or zero column)
alpha_store = tl.where(m_is_1 | is_zero, x0, alpha)
tl.store(H_ptr + pid_b * stride_hb + k * stride_hi + k * stride_hj, alpha_store)
# Store Householder vector v[1:] = x[i] / (x0 - alpha) below diagonal
denom = x0 - alpha
safe_denom = tl.where(is_zero | m_is_1, 1.0, denom)
apply_v = (~is_zero) & (~m_is_1)
for i0 in range(0, n, BLOCK):
i = i0 + offs
mask = (i >= 1) & (i < m)
xi = tl.load(H_ptr + pid_b * stride_hb + (k + i) * stride_hi + k * stride_hj,
mask=mask, other=0.0).to(tl.float32)
vi = tl.where(mask & apply_v, xi / safe_denom, 0.0)
tl.store(H_ptr + pid_b * stride_hb + (k + i) * stride_hi + k * stride_hj,
vi, mask=mask)
# ---------------------------------------------------------------------------
# Kernel 2 — w = Trailing^T @ v (batched matrix–vector)
# ---------------------------------------------------------------------------
@triton.jit
def _compute_w_kernel(
H_ptr, w_ptr, k,
stride_hb, stride_hi, stride_hj,
stride_wb,
m, mp,
n: tl.constexpr,
BLOCK_ROWS: tl.constexpr, BLOCK_COLS: tl.constexpr,
):
pid_b = tl.program_id(0)
pid_c = tl.program_id(1)
col_start = pid_c * BLOCK_COLS
col_off = col_start + tl.arange(0, BLOCK_COLS)
col_mask = col_off < mp
w_acc = tl.zeros([BLOCK_COLS], dtype=tl.float32)
offs = tl.arange(0, BLOCK_ROWS)
for row_start in range(0, n, BLOCK_ROWS):
row_off = row_start + offs
row_mask = row_off < m
v_i = tl.where(row_off == 0, 1.0,
tl.load(H_ptr + pid_b * stride_hb + (k + row_off) * stride_hi + k * stride_hj,
mask=row_mask, other=0.0).to(tl.float32))
a = tl.load(H_ptr + pid_b * stride_hb + (k + row_off)[:, None] * stride_hi
+ (k + 1 + col_off)[None, :] * stride_hj,
mask=row_mask[:, None] & col_mask[None, :], other=0.0).to(tl.float32)
w_acc += tl.sum(v_i[:, None] * a, axis=0)
tl.store(w_ptr + pid_b * stride_wb + col_off, w_acc.to(tl.float32), mask=col_mask)
# ---------------------------------------------------------------------------
# Kernel 3 — trailing -= tau * outer(v, w) (batched rank-1 update)
# ---------------------------------------------------------------------------
@triton.jit
def _apply_hh_kernel(
H_ptr, w_ptr, tau_ptr, k,
stride_hb, stride_hi, stride_hj,
stride_wb, stride_tb,
m, mp,
BLOCK_ROWS: tl.constexpr, BLOCK_COLS: tl.constexpr,
):
pid_b = tl.program_id(0)
pid_r = tl.program_id(1)
pid_c = tl.program_id(2)
row_start = pid_r * BLOCK_ROWS
col_start = pid_c * BLOCK_COLS
row_off = row_start + tl.arange(0, BLOCK_ROWS)
col_off = col_start + tl.arange(0, BLOCK_COLS)
row_mask = row_off < m
col_mask = col_off < mp
v = tl.where(row_off == 0, 1.0,
tl.load(H_ptr + pid_b * stride_hb + (k + row_off) * stride_hi + k * stride_hj,
mask=row_mask, other=0.0).to(tl.float32))
w = tl.load(w_ptr + pid_b * stride_wb + col_off, mask=col_mask, other=0.0).to(tl.float32)
tau = tl.load(tau_ptr + pid_b * stride_tb + k).to(tl.float32)
a = tl.load(H_ptr + pid_b * stride_hb + (k + row_off)[:, None] * stride_hi
+ (k + 1 + col_off)[None, :] * stride_hj,
mask=row_mask[:, None] & col_mask[None, :], other=0.0).to(tl.float32)
a -= tau * v[:, None] * w[None, :]
tl.store(H_ptr + pid_b * stride_hb + (k + row_off)[:, None] * stride_hi
+ (k + 1 + col_off)[None, :] * stride_hj,
a.to(tl.float32), mask=row_mask[:, None] & col_mask[None, :])
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
COMPUTE_BLOCK = 256
W_BLOCK_ROWS = 256
W_BLOCK_COLS = 64
APPLY_BLOCK_ROWS = 64
APPLY_BLOCK_COLS = 64
def custom_kernel(A: torch.Tensor):
batch, n, _ = A.shape
H = A.contiguous().clone()
tau = torch.zeros((batch, n), dtype=torch.float32, device=A.device)
w = torch.empty((batch, n), dtype=torch.float32, device=A.device)
for k in range(n):
m = n - k
mp = m - 1
_hh_compute_kernel[(batch,)](
H, tau, k,
H.stride(0), H.stride(1), H.stride(2),
tau.stride(0),
n=n, BLOCK=COMPUTE_BLOCK,
)
if mp > 0:
_compute_w_kernel[(batch, triton.cdiv(mp, W_BLOCK_COLS))](
H, w, k,
H.stride(0), H.stride(1), H.stride(2),
w.stride(0),
m, mp,
n=n, BLOCK_ROWS=W_BLOCK_ROWS, BLOCK_COLS=W_BLOCK_COLS,
)
_apply_hh_kernel[(batch, triton.cdiv(m, APPLY_BLOCK_ROWS),
triton.cdiv(mp, APPLY_BLOCK_COLS))](
H, w, tau, k,
H.stride(0), H.stride(1), H.stride(2),
w.stride(0), tau.stride(0),
m, mp,
BLOCK_ROWS=APPLY_BLOCK_ROWS, BLOCK_COLS=APPLY_BLOCK_COLS,
)
return H, tau
scrolls · 197 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