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
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.
fp4
NVFP4, tensor-core, or tcgen05 implementation.num-warps = 8
num_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