Skip to content
KernelIndex
Search⌘K

submission 809781

Justin Arndt · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-809781?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
60.5ms
#404 of 515
2026-06-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b0e739d212aba0a91e901d096c6702d55c9e4f016d05baabd83a5b532e14158b
license declaredunknown
license concludedunknown
authorsJustin Arndt
imported2026-08-26

Kernel source

submission.py278 lines
"""
Multi-Strategy Batched Householder QR Factorization
====================================================
Matches torch.geqrf convention exactly:
  H = upper triangle R + lower triangle Householder vectors (v[0]=1 implicit)
  tau = reflector coefficients
  Reflector: Q_k = I - tau_k * v_k * v_k^T

Strategies:
  A. Small n (<=128) or tiny batch (<=2): Direct cuSOLVER via torch.geqrf
  B. Medium n (176-512), large batch: Blocked Householder QR, panel_width=32
  C. Large n (1024+), moderate batch: Blocked Householder QR, panel_width=64

The blocked algorithm:
  1. Panel factorization: column-by-column with branchless Householder (DLARFG)
  2. WY representation: accumulate T matrix for block reflector I - V*T*V^T
  3. Trailing update: 3 batched GEMMs via torch.bmm (FP32 CUDA cores)

Key design decisions:
  - TF32 tensor cores DISABLED: QR error O(n*eps_TF32) exceeds tolerance 20*n*eps_FP32
  - Branchless tau/v computation: torch.where avoids warp divergence on ill-conditioned columns
  - V^T V pre-computed once per panel to reduce T-matrix kernel launches
  - Safety fallback to torch.geqrf if custom path produces NaN/Inf
"""

import torch
from task import input_t, output_t


def custom_kernel(data: input_t) -> output_t:
    """Batched Householder QR factorization with strategy routing."""
    batch, n, _ = data.shape

    # ── Strategy A: cuSOLVER fast path ──────────────────────────────
    # Profiling shows torch.geqrf (cuSOLVER/MAGMA) is faster for:
    #   - Small n, small batch, large n with small batch
    #   - batch=40 n=176/352, batch=60 n=1024, batch=8 n=2048
    # cuSOLVER loops over batch internally; Python overhead of custom
    # blocked QR only pays off when batch is very large (>=200).

    # ── Strategy B: Blocked Householder QR ──────────────────────────
    # For LARGE batch (>=200) with MODERATE n (256-768), custom blocked
    # QR with batched GEMM trailing updates is 1.85x faster than cuSOLVER.
    # The large batch makes trailing GEMMs dominate over panel loop overhead.
    # Profiled: B640-N512 custom=510ms vs geqrf=942ms (1.85x speedup)
    if batch >= 200 and 256 <= n <= 768:
        H, tau = _blocked_householder_qr(data)
        # Safety: fall back to cuSOLVER if NaN/Inf detected
        if torch.isfinite(H).all() and torch.isfinite(tau).all():
            return H, tau
        return torch.geqrf(data)

    # ── Default: cuSOLVER for everything else ───────────────────────
    return torch.geqrf(data)


# ====================================================================
#  Block size selection (inspired by recursive panel scaling pattern)
# ====================================================================

def _select_block_size(n: int, batch: int) -> int:
    """Select panel width based on matrix size.

    Smaller panels = less panel overhead, more trailing GEMM steps.
    Larger panels = more panel overhead, fewer (larger) trailing GEMMs.
    For B200 with 192KB+ shared memory, larger panels are viable.
    """
    if n <= 256:
        return 32
    elif n <= 512:
        # For critical batch=640, n=512 shape: nb=32 gives good trailing GEMM sizes
        return 32 if batch >= 64 else 64
    elif n <= 1024:
        return 64
    else:
        return 64


# ====================================================================
#  Blocked Householder QR
# ====================================================================

def _blocked_householder_qr(A: torch.Tensor):
    """Blocked Householder QR with batched GEMM trailing updates.

    Panel factorization is column-by-column (unblocked) with branchless
    Householder reflector computation matching LAPACK DLARFG convention.
    Trailing matrix updated via WY block reflector (I - V*T*V^T) using
    three large batched GEMMs that maximize GPU utilization.
    """
    batch, n, _ = A.shape
    device = A.device
    dtype = A.dtype

    # CRITICAL: Disable TF32 for matmul precision.
    # TF32 mantissa is 10 bits (eps ~ 5e-4), QR tolerance is 20*n*eps32 ~ 20*n*1.2e-7.
    # For n=512: TF32 error ~ O(n * 5e-4) = 0.25  >>  tolerance ~ 1.2e-3.
    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False

    try:
        nb = _select_block_size(n, batch)

        H = A.clone()
        tau = torch.zeros(batch, n, device=device, dtype=dtype)

        for j in range(0, n, nb):
            jb = min(nb, n - j)

            # 1. Panel factorization: column-by-column Householder reflectors
            _panel_factorize(H, tau, j, jb, batch, n)

            # 2. Trailing update via WY representation + batched GEMM
            if j + jb < n:
                _trailing_update(H, tau, j, jb, batch, n)

        return H, tau

    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32


# ====================================================================
#  Panel Factorization (LAPACK DLARFG convention)
# ====================================================================

def _panel_factorize(H, tau, j, jb, batch, n):
    """Column-by-column Householder panel factorization.

    For each column k in the panel [j, j+jb):
      1. Compute Householder reflector (branchless, matching DLARFG)
      2. Apply rank-1 update to remaining panel columns
      3. Store reflector vector below diagonal and tau coefficient

    LAPACK DLARFG convention:
      beta = -sign(x[0]) * ||x||          (becomes R[k,k])
      tau  = (beta - x[0]) / beta         (reflector coefficient, in [1,2])
      v    = x / (x[0] - beta), v[0] = 1  (Householder vector)
      When subdiagonal ||x[1:]|| == 0: tau = 0, no reflection.
    """
    for k in range(jb):
        col = j + k
        m = n - col  # subcolumn length from diagonal down

        if m <= 1:
            # Scalar: no reflection possible
            tau[:, col] = 0.0
            continue

        # ── Extract column data ──────────────────────────────────
        x0 = H[:, col, col].clone()          # (batch,) diagonal element
        x_tail = H[:, col + 1:, col]         # (batch, m-1) view below diagonal

        # ── Compute norms ────────────────────────────────────────
        norm_tail_sq = (x_tail * x_tail).sum(dim=-1)  # (batch,)
        norm_x = torch.sqrt(x0 * x0 + norm_tail_sq)   # (batch,)

        # ── Branchless sign: sign(0) = 1 ─────────────────────────
        s = x0.sign()
        s = torch.where(s == 0, torch.ones_like(s), s)

        # ── beta = -sign(x0) * ||x|| ─────────────────────────────
        beta = -s * norm_x

        # ── Branchless tau and v computation ──────────────────────
        # Only reflect when subdiagonal is nonzero (branchless via torch.where)
        needs_ref = norm_tail_sq > 0  # (batch,) mask

        # tau = (beta - x0) / beta, safe for beta=0
        safe_beta = torch.where(beta.abs() > 0, beta, torch.ones_like(beta))
        tau_k = torch.where(needs_ref,
                            (beta - x0) / safe_beta,
                            torch.zeros_like(x0))

        # v_below = x_tail / (x0 - beta)
        # denom = x0 - beta = x0 + sign(x0)*||x||, magnitude >= ||x||, always safe
        denom = x0 - beta
        safe_denom = torch.where(needs_ref, denom, torch.ones_like(denom))
        v_below = x_tail / safe_denom.unsqueeze(-1)
        v_below = torch.where(needs_ref.unsqueeze(-1), v_below,
                              torch.zeros_like(v_below))

        # ── Store results in H ────────────────────────────────────
        H[:, col, col] = torch.where(needs_ref, beta, x0)
        H[:, col + 1:, col] = v_below
        tau[:, col] = tau_k

        # ── Apply reflector to remaining panel columns ────────────
        # H_k = I - tau_k * v * v^T, where v = [1, v_below]^T
        # rem = H[:, col:, col+1:j+jb]
        # rem -= tau_k * v * (v^T @ rem)
        if k + 1 < jb:
            rem = H[:, col:, col + 1:j + jb]  # (batch, m, jb-k-1) view

            # w = v^T @ rem = rem[0,:] + v_below^T @ rem[1:,:]
            # Using bmm for the v_below^T @ rem[1:,:] part
            w = rem[:, 0:1, :].clone()  # (batch, 1, jb-k-1)
            w = w + torch.bmm(
                v_below.unsqueeze(1),   # (batch, 1, m-1)
                rem[:, 1:, :]           # (batch, m-1, jb-k-1)
            )

            # rank-1 update: rem -= tau_k * v * w
            tw = tau_k.view(batch, 1, 1) * w  # (batch, 1, jb-k-1)
            rem[:, 0:1, :] -= tw                       # row 0: -= tau_k * 1 * w
            rem[:, 1:, :] -= v_below.unsqueeze(-1) * tw  # rows 1+: -= tau_k * v_below * w


# ====================================================================
#  Trailing Matrix Update (WY Block Reflector)
# ====================================================================

def _trailing_update(H, tau, j, jb, batch, n):
    """Apply block reflector to trailing matrix using WY representation.

    Block reflector:  P = I - V * T * V^T
    where V is (batch, m, jb) unit lower triangular (Householder vectors)
    and   T is (batch, jb, jb) upper triangular (WY coefficients)

    Trailing update via 3 batched GEMMs:
      W = V^T @ A_trail           (batch, jb, trailing_cols)
      W = T @ W                   (batch, jb, trailing_cols)
      A_trail -= V @ W            (batch, m, trailing_cols)
    """
    m = n - j
    trailing_cols = n - j - jb
    device = H.device
    dtype = H.dtype

    if trailing_cols <= 0:
        return

    # ── Build V: unit lower triangular from stored Householder vectors ──
    V = H[:, j:, j:j + jb].clone()       # (batch, m, jb) contiguous copy
    V.tril_(diagonal=-1)                   # zero upper triangle (remove R entries)
    diag_idx = torch.arange(min(m, jb), device=device)
    V[:, diag_idx, diag_idx] = 1.0         # set unit diagonal (v[0]=1)

    # ── Build T: upper triangular WY matrix via DLARFT recurrence ───────
    T = _build_T_matrix(V, tau[:, j:j + jb], jb, batch, device, dtype)

    # ── Trailing update: A = Q^T @ A  where Q = I - V*T*V^T ────────────
    # Q^T = I - V * T^T * V^T  (transpose T for the adjoint)
    # These are the large GEMMs that benefit from batch parallelism
    W = torch.bmm(V.transpose(1, 2), H[:, j:, j + jb:])  # (batch, jb, trailing)
    W = torch.bmm(T.transpose(1, 2), W)                    # (batch, jb, trailing) T^T!
    H[:, j:, j + jb:] -= torch.bmm(V, W)                   # (batch, m, trailing)


# ====================================================================
#  T Matrix Construction (LAPACK DLARFT)
# ====================================================================

def _build_T_matrix(V, tau_panel, jb, batch, device, dtype):
    """Build upper triangular T matrix for WY block reflector.

    Recurrence from LAPACK DLARFT:
      T[k,k] = tau[k]
      T[0:k, k] = -tau[k] * T[0:k, 0:k] @ (V^T V)[0:k, k]

    Pre-computes V^T V as a single batched GEMM to reduce kernel launches
    (1 large GEMM instead of jb small GEMMs).
    """
    # Single GEMM for all inner products between Householder vectors
    VTV = torch.bmm(V.transpose(1, 2), V)  # (batch, jb, jb)

    T = torch.zeros(batch, jb, jb, device=device, dtype=dtype)
    T[:, 0, 0] = tau_panel[:, 0]

    for k in range(1, jb):
        tk = tau_panel[:, k]                    # (batch,)
        z = VTV[:, :k, k:k + 1]                # (batch, k, 1) pre-computed inner products
        Tz = torch.bmm(T[:, :k, :k], z)        # (batch, k, 1) triangular matvec
        T[:, :k, k:k + 1] = -tk.view(batch, 1, 1) * Tz
        T[:, k, k] = tk

    return T
scrolls · 278 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