Skip to content
KernelIndex
Search⌘K

submission 496271

Seraphim · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

helion_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-496271?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
61.4µs
#213 of 310
2026-02-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f772314f0e8a51a38318eb6f510a569ab8ccac962841898cb622a8be5f3756d1
license declaredunknown
license concludedunknown
authorsSeraphim
imported2026-08-15

Techniques

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

fp4NVFP4 block-scaled group GEMM using Triton's tl.dot_scaled for hardware FP4 tensor cores.
stages = 4NUM_STAGES = 4
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128

Kernel source

helion_submission.py152 lines
"""
NVFP4 block-scaled group GEMM using Triton's tl.dot_scaled for hardware FP4 tensor cores.

This kernel uses Triton's block-scaled matmul support (tl.dot_scaled with "e2m1" format)
to leverage B200's native FP4 tensor cores. The scale factors are loaded via TMA in the
preshuffled cuBLAS blocked layout and transposed in-register to match tl.dot_scaled's
expected format.

Adapted from Triton's block-scaled matmul tutorial:
  https://triton-lang.org/main/getting-started/tutorials/10-block-scaled-matmul.html

Key insight: Helion currently compiles hl.dot -> tl.dot, which only supports
float16/bfloat16/float8/int8. Triton's tl.dot_scaled is a separate instruction that
directly uses FP4 tensor cores with block scaling, but Helion doesn't expose it yet.
When Helion adds hl.dot_scaled (or equivalent), this kernel could be written in Helion
instead of raw Triton.

Previous Helion-based approach (software dequant + float16 GEMM) was ~30-50x slower
because it had to dequantize FP4->float16, expand scale factors, and GEMM on 4x larger
data across multiple kernel launches.
"""

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

# NVFP4 block scaling constants
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
VEC_SIZE = 16       # nvfp4: one scale per 16 FP4 elements along K
ELEM_PER_BYTE = 2   # FP4: 2 elements packed per byte
NUM_STAGES = 4


@triton.jit
def nvfp4_gemm_kernel(
    a_desc, a_scale_desc, b_desc, b_scale_desc, c_desc,
    M, N, K,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    VEC_SIZE: tl.constexpr, ELEM_PER_BYTE: tl.constexpr,
    rep_m: tl.constexpr, rep_n: tl.constexpr, rep_k: tl.constexpr,
    NUM_STAGES: tl.constexpr,
):
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    pid_m = pid % num_pid_m
    pid_n = pid // num_pid_m

    offs_am = pid_m * BLOCK_M
    offs_bn = pid_n * BLOCK_N
    offs_k = 0
    offs_scale_m = pid_m * rep_m
    offs_scale_n = pid_n * rep_n
    offs_scale_k = 0

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

    for k in tl.range(0, tl.cdiv(K, BLOCK_K), num_stages=NUM_STAGES):
        # Load packed FP4 tiles via TMA
        a = a_desc.load([offs_am, offs_k])
        b = b_desc.load([offs_bn, offs_k])

        # Load preshuffled scale factors via TMA
        # Scale layout in memory: [1, rest_m, rest_k, 2, 256] (uint8 view of float8)
        # Each [2, 256] block = [32, 4, 4] = 512 bytes of scale data
        scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
        scale_b = b_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])

        # Transpose preshuffled scales to the 2D layout expected by tl.dot_scaled:
        # [rep_m, rep_k, 32, 4, 4] -> [rep_m, 4, 32, rep_k, 4] -> [BLOCK_M, BLOCK_K // VEC_SIZE]
        # See: https://docs.nvidia.com/cuda/cublas/#d-block-scaling-factors-layout
        scale_a = scale_a.reshape(rep_m, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_M, BLOCK_K // VEC_SIZE)
        scale_b = scale_b.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_N, BLOCK_K // VEC_SIZE)

        # Hardware FP4 tensor core matmul with block scaling
        # a: [BLOCK_M, BLOCK_K//2] packed e2m1, b.T: [BLOCK_K//2, BLOCK_N] packed e2m1
        acc = tl.dot_scaled(a, scale_a, "e2m1", b.T, scale_b, "e2m1", acc)

        offs_k += BLOCK_K // ELEM_PER_BYTE
        offs_scale_k += rep_k

    c_desc.store([offs_am, offs_bn], acc.to(tl.float16))


def _prepare_scale_for_tma(sfa_reordered):
    """
    Convert reordered scale factors to 5D TMA format.

    Input:  [32, 4, rest_m, 4, rest_k, L] float8_e4m3fn (on GPU, from generate_input)
    Output: [1, rest_m, rest_k, 2, 256] float8_e4m3fn (contiguous, for TMA descriptor)

    The reordered tensor stores scales in the preshuffled cuBLAS blocked layout:
        reordered[mm32, mm4, block_m, kk4, block_k, l] = original[i, j, l]
    where mm32 = i%32, mm4 = (i%128)//32, block_m = i//128, kk4 = j%4, block_k = j//4.

    We permute to [rest_m, rest_k, 32, 4, 4] then reshape the last 3 dims (512 bytes)
    into [2, 256] for efficient TMA loads.
    """
    s = sfa_reordered[..., 0]                          # [32, 4, rest_m, 4, rest_k]
    rest_m, rest_k = s.shape[2], s.shape[4]
    s = s.permute(2, 4, 0, 1, 3).contiguous()          # [rest_m, rest_k, 32, 4, 4]
    return s.reshape(1, rest_m, rest_k, 2, 256).contiguous()


def custom_kernel(data: input_t) -> output_t:
    abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data

    rep_m = BLOCK_M // 128   # = 1
    rep_n = BLOCK_N // 128   # = 1
    rep_k = BLOCK_K // VEC_SIZE // 4  # = 4

    result_tensors = []
    for (a, b, c), (sfa_reord, sfb_reord), (m, n, k, l) in zip(
        abc_tensors, sfasfb_reordered_tensors, problem_sizes
    ):
        for l_idx in range(l):
            # View FP4 packed tensors as uint8 for TMA
            a_packed = a[:, :, l_idx].contiguous().view(torch.uint8)  # [M, K//2]
            b_packed = b[:, :, l_idx].contiguous().view(torch.uint8)  # [N, K//2]

            # Prepare scale factors: reordered GPU tensors -> 5D TMA format
            a_scale = _prepare_scale_for_tma(sfa_reord)
            b_scale = _prepare_scale_for_tma(sfb_reord)

            # Create TMA descriptors
            a_desc = TensorDescriptor.from_tensor(a_packed, [BLOCK_M, BLOCK_K // ELEM_PER_BYTE])
            b_desc = TensorDescriptor.from_tensor(b_packed, [BLOCK_N, BLOCK_K // ELEM_PER_BYTE])

            c_out = torch.empty(m, n, dtype=torch.float16, device="cuda")
            c_desc = TensorDescriptor.from_tensor(c_out, [BLOCK_M, BLOCK_N])

            a_scale_desc = TensorDescriptor.from_tensor(a_scale, [1, rep_m, rep_k, 2, 256])
            b_scale_desc = TensorDescriptor.from_tensor(b_scale, [1, rep_n, rep_k, 2, 256])

            # Launch Triton kernel
            grid = (triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N), 1)
            nvfp4_gemm_kernel[grid](
                a_desc, a_scale_desc, b_desc, b_scale_desc, c_desc,
                m, n, k,
                BLOCK_M, BLOCK_N, BLOCK_K, VEC_SIZE, ELEM_PER_BYTE,
                rep_m, rep_n, rep_k, NUM_STAGES,
            )

            c[:, :, l_idx] = c_out

        result_tensors.append(c)

    return result_tensors
scrolls · 152 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