submission 844360
heyyowassup3187 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 239 lines, June 9 Researcher Reciprocity License v1.0.
submission_6.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-844360?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:439e21e860f9841ec11d7355faebeee1df8cd0d12d05b34872b2eb5776356714
license declaredunknown
license concludedunknown
authorsheyyowassup3187
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
BLOCK=BLOCK, num_warps=1)Kernel source
submission_6.py239 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ── small matrices (n <= 64) ──────────────────────────────────────────────────
@triton.jit
def qr_kernel_small(
A_ptr, tau_ptr, n,
stride_ab, stride_am, stride_an, stride_tb,
BLOCK: tl.constexpr,
):
bid = tl.program_id(0)
A = A_ptr + bid * stride_ab
T = tau_ptr + bid * stride_tb
ridx = tl.arange(0, BLOCK)
for k in range(n):
rows_below = ridx + k + 1
mask_below = rows_below < n
x0 = tl.load(A + k * stride_am + k * stride_an)
x_below = tl.load(A + rows_below * stride_am + k * stride_an,
mask=mask_below, other=0.0)
norm_sq = x0 * x0 + tl.sum(x_below * x_below, axis=0)
norm = tl.sqrt(norm_sq)
sign = tl.where(x0 >= 0.0, 1.0, -1.0)
alpha = -sign * norm
v0 = x0 - alpha
below_sq = norm_sq - x0 * x0
safe_v0 = tl.where(norm_sq > 0.0, v0, 1.0)
tau_k = tl.where(norm_sq > 0.0, 2.0 * v0 * v0 / (v0 * v0 + below_sq), 0.0)
v_stored = x_below / safe_v0
tl.store(A + k * stride_am + k * stride_an, alpha)
tl.store(A + rows_below * stride_am + k * stride_an, v_stored, mask=mask_below)
tl.store(T + k, tau_k)
for j in range(k + 1, n):
a_kj = tl.load(A + k * stride_am + j * stride_an)
a_below = tl.load(A + rows_below * stride_am + j * stride_an,
mask=mask_below, other=0.0)
dot = a_kj + tl.sum(v_stored * a_below, axis=0)
tl.store(A + k * stride_am + j * stride_an, a_kj - tau_k * dot)
tl.store(A + rows_below * stride_am + j * stride_an,
a_below - tau_k * v_stored * dot, mask=mask_below)
# ── medium matrices (64 < n <= 512): full QR, TILE_J tiling ──────────────────
@triton.jit
def qr_kernel_tiled(
A_ptr, tau_ptr, n,
stride_ab, stride_am, stride_an, stride_tb,
BLOCK_R: tl.constexpr, TILE_J: tl.constexpr,
):
bid = tl.program_id(0)
A = A_ptr + bid * stride_ab
T = tau_ptr + bid * stride_tb
ridx = tl.arange(0, BLOCK_R)
cidx = tl.arange(0, TILE_J)
for k in range(n):
rows_below = ridx + k + 1
mask_below = rows_below < n
x0 = tl.load(A + k * stride_am + k * stride_an)
x_below = tl.load(A + rows_below * stride_am + k * stride_an,
mask=mask_below, other=0.0)
norm_sq = x0 * x0 + tl.sum(x_below * x_below, axis=0)
norm = tl.sqrt(norm_sq)
sign = tl.where(x0 >= 0.0, 1.0, -1.0)
alpha = -sign * norm
v0 = x0 - alpha
below_sq = norm_sq - x0 * x0
safe_v0 = tl.where(norm_sq > 0.0, v0, 1.0)
tau_k = tl.where(norm_sq > 0.0, 2.0 * v0 * v0 / (v0 * v0 + below_sq), 0.0)
v_stored = x_below / safe_v0
tl.store(A + k * stride_am + k * stride_an, alpha)
tl.store(A + rows_below * stride_am + k * stride_an, v_stored, mask=mask_below)
tl.store(T + k, tau_k)
for j_start in range(k + 1, n, TILE_J):
cols = cidx + j_start
col_mask = cols < n
pivot_ptrs = A + k * stride_am + cols * stride_an
pivot = tl.load(pivot_ptrs, mask=col_mask, other=0.0)
tile_ptrs = A + rows_below[:, None] * stride_am + cols[None, :] * stride_an
tile = tl.load(tile_ptrs,
mask=mask_below[:, None] & col_mask[None, :], other=0.0)
dots = pivot + tl.sum(v_stored[:, None] * tile, axis=0)
tl.store(pivot_ptrs, pivot - tau_k * dots, mask=col_mask)
tl.store(tile_ptrs,
tile - tau_k * v_stored[:, None] * dots[None, :],
mask=mask_below[:, None] & col_mask[None, :])
# ── panel kernel (n > 512): T_PANEL steps, applies only within the panel ──────
@triton.jit
def qr_panel_kernel(
A_ptr, tau_ptr, n, ps,
stride_ab, stride_am, stride_an, stride_tb,
BLOCK_R: tl.constexpr, T_PANEL: tl.constexpr,
):
"""
Factorizes panel columns [ps, ps+T_PANEL).
Each reflector is applied ONLY to columns k+1..ps+T_PANEL-1 (within panel).
Trailing columns [ps+T_PANEL, n) are updated externally via WY + torch.bmm.
"""
bid = tl.program_id(0)
A = A_ptr + bid * stride_ab
T = tau_ptr + bid * stride_tb
ridx = tl.arange(0, BLOCK_R)
for kr in range(T_PANEL):
k = ps + kr
rows_below = ridx + k + 1
mask_below = rows_below < n
x0 = tl.load(A + k * stride_am + k * stride_an)
x_below = tl.load(A + rows_below * stride_am + k * stride_an,
mask=mask_below, other=0.0)
norm_sq = x0 * x0 + tl.sum(x_below * x_below, axis=0)
norm = tl.sqrt(norm_sq)
sign = tl.where(x0 >= 0.0, 1.0, -1.0)
alpha = -sign * norm
v0 = x0 - alpha
below_sq = norm_sq - x0 * x0
safe_v0 = tl.where(norm_sq > 0.0, v0, 1.0)
tau_k = tl.where(norm_sq > 0.0, 2.0 * v0 * v0 / (v0 * v0 + below_sq), 0.0)
v_stored = x_below / safe_v0
tl.store(A + k * stride_am + k * stride_an, alpha)
tl.store(A + rows_below * stride_am + k * stride_an, v_stored, mask=mask_below)
tl.store(T + k, tau_k)
# apply reflector only to panel columns [k+1, ps+T_PANEL)
for j in range(k + 1, ps + T_PANEL):
a_kj = tl.load(A + k * stride_am + j * stride_an)
a_below = tl.load(A + rows_below * stride_am + j * stride_an,
mask=mask_below, other=0.0)
dot = a_kj + tl.sum(v_stored * a_below, axis=0)
tl.store(A + k * stride_am + j * stride_an, a_kj - tau_k * dot)
tl.store(A + rows_below * stride_am + j * stride_an,
a_below - tau_k * v_stored * dot, mask=mask_below)
# ── WY trailing update helpers ────────────────────────────────────────────────
def _extract_v(A, ps, T):
"""Build explicit V buffer (batch, n-ps, T) from LAPACK-format A."""
batch, n, _ = A.shape
m = n - ps
V = torch.zeros(batch, m, T, device=A.device, dtype=A.dtype)
for t in range(T):
V[:, t, t] = 1.0
if t + 1 < m:
V[:, t + 1:, t] = A[:, ps + t + 1:, ps + t]
return V
def _build_Tmat(V, tau, ps, T):
"""Build upper-triangular T_mat for WY: H0..H{T-1} = I - V T_mat V^T."""
batch = V.shape[0]
Tm = torch.zeros(batch, T, T, device=V.device, dtype=V.dtype)
for k in range(T):
tk = tau[:, ps + k]
Tm[:, k, k] = tk
if k:
VTv = torch.bmm(V[:, :, :k].transpose(-1, -2).contiguous(),
V[:, :, k:k+1]) # (b, k, 1)
Tm[:, :k, k] = (-tk[:, None] *
torch.bmm(Tm[:, :k, :k], VTv).squeeze(-1))
return Tm
def _wy_update(A, V, Tm, ps, T):
"""A_trail -= V @ Tm^T @ V^T @ A_trail (3 bmm, tensor cores via cuBLAS).
Sequential QR applies H_{T-1}...H_0, which equals (H_0...H_{T-1})^T = I - V Tm^T V^T.
So the trailing update uses Tm^T, not Tm.
"""
At = A[:, ps:, ps + T:].contiguous() # (b, m, n-ps-T)
TmT = Tm.transpose(-1, -2).contiguous() # (b, T, T) lower-tri
W = torch.bmm(V.transpose(-1, -2).contiguous(), At) # (b, T, n-ps-T)
A[:, ps:, ps + T:] -= torch.bmm(V, torch.bmm(TmT, W))
# ── dispatch ──────────────────────────────────────────────────────────────────
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
if n > 1024:
return torch.geqrf(data)
H = data.clone()
tau = torch.zeros(batch, n, device=data.device, dtype=data.dtype)
sa, sm, sn = H.stride(0), H.stride(1), H.stride(2) # row-major strides
if n <= 64:
BLOCK = triton.next_power_of_2(n)
qr_kernel_small[(batch,)](H, tau, n, sa, sm, sn, tau.stride(0),
BLOCK=BLOCK, num_warps=1)
elif n <= 512:
BLOCK_R = triton.next_power_of_2(n)
TILE_J = max(1, 4096 // BLOCK_R)
qr_kernel_tiled[(batch,)](H, tau, n, sa, sm, sn, tau.stride(0),
BLOCK_R=BLOCK_R, TILE_J=TILE_J, num_warps=8)
else:
# panel QR: T columns per panel, WY trailing update via bmm
T = 32
BLOCK_R = triton.next_power_of_2(n)
for ps in range(0, n, T):
Ta = min(T, n - ps)
qr_panel_kernel[(batch,)](
H, tau, n, ps, sa, sm, sn, tau.stride(0),
BLOCK_R=BLOCK_R, T_PANEL=Ta, num_warps=8,
)
if ps + Ta < n:
V = _extract_v(H, ps, Ta)
Tm = _build_Tmat(V, tau, ps, Ta)
_wy_update(H, V, Tm, ps, Ta)
return H, tau
scrolls · 239 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