Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
43.4ms
#391 of 515
2026-06-21

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