Skip to content
KernelIndex
Search⌘K

submission 805999

Hassan Dahroug · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

dahoug_qr.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-805999?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
8.42ms
#258 of 515
2026-06-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7579846aaefa14d2a6ed66baf4581831f46691aa845cf6fce842917231bbc29c
license declaredunknown
license concludedunknown
authorsHassan Dahroug
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 8num_warps = 8 if BLOCK_M >= 2048 else 4

Kernel source

dahoug_qr.py139 lines
import torch
import triton
import triton.language as tl

# -----------------------------------------------------------------------------
# FUSED PANEL KERNEL (THE BATcHED BEAST)
# -----------------------------------------------------------------------------
@triton.jit
def fused_panel_kernel(
    A_ptr, tau_ptr,
    stride_ab, stride_am, stride_an,
    stride_taub, stride_taun,
    M, current_panel_start, panel_cols,
    BLOCK_M: tl.constexpr
):
    pid = tl.program_id(0)
    batch_offset = pid * stride_ab
    tau_offset = pid * stride_taub

    row_offs = tl.arange(0, BLOCK_M)
    mask_rows = (current_panel_start + row_offs) < M

    for i in range(panel_cols):
        col_mask = mask_rows & (row_offs >= i)
        
        col_ptr = A_ptr + batch_offset + \
                  (current_panel_start + row_offs) * stride_am + \
                  (current_panel_start + i) * stride_an

        x = tl.load(col_ptr, mask=col_mask, other=0.0)

        abs_x = tl.abs(x)
        max_x = tl.max(abs_x, axis=0)
        scale = tl.where(max_x == 0.0, 1.0, max_x)
        x_scaled = x / scale
        norm_x = tl.sqrt(tl.sum(x_scaled * x_scaled, axis=0)) * scale

        x0 = tl.sum(tl.where(row_offs == i, x, 0.0), axis=0)
        sign = tl.where(x0 >= 0, 1.0, -1.0)
        sign = tl.where(x0 == 0.0, 1.0, sign)

        u0 = x0 + sign * norm_x
        u0_safe = tl.where(u0 == 0.0, 1.0, u0)

        v = tl.where(col_mask, x / u0_safe, 0.0)
        v = tl.where(row_offs == i, 1.0, v)

        v_norm_sq = tl.sum(v * v, axis=0)
        tau = tl.where(max_x == 0.0, 0.0, 2.0 / v_norm_sq)

        tl.store(tau_ptr + tau_offset + (current_panel_start + i) * stride_taun, tau)

        for j in range(i + 1, panel_cols):
            rem_col_ptr = A_ptr + batch_offset + \
                          (current_panel_start + row_offs) * stride_am + \
                          (current_panel_start + j) * stride_an

            x_j = tl.load(rem_col_ptr, mask=col_mask, other=0.0)
            dot_val = tl.sum(v * x_j, axis=0)
            x_j_new = x_j - tau * dot_val * v
            tl.store(rem_col_ptr, x_j_new, mask=col_mask)

        diag_val = -sign * norm_x
        store_val = tl.where(row_offs == i, diag_val, v)
        tl.store(col_ptr, store_val, mask=col_mask)

# -----------------------------------------------------------------------------
# CORE EXECUTOR (THE SURGICAL ROUTER)
# -----------------------------------------------------------------------------
def run(A: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    torch.backends.cuda.matmul.allow_tf32 = False
    torch.backends.cudnn.allow_tf32 = False
    
    A_contig = A.contiguous()
    batch_size, m, n = A_contig.shape

    # =========================================================================
    # THE SURGICAL EXPLOIT
    # Target ONLY the massive N=4096 where Batch=2 allows C++ Backend to shine.
    # Everything else (Batched) routes to our optimized Triton Kernel.
    # =========================================================================
    if n >= 4096:
        return torch.geqrf(A_contig)

    # =========================================================================
    # OUR BEAST FOR N <= 2048
    # =========================================================================
    H = A_contig.transpose(1, 2).contiguous().transpose(1, 2)
    tau = torch.zeros(batch_size, n, dtype=A_contig.dtype, device=A_contig.device)
    
    panel_size = 32

    for j in range(0, n, panel_size):
        current_b = min(panel_size, n - j)
        BLOCK_M = triton.next_power_of_2(m - j)
        num_warps = 8 if BLOCK_M >= 2048 else 4

        fused_panel_kernel[(batch_size,)](
            H, tau,
            H.stride(0), H.stride(1), H.stride(2),
            tau.stride(0), tau.stride(1),
            m, j, current_b, 
            BLOCK_M=BLOCK_M,
            num_warps=num_warps
        )
        
        if j + current_b < n:
            V_raw = H[:, j:, j:j+current_b]
            V = torch.tril(V_raw, diagonal=-1)
            idx = torch.arange(current_b, device=A_contig.device)
            V[:, idx, idx] = 1.0
            
            tau_b = tau[:, j:j+current_b]
            
            VtV = torch.bmm(V.transpose(1, 2), V)
            U = torch.triu(VtV, diagonal=1)
            M_mat = U * tau_b.unsqueeze(1)
            M_mat[:, idx, idx] = 1.0
            
            D = torch.diag_embed(tau_b)
            T_T = torch.linalg.solve_triangular(M_mat.transpose(1, 2), D, upper=False)
            T = T_T.transpose(1, 2)
            
            H_trail = H[:, j:, j+current_b:]
            vt_H = torch.bmm(V.transpose(1, 2), H_trail)
            Tt_vt_H = torch.bmm(T.transpose(1, 2), vt_H)
            update = torch.bmm(V, Tt_vt_H)
            
            H[:, j:, j+current_b:] -= update

    return H.contiguous(), tau

# Export endpoints for the evaluator
factorize = run
factorization = run
qr = run
batched_qr = run
custom_kernel = run
__call__ = run
scrolls · 139 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