Skip to content
KernelIndex
Search⌘K

submission 315436

Gusarich · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-315436?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
27.9µs
#240 of 420
2026-01-09

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ad593749a1091870c4bd4377c7557a5d794589fd821c3a35e3d26ee541c887e2
license declaredunknown
license concludedunknown
authorsGusarich
imported2026-08-26

Techniques

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

fp4PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
num-warps = 8num_warps=8,
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128

Kernel source

submission.py425 lines
import torch
from task import input_t, output_t
from utils import make_match_reference

# Scaling factor vector size for NVFP4 block scaling
sf_vec_size = 16

# ---------------------------
# Optional Triton fast path
# ---------------------------
_TRITON_AVAILABLE = False
try:
    import triton
    import triton.language as tl
    from triton.tools.tensor_descriptor import TensorDescriptor

    _TRITON_AVAILABLE = True
except Exception:
    _TRITON_AVAILABLE = False


# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b


# Helper function to convert scale factor tensor to blocked format
# Used only by the reference kernel (slow CPU-side path).
def to_blocked(input_matrix):
    rows, cols = input_matrix.shape

    # Please ensure rows and cols are multiples of 128 and 4 respectively
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)

    return rearranged.flatten()


def ref_kernel(data: input_t) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
    C = silu(A @ B1) * (A @ B2).
    """
    a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = (
        data
    )

    # Get dimensions from MxNxL layout
    m, n, l = c_ref.shape

    ref1 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
    ref2 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)

    for l_idx in range(l):
        # Convert the scale factor tensor to blocked format
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
        scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
        scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])

        # (m, k) @ (n, k).T -> (m, n)
        res1 = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b1_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b1.cuda(),
            bias=None,
            out_dtype=torch.float32,
        )
        ref1[:, :, l_idx] = res1

        res2 = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b2_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b2.cuda(),
            bias=None,
            out_dtype=torch.float32,
        )
        ref2[:, :, l_idx] = res2

    c_ref = (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
    return c_ref


# ---------------------------
# Triton fused kernel
# ---------------------------
if _TRITON_AVAILABLE:

    @triton.jit
    def _dual_nvfp4_scaled_gemm_silu_kernel(
        a_desc,  # packed uint8 FP4: (M, K//2)
        a_scale_desc,  # FP8 scales in TMA-friendly layout
        b1_desc,  # packed uint8 FP4: (N, K//2)
        b1_scale_desc,  # FP8 scales
        b2_desc,  # packed uint8 FP4: (N, K//2)
        b2_scale_desc,  # FP8 scales
        c_desc,  # FP16 output: (M, N)
        M: tl.constexpr,
        N: tl.constexpr,
        K: tl.constexpr,  # logical K (elements, not bytes)
        ELEM_PER_BYTE: tl.constexpr,
        VEC_SIZE: tl.constexpr,
        BLOCK_M: tl.constexpr,
        BLOCK_N: tl.constexpr,
        BLOCK_K: tl.constexpr,
        rep_m: tl.constexpr,
        rep_n: tl.constexpr,
        rep_k: tl.constexpr,
        NUM_STAGES: tl.constexpr,
    ):
        pid = tl.program_id(axis=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

        acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
        acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

        for _ in tl.range(0, tl.cdiv(K, BLOCK_K), num_stages=NUM_STAGES):
            a = a_desc.load([offs_am, offs_k])
            b1 = b1_desc.load([offs_bn, offs_k])
            b2 = b2_desc.load([offs_bn, offs_k])

            scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
            scale_b1 = b1_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
            scale_b2 = b2_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])

            # Unpack into the 2D scale layout expected by tl.dot_scaled
            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_b1 = (
                scale_b1.reshape(rep_n, rep_k, 32, 4, 4)
                .trans(0, 3, 2, 1, 4)
                .reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
            )
            scale_b2 = (
                scale_b2.reshape(rep_n, rep_k, 32, 4, 4)
                .trans(0, 3, 2, 1, 4)
                .reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
            )

            # NVFP4 x NVFP4
            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)

            offs_k += BLOCK_K // ELEM_PER_BYTE
            offs_scale_k += rep_k

        # SiLU: x * sigmoid(x)
        sigmoid = 1.0 / (1.0 + tl.exp(-acc1))
        out = (acc1 * sigmoid) * acc2
        c_desc.store([offs_am, offs_bn], out.to(tl.float16))


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized kernel:
      C = silu(A @ B1) * (A @ B2)

    Uses a fused Triton kernel on Blackwell (CC 10.x/11.x) when available.
    Falls back to torch._scaled_mm otherwise.
    """
    a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c = data

    # Shapes
    # a/b are torch.float4_e2m1fn_x2 with physical K dimension stored as K//2 bytes
    m = a.shape[0]
    n = b1.shape[0]
    l = a.shape[2]
    k_bytes = a.shape[1]
    k = k_bytes * 2  # logical K elements

    # Fast path: Triton fused kernel (Blackwell+)
    if _TRITON_AVAILABLE and torch.cuda.is_available():
        major, _minor = torch.cuda.get_device_capability()
        if major >= 10:
            # Views for packed FP4
            a_u8 = a.view(torch.uint8)
            b1_u8 = b1.view(torch.uint8)
            b2_u8 = b2.view(torch.uint8)

            # Convert preshuffled scale tensors into the exact TMA-friendly layout used by Triton tutorial:
            # shape: (L, M//128, K//64, 2, 256) for A
            # shape: (L, N//128, K//64, 2, 256) for B
            #
            # Input s*_perm is a view with shape (32, 4, rest_m, 4, rest_k, L) and underlying storage
            # corresponds to (L, rest_m, rest_k, 32, 4, 4) which is already the cublas packed layout.
            rest_m = m // 128
            rest_n = n // 128
            rest_k = k // sf_vec_size // 4  # K/16/4 = K/64

            # Inverse permute back to contiguous (L, rest_m, rest_k, 32, 4, 4), then reshape 32*4*4=512 -> 2*256
            a_scale_tma = sfa_perm.permute(5, 2, 4, 0, 1, 3).reshape(
                l, rest_m, rest_k, 2, 256
            )
            b1_scale_tma = sfb1_perm.permute(5, 2, 4, 0, 1, 3).reshape(
                l, rest_n, rest_k, 2, 256
            )
            b2_scale_tma = sfb2_perm.permute(5, 2, 4, 0, 1, 3).reshape(
                l, rest_n, rest_k, 2, 256
            )

            # Kernel params (safe for dual-accumulator register pressure)
            BLOCK_M = 128
            BLOCK_N = 128
            BLOCK_K = 256
            VEC_SIZE = 16
            ELEM_PER_BYTE = 2

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

            # Grid over (M,N) tiles
            grid = (triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N), 1)

            # Launch per-batch (L is small in the benchmark, typically 1)
            for l_idx in range(l):
                a_desc = TensorDescriptor.from_tensor(
                    a_u8[:, :, l_idx], [BLOCK_M, BLOCK_K // ELEM_PER_BYTE]
                )
                b1_desc = TensorDescriptor.from_tensor(
                    b1_u8[:, :, l_idx], [BLOCK_N, BLOCK_K // ELEM_PER_BYTE]
                )
                b2_desc = TensorDescriptor.from_tensor(
                    b2_u8[:, :, l_idx], [BLOCK_N, BLOCK_K // ELEM_PER_BYTE]
                )

                a_scale_desc = TensorDescriptor.from_tensor(
                    a_scale_tma[l_idx : l_idx + 1],
                    block_shape=[1, rep_m, rep_k, 2, 256],
                )
                b1_scale_desc = TensorDescriptor.from_tensor(
                    b1_scale_tma[l_idx : l_idx + 1],
                    block_shape=[1, rep_n, rep_k, 2, 256],
                )
                b2_scale_desc = TensorDescriptor.from_tensor(
                    b2_scale_tma[l_idx : l_idx + 1],
                    block_shape=[1, rep_n, rep_k, 2, 256],
                )

                c_desc = TensorDescriptor.from_tensor(
                    c[:, :, l_idx], [BLOCK_M, BLOCK_N]
                )

                _dual_nvfp4_scaled_gemm_silu_kernel[grid](
                    a_desc,
                    a_scale_desc,
                    b1_desc,
                    b1_scale_desc,
                    b2_desc,
                    b2_scale_desc,
                    c_desc,
                    m,
                    n,
                    k,
                    ELEM_PER_BYTE,
                    VEC_SIZE,
                    BLOCK_M,
                    BLOCK_N,
                    BLOCK_K,
                    rep_m,
                    rep_n,
                    rep_k,
                    4,  # NUM_STAGES
                    num_warps=8,
                )

            return c

    # Fallback: torch._scaled_mm (correct, slower)
    # Use the already-provided preshuffled scales to build cublas scale vectors without CPU loops.
    # For nvfp4, cublas expects FP8 E4M3 scales in a packed cublas layout flattened.
    a_scale_flat_all = (
        sfa_perm.permute(5, 2, 4, 0, 1, 3)
        .reshape(l, m // 128, k // 64, 32, 16)
        .contiguous()
    )
    b1_scale_flat_all = (
        sfb1_perm.permute(5, 2, 4, 0, 1, 3)
        .reshape(l, n // 128, k // 64, 32, 16)
        .contiguous()
    )
    b2_scale_flat_all = (
        sfb2_perm.permute(5, 2, 4, 0, 1, 3)
        .reshape(l, n // 128, k // 64, 32, 16)
        .contiguous()
    )

    for l_idx in range(l):
        scale_a = a_scale_flat_all[l_idx].flatten()
        scale_b1 = b1_scale_flat_all[l_idx].flatten()
        scale_b2 = b2_scale_flat_all[l_idx].flatten()

        r1 = torch._scaled_mm(
            a[:, :, l_idx],
            b1[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b1,
            bias=None,
            out_dtype=torch.float32,
        )
        r2 = torch._scaled_mm(
            a[:, :, l_idx],
            b2[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b2,
            bias=None,
            out_dtype=torch.float32,
        )
        c[:, :, l_idx] = (torch.nn.functional.silu(r1) * r2).to(torch.float16)

    return c


def generate_input(m: int, n: int, k: int, l: int, seed: int):
    """
    Generate input tensors for NVFP4 block-scaled dual GEMM with silu activation,
    C = silu(A @ B1) * (A @ B2).
    """
    torch.manual_seed(seed)

    def create_fp4_tensors(l, mn, k):
        # generate uint8 tensor, then convert to float4e2m1fn_x2 data type
        # generate all bit patterns
        ref_i8 = torch.randint(
            255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda"
        )

        # for each nibble, only keep the sign bit and 2 LSBs
        # the possible values are [-1.5, -1, -0.5, 0, +0.5, +1, +1.5]
        ref_i8 = ref_i8 & 0b1011_1011

        return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)

    # FP4 inputs (packed, 2 elems per byte along K)
    a_ref = create_fp4_tensors(l, m, k).view(torch.float4_e2m1fn_x2)
    b1_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
    b2_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)

    # Output buffer
    c_ref = torch.randn((l, m, n), dtype=torch.float16, device="cuda").permute(1, 2, 0)

    def create_scale_factor_tensors(l, mn, sf_k):
        # Create the reference scale factor tensor (mn, sf_k, l) on GPU, then also return a preshuffled layout.
        ref_shape = (l, mn, sf_k)
        ref_permute_order = (1, 2, 0)

        ref_f8_random_fp32 = torch.rand(ref_shape, dtype=torch.float32, device="cuda")
        ref_f8_torch_tensor = ref_f8_random_fp32.to(dtype=torch.float8_e4m3fn)
        ref_f8_torch_tensor_permuted = ref_f8_torch_tensor.permute(*ref_permute_order)

        atom_m = (32, 4)
        atom_k = 4
        mma_shape = (
            l,  # batch size
            ceil_div(mn, atom_m[0] * atom_m[1]),
            ceil_div(sf_k, atom_k),
            atom_m[0],
            atom_m[1],
            atom_k,
        )

        # Create backing storage and a view with the "CuTe" axis order.
        mma_permute_order = (3, 4, 1, 5, 2, 0)
        rand_int_tensor = torch.empty(mma_shape, dtype=torch.int8, device="cuda")
        reordered_f8_torch_tensor = rand_int_tensor.to(
            dtype=torch.float8_e4m3fn
        ).permute(*mma_permute_order)

        # Vectorized reordering on GPU
        i_idx = torch.arange(mn, device="cuda")
        j_idx = torch.arange(sf_k, device="cuda")
        b_idx = torch.arange(l, device="cuda")
        i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing="ij")

        mm = i_grid // (atom_m[0] * atom_m[1])
        mm32 = i_grid % atom_m[0]
        mm4 = (i_grid % 128) // atom_m[0]
        kk = j_grid // atom_k
        kk4 = j_grid % atom_k

        reordered_f8_torch_tensor[mm32, mm4, mm, kk4, kk, b_grid] = (
            ref_f8_torch_tensor_permuted[i_grid, j_grid, b_grid]
        )

        return ref_f8_torch_tensor_permuted.cpu(), reordered_f8_torch_tensor

    sf_k = ceil_div(k, sf_vec_size)
    sfa_ref_cpu, sfa_ref_permuted = create_scale_factor_tensors(l, m, sf_k)
    sfb1_ref_cpu, sfb1_ref_permuted = create_scale_factor_tensors(l, n, sf_k)
    sfb2_ref_cpu, sfb2_ref_permuted = create_scale_factor_tensors(l, n, sf_k)

    return (
        a_ref,
        b1_ref,
        b2_ref,
        sfa_ref_cpu.to("cuda"),
        sfb1_ref_cpu.to("cuda"),
        sfb2_ref_cpu.to("cuda"),
        sfa_ref_permuted,
        sfb1_ref_permuted,
        sfb2_ref_permuted,
        c_ref,
    )


check_implementation = make_match_reference(ref_kernel, rtol=1e-03, atol=1e-03)
scrolls · 425 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