Skip to content
KernelIndex
Search⌘K

submission 826338

arun_k · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 165 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-826338?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
4.51ms
#160 of 515
2026-06-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ef3e13a4cecfae1c8cded583a6ef0123effe10f5962e453c919981b9e71d0dca
license declaredunknown
license concludedunknown
authorsarun_k
imported2026-08-26

Kernel source

submission.py165 lines
import torch
import triton
import triton.language as tl

# =====================================================================
# TRITON KERNEL: Tiled Panel Factorization & In-Kernel WY Matrix Generation
# =====================================================================

@triton.jit
def _panel_kernel(
    P_ptr, TAU_ptr, T_ptr, VOUT_ptr,
    M, IB,
    stride_pb, stride_pr, stride_pc,
    stride_tb, stride_ti,
    stride_Tb, stride_Tr, stride_Tc,
    stride_vb, stride_vr, stride_vc,
    BM: tl.constexpr,   # Static compile-time power-of-2 row bound
    BNB: tl.constexpr  # Static compile-time power-of-2 panel width bound
):
    batch_idx = tl.program_id(0)
    
    # Static 2D index configurations
    r = tl.arange(0, BM)
    c = tl.arange(0, BNB)
    
    row_mask = r < M
    col_mask = c < IB
    
    # Base memory coordinates for this batch slice
    p_panel = P_ptr + batch_idx * stride_pb + r[:, None] * stride_pr + c[None, :] * stride_pc
    tile = tl.load(p_panel, mask=row_mask[:, None] & col_mask[None, :], other=0.0)
    
    tau_vec = tl.zeros((BNB,), dtype=tl.float32)
    
    # Sequential Householder sweep over the panel's internal columns
    for j in range(BNB):
        # Isolate column j cleanly using register masks
        colj = tl.sum(tl.where(c[None, :] == j, tile, 0.0), axis=1)
        
        # Pull out pivot alpha and the norm-squared of elements below it
        alpha = tl.sum(tl.where(r == j, colj, 0.0))
        xn2 = tl.sum(tl.where(r > j, colj * colj, 0.0))
        
        reflect = xn2 > 0.0
        sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
        
        # Calculate Householder scalars stably without triggering branching divergence
        beta = tl.where(reflect, -sgn * tl.sqrt(alpha * alpha + xn2), alpha)
        tau_j = tl.where(reflect, (beta - alpha) / tl.where(reflect, beta, 1.0), 0.0)
        denom = tl.where(reflect, alpha - beta, 1.0)
        
        # Construct Householder vector elements
        vb = colj / denom
        vmask = tl.where(r == j, 1.0, tl.where(r > j, vb, 0.0))
        
        # Local rank-1 update to the remaining column tracks inside this panel register file
        w = tl.sum(tl.where(c[None, :] > j, vmask[:, None] * tile, 0.0), axis=0)
        tile = tile - tau_j * vmask[:, None] * w[None, :]
        
        # Save structural transformations back to our loop state accumulators
        newcol = tl.where(r < j, colj, tl.where(r == j, beta, vb))
        tile = tl.where(c[None, :] == j, newcol[:, None], tile)
        tau_vec = tl.where(c == j, tau_j, tau_vec)

    # Reconstruct the unit-lower triangular reflection matrix V
    V = tl.where(r[:, None] == c[None, :], 1.0, tl.where(r[:, None] > c[None, :], tile, 0.0))
    
    # Store V out to global memory
    VOUT_ptrs = VOUT_ptr + batch_idx * stride_vb + r[:, None] * stride_vr + c[None, :] * stride_vc
    tl.store(VOUT_ptrs, V, mask=row_mask[:, None] & col_mask[None, :])
    
    # Compute the Compact-WY Upper Triangular Block Matrix T entirely in SRAM
    Tt = tl.zeros((BNB, BNB), dtype=tl.float32)
    tau0 = tl.sum(tl.where(c == 0, tau_vec, 0.0))
    Tt = tl.where((c[:, None] == 0) & (c[None, :] == 0), tau0, Tt)
    
    for i in range(1, BNB):
        tau_i = tl.sum(tl.where(c == i, tau_vec, 0.0))
        Vi = tl.sum(tl.where(c[None, :] == i, V, 0.0), axis=1)
        
        # Vectorized internal dot product mapping
        dots = tl.sum(V * Vi[:, None], axis=0)
        z = tl.where(c < i, -tau_i * dots, 0.0)
        Tz = tl.sum(tl.where(c[None, :] < i, Tt * z[None, :], 0.0), axis=1)
        
        newTcol = tl.where(c < i, Tz, tl.where(c == i, tau_i, 0.0))
        Tt = tl.where(c[None, :] == i, newTcol[:, None], Tt)

    # Store computed outputs back to global memory tracks
    T_ptrs = T_ptr + batch_idx * stride_Tb + c[:, None] * stride_Tr + c[None, :] * stride_Tc
    tl.store(T_ptrs, Tt, mask=col_mask[:, None] & col_mask[None, :])
    
    tl.store(p_panel, tile, mask=row_mask[:, None] & col_mask[None, :])
    tl.store(TAU_ptr + batch_idx * stride_tb + c * stride_ti, tau_vec, mask=col_mask)



def qr(A: torch.Tensor, block_size: int, num_warps: int = 8):
    B, m, n = A.shape
    bs = int(block_size)
    BNB = triton.next_power_of_2(bs)
    
    H = A.clone()
    tau = A.new_zeros(B, n)
    
    # Process the entire matrix via sequenced block column partitions
    for k in range(0, n, bs):
        ib = min(bs, n - k)
        BM = triton.next_power_of_2(m - k)
        
        Hv = H[:, k:, k : k + ib]  # In-place tracking slice view
        Tt = A.new_zeros(B, BNB, BNB)
        ts = A.new_zeros(B, BNB)
        Vb = A.new_zeros(B, m - k, ib)
        
        # Spin up the specialized Triton panel kernel configuration
        _panel_kernel[(B,)](
            Hv, ts, Tt, Vb, m - k, ib,
            Hv.stride(0), Hv.stride(1), Hv.stride(2),
            ts.stride(0), ts.stride(1),
            Tt.stride(0), Tt.stride(1), Tt.stride(2),
            Vb.stride(0), Vb.stride(1), Vb.stride(2),
            BM=BM, BNB=BNB,
            num_warps=num_warps
        )
        
        tau[:, k : k + ib] = ts[:, :ib]
        hi = k + ib
        
        # Trailing Submatrix Update: A_trail = A_trail - V @ (T.T @ (V.T @ A_trail))
        # Leverages cuBLAS Tensor Cores natively for maximum performance
        if hi < n:
            V = Vb
            T = Tt[:, :ib, :ib]
            C = H[:, k:, hi:]
            
            W = V.transpose(-1, -2) @ C
            W = T.transpose(-1, -2) @ W
            C.baddbmm_(V, W, beta=1, alpha=-1)
            
    return H, tau


# =====================================================================
# LEADERBOARD ENTRY CONFIGURATION
# =====================================================================

def custom_kernel(A: torch.Tensor):
    n = A.shape[-1]
    
    # Safety fallback bounds for incredibly gargantuan tensor ranges
    if n > 2048:
        return torch.geqrf(A.contiguous())
        
    # Dynamically tune block configurations and warp density depending on matrix layout scale
    if n >= 1024:
        block, nw = 16, 8
    elif n >= 256:
        block, nw = 32, (4 if n == 512 else 8)
    else:
        block, nw = 32, 4
        
    return qr(A.contiguous(), block_size=block, num_warps=nw)

compute = custom_kernel
scrolls · 165 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