Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
22.8ms
#351 of 515
2026-06-29

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 = 1BLOCK=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