Skip to content
KernelIndex
Search⌘K

submission 803280

bobmarleybiceps · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 216 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-803280?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
17.4ms
#337 of 515
2026-06-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f3db1726eb95435f685b1c91331cdce3f83109d04d82db5562ae3a6061c8d0b0
license declaredunknown
license concludedunknown
authorsbobmarleybiceps
imported2026-08-26

Kernel source

submission.py216 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl
from task import input_t, output_t


@triton.jit
def panel_qr_kernel(
    A_ptr, tau_ptr,
    k_start, nb, b,
    stride_Ab, stride_Ar, stride_Ac,
    stride_tb,
    BLOCK_N: tl.constexpr,
):
    """Unblocked Householder QR on the panel A[k_start:n, k_start:k_start+b].

    Updates A in-place (R on diagonal/above, Householder vectors below).
    Writes tau for each of the b reflectors.
    Does NOT update the trailing matrix — that is handled by the block reflector.
    """
    bid = tl.program_id(0)
    A_base = A_ptr + bid * stride_Ab
    tau_base = tau_ptr + bid * stride_tb

    # Local row indices 0..BLOCK_N-1 map to global rows k_start..k_start+nb-1
    rows = tl.arange(0, BLOCK_N)

    for j in range(b):
        valid = (rows >= j) & (rows < nb)

        x = tl.load(
            A_base + (rows + k_start) * stride_Ar + (k_start + j) * stride_Ac,
            mask=valid, other=0.0,
        )

        norm_sq = tl.sum(x * x)
        norm = tl.sqrt(norm_sq)
        x0 = tl.sum(tl.where(rows == j, x, 0.0))
        sign_x0 = tl.where(x0 >= 0.0, 1.0, -1.0)

        # Unnormalized Householder vector
        v0 = x0 + sign_x0 * norm
        v = tl.where(rows == j, v0, x)
        v = tl.where(valid, v, 0.0)

        v_sq = tl.sum(v * v)
        tau_apply = tl.where(v_sq < 1e-30, 0.0, 2.0 / v_sq)
        # LAPACK tau: defined for normalized v where v[j]=1
        tau_lapack = tau_apply * v0 * v0

        tl.store(tau_base + (k_start + j), tau_lapack)

        # R diagonal
        tl.store(
            A_base + (k_start + j) * stride_Ar + (k_start + j) * stride_Ac,
            -sign_x0 * norm,
        )

        # Store normalized Householder vector below diagonal
        safe_v0 = tl.where(v0 * v0 > 1e-60, v0, 1.0)
        v_norm = v / safe_v0
        below = (rows > j) & (rows < nb)
        tl.store(
            A_base + (rows + k_start) * stride_Ar + (k_start + j) * stride_Ac,
            v_norm, mask=below,
        )

        # Apply reflector to the remaining panel columns j+1..b-1 only.
        # The trailing matrix (columns k_start+b and beyond) is updated later
        # via the compact WY block reflector using cuBLAS GEMM.
        for jj in range(j + 1, b):
            col = tl.load(
                A_base + (rows + k_start) * stride_Ar + (k_start + jj) * stride_Ac,
                mask=(rows < nb), other=0.0,
            )
            dot = tl.sum(v * col)
            tl.store(
                A_base + (rows + k_start) * stride_Ar + (k_start + jj) * stride_Ac,
                col - tau_apply * dot * v,
                mask=(rows < nb),
            )


@triton.jit
def build_T_kernel(
    V_ptr, tau_ptr, T_ptr,
    m, b,
    stride_Vba, stride_Vm, stride_Vn,
    stride_tba,
    stride_Tba,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
):
    """Build upper-triangular T for the compact WY representation (LAPACK DLARFT).

    V: batch × m × b (lower unit triangular — unit diagonal already inserted).
    tau: batch × b.
    T: batch × b × b (output, assumed zeroed, contiguous row-major).
    """
    bid = tl.program_id(0)
    V_base = V_ptr + bid * stride_Vba
    tau_base = tau_ptr + bid * stride_tba
    T_base = T_ptr + bid * stride_Tba   # contiguous b×b block

    rows = tl.arange(0, BLOCK_M)   # row index into V (0..m-1)
    bcols = tl.arange(0, BLOCK_B)  # reused for both V columns and T rows/cols

    for j in range(b):
        tau_j = tl.load(tau_base + j)

        # T[j, j] = tau_j
        tl.store(T_base + j * b + j, tau_j)

        # Load vj = V[j:m, j]  (rows >= j are valid; vj[j]==1 already in V)
        valid_rows = (rows >= j) & (rows < m)
        vj = tl.load(
            V_base + rows * stride_Vm + j * stride_Vn,
            mask=valid_rows, other=0.0,
        )

        # Load Vi = V[j:m, 0:j]  — the previous j columns, rows >= j
        valid_cols = bcols < j
        Vi = tl.load(
            V_base + rows[:, None] * stride_Vm + bcols[None, :] * stride_Vn,
            mask=valid_rows[:, None] & valid_cols[None, :],
            other=0.0,
        )  # BLOCK_M × BLOCK_B

        # z[i] = -tau_j * (Vi[:,i] · vj)  for i = 0..j-1
        z = -tau_j * tl.sum(Vi * vj[:, None], axis=0)   # BLOCK_B
        z = tl.where(valid_cols, z, 0.0)

        # Load T[0:j, 0:j]  (written by previous iterations, lives in L2)
        Tj = tl.load(
            T_base + bcols[:, None] * b + bcols[None, :],
            mask=(bcols < j)[:, None] & (bcols < j)[None, :],
            other=0.0,
        )  # BLOCK_B × BLOCK_B

        # result = Tj @ z  → T[0:j, j]
        result = tl.sum(Tj * z[None, :], axis=1)   # BLOCK_B

        tl.store(
            T_base + bcols * b + j,
            result,
            mask=(bcols < j),
        )


def _next_pow2(n: int) -> int:
    p = 1
    while p < n:
        p <<= 1
    return p


def reference_kernel(data: input_t) -> output_t:
    return torch.geqrf(data)


def custom_kernel(data: input_t) -> output_t:
    A = data.clone()   # in-place modifications below; .contiguous() not needed (Triton uses strides)
    batch, n, _ = A.shape
    tau = torch.zeros(batch, n, dtype=A.dtype, device=A.device)

    bs = 32
    BLOCK_N = _next_pow2(n)

    for k in range(0, n, bs):
        b  = min(bs, n - k)
        nb = n - k

        # --- 1. Panel factorization ---
        panel_qr_kernel[(batch,)](
            A, tau,
            k, nb, b,
            A.stride(0), A.stride(1), A.stride(2),
            tau.stride(0),
            BLOCK_N=BLOCK_N,
        )

        if k + b >= n:
            break

        # --- 2. Build V with unit diagonal ---
        # A[:, k:, k:k+b] has row-stride n (not b), so a single .contiguous() gives
        # one compact copy — avoids the previous clone() + contiguous() double-copy.
        V = A[:, k:, k : k + b].contiguous()    # batch × (n-k) × b
        V[:, :b, :].tril_(-1)                   # zero diagonal+above in-place (R entries)
        V.diagonal(dim1=1, dim2=2).fill_(1)     # insert implicit unit diagonal in-place

        # --- 3. Build T (compact WY) ---
        T = torch.zeros(batch, b, b, dtype=A.dtype, device=A.device)
        # Pass tau[:, k:] directly — pointer already offset to column k,
        # stride(0)=n gives tau[bid, k+j] when kernel loads tau_base + j.
        build_T_kernel[(batch,)](
            V, tau[:, k:], T,
            nb, b,
            V.stride(0), V.stride(1), V.stride(2),
            tau.stride(0),
            T.stride(0),
            BLOCK_M=BLOCK_N,
            BLOCK_B=_next_pow2(b),
        )

        # --- 4. Trailing update: A[:, k:, k+b:] -= V @ T^T @ V^T @ A[:, k:, k+b:] ---
        C = A[:, k:, k + b:]
        W = torch.bmm(V.transpose(-1, -2), C)
        W = torch.bmm(T.transpose(-1, -2), W)
        A[:, k:, k + b:] -= torch.bmm(V, W)

    return (A, tau)
scrolls · 216 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