Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
62.6ms
#280 of 286
2026-07-05

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