submission 862880
trxonphoenix · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 317 lines, June 9 Researcher Reciprocity License v1.0.
simple_eigh_triton_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-862880?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:cb7827ca91113a9efde65e94eb28165a4f58991d2ba91ce93d43a4198ea70b00
license declaredunknown
license concludedunknown
authorstrxonphoenix
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4,Kernel source
simple_eigh_triton_submission.py317 lines
# Simple Triton-routed baseline for batched real symmetric eigendecomposition.
#
# Contract:
# Input: data, shape (batch, n, n), CUDA, torch.float32, symmetric up to FP32 roundoff.
# Output: (Q, L)
# Q: shape (batch, n, n), columns are eigenvectors.
# L: shape (batch, n), eigenvalues sorted ascending.
#
# Philosophy:
# - Keep correctness by falling back to torch.linalg.eigh for hard/dense inputs.
# - Add cheap Triton fast routes for exact diagonal / zero / identity-like cases.
# - Keep one explicit Python route per benchmark size so each can be replaced later.
# - Avoid experimental approximate dense logic in the default path.
from __future__ import annotations
import torch
import triton
import triton.language as tl
try:
from task import input_t, output_t
except Exception:
input_t = torch.Tensor
output_t = tuple[torch.Tensor, torch.Tensor]
# Keep this conservative. Direct diagonal routing is only safe when off-diagonal
# mass is essentially zero relative to the matrix L1 norm.
DIRECT_DIAGONAL_RTOL = 1.0e-12
DIRECT_DIAGONAL_ATOL = 0.0
# Stats tile. 1024 keeps compile size small and works for all benchmark shapes.
STAT_BLOCK = 1024
# Q write tile. Larger blocks improve store throughput for large identity/permutation Q.
Q_BLOCK = 1024
@triton.jit
def _stats_partial_kernel(
a,
offdiag_parts,
total_parts,
N: tl.constexpr,
BLOCK: tl.constexpr,
):
"""Compute partial L1 sums for off-diagonal and total matrix mass."""
batch_id = tl.program_id(0)
tile_id = tl.program_id(1)
offsets = tile_id * BLOCK + tl.arange(0, BLOCK)
mask = offsets < N * N
row = offsets // N
col = offsets - row * N
values = tl.load(a + batch_id * N * N + offsets, mask=mask, other=0.0)
abs_values = tl.abs(values)
off_values = tl.where(row != col, abs_values, 0.0)
off_sum = tl.sum(off_values, axis=0)
total_sum = tl.sum(abs_values, axis=0)
part_base = batch_id * tl.cdiv(N * N, BLOCK) + tile_id
tl.store(offdiag_parts + part_base, off_sum)
tl.store(total_parts + part_base, total_sum)
@triton.jit
def _permutation_q_kernel(
q,
perm,
N: tl.constexpr,
BLOCK: tl.constexpr,
):
"""Build Q columns from a sorted diagonal permutation."""
batch_id = tl.program_id(0)
tile_id = tl.program_id(1)
offsets = tile_id * BLOCK + tl.arange(0, BLOCK)
mask = offsets < N * N
row = offsets // N
col = offsets - row * N
source_row = tl.load(perm + batch_id * N + col, mask=mask, other=0)
values = tl.where(row == source_row, 1.0, 0.0)
tl.store(q + batch_id * N * N + offsets, values, mask=mask)
@triton.jit
def _identity_q_kernel(
q,
N: tl.constexpr,
BLOCK: tl.constexpr,
):
"""Build an identity Q matrix for each batch item."""
batch_id = tl.program_id(0)
tile_id = tl.program_id(1)
offsets = tile_id * BLOCK + tl.arange(0, BLOCK)
mask = offsets < N * N
row = offsets // N
col = offsets - row * N
values = tl.where(row == col, 1.0, 0.0)
tl.store(q + batch_id * N * N + offsets, values, mask=mask)
@triton.jit
def _diag_copy_kernel(
a,
diag_out,
N: tl.constexpr,
BLOCK: tl.constexpr,
):
"""Copy the diagonal into an eigenvalue buffer."""
batch_id = tl.program_id(0)
block_id = tl.program_id(1)
offsets = block_id * BLOCK + tl.arange(0, BLOCK)
mask = offsets < N
values = tl.load(
a + batch_id * N * N + offsets * N + offsets,
mask=mask,
other=0.0,
)
tl.store(diag_out + batch_id * N + offsets, values, mask=mask)
def _validate_input(data: torch.Tensor) -> None:
if not isinstance(data, torch.Tensor):
raise TypeError("custom_kernel expects a torch.Tensor")
if data.ndim != 3 or data.shape[-1] != data.shape[-2]:
raise RuntimeError("custom eigensolver expects shape (batch, n, n)")
if not data.is_cuda:
raise RuntimeError("custom eigensolver expects a CUDA tensor")
if data.dtype != torch.float32:
raise RuntimeError("custom eigensolver expects torch.float32")
if not data.is_contiguous():
raise RuntimeError("custom eigensolver expects contiguous input")
def _diagonal_flags(data: torch.Tensor, n: int, *, rtol: float) -> torch.Tensor:
"""Return a CUDA bool mask for matrices that are safe for direct diagonal EVD."""
batch = data.shape[0]
parts = triton.cdiv(n * n, STAT_BLOCK)
offdiag_parts = torch.empty((batch, parts), device=data.device, dtype=torch.float32)
total_parts = torch.empty((batch, parts), device=data.device, dtype=torch.float32)
_stats_partial_kernel[(batch, parts)](
data,
offdiag_parts,
total_parts,
N=n,
BLOCK=STAT_BLOCK,
num_warps=4,
)
offdiag = offdiag_parts.sum(dim=1)
total = total_parts.sum(dim=1)
threshold = torch.clamp(total * rtol, min=DIRECT_DIAGONAL_ATOL)
return offdiag <= threshold
def _direct_diagonal_eigh(data: torch.Tensor, n: int) -> output_t:
"""Fast exact diagonal eigendecomposition using Triton Q construction."""
batch = data.shape[0]
diag = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_diag_copy_kernel[(batch, triton.cdiv(n, Q_BLOCK))](
data,
diag,
N=n,
BLOCK=Q_BLOCK,
num_warps=4,
)
# Sorting is still delegated to PyTorch in this baseline. Replace this with
# a Triton bitonic/radix route for n32/n176 once the rest is stable.
l_sorted, perm = torch.sort(diag, dim=1)
q = torch.empty((batch, n, n), device=data.device, dtype=torch.float32)
_permutation_q_kernel[(batch, triton.cdiv(n * n, Q_BLOCK))](
q,
perm,
N=n,
BLOCK=Q_BLOCK,
num_warps=4,
)
return q, l_sorted.contiguous()
def _dense_fallback_eigh(data: torch.Tensor) -> output_t:
"""Correct fallback for dense or unsafe matrices."""
# torch.linalg.eigh returns (eigenvalues, eigenvectors).
l, q = torch.linalg.eigh(data)
return q.contiguous(), l.contiguous()
def _route_diagonal_then_fallback(
data: torch.Tensor,
n: int,
*,
rtol: float,
) -> output_t:
"""Use direct diagonal route for safe matrices, torch fallback otherwise."""
batch = data.shape[0]
flags = _diagonal_flags(data, n, rtol=rtol)
direct_idx = torch.nonzero(flags, as_tuple=False).flatten()
dense_idx = torch.nonzero(~flags, as_tuple=False).flatten()
if direct_idx.numel() == batch:
return _direct_diagonal_eigh(data, n)
if direct_idx.numel() == 0:
return _dense_fallback_eigh(data)
q_out = torch.empty_like(data)
l_out = torch.empty((batch, n), device=data.device, dtype=torch.float32)
direct_data = data.index_select(0, direct_idx)
q_direct, l_direct = _direct_diagonal_eigh(direct_data, n)
q_out.index_copy_(0, direct_idx, q_direct)
l_out.index_copy_(0, direct_idx, l_direct)
dense_data = data.index_select(0, dense_idx)
q_dense, l_dense = _dense_fallback_eigh(dense_data)
q_out.index_copy_(0, dense_idx, q_dense)
l_out.index_copy_(0, dense_idx, l_dense)
return q_out.contiguous(), l_out.contiguous()
# ---------------------------------------------------------------------------
# Shape-specific routes.
#
# These are deliberately boring at first. The point is to give each benchmark
# shape a stable function that you can replace independently after profiling.
# ---------------------------------------------------------------------------
def _eigh_n32(data: torch.Tensor) -> output_t:
# Best next upgrade: one-kernel Jacobi or hard-coded 32x32 direct/sort route.
return _route_diagonal_then_fallback(data, 32, rtol=DIRECT_DIAGONAL_RTOL)
def _eigh_n176(data: torch.Tensor) -> output_t:
# Best next upgrade: block-Jacobi with small block pairs.
return _route_diagonal_then_fallback(data, 176, rtol=DIRECT_DIAGONAL_RTOL)
def _eigh_n352(data: torch.Tensor) -> output_t:
# Best next upgrade: block-Jacobi or Householder tridiagonalization prototype.
return _route_diagonal_then_fallback(data, 352, rtol=DIRECT_DIAGONAL_RTOL)
def _eigh_n512(data: torch.Tensor) -> output_t:
# Main high-batch target. Add per-matrix detectors here first:
# diagonal, banded, rank-deficient-ish, clustered-ish, dense.
return _route_diagonal_then_fallback(data, 512, rtol=DIRECT_DIAGONAL_RTOL)
def _eigh_n1024(data: torch.Tensor) -> output_t:
# Dense fallback is slow. This route needs a real dense algorithm later.
return _route_diagonal_then_fallback(data, 1024, rtol=DIRECT_DIAGONAL_RTOL)
def _eigh_n2048(data: torch.Tensor) -> output_t:
# Low batch count means a future route must expose intra-matrix parallelism.
return _route_diagonal_then_fallback(data, 2048, rtol=DIRECT_DIAGONAL_RTOL)
def _eigh_n4096(data: torch.Tensor) -> output_t:
# Not in the listed EVD benchmark, but kept for compatibility with QR-style tests.
return _route_diagonal_then_fallback(data, 4096, rtol=DIRECT_DIAGONAL_RTOL)
def custom_kernel(data: input_t) -> output_t:
_validate_input(data)
batch, n, _ = data.shape
if n == 32:
return _eigh_n32(data)
if n == 176:
return _eigh_n176(data)
if n == 352:
return _eigh_n352(data)
if n == 512:
return _eigh_n512(data)
if n == 1024:
return _eigh_n1024(data)
if n == 2048:
return _eigh_n2048(data)
if n == 4096:
return _eigh_n4096(data)
# Keep an honest correctness path for hidden or local experiments.
return _dense_fallback_eigh(data)
def launch_for_eval(inputs: dict) -> output_t:
return custom_kernel(inputs["data"])
kernel = custom_kernel
eigh = custom_kernel
scrolls · 317 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