submission 856420
Sam Ezeh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 95 lines, June 9 Researcher Reciprocity License v1.0.
randomised_subspace_iteration.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-856420?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:0cb749f70deab32183b402d0859c323bd94bc592598e0c382f05f6232648fd23
license declaredunknown
license concludedunknown
authorsSam Ezeh
imported2026-08-26
Kernel source
randomised_subspace_iteration.py95 lines
import torch
from task import input_t, output_t
def custom_kernel(data: input_t) -> output_t:
"""
Highly Optimized Batched Real Symmetric Eigendecomposition.
Uses cuSOLVER for N <= 1024, and Tensor-Core Block Deflation for larger N.
"""
B, N, _ = data.shape
device, dtype = data.device, data.dtype
# 1. Elevate Fast-Path: cuSOLVER dominates everything up to 1024.
if N <= 1024:
L, Q = torch.linalg.eigh(data)
return Q, L
V_total = torch.zeros((B, N, N), device=device, dtype=dtype)
L_total = torch.zeros((B, N), device=device, dtype=dtype)
A_curr = data.clone()
block_size = 512
oversample = 4095
for i in range(0, N, block_size):
current_k = min(block_size, N - i)
remaining_N = N - i
sketch_size = min(N, current_k + oversample)
# Sketching
Omega = torch.randn(B, N, sketch_size, device=device, dtype=dtype)
Y = Omega
# 2. Adaptive Power Iterations
# Only run heavy subspace separation if we are actually truncating the space.
if sketch_size < remaining_N:
with torch.autocast("cuda", dtype=torch.bfloat16):
A_bf16 = A_curr.to(torch.bfloat16)
# Normalize via fast amax
A_max = torch.amax(torch.abs(A_bf16), dim=(-1, -2), keepdim=True).clamp_(min=1e-6)
A_scaled = A_bf16 / A_max
A_sq = A_scaled @ A_scaled
for _ in range(2):
with torch.autocast("cuda", dtype=torch.bfloat16):
Y_bf16 = Y.to(torch.bfloat16)
for _ in range(3):
Y_bf16 = A_sq @ Y_bf16
Y, _ = torch.linalg.qr(Y_bf16.to(dtype))
else:
# If grabbing the entire remaining space, a single QR is sufficient
# to condition the basis for Rayleigh-Ritz.
Y, _ = torch.linalg.qr(Y)
# Exact Rayleigh-Ritz Projection
A_proj = Y.mT @ A_curr @ Y
L_small, V_small = torch.linalg.eigh(A_proj)
# Filter buffer dimensions (keep top current_k by absolute magnitude)
if sketch_size > current_k:
abs_L = torch.abs(L_small)
_, top_idx = torch.sort(abs_L, dim=-1, descending=True)
top_idx = top_idx[:, :current_k]
L_k = torch.gather(L_small, 1, top_idx)
V_k = torch.gather(V_small, 2, top_idx.unsqueeze(1).expand(-1, sketch_size, -1))
else:
L_k, V_k = L_small, V_small
# Lift to full space
V_lifted = Y @ V_k
V_total[:, :, i:i+current_k] = V_lifted
L_total[:, i:i+current_k] = L_k
# 3. In-Place Deflation to save VRAM bandwidth
if i + current_k < N:
deflation = V_lifted @ (L_k.unsqueeze(-1) * V_lifted.mT)
A_curr.sub_(deflation)
# Fast in-place symmetrization
A_curr.add_(A_curr.mT.clone()).mul_(0.5)
# --- Correctness Polish ---
V_total, _ = torch.linalg.qr(V_total)
# Re-extract diagonal in one pass
T = V_total.mT @ data @ V_total
L_total = torch.diagonal(T, dim1=-2, dim2=-1)
L_sorted, indices = torch.sort(L_total, dim=-1)
V_sorted = torch.gather(V_total, 2, indices.unsqueeze(1).expand(-1, N, -1))
return V_sorted, L_sorted
scrolls · 95 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