Skip to content
KernelIndex
Search⌘K

submission 137617

brian4983 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-137617?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
29.7µs
#210 of 369
2025-12-10

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d573b903db6c5e0ed58099dbf5379f50c87a552333903bf642f3bfb350392068
license declaredunknown
license concludedunknown
authorsbrian4983
imported2026-08-26

Kernel source

submission.py168 lines
#!POPCORN leaderboard nvfp4_gemm

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

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

# Ultra-optimized fused kernel with improved memory access patterns
@triton.jit
def fused_to_blocked_kernel(
    src_a_ptr, src_b_ptr,
    dst_a_ptr, dst_b_ptr,
    stride_a_m, stride_a_k,
    stride_b_m, stride_b_k,
    n_rb_a, n_cb_a,
    n_rb_b, n_cb_b,
    num_blocks_a,
    BLOCK_SIZE: tl.constexpr
):
    pid = tl.program_id(0)

    # Determine if processing A or B - use early branching for better performance
    is_a = pid < num_blocks_a
    if is_a:
        block_idx = pid
        rb = block_idx // n_cb_a
        cb = block_idx % n_cb_a
        base_src = src_a_ptr + rb * 128 * stride_a_m + cb * 4 * stride_a_k
        base_dst = dst_a_ptr + block_idx * 512
        stride_m = stride_a_m
        stride_k = stride_a_k
    else:
        block_idx = pid - num_blocks_a
        rb = block_idx // n_cb_b
        cb = block_idx % n_cb_b
        base_src = src_b_ptr + rb * 128 * stride_b_m + cb * 4 * stride_b_k
        base_dst = dst_b_ptr + block_idx * 512
        stride_m = stride_b_m
        stride_k = stride_b_k

    # Process in chunks for better memory coalescing
    # Generate offsets 0..511
    offs = tl.arange(0, 512)

    # Optimized index computation - all in parallel
    inner_row = offs // 16
    rem16 = offs % 16
    outer_row = rem16 // 4
    col = rem16 % 4

    # Compute source offsets - optimize arithmetic
    m_off = (outer_row << 5) + inner_row  # outer_row * 32 + inner_row
    k_off = col

    # Compute all source pointers at once
    src_ptrs = base_src + m_off * stride_m + k_off * stride_k
    src_ptrs = src_ptrs.to(tl.pointer_type(tl.int8))

    # Vectorized load - all 512 elements loaded in parallel
    val = tl.load(src_ptrs)

    # Contiguous store - all elements written in one operation
    dst_ptrs = base_dst + offs
    dst_ptrs = dst_ptrs.to(tl.pointer_type(tl.int8))
    tl.store(dst_ptrs, val)

# Ultra-fast fused transformation - minimal overhead
def fused_to_blocked(sfa_2d, sfb_2d, scale_a_blocked, scale_b_blocked):
    # Compute all metadata in one pass
    rows_a, cols_a = sfa_2d.shape
    rows_b, cols_b = sfb_2d.shape
    n_rb_a = rows_a // 128
    n_cb_a = cols_a // 4
    n_rb_b = rows_b // 128
    n_cb_b = cols_b // 4
    num_blocks_a = n_rb_a * n_cb_a
    num_blocks_b = n_rb_b * n_cb_b
    total_blocks = num_blocks_a + num_blocks_b

    # Get strides once
    stride_a_m, stride_a_k = sfa_2d.stride()
    stride_b_m, stride_b_k = sfb_2d.stride()

    # Launch kernel with minimal overhead
    grid = (total_blocks,)
    fused_to_blocked_kernel[grid](
        sfa_2d, sfb_2d,
        scale_a_blocked, scale_b_blocked,
        stride_a_m, stride_a_k,
        stride_b_m, stride_b_k,
        n_rb_a, n_cb_a,
        n_rb_b, n_cb_b,
        num_blocks_a,
        BLOCK_SIZE=512
    )

# Ultra-optimized implementation with fused kernel launches
def fused_prep_and_gemm(a, b, sfa, sfb, c):
    m, n, l = c.shape

    if l == 1:
        # Maximum performance path for l=1
        # Extract 2D slices once - avoid repeated slicing
        sfa_2d = sfa[:, :, 0]
        sfb_2d = sfb[:, :, 0]
        a_2d = a[:, :, 0]
        b_2d = b[:, :, 0]

        # Compute dimensions once
        rows_a, cols_a = sfa_2d.shape
        rows_b, cols_b = sfb_2d.shape
        n_rb_a = rows_a // 128
        n_cb_a = cols_a // 4
        n_rb_b = rows_b // 128
        n_cb_b = cols_b // 4

        # Allocate buffers with exact sizes
        scale_a_blocked = torch.empty((n_rb_a * n_cb_a * 512,), dtype=sfa.dtype, device=sfa.device)
        scale_b_blocked = torch.empty((n_rb_b * n_cb_b * 512,), dtype=sfb.dtype, device=sfb.device)

        # Single fused kernel launch
        fused_to_blocked(sfa_2d, sfb_2d, scale_a_blocked, scale_b_blocked)

        # Direct GEMM with pre-transposed b
        b_2d_t = b_2d.transpose(0, 1)
        c[:, :, 0] = torch._scaled_mm(
            a_2d,
            b_2d_t,
            scale_a_blocked,
            scale_b_blocked,
            bias=None,
            out_dtype=torch.float16
        )
    else:
        # Multi-batch path - still use fused kernel for each batch
        b_t = b.permute(2, 1, 0)
        for l_idx in range(l):
            rows_a, cols_a = sfa.shape[0], sfa.shape[1]
            rows_b, cols_b = sfb.shape[0], sfb.shape[1]

            n_rb_a = rows_a // 128
            n_cb_a = cols_a // 4
            n_rb_b = rows_b // 128
            n_cb_b = cols_b // 4

            scale_a_blocked = torch.empty((n_rb_a * n_cb_a * 512,), dtype=sfa.dtype, device=sfa.device)
            scale_b_blocked = torch.empty((n_rb_b * n_cb_b * 512,), dtype=sfb.dtype, device=sfb.device)

            # Fused kernel for this batch
            fused_to_blocked(sfa[:, :, l_idx], sfb[:, :, l_idx], scale_a_blocked, scale_b_blocked)

            c[:, :, l_idx] = torch._scaled_mm(
                a[:, :, l_idx],
                b_t[l_idx],
                scale_a_blocked,
                scale_b_blocked,
                bias=None,
                out_dtype=torch.float16
            )
    return c

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa_ref, sfb_ref, _sfa_perm, _sfb_perm, c = data
    return fused_prep_and_gemm(a, b, sfa_ref, sfb_ref, c)
scrolls · 168 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