Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
159.9ms
#507 of 515
2026-06-28

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