Skip to content
KernelIndex
Search⌘K

submission 311717

pagelessuptime · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

nvfp4_dual_gemm_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-311717?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
82.2µs
#385 of 420
2026-01-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:49367b25e3f3e95bbff0f87d593f10c2b217a2b999e5ebc15707e00e74ef6d90
license declaredunknown
license concludedunknown
authorspagelessuptime
imported2026-08-26

Techniques

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

autotunetriton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),
mmaacc1 += tl.dot(a_block, b1_block, out_dtype=tl.float32)
num-warps = 8triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),
stages = 2triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),

Kernel source

nvfp4_dual_gemm_submission.py235 lines
#!POPCORN leaderboard nvfp4_dual_gemm
#!POPCORN gpu B200

import torch
try:
    import triton
    import triton.language as tl

    _has_triton = True
except Exception:
    _has_triton = False

def _ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b


def _to_blocked(scale_matrix: torch.Tensor) -> torch.Tensor:
    """Convert (rows, cols) scale factors to the blocked layout expected by torch._scaled_mm."""
    rows, cols = scale_matrix.shape
    n_row_blocks = _ceil_div(rows, 128)
    n_col_blocks = _ceil_div(cols, 4)

    # Assumes rows are multiple of 128 and cols multiple of 4 (per problem constraints).
    blocks = scale_matrix.reshape(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 _block_all_scales(sf: torch.Tensor) -> torch.Tensor:
    """
    Block all scale tensors in one pass on the GPU.
    sf: [rows, cols, L] -> returns [L, blocked_len]
    """
    # Move batch to front and ensure contiguity for reshapes.
    sf_l = sf.permute(2, 0, 1).contiguous()
    # vmap applies _to_blocked across L.
    return torch.vmap(_to_blocked)(sf_l)


_TRITON_CONFIGS = [
    triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),
    triton.Config({"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 64}, num_warps=4, num_stages=2),
    triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=2),
    triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 128}, num_warps=8, num_stages=2),
]

@triton.autotune(configs=_TRITON_CONFIGS, key=["M", "N", "K"])
@triton.jit
def _dual_gemm_silu_kernel(
    A, B1, B2, C,
    M, N, K,
    stride_aL, stride_am, stride_ak,
    stride_b1L, stride_b1n, stride_b1k,
    stride_b2L, stride_b2n, stride_b2k,
    stride_cL, stride_cm, stride_cn,
    **meta,
):
    BLOCK_M = meta["BLOCK_M"]
    BLOCK_N = meta["BLOCK_N"]
    BLOCK_K = meta["BLOCK_K"]

    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    pid_l = tl.program_id(2)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    a_ptr = A + pid_l * stride_aL
    b1_ptr = B1 + pid_l * stride_b1L
    b2_ptr = B2 + pid_l * stride_b2L
    c_ptr = C + pid_l * stride_cL

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

    k_iter = tl.cdiv(K, BLOCK_K)
    for ki in range(0, k_iter):
        k_idx = ki * BLOCK_K + offs_k
        a_mask = (offs_m[:, None] < M) & (k_idx[None, :] < K)
        b_mask = (offs_n[None, :] < N) & (k_idx[:, None] < K)

        a_block = tl.load(
            a_ptr + offs_m[:, None] * stride_am + k_idx[None, :] * stride_ak,
            mask=a_mask,
            other=0.0,
        )
        b1_block = tl.load(
            b1_ptr + offs_n[None, :] * stride_b1n + k_idx[:, None] * stride_b1k,
            mask=b_mask,
            other=0.0,
        )
        b2_block = tl.load(
            b2_ptr + offs_n[None, :] * stride_b2n + k_idx[:, None] * stride_b2k,
            mask=b_mask,
            other=0.0,
        )

        acc1 += tl.dot(a_block, b1_block, out_dtype=tl.float32)
        acc2 += tl.dot(a_block, b2_block, out_dtype=tl.float32)

    # fused silu(res1) * res2
    silu = acc1 * tl.sigmoid(acc1)
    out_block = silu * acc2

    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(
        c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn,
        out_block.to(tl.float16),
        mask=c_mask,
    )


def _nvfp4_dual_gemm_impl(data):
    """
    Reference-style submission for nvfp4_dual_gemm.
    Accepts the first 7 items of the tuple:
      (a, b1, b2, sfa, sfb1, sfb2, c)
    Extra items (e.g., permuted scales) are ignored to stay compatible with 10-input runners.
    Returns: c_out with shape [M, N, L] in fp16.
    """
    if len(data) < 7:
        raise ValueError(f"Expected at least 7 inputs, got {len(data)}")

    # Only the first 7 are used; any extras (permute variants) are ignored.
    a, b1, b2, sfa, sfb1, sfb2, _c = data[:7]
    device = a.device

    # Ensure inputs are contiguous in their given layout to avoid hidden copies later.
    a = a.contiguous()
    b1 = b1.contiguous()
    b2 = b2.contiguous()
    sfa = sfa.contiguous()
    sfb1 = sfb1.contiguous()
    sfb2 = sfb2.contiguous()

    m, k, l = a.shape
    n, _, _ = b1.shape

    # If triton is available and inputs are already fp16, run fused dual GEMM in one kernel.
    # Otherwise, fall back to the scaled_mm reference path.
    if _has_triton and a.dtype == torch.float16 and b1.dtype == torch.float16 and b2.dtype == torch.float16:
        # Move batch to front: [L, M, K] etc.
        a_b = a.permute(2, 0, 1).contiguous()
        b1_b = b1.permute(2, 0, 1).contiguous()
        b2_b = b2.permute(2, 0, 1).contiguous()

        # Prepare output buffer
        if (
            isinstance(_c, torch.Tensor)
            and _c.shape == (m, n, l)
            and _c.dtype == torch.float16
            and _c.device == device
        ):
            out = _c
        else:
            out = torch.empty((m, n, l), device=device, dtype=torch.float16)

        out_b = out.permute(2, 0, 1).contiguous()

        # Per-shape launch: BLOCK sizes come from autotune.
        grid = (triton.cdiv(m, 128), triton.cdiv(n, 128), l)
        _dual_gemm_silu_kernel[grid](
            a_b, b1_b, b2_b, out_b,
            m, n, k,
            a_b.stride(0), a_b.stride(1), a_b.stride(2),
            b1_b.stride(0), b1_b.stride(1), b1_b.stride(2),
            b2_b.stride(0), b2_b.stride(1), b2_b.stride(2),
            out_b.stride(0), out_b.stride(1), out_b.stride(2),
        )
        return out

    # Fallback: use torch._scaled_mm for correctness (supports FP4 + block scales).
    scale_a_blocked = _block_all_scales(sfa).to(device)
    scale_b1_blocked = _block_all_scales(sfb1).to(device)
    scale_b2_blocked = _block_all_scales(sfb2).to(device)

    a_b = a.permute(2, 0, 1).contiguous()
    b1_b = b1.permute(2, 0, 1).contiguous()
    b2_b = b2.permute(2, 0, 1).contiguous()

    def _per_batch(a_l, b1_l, b2_l, sa_l, sb1_l, sb2_l):
        res1 = torch._scaled_mm(
            a_l,
            b1_l.transpose(0, 1),
            sa_l,
            sb1_l,
            bias=None,
            out_dtype=torch.float32,
        )
        res2 = torch._scaled_mm(
            a_l,
            b2_l.transpose(0, 1),
            sa_l,
            sb2_l,
            bias=None,
            out_dtype=torch.float32,
        )
        return torch.nn.functional.silu(res1) * res2

    if l == 1:
        out_lmn = _per_batch(
            a_b[0], b1_b[0], b2_b[0], scale_a_blocked[0], scale_b1_blocked[0], scale_b2_blocked[0]
        ).unsqueeze(0)
    else:
        out_lmn = torch.vmap(_per_batch)(
            a_b,
            b1_b,
            b2_b,
            scale_a_blocked,
            scale_b1_blocked,
            scale_b2_blocked,
        )

    if (
        isinstance(_c, torch.Tensor)
        and _c.shape == (m, n, l)
        and _c.dtype == torch.float16
        and _c.device == device
    ):
        out = _c
    else:
        out = torch.empty((m, n, l), device=device, dtype=torch.float16)

    out.copy_(out_lmn.permute(1, 2, 0).to(dtype=torch.float16))
    return out


def custom_kernel(data):
    """Popcorn expects a custom_kernel entrypoint."""
    return _nvfp4_dual_gemm_impl(data)
scrolls · 235 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