Skip to content
KernelIndex
Search⌘K

submission 752215

nataliakokoromyti · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1113 lines, June 9 Researcher Reciprocity License v1.0.

aiter_patched_quant_direct_dual_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-752215?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
11.9µs
#357 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:98e9e14cb3da19901b8915c8ee93a91b9259f4d3be1efd891ac30365da653cee
license declaredunknown
license concludedunknown
authorsnataliakokoromyti
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 1num_warps = 1
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
stages = 1num_stages=1,
tile-m = 32BLOCK_M=32,
tile-n = 8BLOCK_N=8,

Kernel source

aiter_patched_quant_direct_dual_submission.py1113 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from task import input_t, output_t
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.quant.quant import _dynamic_mxfp4_quant_kernel
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid


_ASM_OUT_BUFFER_CACHE = {}
_RAW_OUT_BUFFER_CACHE = {}
_DIRECT_QUANT_CACHE = {}

_ASM_KERNEL_OVERRIDES = {
    (64, 7168, 2048): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (256, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
}


def _as_u8_tensor(x):
    import torch

    if x.dtype == torch.uint8:
        return x.contiguous()
    return x.view(torch.uint8).contiguous()


def _get_cached_asm_out(torch_mod, dtypes, device, m: int, n: int):
    key = (device.type, device.index, m, n)
    cached = _ASM_OUT_BUFFER_CACHE.get(key)
    if cached is None:
        padded_m = ((m + 31) // 32) * 32
        cached = torch_mod.empty((padded_m, n), dtype=dtypes.bf16, device=device)
        _ASM_OUT_BUFFER_CACHE[key] = cached
    return cached


def _get_cached_raw_out(torch_mod, device, m: int, n: int):
    key = (device.type, device.index, m, n)
    cached = _RAW_OUT_BUFFER_CACHE.get(key)
    if cached is None:
        cached = torch_mod.empty((m, n), dtype=torch_mod.bfloat16, device=device)
        _RAW_OUT_BUFFER_CACHE[key] = cached
    return cached


def _get_cached_direct_quant(torch_mod, device, m: int, k: int):
    scale_n = (k + 31) // 32
    scale_n_pad = ((scale_n + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256
    key = (device.type, device.index, m, k)
    cached = _DIRECT_QUANT_CACHE.get(key)
    if cached is None:
        a_q = torch_mod.empty((m, k // 2), dtype=torch_mod.uint8, device=device)
        a_scale_raw = torch_mod.empty((m, scale_n), dtype=torch_mod.uint8, device=device)
        a_scale_sh = torch_mod.empty(
            (scale_m_pad, scale_n_pad), dtype=torch_mod.uint8, device=device
        )
        cached = (a_q, a_scale_raw, a_scale_sh)
        _DIRECT_QUANT_CACHE[key] = cached
    return cached


def _get_cached_direct_quant_exact(torch_mod, device):
    key = (device.type, device.index, 256, 1536, "exact")
    cached = _DIRECT_QUANT_CACHE.get(key)
    if cached is None:
        a_q = torch_mod.empty((256, 768), dtype=torch_mod.uint8, device=device)
        a_scale_sh = torch_mod.empty((256, 48), dtype=torch_mod.uint8, device=device)
        cached = (a_q, a_scale_sh)
        _DIRECT_QUANT_CACHE[key] = cached
    return cached


def _get_cached_direct_quant_exact_64x2048(torch_mod, device):
    key = (device.type, device.index, 64, 2048, "exact")
    cached = _DIRECT_QUANT_CACHE.get(key)
    if cached is None:
        a_q = torch_mod.empty((64, 1024), dtype=torch_mod.uint8, device=device)
        a_scale_sh = torch_mod.empty((256, 64), dtype=torch_mod.uint8, device=device)
        cached = (a_q, a_scale_sh)
        _DIRECT_QUANT_CACHE[key] = cached
    return cached


def _pick_config(m: int, n: int, k: int) -> dict | None:
    # AITER ships only one specialized gfx950 A16WFP4_PRESHUFFLED config for the
    # public leaderboard shapes. Everything else falls back to the generic family.
    if (n, k) == (2112, 7168):
        if m <= 8:
            return {
                "BLOCK_SIZE_M": 8,
                "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,
            }
        if m <= 16:
            return {
                "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,
            }
        if m <= 64:
            return {
                "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,
            }
        if m <= 256:
            return {
                "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": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 14,
            }

    # Borrow known-good preshuffled FP4 configs for public shapes that do not
    # have an explicit A16WFP4 specialization in AITER.
    if (n, k) == (4096, 512):
        if m <= 31:
            return {
                "BLOCK_SIZE_M": 8,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 2,
                "num_stages": 1,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 1,
            }
        if m <= 64:
            return {
                "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,
            }
        if m <= 256:
            return {
                "BLOCK_SIZE_M": 256,
                "BLOCK_SIZE_N": 256,
                "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,
            }
        return {
            "BLOCK_SIZE_M": 64,
            "BLOCK_SIZE_N": 256,
            "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,
        }

    if (n, k) == (3072, 1536):
        if m <= 31:
            return {
                "BLOCK_SIZE_M": 8,
                "BLOCK_SIZE_N": 32,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 1,
            }
        if m <= 64:
            return {
                "BLOCK_SIZE_M": 64,
                "BLOCK_SIZE_N": 32,
                "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": 1,
            }
        if m <= 256:
            return {
                "BLOCK_SIZE_M": 128,
                "BLOCK_SIZE_N": 32,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 1,
            }

    return None


def _default_config() -> dict:
    return {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 8,
        "num_stages": 1,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    }


def _pick_raw_config(m: int, n: int, k: int) -> dict:
    if (n, k) == (2112, 7168):
        if m <= 8:
            return {
                "BLOCK_SIZE_M": 8,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 16,
            }
        if m <= 16:
            return {
                "BLOCK_SIZE_M": 16,
                "BLOCK_SIZE_N": 32,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 6,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 8,
            }
        if m <= 32:
            return {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 1024,
                "GROUP_SIZE_M": 1,
                "num_warps": 8,
                "num_stages": 2,
                "waves_per_eu": 8,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 8,
            }
        if m <= 64:
            return {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 32,
                "BLOCK_SIZE_K": 1024,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 1,
            }
        if m <= 128:
            return {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 1024,
                "GROUP_SIZE_M": 1,
                "num_warps": 8,
                "num_stages": 2,
                "waves_per_eu": 6,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 1,
            }
        if m <= 256:
            return {
                "BLOCK_SIZE_M": 128,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 256,
                "GROUP_SIZE_M": 2,
                "num_warps": 4,
                "num_stages": 3,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            }
        return {
            "BLOCK_SIZE_M": 64,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 1024,
            "GROUP_SIZE_M": 1,
            "num_warps": 8,
            "num_stages": 2,
            "waves_per_eu": 4,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        }

    if (n, k) == (3072, 1536):
        if m <= 16:
            return {
                "BLOCK_SIZE_M": 16,
                "BLOCK_SIZE_N": 16,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 3,
                "waves_per_eu": 6,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            }
        if m <= 32:
            return {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 16,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 3,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            }
        if m <= 128:
            return {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 32,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 2,
                "num_warps": 4,
                "num_stages": 3,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            }
        if m <= 256:
            return {
                "BLOCK_SIZE_M": 64,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 2,
                "num_warps": 4,
                "num_stages": 3,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            }
        if m <= 2048:
            return {
                "BLOCK_SIZE_M": 64,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 256,
                "GROUP_SIZE_M": 2,
                "num_warps": 4,
                "num_stages": 3,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": None,
                "NUM_KSPLIT": 1,
            }
        return {
            "BLOCK_SIZE_M": 128,
            "BLOCK_SIZE_N": 256,
            "BLOCK_SIZE_K": 128,
            "GROUP_SIZE_M": 16,
            "num_warps": 4,
            "num_stages": 2,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 32,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        }

    if (n, k) == (7168, 2048):
        if m <= 8:
            return {
                "BLOCK_SIZE_M": 8,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 8,
                "num_stages": 2,
                "waves_per_eu": 1,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 4,
            }
        if m <= 16:
            return {
                "BLOCK_SIZE_M": 16,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 4,
            }
        if m <= 32:
            return {
                "BLOCK_SIZE_M": 16,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 8,
                "num_stages": 2,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 4,
            }
        if m <= 64:
            return {
                "BLOCK_SIZE_M": 16,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 8,
                "num_stages": 2,
                "waves_per_eu": 4,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            }
        if m <= 128:
            return {
                "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,
            }
        return {
            "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,
        }

    # Generic gfx950-GEMM-A16WFP4 config.
    if m <= 16:
        return {
            "BLOCK_SIZE_M": 4,
            "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": ".cg",
            "NUM_KSPLIT": 1,
        }
    return {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 8,
        "num_stages": 1,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg" if m <= 256 else None,
        "NUM_KSPLIT": 1,
    }


def _e8m0_unshuffle(scale_sh, *, rows: int, cols: int):
    restored = scale_sh.view(scale_sh.shape[0] // 32, scale_sh.shape[1] // 8, 4, 16, 2, 2)
    restored = restored.permute(0, 5, 3, 1, 4, 2).contiguous()
    restored = restored.view(scale_sh.shape[0], scale_sh.shape[1])
    return restored[:rows, :cols].contiguous()


@triton.jit
def _e8m0_shuffle_filled_kernel(
    raw_ptr,
    sh_ptr,
    raw_stride_m,
    raw_stride_n,
    sh_stride_m,
    sh_stride_n,
    M,
    M_PAD,
    N_VALID,
    N_PAD,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    linear = offs_m[:, None] * N_PAD + offs_n[None, :]

    out_k = linear & 1
    linear = linear // 2
    out_j = linear & 1
    linear = linear // 2
    out_i = linear % 16
    linear = linear // 16
    out_f = linear % 4
    linear = linear // 4
    out_c = linear % (N_PAD // 8)

    src_m = (offs_m[:, None] // 32) * 32 + out_k * 16 + out_i
    src_n = out_c * 8 + out_j * 4 + out_f

    src_mask = (src_m < M) & (src_n < N_VALID)
    raw = tl.load(
        raw_ptr + src_m * raw_stride_m + src_n * raw_stride_n,
        mask=src_mask,
        other=127,
    )

    dst_mask = (offs_m[:, None] < M_PAD) & (offs_n[None, :] < N_PAD)
    tl.store(
        sh_ptr + offs_m[:, None] * sh_stride_m + offs_n[None, :] * sh_stride_n,
        raw,
        mask=dst_mask,
    )


@triton.jit
def _patched_quant_asm_exact_256x1536_kernel(
    x_ptr,
    x_fp4_ptr,
    bs_sh_ptr,
    stride_x_m,
    stride_x_n,
    stride_x_fp4_m,
    stride_x_fp4_n,
    BLOCK_M: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * 32 + tl.arange(0, 32)
    x_offs = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
    x = tl.load(x_ptr + x_offs).to(tl.float32)

    out_tensor, bs_e8m0 = _mxfp4_quant_op(x, 32, BLOCK_M, 32)

    out_offs_n = pid_n * 16 + tl.arange(0, 16)
    out_offs = offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
    tl.store(x_fp4_ptr + out_offs, out_tensor)

    # Equivalent to fp4_utils.e8m0_shuffle for an exact [256, 48] scale tensor.
    row = offs_m
    col = pid_n
    block_row = row // 32
    rem = row % 32
    b = rem // 16
    c = rem % 16
    d = col // 8
    remc = col % 8
    e = remc // 4
    f = remc % 4
    linear = (((((block_row * 6 + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b)
    tl.store(bs_sh_ptr + linear, bs_e8m0.reshape(BLOCK_M))


@triton.jit
def _patched_quant_asm_exact_64x2048_kernel(
    x_ptr,
    x_fp4_ptr,
    bs_sh_ptr,
    stride_x_m,
    stride_x_n,
    stride_x_fp4_m,
    stride_x_fp4_n,
    BLOCK_M: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * 32 + tl.arange(0, 32)
    x_mask = offs_m[:, None] < 64
    x_offs = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
    x = tl.load(x_ptr + x_offs, mask=x_mask, other=0).to(tl.float32)

    out_tensor, bs_e8m0 = _mxfp4_quant_op(x, 32, BLOCK_M, 32)
    bs_e8m0 = tl.where((offs_m < 64)[:, None], bs_e8m0, 127)

    out_offs_n = pid_n * 16 + tl.arange(0, 16)
    out_offs = offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
    out_mask = offs_m[:, None] < 64
    tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

    row = offs_m
    col = pid_n
    block_row = row // 32
    rem = row % 32
    b = rem // 16
    c = rem % 16
    d = col // 8
    remc = col % 8
    e = remc // 4
    f = remc % 4
    linear = (((((block_row * 8 + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b)
    tl.store(bs_sh_ptr + linear, bs_e8m0.reshape(BLOCK_M))


def _quant_mxfp4_direct(torch_mod, x, a_q, a_scale_raw, a_scale_sh):
    m, k = x.shape
    scale_n = (k + 31) // 32
    scale_n_pad = ((scale_n + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256

    if m <= 32:
        num_iter = 1
        block_size_m = triton.next_power_of_2(m)
        block_size_n = 32
        num_warps = 1
        num_stages_quant = 1
    else:
        num_iter = 4
        block_size_m = 64
        block_size_n = 64
        num_warps = 4
        num_stages_quant = 2
        if k <= 16384:
            block_size_m = 32
            block_size_n = 128

    if k <= 1024:
        num_iter = 1
        num_warps = 4
        num_stages_quant = 1
        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](
        x,
        a_q,
        a_scale_raw,
        *x.stride(),
        *a_q.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_quant,
        num_warps=num_warps,
        waves_per_eu=0,
        num_stages=1,
    )

    grid_shuffle = (
        triton.cdiv(a_scale_sh.shape[0], 32),
        triton.cdiv(a_scale_sh.shape[1], 8),
    )
    _e8m0_shuffle_filled_kernel[grid_shuffle](
        a_scale_raw,
        a_scale_sh,
        a_scale_raw.stride(0),
        a_scale_raw.stride(1),
        a_scale_sh.stride(0),
        a_scale_sh.stride(1),
        m,
        scale_m_pad,
        scale_n,
        scale_n_pad,
        BLOCK_M=32,
        BLOCK_N=8,
    )


def _quant_mxfp4_direct_exact_256x1536(x, a_q, a_scale_sh):
    grid = (2, 48)
    _patched_quant_asm_exact_256x1536_kernel[grid](
        x,
        a_q,
        a_scale_sh,
        x.stride(0),
        x.stride(1),
        a_q.stride(0),
        a_q.stride(1),
        BLOCK_M=128,
        num_warps=4,
        num_stages=1,
    )


def _quant_mxfp4_direct_exact_64x2048(x, a_q, a_scale_sh):
    grid = (1, 64)
    _patched_quant_asm_exact_64x2048_kernel[grid](
        x,
        a_q,
        a_scale_sh,
        x.stride(0),
        x.stride(1),
        a_q.stride(0),
        a_q.stride(1),
        BLOCK_M=64,
        num_warps=4,
        num_stages=2,
    )


@triton.heuristics(
    {
        "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
        * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
    }
)
@triton.jit
def _fixed_a16wfp4_preshuffle_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    b_scales_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_ck,
    stride_cm,
    stride_cn,
    stride_bsn,
    stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr,
    PREQUANT: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid_unified = tl.program_id(axis=0)
    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    SCALE_GROUP_SIZE: tl.constexpr = 32

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)

        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + (
            offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
        )

        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + (
            offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
        )

        offs_bsn = (
            pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
        ) % N
        offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
            0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
        )
        b_scale_ptrs = (
            b_scales_ptr
            + offs_bsn[:, None] * stride_bsn
            + offs_ks[None, :] * stride_bsk
        )

        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for _ in range(0, num_k_iter):
            b_scales = (
                tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .reshape(
                    BLOCK_SIZE_N // 32,
                    BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                    4,
                    16,
                    2,
                    2,
                    1,
                )
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            )

            a_bf16 = tl.load(a_ptrs)
            b_raw = tl.load(b_ptrs, cache_modifier=cache_modifier)
            b = (
                b_raw.reshape(
                    1,
                    BLOCK_SIZE_N // 16,
                    BLOCK_SIZE_K // 64,
                    2,
                    16,
                    16,
                )
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0)
            )

            if PREQUANT:
                a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)

            accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
            b_scale_ptrs += BLOCK_SIZE_K * stride_bsk

        c = accumulator.to(c_ptr.type.element_ty)
        offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        c_ptrs = (
            c_ptr
            + stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        tl.store(c_ptrs, c, mask=c_mask)


def _run_fixed_a16wfp4_preshuffle(x, w, w_scales, *, dtype, config):
    import torch

    M, _ = x.shape
    n_outer, k_outer = w.shape
    N = n_outer * 16
    K = k_outer // 16

    cfg = dict(config)
    if cfg["NUM_KSPLIT"] > 1:
        splitk_block_size, block_size_k, num_ksplit = get_splitk(
            K, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
        )
        cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
        cfg["BLOCK_SIZE_K"] = block_size_k
        cfg["NUM_KSPLIT"] = num_ksplit

    if cfg["BLOCK_SIZE_K"] >= 2 * K:
        cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K
        cfg["NUM_KSPLIT"] = 1
    else:
        cfg.setdefault("SPLITK_BLOCK_SIZE", 2 * K)

    cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)

    if cfg["NUM_KSPLIT"] > 1:
        y_pp = torch.empty((cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device)
        y = torch.empty((M, N), dtype=dtype, device=x.device)
    else:
        y_pp = None
        y = torch.empty((M, N), dtype=dtype, device=x.device)

    grid = lambda META: (  # noqa: E731
        (
            META["NUM_KSPLIT"]
            * triton.cdiv(M, META["BLOCK_SIZE_M"])
            * triton.cdiv(N, META["BLOCK_SIZE_N"])
        ),
    )
    _fixed_a16wfp4_preshuffle_kernel[grid](
        x,
        w,
        y if y_pp is None else y_pp,
        w_scales,
        M,
        N,
        K,
        x.stride(0),
        x.stride(1),
        w.stride(0),
        w.stride(1),
        0 if y_pp is None else y_pp.stride(0),
        y.stride(0) if y_pp is None else y_pp.stride(1),
        y.stride(1) if y_pp is None else y_pp.stride(2),
        w_scales.stride(0),
        w_scales.stride(1),
        PREQUANT=True,
        **cfg,
    )

    if cfg["NUM_KSPLIT"] > 1:
        reduce_block_size_m = 16
        reduce_block_size_n = 64
        actual_ksplit = triton.cdiv(K, (cfg["SPLITK_BLOCK_SIZE"] // 2))
        grid_reduce = (
            triton.cdiv(M, reduce_block_size_m),
            triton.cdiv(N, reduce_block_size_n),
        )
        _gemm_afp4wfp4_reduce_kernel[grid_reduce](
            y_pp,
            y,
            M,
            N,
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            y.stride(0),
            y.stride(1),
            reduce_block_size_m,
            reduce_block_size_n,
            actual_ksplit,
            triton.next_power_of_2(cfg["NUM_KSPLIT"]),
        )

    return y


def custom_kernel(data: input_t) -> output_t:
    import os

    os.environ.setdefault("AITER_LOG_LEVEL", "ERROR")
    import torch
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    def _quant_mxfp4(x):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

    a, _b, b_q, b_shuffle, b_scale_sh = data
    a = a.contiguous()
    b_shuffle = _as_u8_tensor(b_shuffle)
    b_scale_sh = _as_u8_tensor(b_scale_sh)

    m, k_real = a.shape
    n = b_shuffle.shape[0]
    packed_k = k_real // 2
    override = _ASM_KERNEL_OVERRIDES.get((m, n, k_real))

    if override is not None:
        if (m, n, k_real) == (64, 7168, 2048):
            a_q_u8, a_scale_sh_u8 = _get_cached_direct_quant_exact_64x2048(
                torch, a.device
            )
            _quant_mxfp4_direct_exact_64x2048(a, a_q_u8, a_scale_sh_u8)
            a_q = a_q_u8.view(dtypes.fp4x2)
            a_scale_sh = a_scale_sh_u8.view(dtypes.fp8_e8m0)
        elif (m, n, k_real) == (256, 3072, 1536):
            a_q_u8, a_scale_sh_u8 = _get_cached_direct_quant_exact(torch, a.device)
            _quant_mxfp4_direct_exact_256x1536(a, a_q_u8, a_scale_sh_u8)
            a_q = a_q_u8.view(dtypes.fp4x2)
            a_scale_sh = a_scale_sh_u8.view(dtypes.fp8_e8m0)
        else:
            a_q, a_scale_sh = _quant_mxfp4(a)
        kernel_name, split_k = override
        out = _get_cached_asm_out(torch, aiter.dtypes, a.device, m, n)
        aiter.gemm_a4w4_asm(
            a_q.view(m, packed_k),
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            out,
            kernel_name,
            None,
            1.0,
            0.0,
            True,
            split_k,
        )
        return out[:m]

    # The long/wide 7168x2048 family still prefers the older fused-raw path on MI355X.
    if (n, k_real) == (7168, 2048):
        b_q_u8 = _as_u8_tensor(b_q)
        b_scale_raw = _e8m0_unshuffle(b_scale_sh, rows=b_q_u8.shape[0], cols=k_real // 32)
        out = _get_cached_raw_out(torch, a.device, m, b_q_u8.shape[0])
        return gemm_a16wfp4(
            a,
            b_q_u8,
            b_scale_raw,
            atomic_add=False,
            dtype=torch.bfloat16,
            y=out,
            config=_pick_raw_config(m, b_q_u8.shape[0], k_real),
        )

    # GPU MODE's task tensors use the logical task shapes:
    #   B_shuffle   : (N, K/2)
    #   B_scale_sh  : (*, K/32)
    # AITER's preshuffled kernel expects the exact same bytes reinterpreted as:
    #   B           : (N/16, (K/2) * 16)
    #   B_scales    : (* / 32, K)
    # The byte-level round-trip matches the task's public shuffle operators.
    b_phys = b_shuffle.contiguous().view(b_shuffle.numel() // (packed_k * 16), packed_k * 16)
    b_scale_phys = b_scale_sh.contiguous().view(b_scale_sh.numel() // k_real, k_real)
    config = _pick_config(m, n, k_real) or _default_config()

    return _run_fixed_a16wfp4_preshuffle(
        a,
        b_phys,
        b_scale_phys,
        dtype=torch.bfloat16,
        config=config,
    )
scrolls · 1113 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