submission 851948
yellowviper · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 68 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-851948?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:859766f1c9b3612a8c44509ead5ee78675f697d3c1c13da165637cc3b889ec7f
license declaredunknown
license concludedunknown
authorsyellowviper
imported2026-08-26
Kernel source
submission.py68 lines
import torch
from task import input_t, output_t
def custom_kernel(data: input_t) -> output_t:
"""
Optimized Batched Real Symmetric Eigendecomposition.
Implements a zero-allocation fast path for diagonal matrices
and leverages UPLO="L" for dense matrices.
"""
B, N, _ = data.shape
device = data.device
dtype = data.dtype
# 1. Fast diagonal check
nnz_total = torch.count_nonzero(data, dim=(-1, -2))
nnz_diag = torch.count_nonzero(torch.diagonal(data, dim1=-2, dim2=-1), dim=-1)
is_diag = (nnz_total == nnz_diag)
# Branch 1: All matrices are perfectly diagonal
if is_diag.all():
diag_elements = torch.diagonal(data, dim1=-2, dim2=-1)
sorted_vals, indices = torch.sort(diag_elements, dim=-1)
Q = torch.zeros((B, N, N), dtype=dtype, device=device)
b_idx = torch.arange(B, device=device).unsqueeze(1)
j_idx = torch.arange(N, device=device).unsqueeze(0)
Q[b_idx, indices, j_idx] = 1.0
return Q, sorted_vals
# Branch 2: All matrices are dense
elif not is_diag.any():
# UPLO="L" reduces memory bandwidth by only reading the lower triangle
values, vectors = torch.linalg.eigh(data, UPLO="L")
return vectors, values
# Branch 3: Mixed batch
else:
values = torch.empty((B, N), dtype=dtype, device=device)
vectors = torch.empty((B, N, N), dtype=dtype, device=device)
# --- Diagonal Fast Path ---
diag_idx = is_diag.nonzero(as_tuple=True)[0]
diag_elements = torch.diagonal(data[diag_idx], dim1=-2, dim2=-1)
sorted_vals, indices = torch.sort(diag_elements, dim=-1)
# Zero out only the slices for diagonal matrices directly in the output tensor
# This avoids allocating a temporary (num_diag, N, N) tensor like NIAMEY.py did
vectors[diag_idx] = 0.0
# Direct assignment of 1.0s to construct the permutation matrices
b_idx_global = diag_idx.unsqueeze(1)
j_idx = torch.arange(N, device=device).unsqueeze(0)
vectors[b_idx_global, indices, j_idx] = 1.0
values[diag_idx] = sorted_vals
# --- Dense Slow Path ---
not_diag_idx = (~is_diag).nonzero(as_tuple=True)[0]
# Pass only dense matrices to eigh, using UPLO="L" for speed
slow_values, slow_vectors = torch.linalg.eigh(data[not_diag_idx], UPLO="L")
values[not_diag_idx] = slow_values
vectors[not_diag_idx] = slow_vectors
return vectors, values
scrolls · 68 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