Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
50.9ms
#182 of 286
2026-07-08

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 = 4num_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