Skip to content
KernelIndex
Search⌘K

submission 125299

shiyeegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

node_159.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-125299?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.3µs
#171 of 369
2025-12-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d6d555ea8b22d41a99a0a2e5734fe581287583f7e7636116fdae8a277d226d26
license declaredunknown
license concludedunknown
authorsshiyeegao
imported2026-08-15

Techniques

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

num-warps = 8num_warps = 8 if block_n == 256 else 4
persistent-kernelnum_pid = tl.num_programs(axis=0)
split-kSPLIT_K: tl.constexpr,
stages = 5num_stages = 5
tile-m = 128BLOCK_M = 128

Kernel source

node_159.py283 lines
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor


# -------------------------------------------------------------------------
# Kernel Metadata
# -------------------------------------------------------------------------
def _matmul_launch_metadata(grid, kernel, args):
    M, N, K = args["M"], args["N"], args["K"]
    return {
        "name": f"{kernel.name} [M={M}, N={N}, K={K}]",
        "flops": 2.0 * M * N * K,
    }


# -------------------------------------------------------------------------
# Block-Scaled NVFP4 GEMM with TMA (Persistent)
# -------------------------------------------------------------------------
@triton.jit(launch_metadata=_matmul_launch_metadata)
def bmm_fp4_tma_kernel(
    a_desc,         # TMA descriptor for A: [M, L, K/2]
    a_scale_desc,   # TMA descriptor for packed scales of A
    b_desc,         # TMA descriptor for B: [N, L, K/2]
    b_scale_desc,   # TMA descriptor for packed scales of B
    c_ptr,          # Output pointer: [M, N, L]
    # Strides
    stride_cm, stride_cn, stride_cl,
    # Dimensions
    M, N, K, L,
    # Constants
    ELEM_PER_BYTE: tl.constexpr,
    GROUP_SZ: tl.constexpr,
    SPLIT_K: tl.constexpr,
    # Block Tuning
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    REP_M: tl.constexpr,
    REP_N: tl.constexpr,
    REP_K: tl.constexpr,
    # Compiler Hints
    NUM_STAGES: tl.constexpr,
):
    output_dtype: tl.constexpr = tl.float16 if SPLIT_K == 1 else tl.float32
    acc_dtype: tl.constexpr = tl.float32

    BLOCK_K_ELEM_PER_BYTE: tl.constexpr = BLOCK_K // ELEM_PER_BYTE
    BLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZ

    tl.static_assert(BLOCK_M == 128)
    tl.static_assert(BLOCK_N % 128 == 0)
    tl.static_assert(BLOCK_K % (GROUP_SZ * 4) == 0)

    pid = tl.program_id(axis=0)
    num_pid = tl.num_programs(axis=0)

    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    num_tiles_per_batch = num_pid_m * num_pid_n * SPLIT_K
    total_tiles = num_tiles_per_batch * L

    k_tiles_total = tl.cdiv(K, BLOCK_K)
    part_k = tl.cdiv(k_tiles_total, SPLIT_K)

    # Persistent loop keeps blocks active across tiles
    for linear_id in tl.range(pid, total_tiles, num_pid, num_stages=NUM_STAGES):
        tile_split = linear_id % SPLIT_K
        tmp = linear_id // SPLIT_K
        tile_n = tmp % num_pid_n
        tmp = tmp // num_pid_n
        tile_m = tmp % num_pid_m
        tile_l = tmp // num_pid_m

        offs_am = tile_m * BLOCK_M
        offs_bn = tile_n * BLOCK_N
        offs_scale_m = tile_m * REP_M
        offs_scale_n = tile_n * REP_N

        k_start = tile_split * part_k
        k_end = tl.minimum(k_tiles_total, k_start + part_k)

        if k_start < k_end:
            accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)

            for k_idx in tl.range(k_start, k_end):
                offs_k = k_idx * BLOCK_K_ELEM_PER_BYTE
                offs_scale_k = k_idx * REP_K

                # TMA loads for packed FP4 tiles
                a_tile = a_desc.load([offs_am, tile_l, offs_k])
                b_tile = b_desc.load([offs_bn, tile_l, offs_k])

                a_tile = a_tile.reshape(BLOCK_M, BLOCK_K_ELEM_PER_BYTE)
                b_tile = b_tile.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)

                # Load block scales (already swizzled for WGMMA layout)
                scale_a_pack = a_scale_desc.load([tile_l, offs_scale_m, offs_scale_k, 0, 0])
                scale_b_pack = b_scale_desc.load([tile_l, offs_scale_n, offs_scale_k, 0, 0])

                scale_a = (
                    scale_a_pack.reshape(REP_M, REP_K, 32, 4, 4)
                    .trans(0, 3, 2, 1, 4)
                    .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
                )
                scale_b = (
                    scale_b_pack.reshape(REP_N, REP_K, 32, 4, 4)
                    .trans(0, 3, 2, 1, 4)
                    .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
                )

                accumulator = tl.dot_scaled(
                    a_tile,
                    scale_a,
                    "e2m1",
                    b_tile.T,
                    scale_b,
                    "e2m1",
                    accumulator,
                )

            offs_cm = tile_m * BLOCK_M + tl.arange(0, BLOCK_M)
            offs_cn = tile_n * BLOCK_N + tl.arange(0, BLOCK_N)

            c_offset = (
                offs_cm[:, None] * stride_cm
                + offs_cn[None, :] * stride_cn
                + tile_l * stride_cl
            )

            mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)

            if SPLIT_K == 1:
                tl.store(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)
            else:
                tl.atomic_add(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)


# -------------------------------------------------------------------------
# Host API
# -------------------------------------------------------------------------
def _select_config(M: int, N: int, K: int, num_sms: int):
    """Shape-aware tile selection tuned for leaderboard shapes."""
    wide_n = N >= 4096

    block_n = 256 if (wide_n and K < 2048) else 128
    block_k = 256

    if K >= 4096:
        num_stages = 5
    elif K >= 2048:
        num_stages = 4
    else:
        num_stages = 3

    split_k = 1

    num_warps = 8 if block_n == 256 else 4

    # Persistent grid multiplier; slightly conservative to keep occupancy.
    sm_mult = 8

    return {
        "BLOCK_N": block_n,
        "BLOCK_K": block_k,
        "NUM_STAGES": num_stages,
        "NUM_WARPS": num_warps,
        "SM_MULT": sm_mult,
        "SPLIT_K": split_k,
    }


@torch.inference_mode()
def custom_kernel(data):
    """Entry point expected by evaluator."""
    a_tensor, b_tensor, _, _, sfa_permuted, sfb_permuted, c_tensor = data

    BLOCK_M = 128
    GROUP_SZ = 16

    M, K_half, L = a_tensor.shape
    N = b_tensor.shape[0]
    K = 2 * K_half
    ELEM_PER_BYTE = 2

    num_sms = torch.cuda.get_device_properties(a_tensor.device).multi_processor_count
    cfg = _select_config(M, N, K, num_sms)
    BLOCK_N = cfg["BLOCK_N"]
    BLOCK_K = cfg["BLOCK_K"]
    NUM_STAGES = cfg["NUM_STAGES"]
    NUM_WARPS = cfg["NUM_WARPS"]
    SM_MULT = cfg["SM_MULT"]
    split_k = cfg["SPLIT_K"]

    REP_M = BLOCK_M // 128
    REP_N = BLOCK_N // 128
    REP_K = BLOCK_K // GROUP_SZ // 4

    # Reorder A/B for TMA: place K as innermost for contiguous loads
    a_tma = a_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()
    b_tma = b_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()

    a_desc = TensorDescriptor.from_tensor(
        a_tma,
        block_shape=[BLOCK_M, 1, BLOCK_K // ELEM_PER_BYTE],
    )
    b_desc = TensorDescriptor.from_tensor(
        b_tma,
        block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],
    )

    rest_m = M // 128
    rest_n = N // 128
    rest_k = triton.cdiv(K, GROUP_SZ) // 4

    sfa_packed = (
        sfa_permuted.permute(5, 2, 4, 0, 1, 3)
        .contiguous()
        .view(L, rest_m, rest_k, 2, 256)
    )
    sfb_packed = (
        sfb_permuted.permute(5, 2, 4, 0, 1, 3)
        .contiguous()
        .view(L, rest_n, rest_k, 2, 256)
    )

    a_scale_desc = TensorDescriptor.from_tensor(
        sfa_packed,
        block_shape=[1, REP_M, REP_K, 2, 256],
    )
    b_scale_desc = TensorDescriptor.from_tensor(
        sfb_packed,
        block_shape=[1, REP_N, REP_K, 2, 256],
    )

    base_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N) * L
    num_tiles = base_tiles * split_k

    c_buffer = c_tensor
    if split_k > 1:
        c_buffer = torch.zeros_like(c_tensor, dtype=torch.float32)

    stride_cm, stride_cn, stride_cl = c_buffer.stride()

    grid_target = num_sms * SM_MULT
    grid = (max(1, min(num_tiles, grid_target)),)

    bmm_fp4_tma_kernel[grid](
        a_desc,
        a_scale_desc,
        b_desc,
        b_scale_desc,
        c_buffer,
        stride_cm,
        stride_cn,
        stride_cl,
        M,
        N,
        K,
        L,
        ELEM_PER_BYTE,
        GROUP_SZ,
        split_k,
        BLOCK_M,
        BLOCK_N,
        BLOCK_K,
        REP_M,
        REP_N,
        REP_K,
        NUM_STAGES,
        num_warps=NUM_WARPS,
        num_stages=NUM_STAGES,
    )

    if split_k > 1:
        c_tensor.copy_(c_buffer.to(dtype=c_tensor.dtype))

    return c_tensor


__all__ = ["custom_kernel"]
scrolls · 283 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 121413.

- from __future__ import annotations
+ import torch
+ import triton
+ import triton.language as tl
+ from triton.tools.tensor_descriptor import TensorDescriptor
- from typing import Tuple
- import torch
+ # -------------------------------------------------------------------------
+ # Kernel Metadata
+ # -------------------------------------------------------------------------
+ def _matmul_launch_metadata(grid, kernel, args):
+ M, N, K = args["M"], args["N"], args["K"]
+ return {
+ "name": f"{kernel.name} [M={M}, N={N}, K={K}]",
+ "flops": 2.0 * M * N * K,
+ }
- def _prepare_scales_batch(scale_perm: torch.Tensor) -> torch.Tensor:
- permuted = scale_perm.permute(5, 2, 4, 0, 1, 3)
- return permuted.contiguous().view(scale_perm.size(-1), -1)
+ # -------------------------------------------------------------------------
+ # Block-Scaled NVFP4 GEMM with TMA (Persistent)
+ # -------------------------------------------------------------------------
+ @triton.jit(launch_metadata=_matmul_launch_metadata)
+ def bmm_fp4_tma_kernel(
+ a_desc, # TMA descriptor for A: [M, L, K/2]
+ a_scale_desc, # TMA descriptor for packed scales of A
+ b_desc, # TMA descriptor for B: [N, L, K/2]
+ b_scale_desc, # TMA descriptor for packed scales of B
+ c_ptr, # Output pointer: [M, N, L]
+ # Strides
+ stride_cm, stride_cn, stride_cl,
+ # Dimensions
+ M, N, K, L,
+ # Constants
+ ELEM_PER_BYTE: tl.constexpr,
+ GROUP_SZ: tl.constexpr,
+ SPLIT_K: tl.constexpr,
+ # Block Tuning
+ BLOCK_M: tl.constexpr,
+ BLOCK_N: tl.constexpr,
+ BLOCK_K: tl.constexpr,
+ REP_M: tl.constexpr,
+ REP_N: tl.constexpr,
+ REP_K: tl.constexpr,
+ # Compiler Hints
+ NUM_STAGES: tl.constexpr,
+ ):
+ output_dtype: tl.constexpr = tl.float16 if SPLIT_K == 1 else tl.float32
+ acc_dtype: tl.constexpr = tl.float32
+ BLOCK_K_ELEM_PER_BYTE: tl.constexpr = BLOCK_K // ELEM_PER_BYTE
+ BLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZ
+ tl.static_assert(BLOCK_M == 128)
+ tl.static_assert(BLOCK_N % 128 == 0)
+ tl.static_assert(BLOCK_K % (GROUP_SZ * 4) == 0)
+
+ pid = tl.program_id(axis=0)
+ num_pid = tl.num_programs(axis=0)
+
+ num_pid_m = tl.cdiv(M, BLOCK_M)
+ num_pid_n = tl.cdiv(N, BLOCK_N)
+ num_tiles_per_batch = num_pid_m * num_pid_n * SPLIT_K
+ total_tiles = num_tiles_per_batch * L
+
+ k_tiles_total = tl.cdiv(K, BLOCK_K)
+ part_k = tl.cdiv(k_tiles_total, SPLIT_K)
+
+ # Persistent loop keeps blocks active across tiles
+ for linear_id in tl.range(pid, total_tiles, num_pid, num_stages=NUM_STAGES):
+ tile_split = linear_id % SPLIT_K
+ tmp = linear_id // SPLIT_K
+ tile_n = tmp % num_pid_n
+ tmp = tmp // num_pid_n
+ tile_m = tmp % num_pid_m
+ tile_l = tmp // num_pid_m
+
+ offs_am = tile_m * BLOCK_M
+ offs_bn = tile_n * BLOCK_N
+ offs_scale_m = tile_m * REP_M
+ offs_scale_n = tile_n * REP_N
+
+ k_start = tile_split * part_k
+ k_end = tl.minimum(k_tiles_total, k_start + part_k)
+
+ if k_start < k_end:
+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)
+
+ for k_idx in tl.range(k_start, k_end):
+ offs_k = k_idx * BLOCK_K_ELEM_PER_BYTE
+ offs_scale_k = k_idx * REP_K
+
+ # TMA loads for packed FP4 tiles
+ a_tile = a_desc.load([offs_am, tile_l, offs_k])
+ b_tile = b_desc.load([offs_bn, tile_l, offs_k])
+
+ a_tile = a_tile.reshape(BLOCK_M, BLOCK_K_ELEM_PER_BYTE)
+ b_tile = b_tile.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)
+
+ # Load block scales (already swizzled for WGMMA layout)
+ scale_a_pack = a_scale_desc.load([tile_l, offs_scale_m, offs_scale_k, 0, 0])
+ scale_b_pack = b_scale_desc.load([tile_l, offs_scale_n, offs_scale_k, 0, 0])
+
+ scale_a = (
+ scale_a_pack.reshape(REP_M, REP_K, 32, 4, 4)
+ .trans(0, 3, 2, 1, 4)
+ .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
+ )
+ scale_b = (
+ scale_b_pack.reshape(REP_N, REP_K, 32, 4, 4)
+ .trans(0, 3, 2, 1, 4)
+ .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
+ )
+
+ accumulator = tl.dot_scaled(
+ a_tile,
+ scale_a,
+ "e2m1",
+ b_tile.T,
+ scale_b,
+ "e2m1",
+ accumulator,
+ )
+
+ offs_cm = tile_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_cn = tile_n * BLOCK_N + tl.arange(0, BLOCK_N)
+
+ c_offset = (
+ offs_cm[:, None] * stride_cm
+ + offs_cn[None, :] * stride_cn
+ + tile_l * stride_cl
+ )
+
+ mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
+
+ if SPLIT_K == 1:
+ tl.store(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)
+ else:
+ tl.atomic_add(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)
+
+
+ # -------------------------------------------------------------------------
+ # Host API
+ # -------------------------------------------------------------------------
+ def _select_config(M: int, N: int, K: int, num_sms: int):
+ """Shape-aware tile selection tuned for leaderboard shapes."""
+ wide_n = N >= 4096
+
+ block_n = 256 if (wide_n and K < 2048) else 128
+ block_k = 256
+
+ if K >= 4096:
+ num_stages = 5
+ elif K >= 2048:
+ num_stages = 4
+ else:
+ num_stages = 3
+
+ split_k = 1
+
+ num_warps = 8 if block_n == 256 else 4
+
+ # Persistent grid multiplier; slightly conservative to keep occupancy.
+ sm_mult = 8
+
+ return {
+ "BLOCK_N": block_n,
+ "BLOCK_K": block_k,
+ "NUM_STAGES": num_stages,
+ "NUM_WARPS": num_warps,
+ "SM_MULT": sm_mult,
+ "SPLIT_K": split_k,
+ }
+
+
@torch.inference_mode()
- def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
- a, b, _, _, sfa_perm, sfb_perm, c = data
+ def custom_kernel(data):
+ """Entry point expected by evaluator."""
+ a_tensor, b_tensor, _, _, sfa_permuted, sfb_permuted, c_tensor = data
- _, _, l = c.shape
+ BLOCK_M = 128
+ GROUP_SZ = 16
- scale_a_batch = _prepare_scales_batch(sfa_perm)
- scale_b_batch = _prepare_scales_batch(sfb_perm)
+ M, K_half, L = a_tensor.shape
+ N = b_tensor.shape[0]
+ K = 2 * K_half
+ ELEM_PER_BYTE = 2
- for i in range(l):
- res = torch._scaled_mm(
- a[:, :, i],
- b[:, :, i].transpose(0, 1),
- scale_a_batch[i],
- scale_b_batch[i],
- bias=None,
- out_dtype=torch.float16,
- )
- c[:, :, i].copy_(res)
+ num_sms = torch.cuda.get_device_properties(a_tensor.device).multi_processor_count
+ cfg = _select_config(M, N, K, num_sms)
+ BLOCK_N = cfg["BLOCK_N"]
+ BLOCK_K = cfg["BLOCK_K"]
+ NUM_STAGES = cfg["NUM_STAGES"]
+ NUM_WARPS = cfg["NUM_WARPS"]
+ SM_MULT = cfg["SM_MULT"]
+ split_k = cfg["SPLIT_K"]
- return c
+ REP_M = BLOCK_M // 128
+ REP_N = BLOCK_N // 128
+ REP_K = BLOCK_K // GROUP_SZ // 4
+ # Reorder A/B for TMA: place K as innermost for contiguous loads
+ a_tma = a_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()
+ b_tma = b_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()
- __all__ = ["custom_kernel"]
No newline at end of file
+ a_desc = TensorDescriptor.from_tensor(
+ a_tma,
+ block_shape=[BLOCK_M, 1, BLOCK_K // ELEM_PER_BYTE],
+ )
+ b_desc = TensorDescriptor.from_tensor(
+ b_tma,
+ block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],
+ )
+
+ rest_m = M // 128
+ rest_n = N // 128
+ rest_k = triton.cdiv(K, GROUP_SZ) // 4
+
+ sfa_packed = (
+ sfa_permuted.permute(5, 2, 4, 0, 1, 3)
+ .contiguous()
+ .view(L, rest_m, rest_k, 2, 256)
+ )
+ sfb_packed = (
+ sfb_permuted.permute(5, 2, 4, 0, 1, 3)
+ .contiguous()
+ .view(L, rest_n, rest_k, 2, 256)
+ )
+
+ a_scale_desc = TensorDescriptor.from_tensor(
+ sfa_packed,
+ block_shape=[1, REP_M, REP_K, 2, 256],
+ )
+ b_scale_desc = TensorDescriptor.from_tensor(
+ sfb_packed,
+ block_shape=[1, REP_N, REP_K, 2, 256],
+ )
+
+ base_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N) * L
+ num_tiles = base_tiles * split_k
+
+ c_buffer = c_tensor
+ if split_k > 1:
+ c_buffer = torch.zeros_like(c_tensor, dtype=torch.float32)
+
+ stride_cm, stride_cn, stride_cl = c_buffer.stride()
+
+ grid_target = num_sms * SM_MULT
+ grid = (max(1, min(num_tiles, grid_target)),)
+
+ bmm_fp4_tma_kernel[grid](
+ a_desc,
+ a_scale_desc,
+ b_desc,
+ b_scale_desc,
+ c_buffer,
+ stride_cm,
+ stride_cn,
+ stride_cl,
+ M,
+ N,
+ K,
+ L,
+ ELEM_PER_BYTE,
+ GROUP_SZ,
+ split_k,
+ BLOCK_M,
+ BLOCK_N,
+ BLOCK_K,
+ REP_M,
+ REP_N,
+ REP_K,
+ NUM_STAGES,
+ num_warps=NUM_WARPS,
+ num_stages=NUM_STAGES,
+ )
+
+ if split_k > 1:
+ c_tensor.copy_(c_buffer.to(dtype=c_tensor.dtype))
+
+ return c_tensor
+
+
+ __all__ = ["custom_kernel"]
scrolls · 306 diff lines total

Best evidence level for this revision: reported

JSON