submission 639085
Amanpreet Singh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 147 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-639085?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:2d721db00a6e5db1c29f0b0f86fbda1851f041a267877b4e718b447edc09d123
license declaredunknown
license concludedunknown
authorsAmanpreet Singh
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
acc += tl.dot(a_scaled, b_scaled)Kernel source
submission.py147 lines
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
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
def _quant_mxfp4_shuffled(x):
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)
@triton.jit
def _fp4_unpack_and_scale_dot(
a_ptr, a_scale_ptr,
b_ptr, b_scale_ptr,
c_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_blocks = tl.cdiv(K, BLOCK_K)
for k_block in range(num_k_blocks):
k_start = k_block * BLOCK_K
a_offs = offs_m[:, None] * stride_am + ((k_start + offs_k)[None, :] // 2) * stride_ak
mask_a = (offs_m[:, None] < M) & ((k_start + offs_k)[None, :] < K)
a_packed = tl.load(a_ptr + a_offs, mask=mask_a, other=0).to(tl.uint8)
b_offs = ((k_start + offs_k)[:, None] // 2) * stride_bk + offs_n[None, :] * stride_bn
mask_b = ((k_start + offs_k)[:, None] < K) & (offs_n[None, :] < N)
b_packed = tl.load(b_ptr + b_offs, mask=mask_b, other=0).to(tl.uint8)
a_lo = (a_packed & 0x0F).to(tl.float32)
a_hi = ((a_packed >> 4) & 0x0F).to(tl.float32)
b_lo = (b_packed & 0x0F).to(tl.float32)
b_hi = ((b_packed >> 4) & 0x0F).to(tl.float32)
a_f32 = tl.interleave(a_lo, a_hi)
b_f32 = tl.interleave(b_lo, b_hi)
k_scale = (k_start + tl.arange(0, BLOCK_K)) // SCALE_GROUP_SIZE
a_scale_offs = offs_m[:, None] * tl.cdiv(K, SCALE_GROUP_SIZE) + k_scale[None, :]
a_scale_mask = (offs_m[:, None] < M) & (k_scale[None, :] < tl.cdiv(K, SCALE_GROUP_SIZE))
a_scale = tl.load(a_scale_ptr + a_scale_offs, mask=a_scale_mask, other=0).to(tl.uint8)
a_exp = a_scale.to(tl.float32) - 127.0
a_scale_f32 = tl.exp2(a_exp)
b_scale_offs = offs_n[None, :] * tl.cdiv(K, SCALE_GROUP_SIZE) + k_scale[:, None]
b_scale_mask = (offs_n[None, :] < N) & (k_scale[:, None] < tl.cdiv(K, SCALE_GROUP_SIZE))
b_scale = tl.load(b_scale_ptr + b_scale_offs, mask=b_scale_mask, other=0).to(tl.uint8)
b_exp = b_scale.to(tl.float32) - 127.0
b_scale_f32 = tl.exp2(b_exp)
a_scaled = a_f32 * a_scale_f32
b_scaled = b_f32 * b_scale_f32
acc += tl.dot(a_scaled, b_scaled)
c_offs = offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
mask_c = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptr + c_offs, acc.to(tl.bfloat16), mask=mask_c)
_SHAPE_CONFIGS = {
(4, 2880, 512): {"BLOCK_M": 16, "BLOCK_N": 64, "BLOCK_K": 64, "num_warps": 2, "num_stages": 2},
(16, 2112, 7168): {"BLOCK_M": 16, "BLOCK_N": 128, "BLOCK_K": 64, "num_warps": 4, "num_stages": 2},
(32, 4096, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "num_warps": 4, "num_stages": 2},
(32, 2880, 512): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 64, "num_warps": 4, "num_stages": 2},
(64, 7168, 2048): {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "num_warps": 4, "num_stages": 2},
(256, 3072, 1536): {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "num_warps": 4, "num_stages": 3},
}
_DEFAULT_CONFIG = {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 64, "num_warps": 4, "num_stages": 2}
def _get_config(m, n, k):
return _SHAPE_CONFIGS.get((m, n, k), _DEFAULT_CONFIG)
def _custom_quant_gemm(A, B_shuffle, B_scale_sh):
import aiter
A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
out = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0] if hasattr(B_shuffle, 'shape') else B_q.shape[0]
return _custom_quant_gemm(A, B_shuffle, B_scale_sh)
def _quant_mxfp4_noshuf(x):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
def ref_kernel(data: input_t) -> output_t:
import aiter
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
out = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out
check_implementation = make_match_reference(ref_kernel, rtol=1e-02, atol=1e-02)scrolls · 147 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