Skip to content
KernelIndex
Search⌘K

submission 678625

ATIpiu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submit_claude.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-678625?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
23.9µs
#829 of 1143
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4694a445a634d41d2b8542150d0d4796165074b6cd6cf30a42184727fdfb5d21
license declaredunknown
license concludedunknown
authorsATIpiu
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM optimized for AMD MI355X.

Kernel source

submit_claude.py314 lines
"""
FP4 quant + FP4 GEMM optimized for AMD MI355X.
Probe version: collect internal aiter info via logs.
"""
import sys
import os
import glob
import inspect
import pkgutil
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference
from aiter import dtypes
import aiter as aiter_module          # use alias to avoid shadowing
import aiter.ops
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

SCALE_GROUP_SIZE = 32

# ============================================================
# Probe helpers  (all use aiter_module alias, never bare 'aiter')
# ============================================================

def _safe_source(obj):
    try:
        return inspect.getsource(obj)
    except Exception as e:
        return f"<source unavailable: {e}>"

def _safe_dir(obj, label):
    try:
        attrs = sorted(dir(obj))
        print(f"\n[PROBE] dir({label}):", flush=True)
        for a in attrs:
            print(f"  {a}", flush=True)
    except Exception as e:
        print(f"[PROBE] dir({label}) failed: {e}", flush=True)

def _probe_aiter_internals():
    print("\n" + "="*60, flush=True)
    print("[PROBE] aiter internals start", flush=True)
    print("="*60, flush=True)

    # ── 1. Top-level aiter attributes ──────────────────────────
    _safe_dir(aiter_module, "aiter")

    # ── 2. gemm_a4w4 source ────────────────────────────────────
    print("\n[PROBE] gemm_a4w4 file:", flush=True)
    try:
        print(inspect.getsourcefile(aiter_module.gemm_a4w4), flush=True)
        print(_safe_source(aiter_module.gemm_a4w4), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 3. aiter.ops submodule ─────────────────────────────────
    _safe_dir(aiter.ops, "aiter.ops")

    # ── 4. module_gemm_a4w4_asm ───────────────────────────────
    print("\n[PROBE] module_gemm_a4w4_asm:", flush=True)
    try:
        from aiter.jit import module_gemm_a4w4_asm as asm_mod
        _safe_dir(asm_mod, "module_gemm_a4w4_asm")
        # try to print signature of gemm_a4w4_asm
        if hasattr(asm_mod, "gemm_a4w4_asm"):
            print(f"  sig: {inspect.signature(asm_mod.gemm_a4w4_asm)}", flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 5. module_gemm_common ─────────────────────────────────
    print("\n[PROBE] module_gemm_common:", flush=True)
    try:
        from aiter.jit import module_gemm_common as gc_mod
        _safe_dir(gc_mod, "module_gemm_common")
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 6. File system: jit dir ───────────────────────────────
    jit_dir = "/home/runner/aiter/aiter/jit"
    print(f"\n[PROBE] Walking {jit_dir}:", flush=True)
    if os.path.exists(jit_dir):
        for root, dirs, files in os.walk(jit_dir):
            for f in files:
                print(f"  {os.path.join(root, f)}", flush=True)
    else:
        print("  (not found)", flush=True)

    # ── 7. CSV / JSON / tuning / config files ─────────────────
    for pattern in [
        "/home/runner/aiter/**/*.csv",
        "/home/runner/aiter/**/*.json",
        "/home/runner/aiter/**/*tuning*",
        "/home/runner/aiter/**/*config*",
        "/home/runner/aiter/**/*a4w4*",
        "/home/runner/aiter/**/*kernel*name*",
    ]:
        hits = glob.glob(pattern, recursive=True)
        if hits:
            print(f"\n[PROBE] glob '{pattern}':", flush=True)
            for h in hits:
                print(f"  {h}", flush=True)
                # if small text file, dump contents
                if os.path.isfile(h) and os.path.getsize(h) < 32768:
                    try:
                        with open(h) as fh:
                            print(fh.read(), flush=True)
                    except Exception:
                        pass

    # ── 8. aiter.ops.gemm_a4w4 module ────────────────────────
    print("\n[PROBE] aiter.ops.gemm_a4w4:", flush=True)
    try:
        import aiter.ops.gemm_a4w4 as ops_g4
        _safe_dir(ops_g4, "aiter.ops.gemm_a4w4")
        print(_safe_source(ops_g4), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 9. Triton gemm_afp4wfp4 ──────────────────────────────
    print("\n[PROBE] aiter.ops.triton.gemm.basic.gemm_afp4wfp4:", flush=True)
    try:
        import aiter.ops.triton.gemm.basic.gemm_afp4wfp4 as tfp4
        _safe_dir(tfp4, "gemm_afp4wfp4")
        print(_safe_source(tfp4), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 10. aiter.hsa ─────────────────────────────────────────
    print("\n[PROBE] aiter.hsa:", flush=True)
    try:
        import aiter.hsa as hsa_mod
        _safe_dir(hsa_mod, "aiter.hsa")
    except Exception as e:
        print(f"  error: {e}", flush=True)

    print("\n[PROBE] aiter.hsa.codegen:", flush=True)
    try:
        import aiter.hsa.codegen as cg_mod
        print(_safe_source(cg_mod), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 11. Full package walk ─────────────────────────────────
    print("\n[PROBE] aiter package walk:", flush=True)
    try:
        for importer, modname, ispkg in pkgutil.walk_packages(
            path=aiter_module.__path__,
            prefix=aiter_module.__name__ + ".",
            onerror=lambda x: None,
        ):
            tag = "[PKG]" if ispkg else "[MOD]"
            print(f"  {tag} {modname}", flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 12. aiter __init__ source ─────────────────────────────
    print("\n[PROBE] aiter __init__ source:", flush=True)
    print(_safe_source(aiter_module), flush=True)

    # ── 13. gemm_a4w4_asm direct ─────────────────────────────
    print("\n[PROBE] aiter_module.gemm_a4w4_asm:", flush=True)
    try:
        fn = aiter_module.gemm_a4w4_asm
        print(f"  {fn}", flush=True)
        print(f"  sig: {inspect.signature(fn)}", flush=True)
        print(_safe_source(fn), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 14. get_triton_quant source ───────────────────────────
    print("\n[PROBE] aiter.get_triton_quant:", flush=True)
    try:
        fn = aiter_module.get_triton_quant
        print(_safe_source(fn), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 15. dynamic_mxfp4_quant source ───────────────────────
    print("\n[PROBE] dynamic_mxfp4_quant:", flush=True)
    print(_safe_source(dynamic_mxfp4_quant), flush=True)

    # ── 16. e8m0_shuffle source ───────────────────────────────
    print("\n[PROBE] e8m0_shuffle:", flush=True)
    print(_safe_source(e8m0_shuffle), flush=True)

    # ── 17. shuffle_weight source ─────────────────────────────
    print("\n[PROBE] shuffle_weight:", flush=True)
    print(_safe_source(shuffle_weight), flush=True)

    # ── 18. aiter.ops.triton.quant source ────────────────────
    print("\n[PROBE] aiter.ops.triton.quant:", flush=True)
    try:
        import aiter.ops.triton.quant as tq_mod
        print(_safe_source(tq_mod), flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    # ── 19. QuantType ─────────────────────────────────────────
    print("\n[PROBE] aiter.QuantType:", flush=True)
    try:
        from aiter import QuantType
        _safe_dir(QuantType, "QuantType")
        for member in QuantType:
            print(f"  {member.name} = {member.value}", flush=True)
    except Exception as e:
        print(f"  error: {e}", flush=True)

    print("\n" + "="*60, flush=True)
    print("[PROBE] Done", flush=True)
    print("="*60 + "\n", flush=True)
    sys.stdout.flush()


_probe_aiter_internals()


# ============================================================
# Buffer cache
# ============================================================
_out_cache: dict = {}

def _get_out(m, n, device):
    key = (m, n)
    if key not in _out_cache:
        _out_cache[key] = torch.empty((m, n), dtype=torch.bfloat16, device=device)
    return _out_cache[key]


def _quant_mxfp4_shuffled(x: torch.Tensor):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


# ============================================================
# custom_kernel  (logs inputs on first call per shape)
# ============================================================
_logged_shapes: set = set()

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

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

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

    if shape_key not in _logged_shapes:
        _logged_shapes.add(shape_key)
        print(f"\n[custom_kernel] First call M={m} N={n} K={k}", flush=True)
        print(f"  A          : {A.shape}  dtype={A.dtype}  stride={A.stride()}", flush=True)
        print(f"  B_q        : {B_q.shape}  dtype={B_q.dtype}", flush=True)
        print(f"  B_shuffle  : {B_shuffle.shape}  dtype={B_shuffle.dtype}  stride={B_shuffle.stride()}", flush=True)
        print(f"  B_scale_sh : {B_scale_sh.shape}  dtype={B_scale_sh.dtype}", flush=True)

    A_q, A_scale_sh = _quant_mxfp4_shuffled(A)

    if shape_key not in _logged_shapes:
        print(f"  A_q        : {A_q.shape}  dtype={A_q.dtype}", flush=True)
        print(f"  A_scale_sh : {A_scale_sh.shape}  dtype={A_scale_sh.dtype}", flush=True)

    out = aiter_module.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )

    if shape_key not in _logged_shapes:
        print(f"  out        : {out.shape}  dtype={out.dtype}", flush=True)

    return out


# ============================================================
# Support
# ============================================================

def _quant_mxfp4(x, shuffle=True):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def generate_input(m: int, n: int, k: int, seed: int):
    assert k % 64 == 0
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    A = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)
    B_shuffle_t = shuffle_weight(B_q, layout=(16, 16))
    return (A, B, B_q, B_shuffle_t, B_scale_sh)


def ref_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    return aiter_module.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )


check_implementation = make_match_reference(ref_kernel, rtol=1e-02, atol=1e-02)
scrolls · 314 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