Skip to content
KernelIndex
Search⌘K

submission 528029

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_rent_direct_20260311_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-528029?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
#457 of 1143
2026-03-11

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

split-kkernels.add((row["knl_name"], int(row["splitK"])))

Kernel source

submission_rent_direct_20260311_v1.py173 lines
from __future__ import annotations

import csv
from pathlib import Path

from task import input_t, output_t

_AITER = None
_DTYPES = None
_FP4_UTILS = None
_TRITON_QUANT_FUNC = None
_KNOWN_ASM = None

_A_QUANT_BUFFERS = {}
_OUT_BUFFERS = {}

_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, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _KNOWN_ASM
    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

    _AITER = aiter
    _DTYPES = dtypes
    _FP4_UTILS = fp4_utils
    _TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
    _KNOWN_ASM = _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()
    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


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 _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)


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 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])

    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(
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            out,
            preferred[0],
            None,
            bpreshuffle=True,
            log2_k_split=preferred[1],
        )
        return out[:m]

    return _AITER.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=_DTYPES.bf16,
        bpreshuffle=True,
    )
scrolls · 173 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 527938.

⋯ 5 unchanged lines
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
+ _KNOWN_ASM = 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_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,
+ ),
}
- _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
+ global _AITER, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _KNOWN_ASM
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)
- _ASM_KERNELS = _load_asm_kernels(aiter, get_gfx())
+ _KNOWN_ASM = _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 = []
-
+ kernels = set()
with csv_path.open(newline="") as f:
for row in csv.DictReader(f):
- if int(row["bpreshuffle"]) != 1:
- continue
+ if int(row["bpreshuffle"]) == 1:
+ kernels.add((row["knl_name"], int(row["splitK"])))
+ return kernels
- 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
⋯ 12 unchanged lines
return buffers
- def _cached_quantize_a(a):
+ def _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)
+ 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)
- _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)
⋯ 6 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)
- 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()
⋯ 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)
- spec = _BEST_IMPLS.get(shape_key)
- if spec is None:
- spec = _select_best_impl(
+ 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(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
- m,
- n,
- k,
- a.device,
+ out,
+ preferred[0],
+ None,
+ bpreshuffle=True,
+ log2_k_split=preferred[1],
)
- _BEST_IMPLS[shape_key] = spec
+ return out[:m]
- 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]
+ return _AITER.gemm_a4w4(
+ a_q,
+ b_shuffle,
+ a_scale_sh,
+ b_scale_sh,
+ dtype=_DTYPES.bf16,
+ bpreshuffle=True,
+ )
scrolls · 353 diff lines total

Best evidence level for this revision: reported

JSON