submission 746650
guangxiangdebizi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 988 lines, June 9 Researcher Reciprocity License v1.0.
submission_exp17_v05.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-746650?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:f016da420ce7d4ef12098fc742e9884b675934dade3c65c0900a1551f398fed3
license declaredunknown
license concludedunknown
authorsguangxiangdebizi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""Experiment 17 v05 for AMD MXFP4 GEMM.split-k
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)tile-k = 512
- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)tile-m = 16
- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)tile-n = 128
- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)Kernel source
submission_exp17_v05.py988 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""Experiment 17 v05 for AMD MXFP4 GEMM.
Cherry-pick best configs on top of exp15_v01:
- Shape 2 (16,2112,7168): num_stages 1->2 (validated -0.8µs in exp17_v04)
- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)
- Shape 6 (256,3072,1536): keep two-stage
"""
from __future__ import annotations
import torch
from task import input_t, output_t
try:
import triton
import triton.language as tl
except Exception:
triton = None
tl = None
def _cfg(
block_size_m: int,
block_size_n: int,
block_size_k: int,
group_size_m: int,
num_warps: int,
num_stages: int,
waves_per_eu: int,
matrix_instr_nonkdim: int,
cache_modifier: str | None,
num_ksplit: int,
) -> dict[str, int | str | None]:
return {
"BLOCK_SIZE_M": block_size_m,
"BLOCK_SIZE_N": block_size_n,
"BLOCK_SIZE_K": block_size_k,
"GROUP_SIZE_M": group_size_m,
"num_warps": num_warps,
"num_stages": num_stages,
"waves_per_eu": waves_per_eu,
"matrix_instr_nonkdim": matrix_instr_nonkdim,
"cache_modifier": cache_modifier,
"NUM_KSPLIT": num_ksplit,
}
_DEFAULT_CONFIG = _cfg(32, 64, 512, 1, 8, 1, 2, 16, None, 1)
_SMALL_K_TINY_M_CONFIG = _cfg(4, 128, 512, 1, 4, 1, 2, 16, ".cg", 1)
_SMALL_K_SMALL_M_CONFIG = _cfg(16, 128, 512, 1, 4, 1, 2, 16, None, 1)
_SMALL_K_4096_CONFIG = _cfg(8, 128, 512, 1, 8, 1, 2, 16, ".cg", 1)
_LARGE_M_CONFIG = _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)
_TWO_STAGE_64_CONFIG = _cfg(16, 128, 512, 1, 8, 2, 4, 16, None, 1)
_TWO_STAGE_256_CONFIG = _cfg(16, 256, 256, 4, 8, 2, 4, 16, None, 1)
_SPECIALIZED_NK_CONFIGS: dict[
tuple[int, int], list[tuple[int | None, dict[str, int | str | None]]]
] = {
(2112, 7168): [
(8, _cfg(8, 128, 512, 1, 4, 1, 1, 16, ".cg", 14)),
(16, _cfg(16, 128, 512, 1, 4, 1, 1, 16, ".cg", 14)),
(32, _cfg(16, 128, 512, 1, 4, 1, 2, 16, None, 14)),
(64, _cfg(16, 128, 512, 1, 4, 1, 2, 16, None, 14)),
(128, _cfg(32, 128, 512, 1, 4, 1, 2, 16, None, 14)),
(256, _cfg(32, 128, 512, 1, 4, 1, 2, 16, None, 14)),
(None, _cfg(32, 128, 256, 4, 2, 2, 2, 16, None, 1)),
],
(7168, 2048): [
(8, _cfg(8, 128, 512, 1, 8, 2, 1, 16, ".cg", 4)),
(16, _cfg(16, 128, 512, 1, 4, 2, 2, 16, ".cg", 4)),
(32, _cfg(16, 128, 512, 1, 8, 2, 2, 16, ".cg", 4)),
(64, _cfg(16, 128, 512, 1, 8, 2, 4, 16, None, 1)),
(128, _cfg(32, 128, 256, 4, 8, 2, 4, 16, None, 1)),
(256, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
(None, _cfg(32, 256, 256, 1, 8, 2, 1, 16, None, 1)),
],
(3072, 1536): [
(16, _cfg(16, 64, 256, 1, 4, 2, 2, 16, None, 4)),
(64, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
(256, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
(None, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
],
(4096, 512): [
(32, _SMALL_K_4096_CONFIG),
(None, _DEFAULT_CONFIG),
],
}
_FIXED_SHAPE_CONFIGS: dict[tuple[int, int, int], dict[str, int | str | None]] = {
(4, 2880, 512): _SMALL_K_TINY_M_CONFIG,
(16, 2112, 7168): _cfg(16, 128, 512, 1, 4, 2, 1, 16, ".cg", 14),
(32, 4096, 512): _SMALL_K_4096_CONFIG,
(32, 2880, 512): _SMALL_K_SMALL_M_CONFIG,
(64, 7168, 2048): _cfg(16, 128, 512, 1, 8, 2, 4, 16, None, 1),
(256, 3072, 1536): _LARGE_M_CONFIG,
}
_TWO_STAGE_FIXED_SHAPE_CONFIGS: dict[tuple[int, int, int], dict[str, int | str | None]] = {
(256, 3072, 1536): _TWO_STAGE_256_CONFIG,
}
if triton is not None:
@triton.jit
def _pid_grid(
pid: int,
num_pid_m: int,
num_pid_n: int,
GROUP_SIZE_M: tl.constexpr = 1,
):
if GROUP_SIZE_M == 1:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
else:
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
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
tl.assume(group_size_m >= 0)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
return pid_m, pid_n
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
exp_bias_fp32: tl.constexpr = 127
exp_bias_fp4: tl.constexpr = 1
ebits_fp32: tl.constexpr = 8
ebits_fp4: tl.constexpr = 2
mbits_fp32: tl.constexpr = 23
mbits_fp4: tl.constexpr = 1
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, num_quant_blocks, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = (x * quant_scale).to(tl.uint32, bitcast=True)
sign = qx & 0x80000000
qx = qx ^ sign
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
normal_mask = ~(saturate_mask | denormal_mask)
denorm_exp: tl.constexpr = (
(exp_bias_fp32 - exp_bias_fp4) + (mbits_fp32 - mbits_fp4) + 1
)
denorm_mask_int: tl.constexpr = denorm_exp << mbits_fp32
denorm_mask_float: tl.constexpr = tl.cast(
denorm_mask_int, tl.float32, bitcast=True
)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
normal_x = qx
mant_odd = (normal_x >> (mbits_fp32 - mbits_fp4)) & 1
val_to_add = ((exp_bias_fp4 - exp_bias_fp32) << mbits_fp32) + (1 << 21) - 1
normal_x += val_to_add
normal_x += mant_odd
normal_x = (normal_x >> (mbits_fp32 - mbits_fp4)).to(tl.uint8)
e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
sign_lp = sign >> (mbits_fp32 + ebits_fp32 - mbits_fp4 - ebits_fp4)
e2m1_value = e2m1_value | sign_lp.to(tl.uint8)
e2m1_value = tl.reshape(
e2m1_value,
[BLOCK_SIZE_M, num_quant_blocks, MXFP4_QUANT_BLOCK_SIZE // 2, 2],
)
evens, odds = tl.split(e2m1_value)
x_fp4 = (evens | (odds << 4)).reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, num_quant_blocks)
@triton.heuristics({"EVEN_K": lambda args: args["K"] % args["BLOCK_SIZE_K"] == 0})
@triton.jit
def _mxfp4_quant_matrix_kernel(
a_ptr,
a_fp4_ptr,
a_scales_ptr,
M,
K,
stride_am,
stride_ak,
stride_afp4_m,
stride_afp4_k,
stride_asm,
stride_ask,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
EVEN_K: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_k = tl.program_id(axis=1)
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k = pid_k * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
else:
a_bf16 = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K), other=0)
a_fp4, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
offs_k_fp4 = pid_k * (BLOCK_SIZE_K // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
a_fp4_ptrs = (
a_fp4_ptr
+ offs_m[:, None] * stride_afp4_m
+ offs_k_fp4[None, :] * stride_afp4_k
)
tl.store(a_fp4_ptrs, a_fp4, mask=offs_m[:, None] < M)
offs_k_scale = pid_k * (BLOCK_SIZE_K // 32) + tl.arange(0, BLOCK_SIZE_K // 32)
a_scale_ptrs = (
a_scales_ptr
+ offs_m[:, None] * stride_asm
+ offs_k_scale[None, :] * stride_ask
)
tl.store(a_scale_ptrs, a_scales, mask=offs_m[:, None] < M)
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
}
)
@triton.jit
def _gemm_a16wfp4_prequant_kernel(
a_ptr,
a_scales_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_asm,
stride_ask,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_asm > 0)
tl.assume(stride_ask > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = _pid_grid(
pid,
num_pid_m,
num_pid_n,
GROUP_SIZE_M=GROUP_SIZE_M,
)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
tl.assume(pid_k >= 0)
scale_group_size: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_fp4 = tl.arange(0, BLOCK_SIZE_K // 2)
offs_k_split_fp4 = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k_fp4
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_fp4[None, :] * stride_ak
)
offs_asm = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // scale_group_size)) + tl.arange(
0, BLOCK_SIZE_K // scale_group_size
)
a_scale_ptrs = (
a_scales_ptr
+ offs_asm[:, None] * stride_asm
+ offs_ks[None, :] * stride_ask
)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
) % N
offs_ks_b = (pid_k * (SPLITK_BLOCK_SIZE // scale_group_size) * 32) + tl.arange(
0, BLOCK_SIZE_K // scale_group_size * 32
)
b_scale_ptrs = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ offs_ks_b[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
if BLOCK_SIZE_M < 32:
a_scales = tl.load(a_scale_ptrs)
else:
a_scales = (
tl.load(a_scale_ptrs)
.reshape(
BLOCK_SIZE_M // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // scale_group_size)
)
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // scale_group_size)
)
if EVEN_K:
a = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a = tl.load(
a_ptrs,
mask=offs_k_fp4[None, :] < K - k_iter * (BLOCK_SIZE_K // 2),
other=0,
)
b = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :]
< (K - k_iter * (BLOCK_SIZE_K // 2)) * 16,
other=0,
cache_modifier=cache_modifier,
)
b = (
b.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
a_scale_ptrs += (BLOCK_SIZE_K // scale_group_size) * stride_ask
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
}
)
@triton.jit
def _gemm_a16wfp4_preshuffle_kernel(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = _pid_grid(
pid,
num_pid_m,
num_pid_n,
GROUP_SIZE_M=GROUP_SIZE_M,
)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
tl.assume(pid_k >= 0)
scale_group_size: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // scale_group_size) * 32) + tl.arange(
0, BLOCK_SIZE_K // scale_group_size * 32
)
b_scale_ptrs = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // scale_group_size)
)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a_bf16 = tl.load(
a_ptrs,
mask=offs_k_bf16[None, :] < 2 * K - k_iter * BLOCK_SIZE_K,
other=0,
)
b = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :]
< (K - k_iter * (BLOCK_SIZE_K // 2)) * 16,
other=0,
cache_modifier=cache_modifier,
)
b = (
b.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
@triton.jit
def _gemm_reduce_kernel(
c_in_ptr,
c_out_ptr,
M,
N,
stride_c_in_k,
stride_c_in_m,
stride_c_in_n,
stride_c_out_m,
stride_c_out_n,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr,
MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
offs_k = tl.arange(0, MAX_KSPLIT)
c_in_ptrs = (
c_in_ptr
+ (offs_k[:, None, None] * stride_c_in_k)
+ (offs_m[None, :, None] * stride_c_in_m)
+ (offs_n[None, None, :] * stride_c_in_n)
)
if ACTUAL_KSPLIT == MAX_KSPLIT:
c = tl.load(c_in_ptrs)
else:
c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
c = tl.sum(c, axis=0).to(c_out_ptr.type.element_ty)
c_out_ptrs = (
c_out_ptr
+ (offs_m[:, None] * stride_c_out_m)
+ (offs_n[None, :] * stride_c_out_n)
)
tl.store(c_out_ptrs, c)
def _pick_specialized_nk_config(
m: int, n: int, k: int
) -> dict[str, int | str | None] | None:
configs = _SPECIALIZED_NK_CONFIGS.get((n, k))
if configs is None:
return None
for upper_bound, config in configs:
if upper_bound is None or m <= upper_bound:
return config
return None
def _pick_two_stage_config(
m: int, n: int, k: int
) -> dict[str, int | str | None] | None:
config = _TWO_STAGE_FIXED_SHAPE_CONFIGS.get((m, n, k))
if config is None:
return None
return dict(config)
def _pick_config(m: int, n: int, k: int) -> dict[str, int | str | None]:
exact = _FIXED_SHAPE_CONFIGS.get((m, n, k))
if exact is not None:
return dict(exact)
specialized = _pick_specialized_nk_config(m, n, k)
if specialized is not None:
return dict(specialized)
if k == 512 and m <= 8:
return dict(_SMALL_K_TINY_M_CONFIG)
if m >= 128 and n >= 2048 and k >= 1024:
return dict(_LARGE_M_CONFIG)
return dict(_DEFAULT_CONFIG)
def _get_splitk(k_packed: int, block_size_k: int, num_ksplit: int) -> tuple[int, int, int]:
if triton is None:
raise RuntimeError("Triton is required for the fused MXFP4 kernel.")
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k)
* block_size_k
)
while num_ksplit > 1 and block_size_k > 16:
if (
k_packed % (splitk_block_size // 2) == 0
and splitk_block_size % block_size_k == 0
and k_packed % (block_size_k // 2) == 0
):
break
if k_packed % (splitk_block_size // 2) != 0 and num_ksplit > 1:
num_ksplit //= 2
elif splitk_block_size % block_size_k != 0:
if num_ksplit > 1:
num_ksplit //= 2
elif block_size_k > 16:
block_size_k //= 2
elif k_packed % (block_size_k // 2) != 0 and block_size_k > 16:
block_size_k //= 2
else:
break
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k)
* block_size_k
)
return splitk_block_size, block_size_k, num_ksplit
def _reshape_b_shuffle_for_preshuffle(b_shuffle: torch.Tensor) -> torch.Tensor:
n, k_half = b_shuffle.shape
if n % 16 != 0:
raise ValueError(f"Expected N to be divisible by 16, but got N={n}.")
b_in = b_shuffle if b_shuffle.is_contiguous() else b_shuffle.contiguous()
b_bytes = b_in.view(torch.uint8)
return b_bytes.view(n // 16, k_half * 16)
def _reshape_b_scale_for_preshuffle(
b_scale_sh: torch.Tensor,
n: int,
k_bf16: int,
) -> torch.Tensor:
if n % 32 != 0:
raise ValueError(f"Expected N to be divisible by 32, but got N={n}.")
k_scale = k_bf16 // 32
scale_slice = b_scale_sh[:n, :k_scale]
scale_in = scale_slice if scale_slice.is_contiguous() else scale_slice.contiguous()
scale_bytes = scale_in.view(torch.uint8)
return scale_bytes.view(n // 32, k_scale * 32)
def _gemm_a16wfp4_preshuffle(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
if triton is None:
raise RuntimeError("Triton is not available in this environment.")
m, k_bf16 = a.shape
n_blocks, packed_k_x16 = b_shuffle.shape
if packed_k_x16 % 16 != 0:
raise ValueError(
f"Expected preshuffled B second dim to be divisible by 16, got {packed_k_x16}."
)
n = n_blocks * 16
k_packed = packed_k_x16 // 16
if 2 * k_packed != k_bf16:
raise ValueError(
"Unexpected preshuffled B shape: "
f"A has bf16 K={k_bf16}, but B encodes packed K={k_packed}."
)
config = _pick_config(m, n, k_bf16)
num_ksplit = int(config["NUM_KSPLIT"])
block_size_k = int(config["BLOCK_SIZE_K"])
if num_ksplit > 1:
splitk_block_size, block_size_k, num_ksplit = _get_splitk(
k_packed, block_size_k, num_ksplit
)
config["SPLITK_BLOCK_SIZE"] = splitk_block_size
config["BLOCK_SIZE_K"] = block_size_k
config["NUM_KSPLIT"] = num_ksplit
if int(config["BLOCK_SIZE_K"]) >= 2 * k_packed:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
config["NUM_KSPLIT"] = 1
else:
config["SPLITK_BLOCK_SIZE"] = (
int(config["SPLITK_BLOCK_SIZE"])
if "SPLITK_BLOCK_SIZE" in config
else 2 * k_packed
)
config["BLOCK_SIZE_N"] = max(int(config["BLOCK_SIZE_N"]), 32)
if int(config["NUM_KSPLIT"]) > 1:
y_pp = torch.empty(
(int(config["NUM_KSPLIT"]), m, n),
dtype=torch.float32,
device=a.device,
)
y = torch.empty((m, n), dtype=dtype, device=a.device)
else:
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
y_pp = None
y = torch.empty((m, n), dtype=dtype, device=a.device)
grid = lambda meta: ( # noqa: E731
int(meta["NUM_KSPLIT"])
* triton.cdiv(m, int(meta["BLOCK_SIZE_M"]))
* triton.cdiv(n, int(meta["BLOCK_SIZE_N"])),
)
_gemm_a16wfp4_preshuffle_kernel[grid](
a,
b_shuffle,
y if y_pp is None else y_pp,
b_scale_sh,
m,
n,
k_packed,
a.stride(0),
a.stride(1),
b_shuffle.stride(0),
b_shuffle.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
b_scale_sh.stride(0),
b_scale_sh.stride(1),
**config,
)
if y_pp is not None:
reduce_block_m = 16
reduce_block_n = 64
actual_ksplit = triton.cdiv(k_packed, int(config["SPLITK_BLOCK_SIZE"]) // 2)
grid_reduce = (
triton.cdiv(m, reduce_block_m),
triton.cdiv(n, reduce_block_n),
)
_gemm_reduce_kernel[grid_reduce](
y_pp,
y,
m,
n,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
BLOCK_SIZE_M=reduce_block_m,
BLOCK_SIZE_N=reduce_block_n,
ACTUAL_KSPLIT=actual_ksplit,
MAX_KSPLIT=triton.next_power_of_2(int(config["NUM_KSPLIT"])),
)
return y
def _gemm_a16wfp4_two_stage(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
if triton is None:
raise RuntimeError("Triton is not available in this environment.")
m, k_bf16 = a.shape
n_blocks, packed_k_x16 = b_shuffle.shape
if packed_k_x16 % 16 != 0:
raise ValueError(
f"Expected preshuffled B second dim to be divisible by 16, got {packed_k_x16}."
)
n = n_blocks * 16
k_packed = packed_k_x16 // 16
if 2 * k_packed != k_bf16:
raise ValueError(
"Unexpected preshuffled B shape: "
f"A has bf16 K={k_bf16}, but B encodes packed K={k_packed}."
)
config = _pick_two_stage_config(m, n, k_bf16)
if config is None:
raise ValueError(f"Two-stage config not defined for shape {(m, n, k_bf16)}.")
if int(config["NUM_KSPLIT"]) != 1:
raise ValueError("Two-stage prototype currently expects NUM_KSPLIT == 1.")
block_size_k = int(config["BLOCK_SIZE_K"])
config["BLOCK_SIZE_N"] = max(int(config["BLOCK_SIZE_N"]), 32)
a_fp4 = torch.empty((m, k_bf16 // 2), dtype=torch.uint8, device=a.device)
a_scales = torch.empty((m, k_bf16 // 32), dtype=torch.uint8, device=a.device)
y = torch.empty((m, n), dtype=dtype, device=a.device)
grid_quant = (
triton.cdiv(m, int(config["BLOCK_SIZE_M"])),
triton.cdiv(k_bf16, block_size_k),
)
_mxfp4_quant_matrix_kernel[grid_quant](
a,
a_fp4,
a_scales,
m,
k_bf16,
a.stride(0),
a.stride(1),
a_fp4.stride(0),
a_fp4.stride(1),
a_scales.stride(0),
a_scales.stride(1),
BLOCK_SIZE_M=int(config["BLOCK_SIZE_M"]),
BLOCK_SIZE_K=block_size_k,
)
grid = lambda meta: ( # noqa: E731
triton.cdiv(m, int(meta["BLOCK_SIZE_M"]))
* triton.cdiv(n, int(meta["BLOCK_SIZE_N"])),
)
_gemm_a16wfp4_prequant_kernel[grid](
a_fp4,
a_scales,
b_shuffle,
y,
b_scale_sh,
m,
n,
k_packed,
a_fp4.stride(0),
a_fp4.stride(1),
a_scales.stride(0),
a_scales.stride(1),
b_shuffle.stride(0),
b_shuffle.stride(1),
0,
y.stride(0),
y.stride(1),
b_scale_sh.stride(0),
b_scale_sh.stride(1),
BLOCK_SIZE_M=int(config["BLOCK_SIZE_M"]),
BLOCK_SIZE_N=int(config["BLOCK_SIZE_N"]),
BLOCK_SIZE_K=block_size_k,
GROUP_SIZE_M=int(config["GROUP_SIZE_M"]),
NUM_KSPLIT=1,
SPLITK_BLOCK_SIZE=2 * k_packed,
num_warps=int(config["num_warps"]),
num_stages=int(config["num_stages"]),
waves_per_eu=int(config["waves_per_eu"]),
matrix_instr_nonkdim=int(config["matrix_instr_nonkdim"]),
cache_modifier=config["cache_modifier"],
)
return y
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
a, _, _, b_shuffle, b_scale_sh = data
a_in = a if a.is_contiguous() else a.contiguous()
n = b_shuffle.shape[0]
k_bf16 = a_in.shape[1]
b_preshuffle = _reshape_b_shuffle_for_preshuffle(b_shuffle)
b_scale_preshuffle = _reshape_b_scale_for_preshuffle(b_scale_sh, n, k_bf16)
# Only Shape 6 keeps the extra A prequant pass in exp15.
if (a_in.shape[0], n, k_bf16) in _TWO_STAGE_FIXED_SHAPE_CONFIGS:
return _gemm_a16wfp4_two_stage(a_in, b_preshuffle, b_scale_preshuffle)
return _gemm_a16wfp4_preshuffle(a_in, b_preshuffle, b_scale_preshuffle)
scrolls · 988 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