Skip to content
KernelIndex
Search⌘K

submission 272244

boo · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

reference7_fast_v20.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-272244?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
60.4µs
#314 of 420
2026-01-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4b298a4cfbe0440d38e5922e3506e41513d19b19908f5b48db5d24f72485bf92
license declaredunknown
license concludedunknown
authorsboo
imported2026-08-26

Kernel source

reference7_fast_v20.py290 lines

import os
import weakref
import torch
from task import input_t, output_t
from utils import make_match_reference

# -------------------------
# Fast, safe building blocks
# -------------------------

def _cuda_dev_index(t: torch.Tensor) -> int:
    if not t.is_cuda:
        return -1
    return int(t.device.index) if t.device.index is not None else 0

def _tensor_ver(t: torch.Tensor) -> int:
    # PyTorch tensors carry a version counter updated on in-place ops.
    try:
        return int(t._version)  # type: ignore[attr-defined]
    except Exception:
        return 0

# -------------------------
# fp8 scale packing for _scaled_mm
# -------------------------
def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
    """
    Reference mapping used by the baseline:
      input:  [rows, cols] fp8 (e4m3fnuz)
      output: flat tensor in a blocked order expected by torch._scaled_mm
    Assumes rows % 128 == 0 and cols % 4 == 0 for the hot path.
    """
    rows, cols = input_matrix.shape

    # Fallback (should not trigger for contest sizes, but keep it safe).
    if (rows % 128) != 0 or (cols % 4) != 0 or (not input_matrix.is_contiguous()):
        n_row_blocks = (rows + 127) // 128
        n_col_blocks = (cols + 3) // 4
        # NOTE: For contest shapes, rows/cols are divisible and view() is valid.
        blocks = input_matrix.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()

    n_row_blocks = rows // 128
    n_col_blocks = cols // 4

    view5 = input_matrix.as_strided(
        size=(n_row_blocks, n_col_blocks, 32, 4, 4),
        stride=(128 * cols, 4, cols, 32 * cols, 1),
    )
    return view5.contiguous().view(-1)

def to_blocked_out(input_matrix: torch.Tensor, out_flat: torch.Tensor) -> torch.Tensor:
    """Same mapping as to_blocked(), but writes into a preallocated flat buffer."""
    rows, cols = input_matrix.shape

    if (rows % 128) != 0 or (cols % 4) != 0 or (not input_matrix.is_contiguous()):
        tmp = to_blocked(input_matrix)
        out_flat.copy_(tmp)
        return out_flat

    n_row_blocks = rows // 128
    n_col_blocks = cols // 4

    view5 = input_matrix.as_strided(
        size=(n_row_blocks, n_col_blocks, 32, 4, 4),
        stride=(128 * cols, 4, cols, 32 * cols, 1),
    )
    out5 = out_flat.view(n_row_blocks, n_col_blocks, 32, 4, 4)
    out5.copy_(view5)
    return out_flat

# -------------------------
# Buffer reuse (no data_ptr based signatures)
# -------------------------
_SCALE_BUF: dict[tuple, torch.Tensor] = {}
_OUT_BUF: dict[tuple, torch.Tensor] = {}
_MAT_BUF: dict[tuple, torch.Tensor] = {}

def _get_buf(cache: dict, key: tuple, make: callable) -> torch.Tensor:
    buf = cache.get(key)
    if buf is None:
        buf = make()
        cache[key] = buf
        # avoid unbounded growth
        if len(cache) > 64:
            cache.clear()
    return buf

def _get_scale_buf(slot: int, numel: int, dtype: torch.dtype, like: torch.Tensor) -> torch.Tensor:
    dev = _cuda_dev_index(like)
    key = (slot, numel, dtype, dev)
    return _get_buf(_SCALE_BUF, key, lambda: torch.empty((numel,), device=like.device, dtype=dtype))

def _get_out_buf(slot: int, shape: tuple[int, int], dtype: torch.dtype, like: torch.Tensor) -> torch.Tensor:
    dev = _cuda_dev_index(like)
    key = (slot, shape[0], shape[1], dtype, dev)
    return _get_buf(_OUT_BUF, key, lambda: torch.empty(shape, device=like.device, dtype=dtype))

def _get_mat_buf(slot: int, shape: tuple[int, int], dtype: torch.dtype, like: torch.Tensor) -> torch.Tensor:
    dev = _cuda_dev_index(like)
    key = (slot, shape[0], shape[1], dtype, dev)
    return _get_buf(_MAT_BUF, key, lambda: torch.empty(shape, device=like.device, dtype=dtype))

# -------------------------
# _scaled_mm wrapper
# -------------------------
def _scaled_mm_out(mat1, mat2, scale_a_flat, scale_b_flat, out, out_dtype):
    """
    Use aten._scaled_mm.out if available (saves an allocation).
    IMPORTANT: in this contest, mat2 must be transposed: shape (K, N).
    """
    op_out = getattr(torch.ops.aten, "_scaled_mm", None)
    if op_out is not None and hasattr(torch.ops.aten._scaled_mm, "out"):
        # use_fast_accum is not supported for FP4 inputs in the provided environment
        return torch.ops.aten._scaled_mm.out(mat1, mat2, scale_a_flat, scale_b_flat, None, None, out_dtype, False, out=out)

    # fallback
    res = torch._scaled_mm(mat1, mat2, scale_a_flat, scale_b_flat, bias=None, out_dtype=out_dtype)
    out.copy_(res)
    return out

# -------------------------
# Modes
# -------------------------
# 0: safe baseline (2x scaled_mm)
# 1: packed 2N (1x scaled_mm on concatenated B), useful to experiment
_MODE = int(os.environ.get("NVFP4_DUAL_GEMM_MODE", "0"))

# -------------------------
# Kernels
# -------------------------
def _kernel_2x(a, b1, b2, sfa, sfb1, sfb2, c):
    """Fast & stable: two scaled_mm calls; mixed output dtypes."""
    m, n, L = c.shape
    cols_a = sfa.shape[1]
    cols_b = sfb1.shape[1]

    # Hot path for L==1
    if L == 1:
        a0 = a[:, :, 0]
        b10 = b1[:, :, 0].transpose(0, 1)  # (K, N)
        b20 = b2[:, :, 0].transpose(0, 1)  # (K, N)

        scale_a = to_blocked_out(sfa[:, :, 0], _get_scale_buf(0, m * cols_a, sfa.dtype, sfa))
        scale_b1 = to_blocked_out(sfb1[:, :, 0], _get_scale_buf(1, n * cols_b, sfb1.dtype, sfb1))
        scale_b2 = to_blocked_out(sfb2[:, :, 0], _get_scale_buf(2, n * cols_b, sfb2.dtype, sfb2))

        out1 = _get_out_buf(0, (m, n), torch.float32, c)
        out2 = _get_out_buf(1, (m, n), torch.float16, c)

        _scaled_mm_out(a0, b10, scale_a, scale_b1, out=out1, out_dtype=torch.float32)
        _scaled_mm_out(a0, b20, scale_a, scale_b2, out=out2, out_dtype=torch.float16)

        torch.nn.functional.silu(out1, inplace=True)
        torch.mul(out1, out2, out=c[:, :, 0])
        return c

    # Generic L
    out1 = _get_out_buf(0, (m, n), torch.float32, c)
    out2 = _get_out_buf(1, (m, n), torch.float16, c)
    for l_idx in range(L):
        scale_a = to_blocked_out(sfa[:, :, l_idx], _get_scale_buf(0, m * cols_a, sfa.dtype, sfa))
        scale_b1 = to_blocked_out(sfb1[:, :, l_idx], _get_scale_buf(1, n * cols_b, sfb1.dtype, sfb1))
        scale_b2 = to_blocked_out(sfb2[:, :, l_idx], _get_scale_buf(2, n * cols_b, sfb2.dtype, sfb2))

        _scaled_mm_out(a[:, :, l_idx], b1[:, :, l_idx].transpose(0, 1), scale_a, scale_b1, out=out1, out_dtype=torch.float32)
        _scaled_mm_out(a[:, :, l_idx], b2[:, :, l_idx].transpose(0, 1), scale_a, scale_b2, out=out2, out_dtype=torch.float16)

        torch.nn.functional.silu(out1, inplace=True)
        torch.mul(out1, out2, out=c[:, :, l_idx])
    return c

def _kernel_packed_2n(a, b1, b2, sfa, sfb1, sfb2, c):
    """
    One _scaled_mm by concatenating B along N (so output is M x 2N), then SiLU(x)*y.
    FIX: mat2 must be (K, 2N), so we pack B as (2N, K) and pass transpose view.
    """
    m, n, L = c.shape
    k = a.shape[1]
    cols_a = sfa.shape[1]
    cols_b = sfb1.shape[1]

    # Allocate reusable buffers
    bcat = _get_mat_buf(0, (2 * n, k), b1.dtype, b1)
    bcat_t = bcat.transpose(0, 1)  # (K, 2N) view

    scale_a_buf = _get_scale_buf(0, m * cols_a, sfa.dtype, sfa)
    scalecat_buf = _get_scale_buf(3, (2 * n) * cols_b, sfb1.dtype, sfb1)

    tmp = _get_out_buf(2, (m, 2 * n), torch.float32, c)

    # Optional: avoid allocating sfb_cat when rows align to 128-blocks (bench sizes do).
    fast_scale_cat = (n % 128) == 0 and (cols_b % 4) == 0

    if L == 1:
        # pack B in (2N, K) with contiguous copies, then transpose view
        bcat[:n].copy_(b1[:, :, 0])
        bcat[n:].copy_(b2[:, :, 0])

        # scale_a
        scale_a = to_blocked_out(sfa[:, :, 0], scale_a_buf)

        if fast_scale_cat:
            half = n * cols_b
            to_blocked_out(sfb1[:, :, 0], scalecat_buf[:half])
            to_blocked_out(sfb2[:, :, 0], scalecat_buf[half:2 * half])
            scale_b = scalecat_buf
        else:
            # safe path: build (2N, cols_b) then to_blocked_out once
            sfb_cat = _get_mat_buf(1, (2 * n, cols_b), sfb1.dtype, sfb1)
            sfb_cat[:n].copy_(sfb1[:, :, 0])
            sfb_cat[n:].copy_(sfb2[:, :, 0])
            scale_b = to_blocked_out(sfb_cat, scalecat_buf)

        _scaled_mm_out(a[:, :, 0], bcat_t, scale_a, scale_b, out=tmp, out_dtype=torch.float32)

        x = tmp[:, :n]
        y = tmp[:, n:]
        torch.nn.functional.silu(x, inplace=True)
        torch.mul(x, y, out=c[:, :, 0])
        return c

    # Generic L
    for l_idx in range(L):
        bcat[:n].copy_(b1[:, :, l_idx])
        bcat[n:].copy_(b2[:, :, l_idx])

        scale_a = to_blocked_out(sfa[:, :, l_idx], scale_a_buf)

        if fast_scale_cat:
            half = n * cols_b
            to_blocked_out(sfb1[:, :, l_idx], scalecat_buf[:half])
            to_blocked_out(sfb2[:, :, l_idx], scalecat_buf[half:2 * half])
            scale_b = scalecat_buf
        else:
            sfb_cat = _get_mat_buf(1, (2 * n, cols_b), sfb1.dtype, sfb1)
            sfb_cat[:n].copy_(sfb1[:, :, l_idx])
            sfb_cat[n:].copy_(sfb2[:, :, l_idx])
            scale_b = to_blocked_out(sfb_cat, scalecat_buf)

        _scaled_mm_out(a[:, :, l_idx], bcat_t, scale_a, scale_b, out=tmp, out_dtype=torch.float32)

        x = tmp[:, :n]
        y = tmp[:, n:]
        torch.nn.functional.silu(x, inplace=True)
        torch.mul(x, y, out=c[:, :, l_idx])

    return c

def custom_kernel(data: input_t) -> output_t:
    # The runner provides extra tensors; last one is output C.
    a, b1, b2, sfa, sfb1, sfb2 = data[0], data[1], data[2], data[3], data[4], data[5]
    c = data[-1]

    if _MODE == 1:
        return _kernel_packed_2n(a, b1, b2, sfa, sfb1, sfb2, c)
    return _kernel_2x(a, b1, b2, sfa, sfb1, sfb2, c)

# Reference kernel (for correctness checking)
def ref_kernel(data: input_t) -> output_t:
    a, b1, b2, sfa, sfb1, sfb2, _, _, _, c = data
    m, n, L = c.shape
    out1 = torch.empty((m, n, L), device=a.device, dtype=torch.float32)
    out2 = torch.empty((m, n, L), device=a.device, dtype=torch.float32)
    for l_idx in range(L):
        scale_a = to_blocked(sfa[:, :, l_idx])
        scale_b1 = to_blocked(sfb1[:, :, l_idx])
        scale_b2 = to_blocked(sfb2[:, :, l_idx])
        out1[:, :, l_idx] = torch._scaled_mm(
            a[:, :, l_idx],
            b1[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b1,
            bias=None,
            out_dtype=torch.float32,
        )
        out2[:, :, l_idx] = torch._scaled_mm(
            a[:, :, l_idx],
            b2[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b2,
            bias=None,
            out_dtype=torch.float32,
        )
    return (torch.nn.functional.silu(out1) * out2).to(torch.float16)

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