submission 516735
Young Han · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 134 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-516735?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:95c80eea359cbd5911b78d59c87fe98f5ccdd4635380ddf1377b13a0252905f1
license declaredunknown
license concludedunknown
authorsYoung Han
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
sk = config.get("splitK", 0)Kernel source
submission.py134 lines
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
import aiter
import torch
import triton
from aiter import QuantType, dtypes
from aiter.utility import fp4_utils
from task import input_t, output_t
try:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, gemm_a4w4_blockscale, get_GEMM_config
except Exception:
gemm_a4w4_asm = None
gemm_a4w4_blockscale = None
get_GEMM_config = None
_hip_quant = None
try:
_hip_quant = aiter.get_hip_quant(QuantType.per_1x32)
except Exception:
pass
_triton_quant = aiter.get_triton_quant(QuantType.per_1x32)
_quant_kernel = getattr(fp4_utils, "_dynamic_mxfp4_quant_kernel_asm_layout", None)
_ASM_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_64X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
_quant_cache: dict = {}
_out_cache: dict = {}
_dispatch_cache: dict = {}
def _quantize_a_hip(a: torch.Tensor):
return _hip_quant(a, shuffle=True)
def _quantize_a_triton(a: torch.Tensor):
if _quant_kernel is None:
return _triton_quant(a, shuffle=True)
m, k = a.shape
key = (a.device, m, k)
buf = _quant_cache.get(key)
if buf is None:
scale_n_valid = k // 32
scale_n_pad = triton.cdiv(scale_n_valid, 8) * 8
scale_m_pad = triton.cdiv(m, 32) * 32
q_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
scale_u8 = torch.empty(
(triton.cdiv(m, 256) * 256, scale_n_pad),
dtype=torch.uint8, device=a.device,
)
buf = (q_u8, scale_u8, scale_n_valid, scale_n_pad, scale_m_pad)
_quant_cache[key] = buf
q_u8, scale_u8, scale_n_valid, scale_n_pad, scale_m_pad = buf
grid = (triton.cdiv(m, 128), scale_n_pad)
_quant_kernel[grid](
a, q_u8, scale_u8,
*a.stride(), *q_u8.stride(), *scale_u8.stride(),
M=m, N=k, scaleN=scale_n_valid,
scaleM_pad=scale_m_pad, scaleN_pad=scale_n_pad,
BLOCK_SIZE=128, MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0, SHUFFLE=True,
)
return q_u8.view(dtypes.fp4x2), scale_u8.view(dtypes.fp8_e8m0)
_quantize_a = _quantize_a_triton
def _get_out(device: torch.device, m: int, n: int):
key = (device, m, n)
out = _out_cache.get(key)
if out is None:
out = torch.empty((triton.cdiv(m, 32) * 32, n), dtype=dtypes.bf16, device=device)
_out_cache[key] = out
return out
def custom_kernel(data: input_t) -> output_t:
a, _, _, b_shuffle, b_scale_sh = data
m, k = a.shape
n = b_shuffle.shape[0]
a_q, a_scale_sh = _quantize_a(a)
if gemm_a4w4_asm is None or get_GEMM_config is None:
return aiter.gemm_a4w4(
a_q, b_shuffle, a_scale_sh, b_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
out = _get_out(a_q.device, m, n)
shape = (m, n, k)
dispatch = _dispatch_cache.get(shape)
if dispatch is None:
config = get_GEMM_config(m, n, k)
if config is not None:
kn = config["kernelName"]
sk = config.get("splitK", 0)
sk = 0 if sk in (None, "") else int(sk)
if kn and "_ZN" not in kn:
dispatch = ("blockscale", kn, sk)
else:
dispatch = ("asm", kn, sk)
elif k == 512 and m <= 8:
dispatch = ("asm", _ASM_64X128, 0)
elif m <= 64:
dispatch = ("asm", _ASM_32X128, 0)
else:
dispatch = ("asm", "", 0)
_dispatch_cache[shape] = dispatch
kind, kernel_name, split_k = dispatch
if kind == "blockscale":
gemm_a4w4_blockscale(
a_q.view(m, k // 2), b_shuffle, a_scale_sh, b_scale_sh,
out, splitK=split_k,
)
else:
if split_k:
out.zero_()
gemm_a4w4_asm(
a_q.view(m, k // 2), b_shuffle, a_scale_sh, b_scale_sh,
out, kernel_name, None, 1.0, 0.0, True,
log2_k_split=split_k,
)
return out[:m]
scrolls · 134 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