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
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-k
return {"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