submission 710696
mbuchel · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 158 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-710696?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:fd65780c0e044296fae2f6aa5cabcf1a2df306c92980039c933105cdc46ea1c8
license declaredunknown
license concludedunknown
authorsmbuchel
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
_dynamic_mxfp4_quant[(triton.cdiv(M, 128), triton.cdiv(N, 32))](x, x_fp4, bs, *x.stride(), *x_fp4.stride(), *bs.stride(), M, N, 128, 32, num_warps=4, num_stages=3)split-k
SPLIT_K: tl.constexpr,stages = 3
_dynamic_mxfp4_quant[(triton.cdiv(M, 128), triton.cdiv(N, 32))](x, x_fp4, bs, *x.stride(), *x_fp4.stride(), *bs.stride(), M, N, 128, 32, num_warps=4, num_stages=3)tile-m = 512
BK, BN, BM = 512, 256, 32Kernel source
submission.py158 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def _dynamic_mxfp4_quant(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m, stride_x_n, stride_x_fp4_m, stride_x_fp4_n, stride_bs_m, stride_bs_n,
M, N, BLOCK_SIZE: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
pid_m, pid_n = tl.program_id(0), tl.program_id(1)
scaleN_valid = (N + 31) // 32
x_off_m, x_off_n = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE), pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x = tl.load(x_ptr + x_off_m[:, None] * stride_x_m + x_off_n[None, :] * stride_x_n, mask=(x_off_m < M)[:, None] & (x_off_n < N)[None, :]).to(tl.float32)
amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
amax_bits = amax.to(tl.uint32, bitcast=True)
exp_mask: tl.constexpr = 0xFF800000
amax_m = (amax_bits + 0x200000) & exp_mask
su = tl.minimum(tl.maximum((amax_m.to(tl.int32) >> 23) - 129, -127), 127)
exp_val = ((127 - su).to(tl.int32) << 23).to(tl.float32, bitcast=True)
qx = x * exp_val
qb = qx.to(tl.uint32, bitcast=True)
s, qa = qb & 0x80000000, qb ^ (qb & 0x80000000)
dm = (qa.to(tl.float32, bitcast=True) < 1.0)
dvi = 149 << 23
dv = (qa.to(tl.float32, bitcast=True) + tl.cast(dvi, tl.float32, bitcast=True)).to(tl.uint32, bitcast=True) - dvi
qa_i = qa.to(tl.int32)
nv = ((qa_i + ((-126) << 23) + (1 << 21) - 1 + ((qa_i >> 22) & 1)) >> 22).to(tl.uint8)
res = tl.where(qa.to(tl.float32, bitcast=True) >= 6.0, 0x7, tl.where(dm, dv.to(tl.uint8), nv)) | (s >> 28).to(tl.uint8)
res_p = tl.reshape(res, [BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
ev, od = tl.split(res_p)
out = ev | (od << 4)
tl.store(x_fp4_ptr + x_off_m[:, None] * stride_x_fp4_m + (pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2))[None, :] * stride_x_fp4_n, out, mask=(x_off_m < M)[:, None] & ( (pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)) < (N // 2))[None, :])
bs_e8m0 = (su.reshape((BLOCK_SIZE,)) + 127).to(tl.uint8)
b0, b1 = x_off_m // 32, x_off_m % 32
b2, b1_ = b1 % 16, b1 // 16
b3, b4 = pid_n // 8, pid_n % 8
b5, b4_ = b4 % 4, b4 // 4
tl.store(bs_ptr + b1_ + b4_*2 + b2*4 + b5*64 + b3*256 + b0*32*scaleN_valid, bs_e8m0, mask=(x_off_m < M))
@triton.jit
def block_scaled_matmul_kernel_fused(
a_ptr, b_ptr, c_ptr, b_scale_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk, stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
pid_m, pid_n = tl.program_id(0), tl.program_id(1)
num_k_it = tl.cdiv(K, BLOCK_K)
offs_k_a, offs_k_b = tl.arange(0, BLOCK_K), tl.arange(0, BLOCK_K // 2)
offs_am, offs_bn = pid_m * BLOCK_M + tl.arange(0, BLOCK_M), pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_a[None, :] * stride_ak)
b_ptrs = b_ptr + (offs_k_b[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
offs_asn, offs_ks = pid_n * (BLOCK_N // 32) + tl.arange(0, (BLOCK_N // 32)), tl.arange(0, BLOCK_K // 32 * 32)
b_scale_ptrs = b_scale_ptr + offs_asn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, num_k_it):
x = tl.load(a_ptrs, mask=(offs_am[:, None] < M) & (offs_k_a[None, :] < K), other=0).to(tl.float32)
x_g = tl.reshape(x, (BLOCK_M, BLOCK_K // 32, 32))
am = tl.max(tl.abs(x_g), axis=2, keep_dims=True)
am_b = (am.to(tl.uint32, bitcast=True) + 0x200000) & 0xFF800000
su = tl.minimum(tl.maximum((am_b.to(tl.int32) >> 23) - 129, -127), 127)
ev = ((127 - su).to(tl.int32) << 23).to(tl.float32, bitcast=True)
qx = x_g * ev
qb = qx.to(tl.uint32, bitcast=True)
s, qa = qb & 0x80000000, qb ^ (qb & 0x80000000)
dm = (qa.to(tl.float32, bitcast=True) < 1.0)
dvi = 149 << 23
dv = (qa.to(tl.float32, bitcast=True) + tl.cast(dvi, tl.float32, bitcast=True)).to(tl.uint32, bitcast=True) - dvi
qa_i = qa.to(tl.int32)
nv = ((qa_i + ((-126) << 23) + (1 << 21) - 1 + ((qa_i >> 22) & 1)) >> 22).to(tl.uint8)
res = tl.where(qa.to(tl.float32, bitcast=True) >= 6.0, 0x7, tl.where(dm, dv.to(tl.uint8), nv)) | (s >> 28).to(tl.uint8)
res_p = tl.reshape(res, [BLOCK_M, BLOCK_K // 2, 2])
ev_, od_ = tl.split(res_p)
a_mxfp4 = (ev_ | (od_ << 4))
a_s = (su.reshape((BLOCK_M, BLOCK_K // 32)) + 127).to(tl.uint8)
b_s = tl.load(b_scale_ptrs, mask=(offs_asn[:, None] < ((N + 31) // 32) * 32)).reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // 32)
b_raw = tl.load(b_ptrs, mask=(offs_k_b[:, None] < K // 2) & (offs_bn[None, :] < N), other=0)
accumulator += tl.dot_scaled(a_mxfp4, a_s, "e2m1", b_raw, b_s, "e2m1")
a_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * stride_bk
b_scale_ptrs += BLOCK_K * stride_bsk
tl.store(c_ptr + (pid_m * BLOCK_M + tl.arange(0, BLOCK_M))[:, None] * stride_cm + (pid_n * BLOCK_N + tl.arange(0, BLOCK_N))[None, :] * stride_cn, accumulator.to(c_ptr.type.element_ty), mask=(pid_m * BLOCK_M + tl.arange(0, BLOCK_M) < M)[:, None] & (pid_n * BLOCK_N + tl.arange(0, BLOCK_N) < N)[None, :])
@triton.jit
def block_scaled_matmul_kernel_standard(
a_ptr, b_ptr, c_ptr, a_scale_ptr, b_scale_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,
SPLIT_K: tl.constexpr,
):
pid_m, pid_n, pid_k = tl.program_id(0), tl.program_id(1), tl.program_id(2)
num_k_it = tl.cdiv(K, BLOCK_K)
its = tl.cdiv(num_k_it, SPLIT_K)
sk, ek = pid_k * its, tl.minimum((pid_k + 1) * its, num_k_it)
offs_m, offs_n, offs_k = pid_m * BLOCK_M + tl.arange(0, BLOCK_M), pid_n * BLOCK_N + tl.arange(0, BLOCK_N), tl.arange(0, BLOCK_K // 2)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + (sk * BLOCK_K // 2 + offs_k[None, :]) * stride_ak
b_ptrs = b_ptr + (sk * BLOCK_K // 2 + offs_k[:, None]) * stride_bk + offs_n[None, :] * stride_bn
offs_am_s, offs_an_s, offs_ks = pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32), pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32), tl.arange(0, BLOCK_K // 32 * 32)
ase_ptrs = a_scale_ptr + offs_am_s[:, None] * stride_asm + (sk * BLOCK_K + offs_ks[None, :]) * stride_ask
bse_ptrs = b_scale_ptr + offs_an_s[:, None] * stride_bsn + (sk * BLOCK_K + offs_ks[None, :]) * stride_bsk
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(sk, ek):
as_ = tl.load(ase_ptrs, mask=(offs_am_s[:, None] < M)).reshape(BLOCK_M // 32, BLOCK_K // 256, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_M, BLOCK_K // 32)
bs_raw = tl.load(bse_ptrs, mask=(offs_an_s[:, None] < ((N + 31) // 32) * 32))
bs_ = bs_raw.reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // 32)
a, b = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K // 2), other=0), tl.load(b_ptrs, mask=(offs_k[:, None] < K // 2) & (offs_n[None, :] < N), other=0)
acc += tl.dot_scaled(a, as_, "e2m1", b, bs_, "e2m1")
a_ptrs, b_ptrs = a_ptrs + (BLOCK_K // 2) * stride_ak, b_ptrs + (BLOCK_K // 2) * stride_bk
ase_ptrs, bse_ptrs = ase_ptrs + BLOCK_K * stride_ask, bse_ptrs + BLOCK_K * stride_bsk
c_off = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M))[:, None] * stride_cm + (pid_n * BLOCK_N + tl.arange(0, BLOCK_N))[None, :] * stride_cn
c_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M) < M)[:, None] & (pid_n * BLOCK_N + tl.arange(0, BLOCK_N) < N)[None, : ]
if SPLIT_K == 1: tl.store(c_ptr + c_off, acc.to(c_ptr.type.element_ty), mask=c_m)
else: tl.atomic_add(c_ptr + c_off, acc.to(c_ptr.type.element_ty), mask=c_m)
_global_c_cache, _global_q_cache = {}, {}
def triton_dynamic_mxfp4_quant(x):
global _global_q_cache
M, N = x.shape
scaleK_p = triton.cdiv(triton.cdiv(N, 32), 8) * 8
if (M, N) not in _global_q_cache: _global_q_cache[(M, N)] = (torch.empty((M, N // 2), dtype=torch.uint8, device=x.device), torch.empty((triton.cdiv(M, 256) * 256, scaleK_p), dtype=torch.uint8, device=x.device))
x_fp4, bs = _global_q_cache[(M, N)]
_dynamic_mxfp4_quant[(triton.cdiv(M, 128), triton.cdiv(N, 32))](x, x_fp4, bs, *x.stride(), *x_fp4.stride(), *bs.stride(), M, N, 128, 32, num_warps=4, num_stages=3)
return x_fp4, bs
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
M, K, N = A.shape[0], A.shape[1], B.shape[0]
global _global_c_cache
# Best-performing unified adaptive strategy
if M <= 32 and K <= 1024:
if (M, N) not in _global_c_cache: _global_c_cache[(M, N)] = torch.empty((M, N), device=A.device, dtype=torch.bfloat16)
C = _global_c_cache[(M, N)]
BK, BN, BM = 512, 256, 32
grid = (triton.cdiv(M, BM), triton.cdiv(N, BN))
block_scaled_matmul_kernel_fused[grid](A, B_q.view(torch.uint8), C, B_scale_sh.view(torch.uint8), M, N, K, *A.stride(), *B_q.stride(), *C.stride(), triton.cdiv(triton.cdiv(K, 32), 8) * 8 * 32, 1, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, num_warps=8, num_stages=1)
return C
else:
A_q, As = triton_dynamic_mxfp4_quant(A)
if M <= 32:
cfg, sk = {"BM": 32, "BN": 256, "BK": 256, "W": 8}, 4
else:
cfg, sk = {"BM": 128, "BN": 128, "BK": 256, "W": 8}, 1
if (M, N, sk) not in _global_c_cache: _global_c_cache[(M, N, sk)] = torch.empty((M, N), device=A.device, dtype=torch.float32 if sk > 1 else torch.bfloat16)
C = _global_c_cache[(M, N, sk)]
if sk > 1: C.zero_()
block_scaled_matmul_kernel_standard[(triton.cdiv(M, cfg["BM"]), triton.cdiv(N, cfg["BN"]), sk)](A_q, B_q.view(torch.uint8), C, As, B_scale_sh.view(torch.uint8), M, N, K, *A_q.stride(), *B_q.stride(), *C.stride(), triton.cdiv(triton.cdiv(K, 32), 8) * 8 * 32, 1, triton.cdiv(triton.cdiv(K, 32), 8) * 8 * 32, 1, BLOCK_M=cfg["BM"], BLOCK_N=cfg["BN"], BLOCK_K=cfg["BK"], SPLIT_K=sk, num_warps=cfg["W"], num_stages=3)
return C.to(torch.bfloat16) if sk > 1 else C
scrolls · 158 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