Skip to content
KernelIndex
Search⌘K

submission 476068

kaiming-cheng · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

test.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-476068?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 group GEMMsuite of 4 cases
NVIDIA B200
138.1µs
#243 of 310
2026-02-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e575432190ca0c7cc4c1a8bb1d40d41f42fd773fae3605a5a3019746359f132c
license declaredunknown
license concludedunknown
authorskaiming-cheng
imported2026-08-15

Techniques

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

fp4"""Optimized FP4 block-scaled GEMM kernel."""
num-warps = 8num_warps=8,
stages = 4num_stages=4,
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128

Kernel source

test.py195 lines
import triton
import triton.language as tl
import torch


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


@triton.jit
def nvfp4_group_gemm_kernel_optimized(
    # Pointers - a and b are uint8 packed FP4
    a_ptr, b_ptr, c_ptr,
    sfa_ptr, sfb_ptr,
    # Dimensions
    M, N, K,
    # Strides for A [M, K//2] (packed uint8)
    stride_am, stride_ak,
    # Strides for B [N, K//2] (packed uint8)
    stride_bn, stride_bk,
    # Strides for C [M, N]
    stride_cm, stride_cn,
    # Scale strides - reordered format [32, 4, rest_m, 4, rest_k, L]
    stride_sfa_0, stride_sfa_1, stride_sfa_2, stride_sfa_3, stride_sfa_4, stride_sfa_5,
    stride_sfb_0, stride_sfb_1, stride_sfb_2, stride_sfb_3, stride_sfb_4, stride_sfb_5,
    # L index for scale factors
    l_idx,
    # Block sizes
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    SF_VEC_SIZE: tl.constexpr,
):
    """Optimized FP4 block-scaled GEMM kernel."""
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)

    # Enhanced swizzle for better L2 cache locality
    num_pid_in_group = GROUP_SIZE_M * num_pid_n
    group_id = pid // num_pid_in_group
    first_pid_m = group_id * GROUP_SIZE_M
    group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
    pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
    pid_n = (pid % num_pid_in_group) // group_size_m

    # Initialize accumulator
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # Scale block size
    BLOCK_K_SCALE: tl.constexpr = BLOCK_K // SF_VEC_SIZE

    # Number of K iterations
    num_k_iters = tl.cdiv(K, BLOCK_K)

    # Base offsets for this tile - compute once
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k_packed = tl.arange(0, BLOCK_K // 2)
    offs_scale_k = tl.arange(0, BLOCK_K_SCALE)

    # Precompute scale indices for M dimension
    mm32_a = offs_m % 32
    mm4_a = (offs_m % 128) // 32
    mm_a = offs_m // 128

    # Precompute scale indices for N dimension
    mm32_b = offs_n % 32
    mm4_b = (offs_n % 128) // 32
    mm_b = offs_n // 128

    # K scale dimension
    sf_k = tl.cdiv(K, SF_VEC_SIZE)

    # Precompute base pointers
    a_base = a_ptr + offs_m[:, None] * stride_am
    b_base = b_ptr + offs_n[:, None] * stride_bn

    # Precompute scale base pointers
    sfa_base = (sfa_ptr +
                mm32_a[:, None] * stride_sfa_0 +
                mm4_a[:, None] * stride_sfa_1 +
                mm_a[:, None] * stride_sfa_2 +
                l_idx * stride_sfa_5)

    sfb_base = (sfb_ptr +
                mm32_b[:, None] * stride_sfb_0 +
                mm4_b[:, None] * stride_sfb_1 +
                mm_b[:, None] * stride_sfb_2 +
                l_idx * stride_sfb_5)

    # Masks for M and N (reused)
    mask_m = offs_m < M
    mask_n = offs_n < N

    for k_iter in range(num_k_iters):
        # K offset in packed format
        k_offset_packed = k_iter * (BLOCK_K // 2)

        # Load A tile
        a_ptrs = a_base + (offs_k_packed[None, :] + k_offset_packed) * stride_ak
        mask_a = mask_m[:, None] & ((offs_k_packed[None, :] + k_offset_packed) < (K // 2))
        a = tl.load(a_ptrs, mask=mask_a, other=0)

        # Load B tile
        b_ptrs = b_base + (offs_k_packed[None, :] + k_offset_packed) * stride_bk
        mask_b = mask_n[:, None] & ((offs_k_packed[None, :] + k_offset_packed) < (K // 2))
        b = tl.load(b_ptrs, mask=mask_b, other=0)

        # Scale factor K indices
        scale_k_idx = k_iter * BLOCK_K_SCALE
        col_idx = scale_k_idx + offs_scale_k
        kk4 = col_idx % 4
        kk = col_idx // 4

        # Load scales for A
        sfa_ptrs = sfa_base + kk4[None, :] * stride_sfa_3 + kk[None, :] * stride_sfa_4
        mask_sfa = mask_m[:, None] & (col_idx[None, :] < sf_k)
        scale_a = tl.load(sfa_ptrs, mask=mask_sfa, other=1.0)

        # Load scales for B
        sfb_ptrs = sfb_base + kk4[None, :] * stride_sfb_3 + kk[None, :] * stride_sfb_4
        mask_sfb = mask_n[:, None] & (col_idx[None, :] < sf_k)
        scale_b = tl.load(sfb_ptrs, mask=mask_sfb, other=1.0)

        # Scaled dot product for FP4
        acc = tl.dot_scaled(a, scale_a, "e2m1", tl.trans(b), scale_b, "e2m1", acc)

    # Store output
    offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    c_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    mask_c = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)

    c = acc.to(tl.float16)
    tl.store(c_ptrs, c, mask=mask_c)


def kernel_function(abc_tensors, sfasfb_reordered_tensors, problem_sizes):
    """NVFP4 block-scaled group GEMM wrapper using optimized Triton kernel."""
    SF_VEC_SIZE = 16
    results = []

    for group_idx, ((a, b, c), (sfa_reordered, sfb_reordered), (m, n, k, l)) in enumerate(
        zip(abc_tensors, sfasfb_reordered_tensors, problem_sizes)
    ):
        a_uint8 = a.view(torch.uint8)
        b_uint8 = b.view(torch.uint8)

        for l_idx in range(l):
            a_slice = a_uint8[:, :, l_idx].contiguous()
            b_slice = b_uint8[:, :, l_idx].contiguous()
            c_slice = c[:, :, l_idx]

            BLOCK_M = 128
            BLOCK_N = 128
            BLOCK_K = 256
            GROUP_SIZE_M = 8

            grid = (ceil_div(m, BLOCK_M) * ceil_div(n, BLOCK_N),)

            nvfp4_group_gemm_kernel_optimized[grid](
                a_slice, b_slice, c_slice,
                sfa_reordered, sfb_reordered,
                m, n, k,
                a_slice.stride(0), a_slice.stride(1),
                b_slice.stride(0), b_slice.stride(1),
                c_slice.stride(0), c_slice.stride(1),
                sfa_reordered.stride(0), sfa_reordered.stride(1), sfa_reordered.stride(2),
                sfa_reordered.stride(3), sfa_reordered.stride(4), sfa_reordered.stride(5),
                sfb_reordered.stride(0), sfb_reordered.stride(1), sfb_reordered.stride(2),
                sfb_reordered.stride(3), sfb_reordered.stride(4), sfb_reordered.stride(5),
                l_idx,
                BLOCK_M=BLOCK_M,
                BLOCK_N=BLOCK_N,
                BLOCK_K=BLOCK_K,
                GROUP_SIZE_M=GROUP_SIZE_M,
                SF_VEC_SIZE=SF_VEC_SIZE,
                num_warps=8,
                num_stages=4,
            )

        results.append(c)

    return results



# Wrapper for the evaluation system which passes a single data argument
def custom_kernel(data):
    """Wrapper that unpacks data from the evaluation system."""
    abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
    return kernel_function(abc_tensors, sfasfb_reordered_tensors, problem_sizes)
scrolls · 195 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