Skip to content
KernelIndex
Search⌘K

submission 642218

Frank Liu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_mxfp4_mm_smallm_dynamic.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-642218?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
19.9µs
#718 of 1143
2026-03-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8483e06bb0aec136ba839bde9ecd6d50728d88b675147a4a98a34cc288574f25
license declaredunknown
license concludedunknown
authorsFrank Liu
imported2026-08-26

Kernel source

submission_mxfp4_mm_smallm_dynamic.py221 lines
from __future__ import annotations

from typing import Callable, Dict, Tuple

import torch

from task import input_t, output_t

import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

try:
    from aiter import gemm_a4w4
except Exception:  # pragma: no cover
    gemm_a4w4 = getattr(aiter, "gemm_a4w4")

try:
    from aiter.ops.triton.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffled_weight_scales
except Exception:  # pragma: no cover
    gemm_afp4wfp4_preshuffled_weight_scales = None

try:
    from utils import clear_l2_cache_large as clear_l2_cache
except Exception:  # pragma: no cover
    def clear_l2_cache() -> None:
        return None

DTYPE_OUT = torch.bfloat16
RTOL = 1e-2
ATOL = 1e-2

# Cache chosen backend per (device, M, N, K)
_PLAN_CACHE: Dict[Tuple[int, int, int, int], str] = {}
_OUTBUF_CACHE: Dict[Tuple[int, int, int], torch.Tensor] = {}


@torch.no_grad()
def _shape_key(a: torch.Tensor, b_shuffle: torch.Tensor) -> Tuple[int, int, int, int]:
    dev = a.device.index if a.device.index is not None else -1
    return (dev, int(a.shape[0]), int(b_shuffle.shape[0]), int(a.shape[1]))


@torch.no_grad()
def _allclose(a: torch.Tensor, b: torch.Tensor) -> bool:
    return a.shape == b.shape and torch.allclose(
        a.to(torch.float32),
        b.to(torch.float32),
        rtol=RTOL,
        atol=ATOL,
    )


@torch.no_grad()
def _get_outbuf(m: int, n: int, device: torch.device) -> torch.Tensor:
    padded_m = ((m + 31) // 32) * 32
    dev = device.index if device.index is not None else -1
    key = (dev, padded_m, n)
    out = _OUTBUF_CACHE.get(key)
    if out is None or out.shape != (padded_m, n) or out.device != device:
        out = torch.empty((padded_m, n), device=device, dtype=DTYPE_OUT)
        _OUTBUF_CACHE[key] = out
    return out


@torch.no_grad()
def _quant_dynamic_raw(a: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
    # Returns raw packed fp4 bytes + raw E8M0 scale bytes.
    return dynamic_mxfp4_quant(a)


@torch.no_grad()
def _run_wrapper_dynamic(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
    a_q_raw, a_scale_raw = _quant_dynamic_raw(a)
    a_q = a_q_raw.view(dtypes.fp4x2)
    a_scale_sh = e8m0_shuffle(a_scale_raw).view(dtypes.fp8_e8m0)
    return aiter.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


@torch.no_grad()
def _run_outbuf_dynamic(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
    a_q_raw, a_scale_raw = _quant_dynamic_raw(a)
    a_q = a_q_raw.view(dtypes.fp4x2)
    a_scale_sh = e8m0_shuffle(a_scale_raw).view(dtypes.fp8_e8m0)

    out = _get_outbuf(a.shape[0], b_shuffle.shape[0], a.device)
    result = gemm_a4w4(
        a_q,
        b_shuffle.view(a_q.dtype),
        a_scale_sh,
        b_scale_sh.view(a_scale_sh.dtype),
        out,
        bpreshuffle=True,
    )
    if isinstance(result, torch.Tensor):
        return result[: a.shape[0]]
    return out[: a.shape[0]]


@torch.no_grad()
def _run_smallm_dynamic(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
    if gemm_afp4wfp4_preshuffled_weight_scales is None:
        raise RuntimeError("gemm_afp4wfp4_preshuffled_weight_scales is unavailable")

    m = int(a.shape[0])
    n = int(b_shuffle.shape[0])

    a_q_raw, a_scale_raw = _quant_dynamic_raw(a)

    if m >= 32:
        a_scale_sh = e8m0_shuffle(a_scale_raw).contiguous()
        a_scale_kernel = a_scale_sh.view(torch.uint8).view(a_scale_sh.shape[0] // 32, -1)
    else:
        a_scale_kernel = a_scale_raw[:m].contiguous().view(torch.uint8)

    b_q_kernel = b_shuffle.view(torch.uint8).view(n // 16, -1)
    b_scale_kernel = b_scale_sh.view(torch.uint8).view(b_scale_sh.shape[0] // 32, -1)

    y = torch.empty((m, n), device=a.device, dtype=DTYPE_OUT)
    gemm_afp4wfp4_preshuffled_weight_scales(
        a_q_raw.contiguous().view(torch.uint8),
        b_q_kernel,
        a_scale_kernel,
        b_scale_kernel,
        DTYPE_OUT,
        y,
    )
    return y


_BACKEND_FNS: Dict[str, Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]] = {
    "wrapper": _run_wrapper_dynamic,
    "outbuf": _run_outbuf_dynamic,
    "smallm": _run_smallm_dynamic,
}


@torch.no_grad()
def _time_backend(
    fn: Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor],
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    reps: int = 4,
) -> float:
    # One warmup after validation/compile.
    _ = fn(a, b_shuffle, b_scale_sh)
    torch.cuda.synchronize()

    times = []
    for _ in range(reps):
        clear_l2_cache()
        start = torch.cuda.Event(enable_timing=True)
        end = torch.cuda.Event(enable_timing=True)
        start.record()
        out = fn(a, b_shuffle, b_scale_sh)
        end.record()
        torch.cuda.synchronize()
        del out
        times.append(start.elapsed_time(end) * 1e6)
    return float(sum(times) / len(times))


@torch.no_grad()
def _choose_plan(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> str:
    # Always validate against the official reference-style path.
    y_ref = _run_wrapper_dynamic(a, b_shuffle, b_scale_sh)

    valid = {"wrapper": y_ref}

    # Newer AITER builds often support a slightly faster explicit-out GEMM call.
    try:
        y_out = _run_outbuf_dynamic(a, b_shuffle, b_scale_sh)
        if _allclose(y_out, y_ref):
            valid["outbuf"] = y_out
    except Exception:
        pass

    # For M <= 64, vLLM's ROCm Quark path prefers the preshuffled-weight-scales
    # Triton kernel when it is tuned for the shape. We validate numerics first,
    # then benchmark it off the critical path.
    if a.shape[0] <= 64 and gemm_afp4wfp4_preshuffled_weight_scales is not None:
        try:
            y_small = _run_smallm_dynamic(a, b_shuffle, b_scale_sh)
            if _allclose(y_small, y_ref):
                valid["smallm"] = y_small
        except Exception:
            pass

    if len(valid) == 1:
        return "wrapper"

    timings: Dict[str, float] = {}
    for name in valid.keys():
        timings[name] = _time_backend(_BACKEND_FNS[name], a, b_shuffle, b_scale_sh)

    # Choose the empirically fastest valid backend for this (device, M, N, K).
    best_name = min(timings, key=timings.get)
    return best_name


@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
    a, _b, _b_q, b_shuffle, b_scale_sh = data

    key = _shape_key(a, b_shuffle)
    plan = _PLAN_CACHE.get(key)
    if plan is None:
        plan = _choose_plan(a, b_shuffle, b_scale_sh)
        _PLAN_CACHE[key] = plan

    return _BACKEND_FNS[plan](a, b_shuffle, b_scale_sh)
scrolls · 221 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