Skip to content
KernelIndex
Search⌘K

submission 527938

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-rent.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-527938?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
13.6µs
#464 of 1143
2026-03-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8bb1b43b1f9aa75532d568f93ff89c580051c5f4ba856814244176b0cf65ce85
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Kernel source

submission-rent.py302 lines
from __future__ import annotations

import csv
from pathlib import Path

from task import input_t, output_t

_AITER = None
_CLEAR_L2_CACHE = None
_DTYPES = None
_FP4_UTILS = None
_TRITON_QUANT_FUNC = None
_ASM_KERNELS = None

_USE_CACHED_TRITON_QUANT = True

_BEST_IMPLS = {}
_A_QUANT_BUFFERS = {}
_OUT_BUFFERS = {}

_SELECT_ITERS = 5
_SELECT_WARMUP_ITERS = 1
_RTOL = 1e-2
_ATOL = 1e-2

_WRAPPER_ONLY_SHAPES = {
    (64, 7168, 2048),
    (256, 3072, 1536),
}

_PREFERRED_TILE_SPECS = {
    (4, 2880, 512): (64, 128),
    (16, 2112, 7168): (32, 128),
    (32, 4096, 512): (32, 256),
    (32, 2880, 512): (64, 128),
}


def _ceil_div(x: int, y: int) -> int:
    return (x + y - 1) // y


def _lazy_init():
    global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS
    if _AITER is not None:
        return

    import aiter
    from aiter import QuantType, dtypes
    from aiter.jit.utils.chip_info import get_gfx
    from aiter.utility import fp4_utils
    from utils import clear_l2_cache_large

    _AITER = aiter
    _CLEAR_L2_CACHE = clear_l2_cache_large
    _DTYPES = dtypes
    _FP4_UTILS = fp4_utils
    _TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
    _ASM_KERNELS = _load_asm_kernels(aiter, get_gfx())


def _load_asm_kernels(aiter, gfx: str):
    root = Path(aiter.__file__).resolve().parents[1]
    csv_path = root / "hsa" / gfx / "f4gemm" / "f4gemm_bf16_per1x32Fp4.csv"
    kernels = []

    with csv_path.open(newline="") as f:
        for row in csv.DictReader(f):
            if int(row["bpreshuffle"]) != 1:
                continue

            kernels.append(
                (
                    row["knl_name"],
                    int(row["tile_M"]),
                    int(row["tile_N"]),
                )
            )

    return tuple(kernels)


def _spec_for_tiles(tile_m: int, tile_n: int):
    for kernel_name, kernel_tile_m, kernel_tile_n in _ASM_KERNELS:
        if kernel_tile_m == tile_m and kernel_tile_n == tile_n:
            return ("asm", kernel_name, 0)
    return None


def _candidate_specs(m: int, n: int, k: int):
    shape_key = (m, n, k)
    wrapper_spec = ("wrapper", "wrapper", 0)
    if shape_key in _WRAPPER_ONLY_SHAPES:
        return (wrapper_spec,)

    specs = [wrapper_spec]
    seen = {wrapper_spec}

    preferred_tiles = _PREFERRED_TILE_SPECS.get(shape_key)
    if preferred_tiles is not None:
        preferred_spec = _spec_for_tiles(preferred_tiles[0], preferred_tiles[1])
        if preferred_spec is not None:
            specs.append(preferred_spec)
            seen.add(preferred_spec)

    for tile_m in (32, 64):
        for tile_n in (128, 256):
            spec = _spec_for_tiles(tile_m, tile_n)
            if spec is not None and spec not in seen:
                specs.append(spec)
                seen.add(spec)

    return tuple(specs)


def _get_quant_buffers(m: int, k: int, device):
    scale_n = _ceil_div(k, 32)
    scale_n_pad = _ceil_div(scale_n, 8) * 8
    scale_m_pad = _ceil_div(m, 256) * 256
    key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)
    buffers = _A_QUANT_BUFFERS.get(key)

    if buffers is None:
        import torch

        a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
        a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
        buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)
        _A_QUANT_BUFFERS[key] = buffers

    return buffers


def _cached_quantize_a(a):
    m = int(a.shape[0])
    k = int(a.shape[1])
    a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)

    block_size = 128
    grid = (_ceil_div(m, block_size), scale_n_pad)

    _FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](
        a,
        a_q,
        a_scale,
        *a.stride(),
        *a_q.stride(),
        *a_scale.stride(),
        M=m,
        N=k,
        scaleN=scale_n,
        scaleM_pad=scale_m_pad,
        scaleN_pad=scale_n_pad,
        BLOCK_SIZE=block_size,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        SHUFFLE=True,
    )

    return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)


def _quantize_a(a):
    if _USE_CACHED_TRITON_QUANT:
        try:
            return _cached_quantize_a(a)
        except Exception:
            pass

    return _TRITON_QUANT_FUNC(a, shuffle=True)


def _get_out_buffer(m: int, n: int, device):
    rows = _ceil_div(m, 256) * 256
    key = (device.type, device.index, rows, n)
    out = _OUT_BUFFERS.get(key)
    if out is None:
        import torch

        out = torch.empty((rows, n), dtype=_DTYPES.bf16, device=device)
        _OUT_BUFFERS[key] = out
    return out


def _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):
    kind, impl_id, log2_k_split = spec

    if kind == "wrapper":
        return _AITER.gemm_a4w4(
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            dtype=_DTYPES.bf16,
            bpreshuffle=True,
        )

    _AITER.gemm_a4w4_asm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        out,
        impl_id,
        None,
        bpreshuffle=True,
        log2_k_split=log2_k_split,
    )
    return out


def _is_valid(candidate, reference, m: int):
    import torch

    return torch.allclose(candidate[:m], reference[:m], rtol=_RTOL, atol=_ATOL)


def _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):
    import torch

    for _ in range(_SELECT_WARMUP_ITERS):
        _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
    torch.cuda.synchronize()

    times_ms = []
    for _ in range(_SELECT_ITERS):
        _CLEAR_L2_CACHE()
        start = torch.cuda.Event(enable_timing=True)
        end = torch.cuda.Event(enable_timing=True)
        start.record()
        _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
        end.record()
        torch.cuda.synchronize()
        times_ms.append(start.elapsed_time(end))

    times_ms.sort()
    return times_ms[len(times_ms) // 2]


def _select_best_impl(a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, k: int, device):
    wrapper_spec = ("wrapper", "wrapper", 0)
    if (m, n, k) in _WRAPPER_ONLY_SHAPES:
        return wrapper_spec

    out = _get_out_buffer(m, n, device)
    reference = _run_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)

    best_spec = wrapper_spec
    best_time = _time_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)

    for spec in _candidate_specs(m, n, k):
        if spec == wrapper_spec:
            continue

        try:
            candidate = _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
            if not _is_valid(candidate, reference, m):
                continue

            current_time = _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
            if current_time < best_time:
                best_time = current_time
                best_spec = spec
        except Exception:
            continue

    return best_spec


def custom_kernel(data: input_t) -> output_t:
    _lazy_init()

    a, _b, _b_q, b_shuffle, b_scale_sh = data
    a = a.contiguous()

    m = int(a.shape[0])
    k = int(a.shape[1])
    n = int(b_shuffle.shape[0])
    shape_key = (m, n, k)

    a_q, a_scale_sh = _quantize_a(a)

    spec = _BEST_IMPLS.get(shape_key)
    if spec is None:
        spec = _select_best_impl(
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            m,
            n,
            k,
            a.device,
        )
        _BEST_IMPLS[shape_key] = spec

    if spec[0] == "wrapper":
        return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, None)

    out = _get_out_buffer(m, n, a.device)
    return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)[:m]
scrolls · 302 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 527329.

- #!POPCORN leaderboard amd-mxfp4-mm
- #!POPCORN gpu MI355X
+ from __future__ import annotations
- """
- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
- """
+ import csv
+ from pathlib import Path
+
from task import input_t, output_t
+ _AITER = None
+ _CLEAR_L2_CACHE = None
+ _DTYPES = None
+ _FP4_UTILS = None
+ _TRITON_QUANT_FUNC = None
+ _ASM_KERNELS = None
- def custom_kernel(data: input_t) -> output_t:
- """
- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
- gemm_a4w4 with bpreshuffle=True.
- """
+ _USE_CACHED_TRITON_QUANT = True
+
+ _BEST_IMPLS = {}
+ _A_QUANT_BUFFERS = {}
+ _OUT_BUFFERS = {}
+
+ _SELECT_ITERS = 5
+ _SELECT_WARMUP_ITERS = 1
+ _RTOL = 1e-2
+ _ATOL = 1e-2
+
+ _WRAPPER_ONLY_SHAPES = {
+ (64, 7168, 2048),
+ (256, 3072, 1536),
+ }
+
+ _PREFERRED_TILE_SPECS = {
+ (4, 2880, 512): (64, 128),
+ (16, 2112, 7168): (32, 128),
+ (32, 4096, 512): (32, 256),
+ (32, 2880, 512): (64, 128),
+ }
+
+
+ def _ceil_div(x: int, y: int) -> int:
+ return (x + y - 1) // y
+
+
+ def _lazy_init():
+ global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS
+ if _AITER is not None:
+ return
+
import aiter
from aiter import QuantType, dtypes
+ from aiter.jit.utils.chip_info import get_gfx
+ from aiter.utility import fp4_utils
+ from utils import clear_l2_cache_large
- A, B, B_q, B_shuffle, B_scale_sh = data
- A = A.contiguous()
- B = B.contiguous()
- m, k = A.shape
- n, _ = B.shape
+ _AITER = aiter
+ _CLEAR_L2_CACHE = clear_l2_cache_large
+ _DTYPES = dtypes
+ _FP4_UTILS = fp4_utils
+ _TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
+ _ASM_KERNELS = _load_asm_kernels(aiter, get_gfx())
- quant_func = aiter.get_triton_quant(QuantType.per_1x32)
- A_q, A_scale_sh = quant_func(A, shuffle=True)
- out_gemm = aiter.gemm_a4w4(
- A_q,
- B_shuffle,
- A_scale_sh,
- B_scale_sh,
- dtype=dtypes.bf16,
+
+ def _load_asm_kernels(aiter, gfx: str):
+ root = Path(aiter.__file__).resolve().parents[1]
+ csv_path = root / "hsa" / gfx / "f4gemm" / "f4gemm_bf16_per1x32Fp4.csv"
+ kernels = []
+
+ with csv_path.open(newline="") as f:
+ for row in csv.DictReader(f):
+ if int(row["bpreshuffle"]) != 1:
+ continue
+
+ kernels.append(
+ (
+ row["knl_name"],
+ int(row["tile_M"]),
+ int(row["tile_N"]),
+ )
+ )
+
+ return tuple(kernels)
+
+
+ def _spec_for_tiles(tile_m: int, tile_n: int):
+ for kernel_name, kernel_tile_m, kernel_tile_n in _ASM_KERNELS:
+ if kernel_tile_m == tile_m and kernel_tile_n == tile_n:
+ return ("asm", kernel_name, 0)
+ return None
+
+
+ def _candidate_specs(m: int, n: int, k: int):
+ shape_key = (m, n, k)
+ wrapper_spec = ("wrapper", "wrapper", 0)
+ if shape_key in _WRAPPER_ONLY_SHAPES:
+ return (wrapper_spec,)
+
+ specs = [wrapper_spec]
+ seen = {wrapper_spec}
+
+ preferred_tiles = _PREFERRED_TILE_SPECS.get(shape_key)
+ if preferred_tiles is not None:
+ preferred_spec = _spec_for_tiles(preferred_tiles[0], preferred_tiles[1])
+ if preferred_spec is not None:
+ specs.append(preferred_spec)
+ seen.add(preferred_spec)
+
+ for tile_m in (32, 64):
+ for tile_n in (128, 256):
+ spec = _spec_for_tiles(tile_m, tile_n)
+ if spec is not None and spec not in seen:
+ specs.append(spec)
+ seen.add(spec)
+
+ return tuple(specs)
+
+
+ def _get_quant_buffers(m: int, k: int, device):
+ scale_n = _ceil_div(k, 32)
+ scale_n_pad = _ceil_div(scale_n, 8) * 8
+ scale_m_pad = _ceil_div(m, 256) * 256
+ key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)
+ buffers = _A_QUANT_BUFFERS.get(key)
+
+ if buffers is None:
+ import torch
+
+ a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
+ a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
+ buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)
+ _A_QUANT_BUFFERS[key] = buffers
+
+ return buffers
+
+
+ def _cached_quantize_a(a):
+ m = int(a.shape[0])
+ k = int(a.shape[1])
+ a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)
+
+ block_size = 128
+ grid = (_ceil_div(m, block_size), scale_n_pad)
+
+ _FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](
+ a,
+ a_q,
+ a_scale,
+ *a.stride(),
+ *a_q.stride(),
+ *a_scale.stride(),
+ M=m,
+ N=k,
+ scaleN=scale_n,
+ scaleM_pad=scale_m_pad,
+ scaleN_pad=scale_n_pad,
+ BLOCK_SIZE=block_size,
+ MXFP4_QUANT_BLOCK_SIZE=32,
+ SCALING_MODE=0,
+ SHUFFLE=True,
+ )
+
+ return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)
+
+
+ def _quantize_a(a):
+ if _USE_CACHED_TRITON_QUANT:
+ try:
+ return _cached_quantize_a(a)
+ except Exception:
+ pass
+
+ return _TRITON_QUANT_FUNC(a, shuffle=True)
+
+
+ def _get_out_buffer(m: int, n: int, device):
+ rows = _ceil_div(m, 256) * 256
+ key = (device.type, device.index, rows, n)
+ out = _OUT_BUFFERS.get(key)
+ if out is None:
+ import torch
+
+ out = torch.empty((rows, n), dtype=_DTYPES.bf16, device=device)
+ _OUT_BUFFERS[key] = out
+ return out
+
+
+ def _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):
+ kind, impl_id, log2_k_split = spec
+
+ if kind == "wrapper":
+ return _AITER.gemm_a4w4(
+ a_q,
+ b_shuffle,
+ a_scale_sh,
+ b_scale_sh,
+ dtype=_DTYPES.bf16,
+ bpreshuffle=True,
+ )
+
+ _AITER.gemm_a4w4_asm(
+ a_q,
+ b_shuffle,
+ a_scale_sh,
+ b_scale_sh,
+ out,
+ impl_id,
+ None,
bpreshuffle=True,
+ log2_k_split=log2_k_split,
)
- return out_gemm
+ return out
+
+
+ def _is_valid(candidate, reference, m: int):
+ import torch
+
+ return torch.allclose(candidate[:m], reference[:m], rtol=_RTOL, atol=_ATOL)
+
+
+ def _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):
+ import torch
+
+ for _ in range(_SELECT_WARMUP_ITERS):
+ _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
+ torch.cuda.synchronize()
+
+ times_ms = []
+ for _ in range(_SELECT_ITERS):
+ _CLEAR_L2_CACHE()
+ start = torch.cuda.Event(enable_timing=True)
+ end = torch.cuda.Event(enable_timing=True)
+ start.record()
+ _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
+ end.record()
+ torch.cuda.synchronize()
+ times_ms.append(start.elapsed_time(end))
+
+ times_ms.sort()
+ return times_ms[len(times_ms) // 2]
+
+
+ def _select_best_impl(a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, k: int, device):
+ wrapper_spec = ("wrapper", "wrapper", 0)
+ if (m, n, k) in _WRAPPER_ONLY_SHAPES:
+ return wrapper_spec
+
+ out = _get_out_buffer(m, n, device)
+ reference = _run_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
+
+ best_spec = wrapper_spec
+ best_time = _time_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
+
+ for spec in _candidate_specs(m, n, k):
+ if spec == wrapper_spec:
+ continue
+
+ try:
+ candidate = _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
+ if not _is_valid(candidate, reference, m):
+ continue
+
+ current_time = _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
+ if current_time < best_time:
+ best_time = current_time
+ best_spec = spec
+ except Exception:
+ continue
+
+ return best_spec
+
+
+ def custom_kernel(data: input_t) -> output_t:
+ _lazy_init()
+
+ a, _b, _b_q, b_shuffle, b_scale_sh = data
+ a = a.contiguous()
+
+ m = int(a.shape[0])
+ k = int(a.shape[1])
+ n = int(b_shuffle.shape[0])
+ shape_key = (m, n, k)
+
+ a_q, a_scale_sh = _quantize_a(a)
+
+ spec = _BEST_IMPLS.get(shape_key)
+ if spec is None:
+ spec = _select_best_impl(
+ a_q,
+ b_shuffle,
+ a_scale_sh,
+ b_scale_sh,
+ m,
+ n,
+ k,
+ a.device,
+ )
+ _BEST_IMPLS[shape_key] = spec
+
+ if spec[0] == "wrapper":
+ return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, None)
+
+ out = _get_out_buffer(m, n, a.device)
+ return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)[:m]
scrolls · 326 diff lines total

Best evidence level for this revision: reported

JSON