Skip to content
KernelIndex
Search⌘K

submission 327716

HayatoFujihara · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

nvfp4_dual_gemm_v12.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-327716?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 dual GEMMsuite of 4 cases
NVIDIA B200
23.5µs
#196 of 420
2026-01-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f7405e11b6e8341075a2dd1d45ffd98c5511eaf1eebb4fd1475dde646d67edb6
license declaredunknown
license concludedunknown
authorsHayatoFujihara
imported2026-08-26

Techniques

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

fp4NVFP4 Dual GEMM + SiLU Fusion Kernel v12
fused-epilogue- SiLU fused in epilogue
stages = 4NUM_STAGES = 4 # Pipeline stages (optimal)
tile-k = 256BLOCK_K = 256 # K-dimension tile size (FP4 native K)
tile-m = 128BLOCK_M = 128 # M-dimension tile size
tile-n = 128BLOCK_N = 128 # N-dimension tile size

Kernel source

nvfp4_dual_gemm_v12.py268 lines
"""
NVFP4 Dual GEMM + SiLU Fusion Kernel v12
============================================================================
Triton tl.dot_scaled Implementation - Host-side Optimization

Based on v11 (24.8us). Changes:
  1. 2D grid instead of 1D (removes div/mod overhead in kernel)
  2. L=1 specialization (removes batch loop overhead)
  3. Inline scale conversion (reduces function call overhead)
  4. Pre-computed constants outside batch loop

Target: 13.915us (1st place)
============================================================================
"""
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
from task import input_t, output_t


# ============================================================================
# Configuration - Tuned for Blackwell B200
# ============================================================================
BLOCK_M = 128       # M-dimension tile size
BLOCK_N = 128       # N-dimension tile size
BLOCK_K = 256       # K-dimension tile size (FP4 native K)
VEC_SIZE = 16       # NVFP4: 16 elements per E4M3 scale
ELEM_PER_BYTE = 2   # FP4 packs 2 elements per byte
NUM_STAGES = 4      # Pipeline stages (optimal)


# ============================================================================
# Triton Kernel - Fused Dual GEMM + SiLU
# ============================================================================
@triton.jit
def fused_dual_gemm_silu_kernel(
    # TMA descriptors
    a_desc,
    a_scale_desc,
    b1_desc,
    b1_scale_desc,
    b2_desc,
    b2_scale_desc,
    c_desc,
    # Shape parameters
    M: tl.constexpr,
    N: tl.constexpr,
    K: tl.constexpr,
    # Tile sizes
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    VEC_SIZE: tl.constexpr,
    # Iteration parameters
    rep_m: tl.constexpr,
    rep_n: tl.constexpr,
    rep_k: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    ELEM_PER_BYTE: tl.constexpr,
):
    """
    Fused Dual GEMM + SiLU kernel using TMA and tl.dot_scaled.

    Computation: C = silu(A @ B1.T) * (A @ B2.T)

    Optimizations:
    - TMA for efficient memory access
    - A matrix loaded once, reused for both B1 and B2 products
    - Intermediate results kept in registers
    - SiLU fused in epilogue
    - 2D grid for direct program ID access (no div/mod)
    """
    # Program ID (2D grid - direct access, no div/mod)
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    # Offset calculations
    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

    # Initialize accumulators (FP32 for precision)
    acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # K-dimension loop with pipelining
    num_k_iters = tl.cdiv(K, BLOCK_K)
    for _ in tl.range(0, num_k_iters, num_stages=NUM_STAGES):
        # Load A tile via TMA
        a = a_desc.load([offs_am, offs_k])

        # Load and transform A scale factors
        raw_scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
        scale_a = raw_scale_a.reshape(rep_m, rep_k, 32, 4, 4)
        scale_a = scale_a.trans(0, 3, 2, 1, 4)
        scale_a = scale_a.reshape(BLOCK_M, BLOCK_K // VEC_SIZE)

        # Load B1 tile
        b1 = b1_desc.load([offs_bn, offs_k])

        # Load and transform B1 scale factors
        raw_scale_b1 = b1_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
        scale_b1 = raw_scale_b1.reshape(rep_n, rep_k, 32, 4, 4)
        scale_b1 = scale_b1.trans(0, 3, 2, 1, 4)
        scale_b1 = scale_b1.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)

        # Load B2 tile
        b2 = b2_desc.load([offs_bn, offs_k])

        # Load and transform B2 scale factors
        raw_scale_b2 = b2_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
        scale_b2 = raw_scale_b2.reshape(rep_n, rep_k, 32, 4, 4)
        scale_b2 = scale_b2.trans(0, 3, 2, 1, 4)
        scale_b2 = scale_b2.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)

        # Dual GEMM with A reuse
        # tl.dot_scaled handles FP4 E2M1 with E4M3 block scales natively
        acc1 = tl.dot_scaled(a, scale_a, "e2m1", b1.T, scale_b1, "e2m1", acc1)
        acc2 = tl.dot_scaled(a, scale_a, "e2m1", b2.T, scale_b2, "e2m1", acc2)

        # Update offsets for next iteration
        offs_k += BLOCK_K // ELEM_PER_BYTE
        offs_scale_k += rep_k

    # Fused SiLU + multiply epilogue
    # silu(x) = x * sigmoid(x)
    result = (acc1 * tl.sigmoid(acc1)) * acc2

    # Convert to FP16 for output
    result = result.to(tl.float16)

    # Store result via TMA
    c_desc.store([offs_am, offs_bn], result)


# ============================================================================
# Host-side Entry Point
# ============================================================================
def custom_kernel(data: input_t) -> output_t:
    """
    Triton Fused Dual GEMM + SiLU (TMA + tl.dot_scaled)

    Computation: C = silu(A @ B1.T) * (A @ B2.T)

    Optimizations applied:
    - L=1 specialization: no batch loop overhead
    - Inline scale conversion: no function call overhead
    - Pre-computed shapes: avoid repeated calculations
    """
    # Unpack inputs (use permuted scales for TMA compatibility)
    a, b1, b2, _, _, _, sfa_perm, sfb1_perm, sfb2_perm, c = data

    # Get dimensions
    M, N, L = c.shape
    K = a.shape[1] * 2  # FP4 packs 2 elements per byte

    # Pre-computed constants (avoid recalculation in loop)
    rep_m = BLOCK_M // 128  # = 1
    rep_n = BLOCK_N // 128  # = 1
    rep_k = BLOCK_K // VEC_SIZE // 4  # = 4
    k_bytes = BLOCK_K // ELEM_PER_BYTE  # = 128

    # TMA block shapes (computed once)
    a_block_shape = [BLOCK_M, k_bytes]
    b_block_shape = [BLOCK_N, k_bytes]
    c_block_shape = [BLOCK_M, BLOCK_N]
    scale_block_shape = [1, rep_m, rep_k, 2, 256]

    # Grid configuration
    grid = (M // BLOCK_M, N // BLOCK_N)

    # L=1 specialization: inline everything for minimal overhead
    if L == 1:
        # Direct slice without loop variable
        a_u8 = a[:, :, 0].view(torch.uint8)
        b1_u8 = b1[:, :, 0].view(torch.uint8)
        b2_u8 = b2[:, :, 0].view(torch.uint8)
        c_l = c[:, :, 0]

        # Inline scale conversion for sfa
        # (32, 4, rest_m, 4, rest_k, L) -> (1, rest_m, rest_k, 2, 256)
        rest_m = sfa_perm.shape[2]
        rest_k = sfa_perm.shape[4]
        sfa_tma = sfa_perm[:, :, :, :, :, 0].permute(2, 4, 0, 1, 3).reshape(
            1, rest_m, rest_k, 2, 256
        ).contiguous()

        # Inline scale conversion for sfb1
        rest_n = sfb1_perm.shape[2]
        sfb1_tma = sfb1_perm[:, :, :, :, :, 0].permute(2, 4, 0, 1, 3).reshape(
            1, rest_n, rest_k, 2, 256
        ).contiguous()

        # Inline scale conversion for sfb2
        sfb2_tma = sfb2_perm[:, :, :, :, :, 0].permute(2, 4, 0, 1, 3).reshape(
            1, rest_n, rest_k, 2, 256
        ).contiguous()

        # Create TMA tensor descriptors
        a_desc = TensorDescriptor.from_tensor(a_u8, a_block_shape)
        b1_desc = TensorDescriptor.from_tensor(b1_u8, b_block_shape)
        b2_desc = TensorDescriptor.from_tensor(b2_u8, b_block_shape)
        c_desc = TensorDescriptor.from_tensor(c_l, c_block_shape)
        a_scale_desc = TensorDescriptor.from_tensor(sfa_tma, scale_block_shape)
        b1_scale_desc = TensorDescriptor.from_tensor(sfb1_tma, scale_block_shape)
        b2_scale_desc = TensorDescriptor.from_tensor(sfb2_tma, scale_block_shape)

        # Launch kernel
        fused_dual_gemm_silu_kernel[grid](
            a_desc, a_scale_desc,
            b1_desc, b1_scale_desc,
            b2_desc, b2_scale_desc,
            c_desc,
            M=M, N=N, K=K,
            BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
            VEC_SIZE=VEC_SIZE,
            rep_m=rep_m, rep_n=rep_n, rep_k=rep_k,
            NUM_STAGES=NUM_STAGES, ELEM_PER_BYTE=ELEM_PER_BYTE,
        )
    else:
        # Generic batch loop for L > 1
        for l_idx in range(L):
            a_u8 = a[:, :, l_idx].view(torch.uint8)
            b1_u8 = b1[:, :, l_idx].view(torch.uint8)
            b2_u8 = b2[:, :, l_idx].view(torch.uint8)
            c_l = c[:, :, l_idx]

            # Scale conversion
            rest_m = sfa_perm.shape[2]
            rest_k = sfa_perm.shape[4]
            rest_n = sfb1_perm.shape[2]

            sfa_tma = sfa_perm[:, :, :, :, :, l_idx].permute(2, 4, 0, 1, 3).reshape(
                1, rest_m, rest_k, 2, 256
            ).contiguous()
            sfb1_tma = sfb1_perm[:, :, :, :, :, l_idx].permute(2, 4, 0, 1, 3).reshape(
                1, rest_n, rest_k, 2, 256
            ).contiguous()
            sfb2_tma = sfb2_perm[:, :, :, :, :, l_idx].permute(2, 4, 0, 1, 3).reshape(
                1, rest_n, rest_k, 2, 256
            ).contiguous()

            a_desc = TensorDescriptor.from_tensor(a_u8, a_block_shape)
            b1_desc = TensorDescriptor.from_tensor(b1_u8, b_block_shape)
            b2_desc = TensorDescriptor.from_tensor(b2_u8, b_block_shape)
            c_desc = TensorDescriptor.from_tensor(c_l, c_block_shape)
            a_scale_desc = TensorDescriptor.from_tensor(sfa_tma, scale_block_shape)
            b1_scale_desc = TensorDescriptor.from_tensor(sfb1_tma, scale_block_shape)
            b2_scale_desc = TensorDescriptor.from_tensor(sfb2_tma, scale_block_shape)

            fused_dual_gemm_silu_kernel[grid](
                a_desc, a_scale_desc,
                b1_desc, b1_scale_desc,
                b2_desc, b2_scale_desc,
                c_desc,
                M=M, N=N, K=K,
                BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
                VEC_SIZE=VEC_SIZE,
                rep_m=rep_m, rep_n=rep_n, rep_k=rep_k,
                NUM_STAGES=NUM_STAGES, ELEM_PER_BYTE=ELEM_PER_BYTE,
            )

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