Skip to content
KernelIndex
Search⌘K

submission 690209

Tecahens · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submissionv1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-690209?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
#283 of 1143
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:65198b9188f95bb5b53d220032f41ce08de0bb1ce6736e70c24b9f724c68fb05
license declaredunknown
license concludedunknown
authorsTecahens
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM: A16 preshuffle with reused config strings and workspace buffers.
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

Kernel source

submissionv1.py277 lines
"""
MXFP4 GEMM: A16 preshuffle with reused config strings and workspace buffers.
"""
import torch
import triton

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,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle_
from aiter.ops.triton.utils.common_utils import serialize_dict

_DIRECT_CFG = {
    (4, 2880, 512): {
        "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": 1,
    },
    (32, 4096, 512): {
        "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": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (32, 2880, 512): {
        "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": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 14,
    },
    (64, 7168, 2048): {
        "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": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (256, 3072, 1536): {
        "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": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
}
_A16_CONFIGS = {
    (7168, 2048): {
        "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": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (3072, 1536): {
        "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": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
}

_DIRECT_STATE = {}
_A16_WORKSPACE = {}
_DIRECT_REDUCE_CFG = {
    (16, 2112, 7168): (16, 16),
}


def _workspace_key(tensor, n, k, config_override):
    cfg_key = tuple(sorted(config_override.items())) if config_override else None
    return (str(tensor.device), tensor.shape[0], n, k, cfg_key)


def _get_a16_workspace(a, n, k, config_override):
    key = _workspace_key(a, n, k, config_override)
    workspace = _A16_WORKSPACE.get(key)
    if workspace is not None:
        return workspace

    workspace = {
        "config": serialize_dict(config_override) if config_override is not None else None,
        "out": torch.empty((a.shape[0], n), dtype=torch.bfloat16, device=a.device),
    }
    _A16_WORKSPACE[key] = workspace
    return workspace


def _prepare_direct_entry(m, n, k, config):
    logical_k = k >> 1
    direct_config = dict(config)
    if direct_config["NUM_KSPLIT"] > 1:
        splitk_block_size, block_size_k, num_ksplit = get_splitk(
            logical_k, direct_config["BLOCK_SIZE_K"], direct_config["NUM_KSPLIT"]
        )
        direct_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
        direct_config["BLOCK_SIZE_K"] = block_size_k
        direct_config["NUM_KSPLIT"] = num_ksplit
    if direct_config["BLOCK_SIZE_K"] >= 2 * logical_k:
        direct_config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * logical_k)
        direct_config["NUM_KSPLIT"] = 1
    direct_config["BLOCK_SIZE_N"] = max(direct_config["BLOCK_SIZE_N"], 32)
    if direct_config["BLOCK_SIZE_N"] % 32 != 0 or direct_config["BLOCK_SIZE_K"] % 256 != 0:
        return None
    direct_config["SPLITK_BLOCK_SIZE"] = (
        2 * logical_k
        if direct_config["NUM_KSPLIT"] == 1
        else direct_config["SPLITK_BLOCK_SIZE"]
    )
    grid = (
        direct_config["NUM_KSPLIT"]
        * triton.cdiv(m, direct_config["BLOCK_SIZE_M"])
        * triton.cdiv(n, direct_config["BLOCK_SIZE_N"]),
    )
    return {
        "config": direct_config,
        "grid": (grid,),
    }


def _get_direct_state(a, n, k):
    key = (str(a.device), a.shape[0], n, k)
    state = _DIRECT_STATE.get(key)
    if state is not None:
        return state

    config = _DIRECT_CFG.get((a.shape[0], n, k))
    if config is None:
        return None

    state = _prepare_direct_entry(a.shape[0], n, k, config)
    if state is None:
        return None
    state["out"] = torch.empty((a.shape[0], n), dtype=torch.bfloat16, device=a.device)
    if state["config"]["NUM_KSPLIT"] > 1:
        state["y_pp"] = torch.empty(
            (state["config"]["NUM_KSPLIT"], a.shape[0], n),
            dtype=torch.float32,
            device=a.device,
        )
    _DIRECT_STATE[key] = state
    return state


def _run_direct_shape(a, w, w_scales, n, k):
    state = _get_direct_state(a, n, k)
    if state is None:
        return None

    logical_k = k >> 1
    out = state.get("y_pp", state["out"])

    _gemm_a16wfp4_preshuffle_kernel[state["grid"][0]](
        a,
        w,
        out,
        w_scales,
        a.shape[0],
        n,
        logical_k,
        a.stride(0),
        a.stride(1),
        w.stride(0),
        w.stride(1),
        0 if "y_pp" not in state else state["y_pp"].stride(0),
        state["out"].stride(0) if "y_pp" not in state else state["y_pp"].stride(1),
        state["out"].stride(1) if "y_pp" not in state else state["y_pp"].stride(2),
        w_scales.stride(0),
        w_scales.stride(1),
        PREQUANT=True,
        **state["config"],
    )

    if "y_pp" in state:
        reduce_m, reduce_n = _DIRECT_REDUCE_CFG.get((a.shape[0], n, k), (16, 64))
        actual_ksplit = triton.cdiv(
            logical_k, state["config"]["SPLITK_BLOCK_SIZE"] >> 1
        )
        _gemm_afp4wfp4_reduce_kernel[
            (triton.cdiv(a.shape[0], reduce_m), triton.cdiv(n, reduce_n))
        ](
            state["y_pp"],
            state["out"],
            a.shape[0],
            n,
            state["y_pp"].stride(0),
            state["y_pp"].stride(1),
            state["y_pp"].stride(2),
            state["out"].stride(0),
            state["out"].stride(1),
            reduce_m,
            reduce_n,
            actual_ksplit,
            triton.next_power_of_2(state["config"]["NUM_KSPLIT"]),
        )

    return state["out"]
def custom_kernel(data):
    a, _, _, b_shuffle, b_scale_sh = data
    k = a.shape[1]
    n = b_shuffle.shape[0]

    w = b_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
    w_scales = b_scale_sh.view(torch.uint8).reshape(b_scale_sh.shape[0] // 32, k)

    out = _run_direct_shape(a, w, w_scales, n, k)
    if out is not None:
        return out

    workspace = _get_a16_workspace(a, n, k, _A16_CONFIGS.get((n, k)))

    return gemm_a16wfp4_preshuffle_(
        a,
        w,
        w_scales,
        True,
        torch.bfloat16,
        workspace["out"],
        workspace["config"],
        False,
    )
scrolls · 277 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