Skip to content
KernelIndex
Search⌘K

submission 797501

msaroufim · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

modest_attempt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-797501?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
18.4ms
#338 of 515
2026-06-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3312b787ced4e308393886165ad32222488ef5d1ea55feace8ce91420c2d3f0f
license declaredunknown
license concludedunknown
authorsmsaroufim
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4NVFP4, tensor-core, or tcgen05 implementation.
num-warps = 8num_warps=8,

Kernel source

modest_attempt.py464 lines
"""Hybrid CoreEvolve-derived Triton compact-Householder GEQRF candidate.

This is the private-repo copy of the experimental square fp32 compact
Householder artifact from CoreEvolve/CoreKernels. It is not a low-bit,
NVFP4, tensor-core, or tcgen05 implementation.

The Triton panel-4 path wins most Shampoo benchmark shapes, but the B200
measurements show PyTorch's native GEQRF is faster for the 4096 case. This
variant keeps the Triton path below 4096 and falls back to torch.geqrf at 4096.
"""

from __future__ import annotations

from typing import Any, cast

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

_PANEL = 4
_SMALL_PANEL = 4
_FACTOR_ROWS = 256
_APPLY_ROWS = 64
_APPLY_COLS = 32


@triton.jit
def _qr_panel_factor_kernel(
  a,
  tau,
  panel_start,
  panel_end,
  n,
  BLOCK_R: tl.constexpr,
  GROUPED: tl.constexpr,
  SKETCHED: tl.constexpr,
):
  b = tl.program_id(0)
  offs = tl.arange(0, BLOCK_R)
  base = b * n * n
  k = panel_start
  while k < panel_end:
    alpha = tl.load(a + base + k * n + k).to(tl.float32)
    if SKETCHED:
      sketch_sigma = tl.full((), 0.0, tl.float32)
      resid_sigma = tl.full((), 0.0, tl.float32)
      r = k + 1
      while r < n:
        rows = r + offs
        mask = rows < n
        x = tl.load(a + base + rows * n + k, mask=mask, other=0.0).to(tl.float32)
        sq = x * x
        in_sketch = (rows & 3) == 0
        sketch_part = tl.where(in_sketch, sq, 0.0)
        sketch_sigma += tl.sum(sketch_part, axis=0)
        resid_sigma += tl.sum(sq - sketch_part, axis=0)
        r += BLOCK_R
      sigma = sketch_sigma + resid_sigma
    else:
      sigma = tl.full((), 0.0, tl.float32)
      r = k + 1
      while r < n:
        rows = r + offs
        mask = rows < n
        x = tl.load(a + base + rows * n + k, mask=mask, other=0.0).to(tl.float32)
        sigma += tl.sum(x * x, axis=0)
        r += BLOCK_R

    mu = tl.sqrt(alpha * alpha + sigma)
    beta = tl.where(sigma == 0.0, alpha, tl.where(alpha < 0.0, mu, -mu))
    denom = tl.where(sigma == 0.0, 1.0, alpha - beta)
    v_norm2 = 1.0 + sigma / (denom * denom)
    tau_k = tl.where(sigma == 0.0, 0.0, 2.0 / v_norm2)
    tl.store(a + base + k * n + k, beta)
    tl.store(tau + b * n + k, tau_k)

    r = k + 1
    while r < n:
      rows = r + offs
      mask = rows < n
      x = tl.load(a + base + rows * n + k, mask=mask, other=0.0).to(tl.float32)
      tl.store(a + base + rows * n + k, x / denom, mask=mask)
      r += BLOCK_R

    if GROUPED:
      tl.debug_barrier()

    col = k + 1
    while col < panel_end:
      top_dot = tl.load(a + base + k * n + col).to(tl.float32)
      if SKETCHED:
        sketch_dot = tl.full((), 0.0, tl.float32)
        resid_dot = tl.full((), 0.0, tl.float32)
        r = k + 1
        while r < n:
          rows = r + offs
          mask = rows < n
          v = tl.load(a + base + rows * n + k, mask=mask, other=0.0).to(tl.float32)
          x = tl.load(a + base + rows * n + col, mask=mask, other=0.0).to(tl.float32)
          prod = v * x
          in_sketch = (rows & 3) == 0
          sketch_part = tl.where(in_sketch, prod, 0.0)
          sketch_dot += tl.sum(sketch_part, axis=0)
          resid_dot += tl.sum(prod - sketch_part, axis=0)
          r += BLOCK_R
        dot = top_dot + sketch_dot + resid_dot
      else:
        dot = top_dot
        r = k + 1
        while r < n:
          rows = r + offs
          mask = rows < n
          v = tl.load(a + base + rows * n + k, mask=mask, other=0.0).to(tl.float32)
          x = tl.load(a + base + rows * n + col, mask=mask, other=0.0).to(tl.float32)
          dot += tl.sum(v * x, axis=0)
          r += BLOCK_R

      tl.store(a + base + k * n + col, top_dot - tau_k * dot)
      r = k + 1
      while r < n:
        rows = r + offs
        mask = rows < n
        v = tl.load(a + base + rows * n + k, mask=mask, other=0.0).to(tl.float32)
        x = tl.load(a + base + rows * n + col, mask=mask, other=0.0).to(tl.float32)
        tl.store(a + base + rows * n + col, x - tau_k * v * dot, mask=mask)
        r += BLOCK_R
      col += 1

    if GROUPED:
      tl.debug_barrier()
    k += 1


@triton.jit
def _qr_panel_apply_kernel(
  a,
  tau,
  panel_start,
  panel_end,
  n,
  BLOCK_R: tl.constexpr,
  BLOCK_C: tl.constexpr,
  GROUPED: tl.constexpr,
):
  col_tile = tl.program_id(0)
  b = tl.program_id(1)
  offs_r = tl.arange(0, BLOCK_R)
  offs_c = tl.arange(0, BLOCK_C)
  cols = panel_end + col_tile * BLOCK_C + offs_c
  col_mask = cols < n
  base = b * n * n
  k = panel_start
  while k < panel_end:
    tau_k = tl.load(tau + b * n + k).to(tl.float32)
    dots = tl.load(a + base + k * n + cols, mask=col_mask, other=0.0).to(tl.float32)

    r = k + 1
    while r < n:
      rows = r + offs_r
      row_mask = rows < n
      v = tl.load(a + base + rows * n + k, mask=row_mask, other=0.0).to(tl.float32)
      x = tl.load(
        a + base + rows[:, None] * n + cols[None, :],
        mask=row_mask[:, None] & col_mask[None, :],
        other=0.0,
      ).to(tl.float32)
      dots += tl.sum(v[:, None] * x, axis=0)
      r += BLOCK_R

    top = tl.load(a + base + k * n + cols, mask=col_mask, other=0.0).to(tl.float32)
    tl.store(a + base + k * n + cols, top - tau_k * dots, mask=col_mask)

    r = k + 1
    while r < n:
      rows = r + offs_r
      row_mask = rows < n
      v = tl.load(a + base + rows * n + k, mask=row_mask, other=0.0).to(tl.float32)
      x = tl.load(
        a + base + rows[:, None] * n + cols[None, :],
        mask=row_mask[:, None] & col_mask[None, :],
        other=0.0,
      ).to(tl.float32)
      x = x - tau_k * v[:, None] * dots[None, :]
      tl.store(
        a + base + rows[:, None] * n + cols[None, :],
        x,
        mask=row_mask[:, None] & col_mask[None, :],
      )
      r += BLOCK_R

    if GROUPED:
      tl.debug_barrier()
    k += 1


@triton.jit
def _qr_panel2_apply_kernel(a, tau, panel_start, n, BLOCK_R: tl.constexpr, BLOCK_C: tl.constexpr):
  col_tile = tl.program_id(0)
  b = tl.program_id(1)
  offs_r = tl.arange(0, BLOCK_R)
  offs_c = tl.arange(0, BLOCK_C)
  k0 = panel_start
  k1 = panel_start + 1
  cols = k1 + 1 + col_tile * BLOCK_C + offs_c
  col_mask = cols < n
  base = b * n * n

  tau0 = tl.load(tau + b * n + k0).to(tl.float32)
  tau1 = tl.load(tau + b * n + k1).to(tl.float32)
  top0 = tl.load(a + base + k0 * n + cols, mask=col_mask, other=0.0).to(tl.float32)
  row1 = tl.load(a + base + k1 * n + cols, mask=col_mask, other=0.0).to(tl.float32)
  v0_k1 = tl.load(a + base + k1 * n + k0).to(tl.float32)

  dot0 = top0 + v0_k1 * row1
  dot1 = row1
  cross = v0_k1

  r = k1 + 1
  while r < n:
    rows = r + offs_r
    row_mask = rows < n
    v0 = tl.load(a + base + rows * n + k0, mask=row_mask, other=0.0).to(tl.float32)
    v1 = tl.load(a + base + rows * n + k1, mask=row_mask, other=0.0).to(tl.float32)
    x = tl.load(
      a + base + rows[:, None] * n + cols[None, :],
      mask=row_mask[:, None] & col_mask[None, :],
      other=0.0,
    ).to(tl.float32)
    dot0 += tl.sum(v0[:, None] * x, axis=0)
    dot1 += tl.sum(v1[:, None] * x, axis=0)
    cross += tl.sum(v0 * v1, axis=0)
    r += BLOCK_R

  dot1 = dot1 - tau0 * cross * dot0
  tl.store(a + base + k0 * n + cols, top0 - tau0 * dot0, mask=col_mask)
  tl.store(
    a + base + k1 * n + cols,
    row1 - tau0 * v0_k1 * dot0 - tau1 * dot1,
    mask=col_mask,
  )

  r = k1 + 1
  while r < n:
    rows = r + offs_r
    row_mask = rows < n
    v0 = tl.load(a + base + rows * n + k0, mask=row_mask, other=0.0).to(tl.float32)
    v1 = tl.load(a + base + rows * n + k1, mask=row_mask, other=0.0).to(tl.float32)
    x = tl.load(
      a + base + rows[:, None] * n + cols[None, :],
      mask=row_mask[:, None] & col_mask[None, :],
      other=0.0,
    ).to(tl.float32)
    x = x - tau0 * v0[:, None] * dot0[None, :] - tau1 * v1[:, None] * dot1[None, :]
    tl.store(
      a + base + rows[:, None] * n + cols[None, :],
      x,
      mask=row_mask[:, None] & col_mask[None, :],
    )
    r += BLOCK_R


@triton.jit
def _qr_panel4_apply_kernel(a, tau, panel_start, n, BLOCK_R: tl.constexpr, BLOCK_C: tl.constexpr):
  col_tile = tl.program_id(0)
  b = tl.program_id(1)
  offs_r = tl.arange(0, BLOCK_R)
  offs_c = tl.arange(0, BLOCK_C)
  k0 = panel_start
  k1 = panel_start + 1
  k2 = panel_start + 2
  k3 = panel_start + 3
  cols = k3 + 1 + col_tile * BLOCK_C + offs_c
  col_mask = cols < n
  base = b * n * n

  tau0 = tl.load(tau + b * n + k0).to(tl.float32)
  tau1 = tl.load(tau + b * n + k1).to(tl.float32)
  tau2 = tl.load(tau + b * n + k2).to(tl.float32)
  tau3 = tl.load(tau + b * n + k3).to(tl.float32)

  row0 = tl.load(a + base + k0 * n + cols, mask=col_mask, other=0.0).to(tl.float32)
  row1 = tl.load(a + base + k1 * n + cols, mask=col_mask, other=0.0).to(tl.float32)
  row2 = tl.load(a + base + k2 * n + cols, mask=col_mask, other=0.0).to(tl.float32)
  row3 = tl.load(a + base + k3 * n + cols, mask=col_mask, other=0.0).to(tl.float32)

  v0_1 = tl.load(a + base + k1 * n + k0).to(tl.float32)
  v0_2 = tl.load(a + base + k2 * n + k0).to(tl.float32)
  v0_3 = tl.load(a + base + k3 * n + k0).to(tl.float32)
  v1_2 = tl.load(a + base + k2 * n + k1).to(tl.float32)
  v1_3 = tl.load(a + base + k3 * n + k1).to(tl.float32)
  v2_3 = tl.load(a + base + k3 * n + k2).to(tl.float32)

  dot0 = row0 + v0_1 * row1 + v0_2 * row2 + v0_3 * row3
  dot1 = row1 + v1_2 * row2 + v1_3 * row3
  dot2 = row2 + v2_3 * row3
  dot3 = row3

  c10 = v0_1 + v1_2 * v0_2 + v1_3 * v0_3
  c20 = v0_2 + v2_3 * v0_3
  c21 = v1_2 + v2_3 * v1_3
  c30 = v0_3
  c31 = v1_3
  c32 = v2_3

  r = k3 + 1
  while r < n:
    rows = r + offs_r
    row_mask = rows < n
    v0 = tl.load(a + base + rows * n + k0, mask=row_mask, other=0.0).to(tl.float32)
    v1 = tl.load(a + base + rows * n + k1, mask=row_mask, other=0.0).to(tl.float32)
    v2 = tl.load(a + base + rows * n + k2, mask=row_mask, other=0.0).to(tl.float32)
    v3 = tl.load(a + base + rows * n + k3, mask=row_mask, other=0.0).to(tl.float32)
    x = tl.load(
      a + base + rows[:, None] * n + cols[None, :],
      mask=row_mask[:, None] & col_mask[None, :],
      other=0.0,
    ).to(tl.float32)
    dot0 += tl.sum(v0[:, None] * x, axis=0)
    dot1 += tl.sum(v1[:, None] * x, axis=0)
    dot2 += tl.sum(v2[:, None] * x, axis=0)
    dot3 += tl.sum(v3[:, None] * x, axis=0)
    c10 += tl.sum(v1 * v0, axis=0)
    c20 += tl.sum(v2 * v0, axis=0)
    c21 += tl.sum(v2 * v1, axis=0)
    c30 += tl.sum(v3 * v0, axis=0)
    c31 += tl.sum(v3 * v1, axis=0)
    c32 += tl.sum(v3 * v2, axis=0)
    r += BLOCK_R

  work0 = dot0
  work1 = dot1 - tau0 * c10 * work0
  work2 = dot2 - tau0 * c20 * work0 - tau1 * c21 * work1
  work3 = dot3 - tau0 * c30 * work0 - tau1 * c31 * work1 - tau2 * c32 * work2

  tl.store(a + base + k0 * n + cols, row0 - tau0 * work0, mask=col_mask)
  tl.store(
    a + base + k1 * n + cols,
    row1 - tau0 * v0_1 * work0 - tau1 * work1,
    mask=col_mask,
  )
  tl.store(
    a + base + k2 * n + cols,
    row2 - tau0 * v0_2 * work0 - tau1 * v1_2 * work1 - tau2 * work2,
    mask=col_mask,
  )
  tl.store(
    a + base + k3 * n + cols,
    row3 - tau0 * v0_3 * work0 - tau1 * v1_3 * work1 - tau2 * v2_3 * work2 - tau3 * work3,
    mask=col_mask,
  )

  r = k3 + 1
  while r < n:
    rows = r + offs_r
    row_mask = rows < n
    v0 = tl.load(a + base + rows * n + k0, mask=row_mask, other=0.0).to(tl.float32)
    v1 = tl.load(a + base + rows * n + k1, mask=row_mask, other=0.0).to(tl.float32)
    v2 = tl.load(a + base + rows * n + k2, mask=row_mask, other=0.0).to(tl.float32)
    v3 = tl.load(a + base + rows * n + k3, mask=row_mask, other=0.0).to(tl.float32)
    x = tl.load(
      a + base + rows[:, None] * n + cols[None, :],
      mask=row_mask[:, None] & col_mask[None, :],
      other=0.0,
    ).to(tl.float32)
    x = (
      x
      - tau0 * v0[:, None] * work0[None, :]
      - tau1 * v1[:, None] * work1[None, :]
      - tau2 * v2[:, None] * work2[None, :]
      - tau3 * v3[:, None] * work3[None, :]
    )
    tl.store(
      a + base + rows[:, None] * n + cols[None, :],
      x,
      mask=row_mask[:, None] & col_mask[None, :],
    )
    r += BLOCK_R


def qr_geqrf(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
  """Return compact Householder QR factors for a square fp32 CUDA batch."""
  if a.ndim != 3 or a.shape[-1] != a.shape[-2]:
    raise ValueError("qr_geqrf expects a square rank-3 tensor")
  if not a.is_cuda or a.dtype != torch.float32:
    raise TypeError("qr_geqrf expects a fp32 CUDA tensor")
  if a.shape[-1] >= 4096:
    return torch.geqrf(a)

  h = a.contiguous().clone()
  batch = h.shape[0]
  n = h.shape[-1]
  tau = torch.empty((batch, n), device=h.device, dtype=h.dtype)

  panel = _SMALL_PANEL if n <= 512 else _PANEL
  grouped_panel = panel > 1
  factor_rows = 2048 if n > 512 else _FACTOR_ROWS
  apply_cols = 32 if n > 512 else _APPLY_COLS

  # Triton launch objects accept constexpr meta-arguments that ty does not model.
  qr_panel_factor_kernel = cast(Any, _qr_panel_factor_kernel)
  qr_panel4_apply_kernel = cast(Any, _qr_panel4_apply_kernel)
  qr_panel2_apply_kernel = cast(Any, _qr_panel2_apply_kernel)
  qr_panel_apply_kernel = cast(Any, _qr_panel_apply_kernel)

  for panel_start in range(0, n, panel):
    panel_end = min(panel_start + panel, n)
    qr_panel_factor_kernel[(batch,)](
      h,
      tau,
      panel_start,
      panel_end,
      n,
      BLOCK_R=factor_rows,
      GROUPED=grouped_panel,
      SKETCHED=panel_start == 0 and n <= 512,
      num_warps=8,
    )
    if panel_end < n:
      grid = (triton.cdiv(n - panel_end, apply_cols), batch)
      if panel == 4 and panel_end == panel_start + 4:
        qr_panel4_apply_kernel[grid](
          h,
          tau,
          panel_start,
          n,
          BLOCK_R=_APPLY_ROWS,
          BLOCK_C=apply_cols,
          num_warps=4,
        )
      elif panel == 2 and panel_end == panel_start + 2:
        qr_panel2_apply_kernel[grid](
          h,
          tau,
          panel_start,
          n,
          BLOCK_R=_APPLY_ROWS,
          BLOCK_C=apply_cols,
          num_warps=4,
        )
      else:
        qr_panel_apply_kernel[grid](
          h,
          tau,
          panel_start,
          panel_end,
          n,
          BLOCK_R=_APPLY_ROWS,
          BLOCK_C=apply_cols,
          GROUPED=grouped_panel,
          num_warps=4,
        )
  return h, tau


def run(a: torch.Tensor) -> torch.Tensor:
  """Compatibility alias for evaluator harnesses that look for `run`."""
  h, tau = qr_geqrf(a)
  return torch.linalg.householder_product(h, tau)


def custom_kernel(data: input_t) -> output_t:
  return qr_geqrf(data)
scrolls · 464 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