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
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.
fp4
FP4 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