Skip to content
KernelIndex
Search⌘K

submission 183549

pongtsu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-183549?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 GEMMsuite of 3 cases
NVIDIA B200
17.7µs
#181 of 369
2025-12-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1755c810e372e580d275adc99ed6d7a47270723bf1093adef79cecca73b9c063
license declaredunknown
license concludedunknown
authorspongtsu
imported2026-08-26

Techniques

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

fp4Optimized NVFP4 block-scaled GEMM for NVIDIA B200.

Kernel source

v3.py140 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t


@triton.jit
def to_blocked_kernel(
    # Input: pre-permuted scale factors [32, 4, rb, 4, cb, L]
    sf_in_ptr,
    # Output: blocked format [packed, L]
    sf_out_ptr,
    # Dimensions
    rb,
    cb,
    L,
    # Strides for input [32, 4, rb, 4, cb, L]
    stride_mm32,
    stride_mm4,
    stride_mm,
    stride_kk4,
    stride_kk,
    stride_l,
    # Output stride
    out_stride_packed,
    out_stride_l,
    BLOCK_SIZE: tl.constexpr,
):
    """
    Convert pre-permuted scale factors to cuBLAS blocked format.

    Pre-permuted: [32, 4, rb, 4, cb, L] indexed as (mm32, mm4, mm, kk4, kk, l)
    Blocked output index: kk4 + 4*mm4 + 16*mm32 + 512*kk + 512*cb*mm
    """
    pid = tl.program_id(0)
    l_idx = tl.program_id(1)

    offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    total = 512 * rb * cb
    mask = offs < total

    # Decode output position to (mm, kk, mm32, mm4, kk4)
    tmp = offs
    kk4 = tmp % 4
    tmp = tmp // 4
    mm4 = tmp % 4
    tmp = tmp // 4
    mm32 = tmp % 32
    tmp = tmp // 32
    kk = tmp % cb
    mm = tmp // cb

    # Compute input offset (treating as uint8 for fp8 compatibility)
    in_offset = (
        mm32 * stride_mm32
        + mm4 * stride_mm4
        + mm * stride_mm
        + kk4 * stride_kk4
        + kk * stride_kk
        + l_idx * stride_l
    )

    val = tl.load(sf_in_ptr + in_offset, mask=mask)
    out_offset = offs * out_stride_packed + l_idx * out_stride_l
    tl.store(sf_out_ptr + out_offset, val, mask=mask)


def convert_scale_factors_triton(sf_perm: torch.Tensor) -> torch.Tensor:
    """Convert pre-permuted [32, 4, rb, 4, cb, L] to blocked [packed, L] using Triton."""
    d0, d1, rb, d3, cb, L = sf_perm.shape

    packed = 512 * rb * cb
    # View as uint8 for Triton compatibility, will view back after
    sf_u8 = sf_perm.view(torch.uint8)
    out_u8 = torch.empty((packed, L), dtype=torch.uint8, device=sf_perm.device)

    BLOCK_SIZE = 1024
    grid = (triton.cdiv(packed, BLOCK_SIZE), L)

    to_blocked_kernel[grid](
        sf_u8,
        out_u8,
        rb,
        cb,
        L,
        sf_u8.stride(0),
        sf_u8.stride(1),
        sf_u8.stride(2),
        sf_u8.stride(3),
        sf_u8.stride(4),
        sf_u8.stride(5),
        out_u8.stride(0),
        out_u8.stride(1),
        BLOCK_SIZE=BLOCK_SIZE,
    )

    return out_u8.view(sf_perm.dtype)


def convert_scale_factors_pytorch(sf_perm: torch.Tensor) -> torch.Tensor:
    """
    Convert pre-permuted [32, 4, rb, 4, cb, L] to blocked [packed, L] using PyTorch.
    Single permute + reshape - much faster than the to_blocked chain.
    """
    # Pre-permuted: [32, 4, rb, 4, cb, L] = (mm32, mm4, mm, kk4, kk, l)
    # Need: [rb, cb, 32, 4, 4, L] = (mm, kk, mm32, mm4, kk4, l)
    # This gives correct blocked order when flattened
    return sf_perm.permute(2, 4, 0, 1, 3, 5).reshape(-1, sf_perm.shape[-1]).contiguous()


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized NVFP4 block-scaled GEMM for NVIDIA B200.

    Key optimizations:
    1. Uses pre-permuted scale factors (avoids expensive to_blocked chain)
    2. Single permute+reshape vs view/permute/reshape/transpose chain
    3. Direct torch._scaled_mm for optimized tensor core GEMM
    """
    a, b, _, _, sfa_perm, sfb_perm, c = data

    _, _, L = c.shape

    # Convert pre-permuted scale factors to blocked format
    # Use PyTorch permute (single op) - faster for typical sizes
    blocked_a = convert_scale_factors_pytorch(sfa_perm)
    blocked_b = convert_scale_factors_pytorch(sfb_perm)

    for l_idx in range(L):
        c[:, :, l_idx] = torch._scaled_mm(
            a[:, :, l_idx],
            b[:, :, l_idx].transpose(0, 1),
            blocked_a[:, l_idx],
            blocked_b[:, l_idx],
            bias=None,
            out_dtype=torch.float16,
        )

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