Skip to content
KernelIndex
Search⌘K

submission 189474

francescagreco · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

chew2ULTRAsolgems.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-189474?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
43.1µs
#289 of 420
2025-12-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fe61da1c03534ffc42cdfaae3a01e7553f7719466df8a84f685885d0c55ebb57
license declaredunknown
license concludedunknown
authorsfrancescagreco
imported2026-08-26

Techniques

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

fp4128x4 blocked layout expected by Blackwell nvFP4 block-scaled GEMM.

Kernel source

chew2ULTRAsolgems.py98 lines
import torch
import torch.nn.functional as F
from task import input_t, output_t

# From spec
sf_vec_size = 16

def ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b

def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
    """
    Convert a single scale factor matrix [rows, cols] into the
    128x4 blocked layout expected by Blackwell nvFP4 block-scaled GEMM.
    """
    rows, cols = input_matrix.shape
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = (
        blocks
        .reshape(-1, 4, 32, 4)
        .transpose(1, 2)
        .reshape(-1, 32, 16)
    )
    return rearranged.flatten()

def blocked_all_l(s: torch.Tensor) -> torch.Tensor:
    """Processes scales for all layers into hardware-aligned layout."""
    rows, cols, L = s.shape
    blocked = [to_blocked(s[:, :, l]) for l in range(L)]
    return torch.stack(blocked, dim=0).contiguous()

_WARMUP_DONE = False

def custom_kernel(data: input_t) -> output_t:
    """
    Dual nvFP4 GEMM + SwiGLU optimized for Blackwell (B200).
    Corrected to handle 10-way data unpacking.
    """
    global _WARMUP_DONE

    # Correct unpacking of 10 elements provided by the environment
    (
        a, b1, b2, 
        sfa, sfb1, sfb2, 
        sfa_p, sfb1_p, sfb2_p, 
        c
    ) = data
    
    M, N, L = c.shape
    device = a.device

    with torch.inference_mode():
        # 1. Scale preparation: Use permuted scales if available (L2-friendly)
        if sfa_p is not None and sfa_p.numel() > 0:
            scale_a  = sfa_p.permute(5, 2, 4, 0, 1, 3).reshape(L, -1).contiguous()
            scale_b1 = sfb1_p.permute(5, 2, 4, 0, 1, 3).reshape(L, -1).contiguous()
            scale_b2 = sfb2_p.permute(5, 2, 4, 0, 1, 3).reshape(L, -1).contiguous()
        else:
            scale_a  = blocked_all_l(sfa).to(device)
            scale_b1 = blocked_all_l(sfb1).to(device)
            scale_b2 = blocked_all_l(sfb2).to(device)

        # 2. Warmup
        if not _WARMUP_DONE:
            torch._scaled_mm(
                a[:, :, 0], b1[:, :, 0].t(),
                scale_a[0], scale_b1[0],
                out_dtype=torch.float32,
            )
            _WARMUP_DONE = True

        # 3. Memory Allocation (Reused for all layers to keep L2 cache hot)
        g1 = torch.empty((M, N), dtype=torch.float32, device=device)
        g2 = torch.empty((M, N), dtype=torch.float32, device=device)

        # 4. Main Loop
        for l in range(L):
            A_l  = a[:, :, l]
            # .t() is metadata-only; avoids copying K-major inputs
            B1_l = b1[:, :, l].t()
            B2_l = b2[:, :, l].t()

            # Dual GEMM
            torch._scaled_mm(A_l, B1_l, scale_a[l], scale_b1[l], out=g1)
            torch._scaled_mm(A_l, B2_l, scale_a[l], scale_b2[l], out=g2)

            # In-place SwiGLU Epilogue
            # This sequence minimizes HBM write-backs
            res = F.silu(g1) 
            res.mul_(g2)
            
            c[:, :, l].copy_(res.to(torch.float16))

    return c
scrolls · 98 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