submission 754001
Navdeep Singh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 200 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754001?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:080933fe27ec1af851bd56aa8db08e595702acf1cff5a246174e2aba03e8e4d8
license declaredunknown
license concludedunknown
authorsNavdeep Singh
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
return e5m2_bytes.to(tl.int8).to(tl.float8e5, bitcast=True)mma
inner_acc = tl.dot(a_fp8, tl.trans(b_fp8), out_dtype=tl.float32)num-warps = 16
block_m, block_n, block_k, num_warps = 16, 64, 32, 4split-k
GROUP_M: tl.constexpr, SPLIT_K: tl.constexpr, EVEN_K: tl.constexpr,Kernel source
submission.py200 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
def _inverse_e8m0_shuffle(scale_shuffled: torch.Tensor, logical_m: int) -> torch.Tensor:
scale_u8 = scale_shuffled.view(torch.uint8)
sm, sn = scale_u8.shape
scale = (
scale_u8.view(sm // 32, sn // 8, 4, 16, 2, 2)
.permute(0, 5, 3, 1, 4, 2)
.contiguous()
.view(sm, sn)
)
return scale[:logical_m].contiguous()
@triton.jit
def _pid_grid(pid: int, num_pid_m: int, num_pid_n: int, group_size_m: tl.constexpr):
num_pid_in_group = group_size_m * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * group_size_m
actual_group_size_m = tl.minimum(num_pid_m - first_pid_m, group_size_m)
pid_m = first_pid_m + (pid % actual_group_size_m)
pid_n = (pid % num_pid_in_group) // actual_group_size_m
return pid_m, pid_n
@triton.jit
def decode_e2m1_to_fp8(packed_bytes, is_odd):
# Natively transforms MXFP4 e2m1 into standard fp8 native types using pure logic
nibbles = tl.where(is_odd, packed_bytes >> 4, packed_bytes & 0xF).to(tl.int32)
abs_val = nibbles & 0x7
base = 56 + (abs_val << 1)
res = tl.where(abs_val == 1, 56, base)
res = tl.where(abs_val == 0, 0, res)
e5m2_bytes = res | ((nibbles & 0x8) << 4)
return e5m2_bytes.to(tl.int8).to(tl.float8e5, bitcast=True)
@triton.heuristics({
"EVEN_K": lambda args: args["K"] % args["BLOCK_K"] == 0,
})
@triton.jit
def _mxfp4_gemm_kernel(
a_ptr, b_ptr, c_ptr,
a_scales_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
stride_asm, stride_ask,
stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr, SPLIT_K: tl.constexpr, EVEN_K: tl.constexpr,
):
pid = tl.program_id(axis=0)
pid_k = tl.program_id(axis=1)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m, pid_n = _pid_grid(pid, num_pid_m, num_pid_n, GROUP_M)
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)
byte_offs_k = offs_k // 2
is_odd = (offs_k % 2) == 1
a_ptrs = a_ptr + offs_m[:, None] * stride_am + byte_offs_k[None, :] * stride_ak
b_ptrs = b_ptr + offs_n[:, None] * stride_bn + byte_offs_k[None, :] * stride_bk
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_tiles_per_split = tl.cdiv(K, BLOCK_K)
if SPLIT_K > 1:
tiles_per_split = tl.cdiv(k_tiles_per_split, SPLIT_K)
start_k = pid_k * tiles_per_split
end_k = tl.minimum((pid_k + 1) * tiles_per_split, k_tiles_per_split)
else:
start_k = 0
end_k = k_tiles_per_split
a_ptrs += start_k * (BLOCK_K // 2) * stride_ak
b_ptrs += start_k * (BLOCK_K // 2) * stride_bk
for k_i in range(start_k, end_k):
# Scale loading (BLOCK_K is exactly 32)
# Therefore each iteration handles exactly ONE scale
a_scale_ptrs = a_scales_ptr + offs_m[:, None] * stride_asm + k_i * stride_ask
b_scale_ptrs = b_scales_ptr + offs_n[:, None] * stride_bsn + k_i * stride_bsk
if EVEN_K:
a_bytes = tl.load(a_ptrs, mask=(offs_m[:, None] < M), other=0)
b_bytes = tl.load(b_ptrs, mask=(offs_n[:, None] < N), other=0)
a_scale_val = tl.load(a_scale_ptrs, mask=(offs_m[:, None] < M), other=0)
b_scale_val = tl.load(b_scale_ptrs, mask=(offs_n[:, None] < N), other=0)
else:
k_mask_byte = (k_i * (BLOCK_K // 2) + byte_offs_k) < (K // 2)
a_bytes = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & k_mask_byte[None, :], other=0)
b_bytes = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & k_mask_byte[None, :], other=0)
a_scale_val = tl.load(a_scale_ptrs, mask=(offs_m[:, None] < M), other=0)
b_scale_val = tl.load(b_scale_ptrs, mask=(offs_n[:, None] < N), other=0)
# Decode via bitcast algebraic logic perfectly fitting into registers!
a_fp8 = decode_e2m1_to_fp8(a_bytes, is_odd[None, :])
b_fp8 = decode_e2m1_to_fp8(b_bytes, is_odd[None, :])
inner_acc = tl.dot(a_fp8, tl.trans(b_fp8), out_dtype=tl.float32)
# Scale outer product algebraically
a_sc_f32 = tl.exp2(a_scale_val.to(tl.float32) - 127.0) # [BLOCK_M, 1]
b_sc_f32 = tl.exp2(b_scale_val.to(tl.float32) - 127.0) # [BLOCK_N, 1]
acc += inner_acc * a_sc_f32 * tl.trans(b_sc_f32)
a_ptrs += (BLOCK_K // 2) * stride_ak
b_ptrs += (BLOCK_K // 2) * stride_bk
acc_out = acc.to(c_ptr.type.element_ty)
if SPLIT_K > 1:
c_ptrs = c_ptr + pid_k * (M * N) + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
else:
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, acc_out, mask=c_mask)
def get_tune_config(M, N, K):
if M <= 16:
block_m, block_n, block_k, num_warps = 16, 64, 32, 4
split_k = 16 if K >= 4096 else 8
elif M <= 32:
block_m, block_n, block_k, num_warps = 32, 64, 32, 4
split_k = 8 if K >= 2048 else 4
elif M <= 64:
block_m, block_n, block_k, num_warps = 64, 64, 32, 4
split_k = 4 if K >= 2048 else 2
else:
block_m, block_n, block_k, num_warps = 128, 128, 32, 8
split_k = 1
k_tiles = triton.cdiv(K, block_k)
if k_tiles < split_k:
split_k = max(1, k_tiles)
return block_m, block_n, block_k, split_k, num_warps, 2
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_q.shape[0]
A_q_u8, A_scale_raw = dynamic_mxfp4_quant(A)
# Native Linear Math arrays
A_q = A_q_u8[:M, :K // 2].contiguous().view(torch.uint8)
A_scale = A_scale_raw[:M, :K // 32].contiguous().view(torch.uint8)
B_q_u8 = B_q.contiguous().view(torch.uint8)
B_scale = _inverse_e8m0_shuffle(B_scale_sh, N).view(torch.uint8).contiguous()
block_m, block_n, block_k, split_k, num_warps, num_stages = get_tune_config(M, N, K)
if split_k > 1:
out = torch.empty((split_k, M, N), device=A.device, dtype=torch.float32)
else:
out = torch.empty((M, N), device=A.device, dtype=torch.bfloat16)
grid = (triton.cdiv(M, block_m) * triton.cdiv(N, block_n), split_k)
_mxfp4_gemm_kernel[grid](
A_q, B_q_u8, out,
A_scale, B_scale,
M, N, K,
A_q.stride(0), A_q.stride(1),
B_q_u8.stride(0), B_q_u8.stride(1),
out.stride(-2), out.stride(-1),
A_scale.stride(0), A_scale.stride(1),
B_scale.stride(0), B_scale.stride(1),
BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k,
GROUP_M=4, SPLIT_K=split_k,
num_warps=num_warps, num_stages=num_stages
)
if split_k > 1:
return out.sum(dim=0).to(torch.bfloat16)
return out
scrolls · 200 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