submission 550927
oofbaroomf · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1891 lines, June 9 Researcher Reciprocity License v1.0.
amd_mxfp4_mm_hybrid_cfgsearch_aw.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-550927?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:8c1b00320a09f374ecc6c704c9446c77e6d7faa6e963db5a4f173954fd284ac4
license declaredunknown
license concludedunknown
authorsoofbaroomf
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
num_warps = 1split-k
f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")stages = 1
num_stages = 1tile-m = 32
BLOCK_SIZE_M=32,tile-n = 128
BLOCK_SIZE_N=128,Kernel source
amd_mxfp4_mm_hybrid_cfgsearch_aw.py1891 lines
import importlib.util
import os
from pathlib import Path
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_VIEW_CACHE: dict[tuple[int, int, int], tuple[torch.Tensor, torch.Tensor]] = {}
_OUT_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_PARTIAL_CACHE: dict[tuple[int, int, int, int], torch.Tensor] = {}
_DEFAULT_CFG_CACHE: dict[tuple[int, int, int], dict] = {}
_B_Q_U8_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_UNSHUFFLED_SCALE_CACHE: dict[tuple[int, int, int, int], torch.Tensor] = {}
_A_Q_RAW_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_A_SCALE_RAW_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_A_SCALE_PAD_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_A_SCALE_SH_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_A_SCALE_SH256_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_OUT_PAD32_CACHE: dict[tuple[int, int, int], torch.Tensor] = {}
_LARGE_ROUTE_CACHE: dict[tuple[int, int, int, int], tuple[str, object | None]] = {}
_LARGE_ROUTE_CACHE: dict[tuple[int, int, int, int], tuple[str, str | None]] = {}
_A4W4_CFG_PATH = Path.home() / ".cache" / "oof_a4w4_task_tuned_cachedwrap.csv"
def _write_a4w4_cfg() -> None:
_A4W4_CFG_PATH.parent.mkdir(parents=True, exist_ok=True)
rows = [
(
256,
16,
3072,
1536,
21,
0,
6.0090,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
25.13,
411.03,
0.0,
),
(
256,
32,
3072,
1536,
29,
0,
6.1627,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
49.00,
418.73,
0.0,
),
(
256,
64,
3072,
1536,
21,
0,
6.1490,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
98.22,
455.63,
0.0,
),
(
256,
128,
3072,
1536,
21,
0,
6.1683,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
195.83,
525.92,
0.0,
),
(
256,
256,
3072,
1536,
21,
0,
6.1771,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
391.11,
668.40,
0.0,
),
]
with _A4W4_CFG_PATH.open("w", encoding="ascii") as f:
f.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")
for row in rows:
f.write(",".join(str(x) for x in row) + "\n")
_write_a4w4_cfg()
def _set_a4w4_cfg_env() -> None:
spec = importlib.util.find_spec("aiter")
if spec is None or spec.origin is None:
os.environ["AITER_CONFIG_GEMM_A4W4"] = str(_A4W4_CFG_PATH)
return
stock_cfg = Path(spec.origin).resolve().parent / "configs" / "a4w4_blockscale_tuned_gemm.csv"
if stock_cfg.exists():
os.environ["AITER_CONFIG_GEMM_A4W4"] = os.pathsep.join(
[str(stock_cfg), str(_A4W4_CFG_PATH)]
)
else:
os.environ["AITER_CONFIG_GEMM_A4W4"] = str(_A4W4_CFG_PATH)
_set_a4w4_cfg_env()
_CFG_2880_512_M_LEQ_8 = {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_CFG_2880_512_M_LEQ_4 = {
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
_CFG_2880_512_M32 = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_CFG_2112_7168_M_LEQ_16 = {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
}
_CFG_2112_7168_STOCK_M16 = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 14,
}
_CFG_2112_7168_M16_N64_S2 = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
}
_CFG_2112_7168_STOCK_M32 = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 14,
}
_CFG_2112_7168_STOCK_ANY = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_CFG_7168_2048_A16_M128 = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_CFG_7168_2048_A16_M256 = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_CFG_7168_2048_A16_M256_W1 = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_CFG_4096_512_M32 = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
_QCFG_CURRENT = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 4,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
}
_QCFG_LARGE_64_2048 = [
{
"BLOCK_SIZE_M": 1,
"BLOCK_SIZE_N": 64,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 64,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 64,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 1,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 8,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
_QCFG_CURRENT,
]
_QCFG_LARGE_256_1536 = [
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 64,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 64,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 2,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 8,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 2,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
_QCFG_CURRENT,
]
_ASMQCFG_LARGE_64_2048 = [
{"BLOCK_SIZE": 16, "num_warps": 2},
{"BLOCK_SIZE": 32, "num_warps": 2},
{"BLOCK_SIZE": 64, "num_warps": 4},
{"BLOCK_SIZE": 128, "num_warps": 4},
]
_ASMQCFG_LARGE_256_1536 = [
{"BLOCK_SIZE": 16, "num_warps": 2},
{"BLOCK_SIZE": 32, "num_warps": 4},
{"BLOCK_SIZE": 64, "num_warps": 4},
{"BLOCK_SIZE": 128, "num_warps": 4},
]
_ASMQCFG_LARGE_64_1536 = [
{"BLOCK_SIZE": 16, "num_warps": 2},
{"BLOCK_SIZE": 32, "num_warps": 4},
{"BLOCK_SIZE": 64, "num_warps": 4},
]
_A4W4_QUANT_CFGS = {
"bm32_bn128_i4_w4": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 4,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm32_bn256_i2_w4": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm64_bn128_i4_w4": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"NUM_ITER": 4,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm64_bn256_i2_w4": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm32_bn512_i1_w4": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm64_bn512_i1_w4": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 512,
"NUM_ITER": 1,
"NUM_STAGES": 2,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm32_bn256_i2_w8": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 2,
"num_warps": 8,
"waves_per_eu": 0,
},
"bm64_bn256_i2_w8": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 2,
"num_warps": 8,
"waves_per_eu": 0,
},
"bm32_bn256_i2s1_w4": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
"bm64_bn256_i2s1_w4": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 256,
"NUM_ITER": 2,
"NUM_STAGES": 1,
"num_warps": 4,
"waves_per_eu": 0,
},
}
_A4W4_3072_CANDIDATES = (
"bm32_bn128_i4_w4",
"bm32_bn256_i2_w4",
"bm64_bn128_i4_w4",
"bm64_bn256_i2_w4",
"bm32_bn512_i1_w4",
"bm64_bn512_i1_w4",
"bm32_bn256_i2_w8",
"bm64_bn256_i2_w8",
"bm32_bn256_i2s1_w4",
"bm64_bn256_i2s1_w4",
)
_A4W4_7168_CANDIDATES = (
"bm32_bn128_i4_w4",
"bm32_bn256_i2_w4",
"bm64_bn128_i4_w4",
"bm64_bn256_i2_w4",
"bm32_bn512_i1_w4",
"bm64_bn512_i1_w4",
"bm32_bn256_i2_w8",
"bm64_bn256_i2_w8",
"bm32_bn256_i2s1_w4",
"bm64_bn256_i2s1_w4",
)
@triton.jit
def _mxfp4_quant_op_shuffled(
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
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
normal_mask = not (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)
normal_x = normal_x.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 = s >> (mbits_fp32 + ebits_fp32 - mbits_fp4 - ebits_fp4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
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.jit
def _dynamic_mxfp4_quant_kernel_asm_layout_cfg(
x_ptr,
x_fp4_ptr,
bs_ptr,
stride_x_m,
stride_x_n,
stride_x_fp4_m,
stride_x_fp4_n,
M: tl.constexpr,
N: tl.constexpr,
SCALE_N_VALID: tl.constexpr,
SCALE_M_PAD: tl.constexpr,
SCALE_N_PAD: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
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)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
mant = qx & 0x7FFFFF
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
mant = tl.where(e < E8_BIAS, (0x400000 | (mant >> 1)) >> adjusted_exponents, mant)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (mant >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
out_tensor = evens | (odds << 4)
out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
0, MXFP4_QUANT_BLOCK_SIZE // 2
)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
)
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
bs_offs_n = pid_n
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 4
+ bs_offs_5 * 64
+ bs_offs_3 * 256
+ bs_offs_0 * 32 * SCALE_N_VALID
)
bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < SCALE_N_VALID)[None, :]
bs_mask2 = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[None, :]
bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
@triton.jit
def _shuffle_e8m0_kernel(
src_ptr,
dst_ptr,
stride_src_m,
stride_src_n,
stride_dst_m,
stride_dst_n,
M,
SCALE_N,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_src_m = tl.cast(stride_src_m, tl.int64)
stride_src_n = tl.cast(stride_src_n, tl.int64)
stride_dst_m = tl.cast(stride_dst_m, tl.int64)
stride_dst_n = tl.cast(stride_dst_n, tl.int64)
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
mask = (offs_m < M)[:, None] & (offs_n < SCALE_N)[None, :]
vals = tl.load(
src_ptr + offs_m[:, None] * stride_src_m + offs_n[None, :] * stride_src_n,
mask=mask,
other=127,
)
bs_offs_0 = offs_m[:, None] // 32
bs_offs_1 = offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = offs_n[None, :] // 8
bs_offs_4 = offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
flat = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 4
+ bs_offs_5 * 64
+ bs_offs_3 * 256
+ bs_offs_0 * 32 * SCALE_N
)
dst_rows = flat // SCALE_N
dst_cols = flat % SCALE_N
tl.store(
dst_ptr + dst_rows * stride_dst_m + dst_cols * stride_dst_n,
vals,
mask=mask,
)
@triton.heuristics(
{
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
}
)
@triton.jit
def _dynamic_mxfp4_quant_kernel_shuffled(
x_ptr,
x_fp4_ptr,
bs_sh_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
stride_bs_m_in,
stride_bs_n_in,
M,
N,
SCALE_N,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
stride_bs_n = tl.cast(stride_bs_n_in, tl.int64)
num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
tl.float32
)
out_tensor, bs_e8m0 = _mxfp4_quant_op_shuffled(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
)
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * num_quant_blocks + tl.arange(0, num_quant_blocks)
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_flat = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 4
+ bs_offs_5 * 64
+ bs_offs_3 * 256
+ bs_offs_0 * 32 * SCALE_N
)
bs_rows = bs_flat // SCALE_N
bs_cols = bs_flat % SCALE_N
bs_ptrs = bs_sh_ptr + bs_rows * stride_bs_m + bs_cols * stride_bs_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < SCALE_N)[None, :]
tl.store(bs_ptrs, bs_e8m0, mask=bs_mask)
def _unshuffle_scales(scales_shuffled: torch.Tensor) -> torch.Tensor:
sm, sn = scales_shuffled.shape
scales = scales_shuffled.view(sm // 32, sn // 8, 4, 16, 2, 2)
scales = scales.permute(0, 5, 3, 1, 4, 2).contiguous()
return scales.view(sm, sn)
def _get_views(
b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), b_shuffle.device.index or 0)
cached = _VIEW_CACHE.get(key)
if cached is None:
w = b_shuffle.view(torch.uint8).view(b_shuffle.shape[0] // 16, -1)
w_scale = b_scale_sh.view(torch.uint8).view(b_scale_sh.shape[0] // 32, -1)
cached = (w, w_scale)
_VIEW_CACHE[key] = cached
return cached
def _get_b_q_u8(b_q: torch.Tensor) -> torch.Tensor:
key = (b_q.data_ptr(), b_q.device.index or 0, b_q.shape[0])
cached = _B_Q_U8_CACHE.get(key)
if cached is None:
cached = b_q.view(torch.uint8).contiguous()
_B_Q_U8_CACHE[key] = cached
return cached
def _get_unshuffled_b_scale(
b_q: torch.Tensor, b_scale_sh: torch.Tensor, k: int
) -> torch.Tensor:
key = (b_q.data_ptr(), b_scale_sh.data_ptr(), b_q.device.index or 0, k)
cached = _UNSHUFFLED_SCALE_CACHE.get(key)
if cached is None:
k_scale = k // 32
cached = _unshuffle_scales(b_scale_sh)[: b_q.shape[0], :k_scale].view(
torch.uint8
).contiguous()
_UNSHUFFLED_SCALE_CACHE[key] = cached
return cached
def _get_default_config(m: int, n: int, k: int) -> dict:
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
key = (m, n, k)
cached = _DEFAULT_CFG_CACHE.get(key)
if cached is None:
cached, _ = _get_config(m, n, k // 2, True)
_DEFAULT_CFG_CACHE[key] = dict(cached)
return dict(cached)
def _normalize_config(config: dict, k: int) -> dict:
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
config = dict(config)
if config["NUM_KSPLIT"] > 1:
splitk_block_size, block_size_k, num_ksplit = get_splitk(
k, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = splitk_block_size
config["BLOCK_SIZE_K"] = block_size_k
config["NUM_KSPLIT"] = num_ksplit
if config["BLOCK_SIZE_K"] >= 2 * k:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k)
config["SPLITK_BLOCK_SIZE"] = 2 * k
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
if config["NUM_KSPLIT"] == 1:
config["SPLITK_BLOCK_SIZE"] = 2 * k
return config
def _quant_a4w4(a: torch.Tensor):
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
a_q_raw, a_scale_raw = dynamic_mxfp4_quant(a)
return a_q_raw.view(dtypes.fp4x2), e8m0_shuffle(a_scale_raw).view(dtypes.fp8_e8m0)[
: a.shape[0]
]
def _get_a4w4_quant_buffers(m: int, k: int, device: torch.device):
key = (device.index or 0, m, k)
a_q_raw = _A_Q_RAW_CACHE.get(key)
a_scale_raw = _A_SCALE_RAW_CACHE.get(key)
a_scale_pad = _A_SCALE_PAD_CACHE.get(key)
a_scale_sh = _A_SCALE_SH_CACHE.get(key)
m_pad = triton.cdiv(m, 32) * 32
n_pad = triton.cdiv(k // 32, 8) * 8
if a_q_raw is None:
a_q_raw = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
_A_Q_RAW_CACHE[key] = a_q_raw
if a_scale_raw is None:
a_scale_raw = torch.empty((m, k // 32), dtype=torch.uint8, device=device)
_A_SCALE_RAW_CACHE[key] = a_scale_raw
if a_scale_pad is None:
a_scale_pad = torch.empty((m_pad, n_pad), dtype=torch.uint8, device=device)
_A_SCALE_PAD_CACHE[key] = a_scale_pad
if a_scale_sh is None:
a_scale_sh = torch.empty((m_pad, n_pad), dtype=torch.uint8, device=device)
_A_SCALE_SH_CACHE[key] = a_scale_sh
return a_q_raw, a_scale_raw, a_scale_pad, a_scale_sh
def _get_a4w4_asm_quant_buffers(m: int, k: int, device: torch.device):
key = (device.index or 0, m, k)
a_q_raw = _A_Q_RAW_CACHE.get(key)
a_scale_sh = _A_SCALE_SH256_CACHE.get(key)
if a_q_raw is None:
a_q_raw = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
_A_Q_RAW_CACHE[key] = a_q_raw
if a_scale_sh is None:
m_pad = triton.cdiv(m, 256) * 256
n_pad = triton.cdiv(k // 32, 8) * 8
a_scale_sh = torch.empty((m_pad, n_pad), dtype=torch.uint8, device=device)
_A_SCALE_SH256_CACHE[key] = a_scale_sh
return a_q_raw, a_scale_sh
def _quant_a4w4_cached(a: torch.Tensor):
from aiter import dtypes
from aiter.ops.triton.quant.quant import _dynamic_mxfp4_quant_kernel
m, k = a.shape
a_q_raw, a_scale_raw, a_scale_pad, a_scale_sh = _get_a4w4_quant_buffers(
m, k, a.device
)
if m <= 32:
num_iter = 1
block_size_m = triton.next_power_of_2(m)
block_size_n = 32
num_warps = 1
num_stages = 1
else:
num_iter = 4
block_size_m = 64
block_size_n = 64
num_warps = 4
num_stages = 2
if k <= 16384:
block_size_m = 32
block_size_n = 128
if k <= 1024:
num_iter = 1
num_stages = 1
num_warps = 4
block_size_n = min(256, triton.next_power_of_2(k))
block_size_n = max(32, block_size_n)
block_size_m = min(8, triton.next_power_of_2(m))
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(k, block_size_n * num_iter),
)
_dynamic_mxfp4_quant_kernel[grid](
a,
a_q_raw,
a_scale_raw,
*a.stride(),
*a_q_raw.stride(),
*a_scale_raw.stride(),
M=m,
N=k,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=num_iter,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
NUM_STAGES=num_stages,
num_warps=num_warps,
waves_per_eu=0,
num_stages=1,
)
sm, sn = a_scale_pad.shape
a_scale_pad[:m, : k // 32] = a_scale_raw
a_scale_sh.view(sm // 32, sn // 8, 4, 16, 2, 2).copy_(
a_scale_pad.view(sm // 32, 2, 16, sn // 8, 2, 4).permute(0, 3, 5, 2, 4, 1)
)
return a_q_raw.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)[:m]
def _quant_a4w4_cached_fast_m32_k512(a: torch.Tensor):
from aiter import dtypes
m, k = a.shape
a_q_raw, _a_scale_raw, _a_scale_pad, a_scale_sh = _get_a4w4_quant_buffers(
m, k, a.device
)
grid = (triton.cdiv(m, 32), triton.cdiv(k, 128 * 4))
_dynamic_mxfp4_quant_kernel_shuffled[grid](
a,
a_q_raw,
a_scale_sh,
*a.stride(),
*a_q_raw.stride(),
*a_scale_sh.stride(),
M=m,
N=k,
SCALE_N=k // 32,
BLOCK_SIZE_M=32,
BLOCK_SIZE_N=128,
NUM_ITER=4,
NUM_STAGES=2,
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=4,
waves_per_eu=0,
num_stages=1,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)
def _quant_a4w4_cached_large_cfg(a: torch.Tensor, cfg: dict):
from aiter import dtypes
m, k = a.shape
a_q_raw, _a_scale_raw, _a_scale_pad, a_scale_sh = _get_a4w4_quant_buffers(
m, k, a.device
)
grid = (
triton.cdiv(m, cfg["BLOCK_SIZE_M"]),
triton.cdiv(k, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
)
_dynamic_mxfp4_quant_kernel_shuffled[grid](
a,
a_q_raw,
a_scale_sh,
*a.stride(),
*a_q_raw.stride(),
*a_scale_sh.stride(),
M=m,
N=k,
SCALE_N=k // 32,
BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
NUM_ITER=cfg["NUM_ITER"],
NUM_STAGES=cfg["NUM_STAGES"],
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=cfg["num_warps"],
waves_per_eu=cfg["waves_per_eu"],
num_stages=1,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)
def _quant_a4w4_cached_rawshuf_cfg(a: torch.Tensor, cfg: dict):
from aiter import dtypes
from aiter.ops.triton.quant.quant import _dynamic_mxfp4_quant_kernel
m, k = a.shape
a_q_raw, a_scale_raw, _a_scale_pad, a_scale_sh = _get_a4w4_quant_buffers(
m, k, a.device
)
grid = (
triton.cdiv(m, cfg["BLOCK_SIZE_M"]),
triton.cdiv(k, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
)
_dynamic_mxfp4_quant_kernel[grid](
a,
a_q_raw,
a_scale_raw,
*a.stride(),
*a_q_raw.stride(),
*a_scale_raw.stride(),
M=m,
N=k,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=cfg["NUM_ITER"],
BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
NUM_STAGES=cfg["NUM_STAGES"],
num_warps=cfg["num_warps"],
waves_per_eu=cfg["waves_per_eu"],
num_stages=1,
)
grid_shuf = (triton.cdiv(m, 32), triton.cdiv(k // 32, 8))
_shuffle_e8m0_kernel[grid_shuf](
a_scale_raw,
a_scale_sh,
*a_scale_raw.stride(),
*a_scale_sh.stride(),
M=m,
SCALE_N=k // 32,
BLOCK_SIZE_M=32,
BLOCK_SIZE_N=8,
num_warps=4,
num_stages=1,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)
def _quant_a4w4_cached_asm_cfg(a: torch.Tensor, cfg: dict):
from aiter import dtypes
m, k = a.shape
a_q_raw, a_scale_sh = _get_a4w4_asm_quant_buffers(m, k, a.device)
grid = (triton.cdiv(m, cfg["BLOCK_SIZE"]), k // 32)
_dynamic_mxfp4_quant_kernel_asm_layout_cfg[grid](
a,
a_q_raw,
a_scale_sh,
*a.stride(),
*a_q_raw.stride(),
M=m,
N=k,
SCALE_N_VALID=k // 32,
SCALE_M_PAD=a_scale_sh.shape[0],
SCALE_N_PAD=a_scale_sh.shape[1],
BLOCK_SIZE=cfg["BLOCK_SIZE"],
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=cfg["num_warps"],
num_stages=1,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)
def _benchmark_ms(fn, warmup: int = 2, iters: int = 5) -> float:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
for _ in range(warmup):
fn()
torch.cuda.synchronize()
samples: list[float] = []
for _ in range(iters):
start.record()
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
samples.sort()
return samples[len(samples) // 2]
def _quant_a4w4_cached_shuffled_cfg(a: torch.Tensor, cfg_name: str):
from aiter import dtypes
cfg = _A4W4_QUANT_CFGS[cfg_name]
m, k = a.shape
a_q_raw, _a_scale_raw, _a_scale_pad, a_scale_sh = _get_a4w4_quant_buffers(
m, k, a.device
)
grid = (
triton.cdiv(m, cfg["BLOCK_SIZE_M"]),
triton.cdiv(k, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
)
_dynamic_mxfp4_quant_kernel_shuffled[grid](
a,
a_q_raw,
a_scale_sh,
*a.stride(),
*a_q_raw.stride(),
*a_scale_sh.stride(),
M=m,
N=k,
SCALE_N=k // 32,
BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
NUM_ITER=cfg["NUM_ITER"],
NUM_STAGES=cfg["NUM_STAGES"],
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=cfg["num_warps"],
waves_per_eu=cfg["waves_per_eu"],
num_stages=1,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)
def _get_out(m: int, n: int, device: torch.device) -> torch.Tensor:
key = (device.index or 0, m, n)
out = _OUT_CACHE.get(key)
if out is None:
out = torch.empty((m, n), dtype=torch.bfloat16, device=device)
_OUT_CACHE[key] = out
return out
def _get_out_pad32(m: int, n: int, device: torch.device) -> torch.Tensor:
m_pad = triton.cdiv(m, 32) * 32
key = (device.index or 0, m_pad, n)
out = _OUT_PAD32_CACHE.get(key)
if out is None:
out = torch.empty((m_pad, n), dtype=torch.bfloat16, device=device)
_OUT_PAD32_CACHE[key] = out
return out
def _get_partial(num_ksplit: int, m: int, n: int, device: torch.device) -> torch.Tensor:
key = (device.index or 0, num_ksplit, m, n)
partial = _PARTIAL_CACHE.get(key)
if partial is None:
partial = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)
_PARTIAL_CACHE[key] = partial
return partial
def _run_preshuffle(
a: torch.Tensor,
w: torch.Tensor,
w_scales: torch.Tensor,
config: dict,
) -> torch.Tensor:
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
m, k = a.shape
n, k_packed = w.shape
n *= 16
k = k_packed // 16
config = _normalize_config(config, k)
out = _get_out(m, n, a.device)
if config["NUM_KSPLIT"] > 1:
partial = _get_partial(config["NUM_KSPLIT"], m, n, a.device)
else:
partial = None
grid = lambda meta: ( # noqa: E731
meta["NUM_KSPLIT"]
* triton.cdiv(m, meta["BLOCK_SIZE_M"])
* triton.cdiv(n, meta["BLOCK_SIZE_N"]),
)
_gemm_a16wfp4_preshuffle_kernel[grid](
a,
w,
out if partial is None else partial,
w_scales,
m,
n,
k,
a.stride(0),
a.stride(1),
w.stride(0),
w.stride(1),
0 if partial is None else partial.stride(0),
out.stride(0) if partial is None else partial.stride(1),
out.stride(1) if partial is None else partial.stride(2),
w_scales.stride(0),
w_scales.stride(1),
PREQUANT=True,
**config,
)
if partial is None:
return out
actual_ksplit = triton.cdiv(k, config["SPLITK_BLOCK_SIZE"] // 2)
grid_reduce = (triton.cdiv(m, 16), triton.cdiv(n, 64))
_gemm_afp4wfp4_reduce_kernel[grid_reduce](
partial,
out,
m,
n,
partial.stride(0),
partial.stride(1),
partial.stride(2),
out.stride(0),
out.stride(1),
16,
64,
actual_ksplit,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return out
def _run_default_preshuffle(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
out = _get_out(a.shape[0], b_shuffle.shape[0], a.device)
w = b_shuffle.view(torch.uint8).contiguous().view(b_shuffle.shape[0] // 16, -1)
w_scales = b_scale_sh.view(torch.uint8).contiguous().view(
b_scale_sh.shape[0] // 32, -1
)
return gemm_a16wfp4_preshuffle(a, w, w_scales, True, torch.bfloat16, out)
def _run_direct_a16wfp4(
a: torch.Tensor,
b_q: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
m, k = a.shape
n = b_q.shape[0]
out = _get_out(m, n, a.device)
return gemm_a16wfp4(
a,
_get_b_q_u8(b_q),
_get_unshuffled_b_scale(b_q, b_scale_sh, k),
False,
torch.bfloat16,
out,
)
def _run_direct_a4w4(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4
m, k = a.shape
a = a.contiguous()
if m % 32 == 0 and k % 512 == 0:
a_q, a_scale = _quant_a4w4_cached_fast_m32_k512(a)
else:
a_q, a_scale = _quant_a4w4_cached(a)
return gemm_a4w4(
a_q.view(m, k // 2),
b_shuffle,
a_scale,
b_scale_sh,
dtype=torch.bfloat16,
bpreshuffle=True,
)
def _run_direct_a4w4_large_cfg(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
cfg: dict,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4
m, k = a.shape
a_q, a_scale = _quant_a4w4_cached_large_cfg(a.contiguous(), cfg)
return gemm_a4w4(
a_q.view(m, k // 2),
b_shuffle,
a_scale,
b_scale_sh,
dtype=torch.bfloat16,
bpreshuffle=True,
)
def _run_direct_a4w4_rawshuf_cfg(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
cfg: dict,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4
m, k = a.shape
a_q, a_scale = _quant_a4w4_cached_rawshuf_cfg(a.contiguous(), cfg)
return gemm_a4w4(
a_q.view(m, k // 2),
b_shuffle,
a_scale,
b_scale_sh,
dtype=torch.bfloat16,
bpreshuffle=True,
)
def _run_direct_a4w4_asm_cfg(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
cfg: dict,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4
m, k = a.shape
a_q, a_scale = _quant_a4w4_cached_asm_cfg(a.contiguous(), cfg)
return gemm_a4w4(
a_q.view(m, k // 2),
b_shuffle,
a_scale,
b_scale_sh,
dtype=torch.bfloat16,
bpreshuffle=True,
)
def _run_direct_a4w4_cfg(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
cfg_name: str,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4
m, k = a.shape
a_q, a_scale = _quant_a4w4_cached_shuffled_cfg(a.contiguous(), cfg_name)
return gemm_a4w4(
a_q.view(m, k // 2),
b_shuffle,
a_scale,
b_scale_sh,
dtype=torch.bfloat16,
bpreshuffle=True,
)
def _measure_route_us(fn, warmup: int = 1, reps: int = 3) -> float:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
samples: list[float] = []
for _ in range(reps):
start.record()
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000.0)
samples.sort()
return samples[len(samples) // 2]
def _pick_large_route(
a: torch.Tensor,
b_q: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> tuple[str, str | None]:
m, k = a.shape
n = b_shuffle.shape[0]
shape_key = (a.device.index or 0, m, n, k)
cached = _LARGE_ROUTE_CACHE.get(shape_key)
if cached is not None:
return cached
candidates: list[tuple[str, str | None, callable]] = []
if (n, k) == (7168, 2048):
candidates.append(
("a16_direct", None, lambda: _run_direct_a16wfp4(a, b_q, b_scale_sh))
)
for cfg_name in _A4W4_7168_CANDIDATES:
candidates.append(
(
"a4w4_direct",
cfg_name,
lambda cfg_name=cfg_name: _run_direct_a4w4_cfg(
a, b_shuffle, b_scale_sh, cfg_name
),
)
)
elif (n, k) == (3072, 1536):
for cfg_name in _A4W4_3072_CANDIDATES:
candidates.append(
(
"a4w4_direct",
cfg_name,
lambda cfg_name=cfg_name: _run_direct_a4w4_cfg(
a, b_shuffle, b_scale_sh, cfg_name
),
)
)
else:
raise AssertionError(f"unexpected large shape {(m, n, k)}")
best_kind, best_cfg, best_score = "", None, float("inf")
for kind, cfg_name, fn in candidates:
score = _measure_route_us(fn)
if score < best_score:
best_kind, best_cfg, best_score = kind, cfg_name, score
cached = (best_kind, best_cfg)
print(
f"[route] shape={(m, n, k)} kind={best_kind} cfg={best_cfg} median_us={best_score:.3f}"
)
_LARGE_ROUTE_CACHE[shape_key] = cached
return cached
def _run_direct_a4w4_uncached(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4
a_q, a_scale = _quant_a4w4(a.contiguous())
return gemm_a4w4(
a_q,
b_shuffle,
a_scale,
b_scale_sh,
dtype=torch.bfloat16,
bpreshuffle=True,
)
def _run_direct_a4w4_asm(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
kernel_name: str,
splitk: int,
) -> torch.Tensor:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
m, k = a.shape
n = b_shuffle.shape[0]
a = a.contiguous()
if m % 32 == 0 and k % 512 == 0:
a_q, a_scale = _quant_a4w4_cached_fast_m32_k512(a)
else:
a_q, a_scale = _quant_a4w4_cached(a)
out = _get_out_pad32(m, n, a.device)
gemm_a4w4_asm(
a_q.view(m, k // 2),
b_shuffle,
a_scale,
b_scale_sh,
out,
kernel_name,
None,
1.0,
0.0,
True,
log2_k_split=splitk,
)
return out[:m]
def _get_large_q_candidates(m: int, k: int) -> list[dict]:
if (m, k) == (64, 2048):
return _QCFG_LARGE_64_2048
if (m, k) == (256, 1536):
return _QCFG_LARGE_256_1536
return [_QCFG_CURRENT]
def _get_large_asm_q_candidates(m: int, k: int) -> list[dict]:
return []
def _select_large_route(
a: torch.Tensor,
b_q: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> tuple[str, dict | None]:
m, k = a.shape
n = b_shuffle.shape[0]
key = (a.device.index or 0, m, n, k)
cached = _LARGE_ROUTE_CACHE.get(key)
if cached is not None:
return cached
candidates: list[tuple[str, dict | None]] = []
if (n, k) == (7168, 2048):
candidates.append(("a16", None))
for cfg in _get_large_q_candidates(m, k):
candidates.append(("a4w4_shuf", cfg))
for kind, cfg in candidates:
if kind == "a16":
_run_direct_a16wfp4(a, b_q, b_scale_sh)
else:
_run_direct_a4w4_large_cfg(a, b_shuffle, b_scale_sh, cfg)
torch.cuda.synchronize()
best = candidates[0]
best_ms = float("inf")
scores: list[tuple[str, str, float]] = []
for kind, cfg in candidates:
if kind == "a16":
ms = _benchmark_ms(lambda: _run_direct_a16wfp4(a, b_q, b_scale_sh))
else:
ms = _benchmark_ms(
lambda cfg=cfg: _run_direct_a4w4_large_cfg(
a, b_shuffle, b_scale_sh, cfg
)
)
label = "a16" if cfg is None else (
f"bm{cfg['BLOCK_SIZE_M']}_bn{cfg['BLOCK_SIZE_N']}_i{cfg['NUM_ITER']}_w{cfg['num_warps']}"
)
scores.append((kind, label, ms))
if ms < best_ms:
best_ms = ms
best = (kind, cfg)
print(f"route {(m, n, k)} -> {scores} -> best {best_ms:.3f} ms", flush=True)
_LARGE_ROUTE_CACHE[key] = best
return best
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 (n, k) == (2880, 512) and m <= 4:
w, w_scales = _get_views(b_shuffle, b_scale_sh)
return _run_preshuffle(a, w, w_scales, _CFG_2880_512_M_LEQ_4)
if (n, k) == (2880, 512) and m <= 8:
w, w_scales = _get_views(b_shuffle, b_scale_sh)
return _run_preshuffle(a, w, w_scales, _CFG_2880_512_M_LEQ_8)
if (n, k) == (2880, 512) and m >= 32:
return _run_default_preshuffle(a, b_shuffle, b_scale_sh)
if (n, k) == (2112, 7168) and m <= 16:
w, w_scales = _get_views(b_shuffle, b_scale_sh)
return _run_preshuffle(a, w, w_scales, _CFG_2112_7168_M16_N64_S2)
if (n, k) == (4096, 512):
return _run_default_preshuffle(a, b_shuffle, b_scale_sh)
if (n, k) in {(7168, 2048), (3072, 1536)}:
route_kind, route_cfg = _select_large_route(a, b_q, b_shuffle, b_scale_sh)
if route_kind == "a16":
return _run_direct_a16wfp4(a, b_q, b_scale_sh)
return _run_direct_a4w4_large_cfg(a, b_shuffle, b_scale_sh, route_cfg)
return _run_default_preshuffle(a, b_shuffle, b_scale_sh)
scrolls · 1891 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