submission 840565
HankBO · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 126 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-840565?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:d79e89c709d541627678def417a9077c67c2674f7c023b11053cfc167fe7af6f
license declaredunknown
license concludedunknown
authorsHankBO
imported2026-08-26
Kernel source
submission.py126 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def qr_kernel_s_mat_l_batch(a_ptr,
output_ptr,
tau_ptr,
stride_a_batch, stride_a_row, stride_a_col,
stride_tau_batch, stride_tau_col,
n: tl.constexpr,
BLOCK_SIZE: tl.constexpr):
"""
Householder QR decomposition with small matrix size and large batch size.
assume a shape of n x n.
Load all data once to sram, updating interatively. Store to HBM finally.
"""
pid = tl.program_id(axis=0)
a_batch_ptr = a_ptr + pid * stride_a_batch
out_batch_ptr = output_ptr + pid * stride_a_batch
tau_batch_ptr = tau_ptr + pid * stride_tau_batch
rows = tl.arange(0, BLOCK_SIZE)[:, None] # shape: (N, 1)
cols = tl.arange(0, BLOCK_SIZE)[None, :] # shape: (1, N)
a_offsets = rows * stride_a_row + cols * stride_a_col
valid_mask = (rows < n) & (cols < n)
a_sram = tl.load(a_batch_ptr + a_offsets, mask=valid_mask, other=0.0)
tau_sram = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
for k in range(n):
mask_k_col = (cols == k)
# extract col k, lower dim as (BLOCK_SIZE, 1), block data above row k
v_col = tl.sum(tl.where(mask_k_col, a_sram, 0.0), axis=1)[:, None]
alpha = tl.sum(tl.where(rows == k, v_col, 0.0), axis=0) # shape: (1,)
mask_tail = (rows > k) & (rows < n)
tail_norm_sq = tl.sum(tl.where(mask_tail, v_col * v_col, 0.0), axis=0)
norm_x = tl.sqrt(alpha * alpha + tail_norm_sq)
sign_alpha = tl.where(alpha >= 0, 1.0, -1.0)
beta_standard = -sign_alpha * norm_x
is_reflection_needed = (tail_norm_sq > 0.0)
beta = tl.where(is_reflection_needed, beta_standard, alpha)
beta_safe = tl.where(beta == 0.0, 1.0, beta)
tau = tl.where(tail_norm_sq == 0.0, 0.0, (beta - alpha) / beta_safe)
v_0 = alpha - beta
v_0_safe = tl.where(v_0 == 0.0, 1.0, v_0)
v = tl.where(rows == k, 1.0, v_col / v_0_safe)
v = tl.where(rows >= k, v, 0.0)
# update sub matrix
mask_A_sub = (rows >= k) & (rows < n) & (cols > k) & (cols < n)
A_sub = tl.where(mask_A_sub, a_sram, 0.0)
# 计算 v^T * A (矩阵乘法转为 Element-wise 乘法加规约)
# v: (BLOCK_SIZE, 1), A_sub: (BLOCK_SIZE, BLOCK_SIZE)
# axis=0 规约后得到 (BLOCK_SIZE,),升维为 (1, BLOCK_SIZE)
v_T_A = tl.sum(v * A_sub, axis=0)[None, :]
update = tau * v * v_T_A
a_sram = tl.where(mask_A_sub, a_sram - update, a_sram)
# write back
mask_beta = (rows == k) & (cols == k)
mask_v_store = (rows > k) & (rows < n) & (cols == k)
a_sram = tl.where(mask_beta, beta, a_sram)
a_sram = tl.where(mask_v_store, v, a_sram)
idx_1d = tl.arange(0, BLOCK_SIZE)
tau_sram = tl.where(idx_1d == k, tau, tau_sram)
tl.store(out_batch_ptr + a_offsets, a_sram, mask=valid_mask)
tau_offsets = tl.arange(0, BLOCK_SIZE) * stride_tau_col
mask_tau = tl.arange(0, BLOCK_SIZE) < n
tl.store(tau_batch_ptr + tau_offsets, tau_sram, mask=mask_tau)
def triton_qr(a: torch.Tensor):
"""
shape of a: b x n x n
Grid allocation
1. n <= 512, qr_kernel_s_mat_l_batch: put the whole matrix in a single SM
2. n > 512, qr_kernel_l_mat_s_batch: distribute n/b panels to blocks, do grid-level synchronization
"""
a = a.contiguous()
b, n, _ = a.shape
output = torch.empty_like(a)
tau = torch.empty((b, n), dtype=a.dtype, device=a.device)
stride_a_batch = a.stride(0)
stride_a_row = a.stride(1)
stride_a_col = a.stride(2)
stride_tau_batch = tau.stride(0)
stride_tau_col = tau.stride(1)
BLOCK_SIZE = triton.next_power_of_2(n)
if n <= 512:
grid = lambda meta: (b,)
qr_kernel_s_mat_l_batch[grid](a, output, tau,
stride_a_batch, stride_a_row, stride_a_col,
stride_tau_batch, stride_tau_col,
n, BLOCK_SIZE=BLOCK_SIZE)
return output, tau
def custom_kernel(data: input_t) -> output_t:
if len(data[0]) <=512:
return triton_qr(data)
else:
return torch.geqrf(data)
scrolls · 126 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