Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
10.7µs
#282 of 1143
2026-03-14

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 = 1num_warps = 1
split-kf.write("cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n")
stages = 1num_stages = 1
tile-m = 32BLOCK_SIZE_M=32,
tile-n = 128BLOCK_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