submission 585823
jefflyu_47387 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 168 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-585823?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:57844e5c30e1210ed36032a1764a9009250f077ec8bcdebf23746b9522c040e5
license declaredunknown
license concludedunknown
authorsjefflyu_47387
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
- B uses provided MXFP4 packed tensor (B_q), and we reconstruct raw scales from B.stages = 2
num_stages = 2Kernel source
submission.py168 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Phase 1 prototype: TileLang fused dequant(B)+GEMM kernel.
Numerical path:
- A is consumed directly as bf16 (no A quantization in this prototype).
- B uses provided MXFP4 packed tensor (B_q), and we reconstruct raw scales from B.
"""
from task import input_t, output_t
_KERNEL_CACHE = {}
def _aiter_fallback(data: input_t) -> output_t:
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
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)
A, B, B_q, B_shuffle, B_scale_sh = data
del B_q
A = A.contiguous()
B = B.contiguous()
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def _build_mxfp4_dequant_gemm_kernel(m: int, n: int, k: int):
import tilelang
import tilelang.language as T
from tvm import tir
block_M = 64
block_N = 64
block_K = 64
num_stages = 2
threads = 128
num_bits = 4
scale_size = 32
num_elems_per_byte = 8 // num_bits
qk = k // num_elems_per_byte
block_qk = block_K // num_elems_per_byte
def _tir_u8_to_f4_to_bf16(val: tir.PrimExpr, pos: tir.PrimExpr):
mask = tir.const((1 << num_bits) - 1, T.uint16)
f4 = (val >> (pos.astype(T.uint16) * tir.const(num_bits, T.uint16))) & mask
s = f4 >> tir.const(3, T.uint16)
e_f4 = (f4 & tir.const(6, T.uint16)) >> tir.const(1, T.uint16)
e_bf16 = e_f4 + tir.const(126, T.uint16)
m_f4 = f4 & tir.const(1, T.uint16)
return tir.reinterpret(
T.bfloat16,
((((s << tir.const(8, T.uint16)) | e_bf16) << tir.const(7, T.uint16))
| (m_f4 << tir.const(6, T.uint16))).astype(T.uint16),
)
@T.macro
def _simple_dequant_bf16_fp4(B_shared, B_dequantize_shared, Scale, k_block):
B_local = T.alloc_fragment((block_N, block_qk), "uint8")
B_dequantize_local = T.alloc_fragment((block_N, block_K), "bfloat16")
bx = T.get_block_binding(0)
T.copy(B_shared, B_local)
for i, j in T.Parallel(block_N, block_K):
fp4_as_bf16 = _tir_u8_to_f4_to_bf16(
B_local[i, j // num_elems_per_byte],
j % num_elems_per_byte,
)
# E8M0 scale is exponent-only, applying power-of-two multiplication.
scale_exp = Scale[
bx * block_N + i,
k_block * block_K // scale_size + j // scale_size,
]
B_dequantize_local[i, j] = fp4_as_bf16 * T.shift_left(1, scale_exp)
T.copy(B_dequantize_local, B_dequantize_shared)
@tilelang.jit(out_idx=[-1])
def _mxfp4_dequant_gemm(M, N, K):
@T.prim_func
def main(
A: T.Tensor((M, K), "bfloat16"),
B_q_u8: T.Tensor((N, qk), "uint8"),
B_scale_u8: T.Tensor((N, K // scale_size), "uint8"),
C: T.Tensor((M, N), "bfloat16"),
):
with T.Kernel(
T.ceildiv(N, block_N),
T.ceildiv(M, block_M),
threads=threads,
) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), "bfloat16")
B_shared = T.alloc_shared((block_N, block_qk), "uint8")
B_dequantize_shared = T.alloc_shared((block_N, block_K), "bfloat16")
C_local = T.alloc_fragment((block_M, block_N), "float32")
T.clear(C_local)
for ko in T.Pipelined(K // block_K, num_stages=num_stages):
T.copy(A[by * block_M, ko * block_K], A_shared)
T.copy(B_q_u8[bx * block_N, ko * block_qk], B_shared)
_simple_dequant_bf16_fp4(B_shared, B_dequantize_shared, B_scale_u8, ko)
T.gemm(A_shared, B_dequantize_shared, C_local, transpose_B=True)
T.copy(C_local, C[by * block_M, bx * block_N])
return main
return _mxfp4_dequant_gemm(m, n, k)
def custom_kernel(data: input_t) -> output_t:
import torch
from aiter.ops.triton.quant import dynamic_mxfp4_quant
A, B, B_q, B_shuffle, B_scale_sh = data
try:
import tilelang # noqa: F401
except ModuleNotFoundError:
return _aiter_fallback(data)
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n, _ = B.shape
# Task constraints guarantee K % 64 == 0.
assert k % 64 == 0
# B_q is fp4x2-packed; reinterpret as raw packed bytes for TileLang dequant macro.
B_q_u8 = B_q.view(torch.uint8).contiguous()
# Task input provides shuffled scales; dequant macro here expects unshuffled scales.
_, B_scale_raw = dynamic_mxfp4_quant(B)
B_scale_u8 = B_scale_raw.view(torch.uint8)[:n, : (k // 32)].contiguous()
key = (m, n, k)
kernel = _KERNEL_CACHE.get(key)
if kernel is None:
try:
kernel = _build_mxfp4_dequant_gemm_kernel(m, n, k)
_KERNEL_CACHE[key] = kernel
except ModuleNotFoundError:
return _aiter_fallback(data)
return kernel(A, B_q_u8, B_scale_u8)
scrolls · 168 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