Skip to content
KernelIndex
Search⌘K

submission 528216

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_rent2_exhaustive_20260311_v3-1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-528216?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.5µs
#446 of 1143
2026-03-11

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

split-kint(row["splitK"]),

Kernel source

submission_rent2_exhaustive_20260311_v3-1.py317 lines
from __future__ import annotations

import csv
from pathlib import Path

import torch
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

_PREFERRED_KERNELS = {
    (4, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (16, 2112, 7168): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (32, 4096, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
        0,
    ),
    (32, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
        0,
    ),
    (64, 7168, 2048): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (256, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
}


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"]),
                    int(row["splitK"]),
                )
            )

    return tuple(kernels)


def _kernel_spec(kernel_name: str, split_k: int):
    return ("asm", kernel_name, split_k)


def _candidate_specs(m: int, n: int, k: int):
    shape_key = (m, n, k)
    wrapper_spec = ("wrapper", "wrapper", 0)
    specs = [wrapper_spec]
    seen = {wrapper_spec}

    preferred = _PREFERRED_KERNELS.get(shape_key)
    if preferred is not None:
        spec = _kernel_spec(*preferred)
        specs.append(spec)
        seen.add(spec)

    padded_m = _ceil_div(m, 32) * 32

    def _kernel_rank(entry):
        kernel_name, tile_m, tile_n, split_k = entry
        m_penalty = abs(tile_m - min(max(32, padded_m), 256))
        n_penalty = abs(tile_n - min(max(128, n), 1024))
        split_penalty = split_k
        return (m_penalty, n_penalty, split_penalty, tile_m, tile_n, kernel_name)

    for kernel_name, _tile_m, _tile_n, split_k in sorted(_ASM_KERNELS, key=_kernel_rank):
        spec = _kernel_spec(kernel_name, split_k)
        if 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, 32) * 32
    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)
    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


@torch.no_grad()
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 · 317 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 528029.

⋯ 2 unchanged lines
import csv
from pathlib import Path
+ import torch
from task import input_t, output_t
_AITER = None
+ _CLEAR_L2_CACHE = None
_DTYPES = None
_FP4_UTILS = None
_TRITON_QUANT_FUNC = None
- _KNOWN_ASM = 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
+
_PREFERRED_KERNELS = {
(4, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
⋯ 27 unchanged lines
def _lazy_init():
- global _AITER, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _KNOWN_ASM
+ global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS
if _AITER is not None:
return
⋯ 1 unchanged lines
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)
- _KNOWN_ASM = _load_asm_kernels(aiter, get_gfx())
+ _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 = set()
+ kernels = []
+
with csv_path.open(newline="") as f:
for row in csv.DictReader(f):
- if int(row["bpreshuffle"]) == 1:
- kernels.add((row["knl_name"], int(row["splitK"])))
- return kernels
+ if int(row["bpreshuffle"]) != 1:
+ continue
+ kernels.append(
+ (
+ row["knl_name"],
+ int(row["tile_M"]),
+ int(row["tile_N"]),
+ int(row["splitK"]),
+ )
+ )
+ return tuple(kernels)
+
+ def _kernel_spec(kernel_name: str, split_k: int):
+ return ("asm", kernel_name, split_k)
+
+
+ def _candidate_specs(m: int, n: int, k: int):
+ shape_key = (m, n, k)
+ wrapper_spec = ("wrapper", "wrapper", 0)
+ specs = [wrapper_spec]
+ seen = {wrapper_spec}
+
+ preferred = _PREFERRED_KERNELS.get(shape_key)
+ if preferred is not None:
+ spec = _kernel_spec(*preferred)
+ specs.append(spec)
+ seen.add(spec)
+
+ padded_m = _ceil_div(m, 32) * 32
+
+ def _kernel_rank(entry):
+ kernel_name, tile_m, tile_n, split_k = entry
+ m_penalty = abs(tile_m - min(max(32, padded_m), 256))
+ n_penalty = abs(tile_n - min(max(128, n), 1024))
+ split_penalty = split_k
+ return (m_penalty, n_penalty, split_penalty, tile_m, tile_n, kernel_name)
+
+ for kernel_name, _tile_m, _tile_n, split_k in sorted(_ASM_KERNELS, key=_kernel_rank):
+ spec = _kernel_spec(kernel_name, split_k)
+ if 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
⋯ 12 unchanged lines
return buffers
- def _quantize_a(a):
+ 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)
- try:
- 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)
- except Exception:
- return _TRITON_QUANT_FUNC(a, shuffle=True)
+ 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
+ rows = _ceil_div(m, 32) * 32
key = (device.type, device.index, rows, n)
out = _OUT_BUFFERS.get(key)
if out is None:
⋯ 4 unchanged lines
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)
+ 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
+
+
+ @torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
_lazy_init()
⋯ 3 unchanged lines
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)
- preferred = _PREFERRED_KERNELS.get((m, n, k))
- if preferred is not None and preferred in _KNOWN_ASM:
- out = _get_out_buffer(m, n, a.device)
- _AITER.gemm_a4w4_asm(
+ spec = _BEST_IMPLS.get(shape_key)
+ if spec is None:
+ spec = _select_best_impl(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
- out,
- preferred[0],
- None,
- bpreshuffle=True,
- log2_k_split=preferred[1],
+ m,
+ n,
+ k,
+ a.device,
)
- return out[:m]
+ _BEST_IMPLS[shape_key] = spec
- return _AITER.gemm_a4w4(
- a_q,
- b_shuffle,
- a_scale_sh,
- b_scale_sh,
- dtype=_DTYPES.bf16,
- bpreshuffle=True,
- )
+ 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 · 323 diff lines total

Best evidence level for this revision: reported

JSON