Skip to content
KernelIndex
Search⌘K

submission 565041

wuxin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v50.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-565041?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
22.2µs
#745 of 1143
2026-03-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bbb33142ba927638fe62e6d9ef99e5417ea7f2fb08102db4e4706d620cb8fe2a
license declaredunknown
license concludedunknown
authorswuxin
imported2026-08-26

Techniques

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

split-kreturn {"splitK": 2, "kernelName": KERNEL_32}

Kernel source

v50.py139 lines
from task import input_t, output_t

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

go = importlib.import_module("aiter.ops.gemm_op_a4w4")

KERNEL_32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
TARGET_SHAPE = (16, 2112, 7168)
_PROBED = False

def _p(msg: str):
    sys.stderr.write(msg + "\n")

def _patched_get_GEMM_config(m, n, k):
    # 当前 best: v21
    if k == 7168:
        return {"splitK": 2, "kernelName": KERNEL_32}
    if k == 512 and m <= 4:
        return {"splitK": 0, "kernelName": KERNEL_32}
    if k == 512:
        return {"splitK": 1, "kernelName": KERNEL_32}
    return {"splitK": 0, "kernelName": KERNEL_32}

go.get_GEMM_config = _patched_get_GEMM_config

def _probe_alignment(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k):
    global _PROBED
    if _PROBED or (m, n, k) != TARGET_SHAPE:
        return
    _PROBED = True

    _p(f"[align] ==== target_shape={(m, n, k)} ====")

    # 1) v21 asm baseline
    out_asm = torch.empty((((m + 31) // 32) * 32, n), dtype=torch.bfloat16, device=A_q.device)
    out_asm = aiter.gemm_a4w4_asm(
        A_q.view(m, k // 2),
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        out_asm,
        KERNEL_32,
        None,
        1.0,
        0.0,
        True,
        2,
    )[:m]

    # 2) blockscale_tune candidate
    out_bs = torch.empty((((m + 31) // 32) * 32, n), dtype=torch.bfloat16, device=A_q.device)
    out_bs = aiter.gemm_a4w4_blockscale_tune(
        A_q.view(m, k // 2),
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        out_bs,
        14,
        2,
    )[:m]

    # 转成 fp32 比较,避免 bf16 比较太粗
    ref = out_asm.float()
    got = out_bs.float()

    diff = (got - ref).abs()
    denom = ref.abs().clamp_min(1e-6)
    rel = diff / denom

    max_abs = diff.max().item()
    mean_abs = diff.mean().item()
    max_rel = rel.max().item()
    mean_rel = rel.mean().item()

    _p(f"[align:stats] max_abs={max_abs:.6f} mean_abs={mean_abs:.6f} max_rel={max_rel:.6f} mean_rel={mean_rel:.6f}")

    # 整体尺度对比
    ref_abs_mean = ref.abs().mean().item()
    got_abs_mean = got.abs().mean().item()
    ratio = got_abs_mean / max(ref_abs_mean, 1e-12)
    _p(f"[align:scale] ref_abs_mean={ref_abs_mean:.6f} got_abs_mean={got_abs_mean:.6f} abs_mean_ratio={ratio:.6f}")

    # 取最大误差的前 8 个位置
    flat_diff = diff.flatten()
    topk = min(8, flat_diff.numel())
    vals, idxs = torch.topk(flat_diff, k=topk)

    for rank, (v, idx) in enumerate(zip(vals.tolist(), idxs.tolist()), start=1):
        i = idx // n
        j = idx % n
        r = ref[i, j].item()
        g = got[i, j].item()
        rr = abs(g - r) / max(abs(r), 1e-6)
        _p(f"[align:topdiff] rank={rank} i={i} j={j} ref={r:.6f} got={g:.6f} abs={abs(g-r):.6f} rel={rr:.6f}")

    # 行统计,判断是否像布局/scale 问题
    row_abs = diff.mean(dim=1)
    for i in range(min(4, m)):
        _p(f"[align:row_mean_abs] row={i} mean_abs={row_abs[i].item():.6f}")

    # 列采样统计
    for j in [0, 1, 2, 3, 63, 127, 255, 511, 1023, 2047]:
        if j < n:
            col_mean = diff[:, j].mean().item()
            _p(f"[align:col_mean_abs] col={j} mean_abs={col_mean:.6f}")

def custom_kernel(data: input_t) -> output_t:
    A, _, _, B_shuffle, B_scale_sh = data

    if not A.is_contiguous():
        A = A.contiguous()

    m, k = A.shape
    n = B_shuffle.shape[0]

    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(A)
    bs_e8m0 = e8m0_shuffle(bs_e8m0)

    A_q = x_fp4.view(dtypes.fp4x2)
    A_scale_sh = bs_e8m0.view(dtypes.fp8_e8m0)

    # 只做一次对齐探针
    _probe_alignment(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)

    # 真正返回仍然走当前 best
    return aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
scrolls · 139 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