Skip to content
KernelIndex
Search⌘K

submission 735371

zx199303 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v14.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-735371?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
184.7µs
#559 of 782
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:deac9cb7880fd129a41a068188c66891e84e874720ea26112610acf1fc093662
license declaredunknown
license concludedunknown
authorszx199303
imported2026-08-26

Kernel source

submission_v14.py333 lines
from task import input_t, output_t


_FUSED_MOE_PARAM_CACHE = None
_OUTPUT_CACHE = {}
_PAD_CACHE = {}
_INJECTED_FMOE_CONFIG = False


# Shape-aware policy inspired by mxfp4-mm exact-shape dispatch style:
# - EP-off (d_expert=256): preserve small-batch latency
# - EP-on  (d_expert=2048): push contiguous earlier
_SMALL_BS = 32
_EP_ON_CONTIGUOUS_THRESHOLD = 64
_EP_OFF_CONTIGUOUS_THRESHOLD = 128


def _get_fused_moe_params(fused_moe_fn):
    global _FUSED_MOE_PARAM_CACHE
    if _FUSED_MOE_PARAM_CACHE is None:
        import inspect

        _FUSED_MOE_PARAM_CACHE = set(inspect.signature(fused_moe_fn).parameters.keys())
    return _FUSED_MOE_PARAM_CACHE


def _get_output_buffer(m: int, h: int, device, dtype):
    key = (int(m), int(h), str(device), str(dtype))
    out = _OUTPUT_CACHE.get(key)
    if out is None:
        import torch

        out = torch.empty((m, h), device=device, dtype=dtype)
        _OUTPUT_CACHE[key] = out
    return out


def _get_pads(config):
    key = (
        int(config["d_hidden"]),
        int(config["d_hidden_pad"]),
        int(config["d_expert"]),
        int(config["d_expert_pad"]),
    )
    cached = _PAD_CACHE.get(key)
    if cached is None:
        cached = (
            config["d_hidden_pad"] - config["d_hidden"],
            config["d_expert_pad"] - config["d_expert"],
        )
        _PAD_CACHE[key] = cached
    return cached


def _maybe_contiguous(bs: int, d_expert: int, tensors):
    # keep very small batches untouched
    if bs <= _SMALL_BS:
        return tensors

    threshold = (
        _EP_ON_CONTIGUOUS_THRESHOLD if d_expert >= 2048 else _EP_OFF_CONTIGUOUS_THRESHOLD
    )
    if bs < threshold:
        return tensors

    converted = []
    for x in tensors:
        converted.append(x if x.is_contiguous() else x.contiguous())
    return tuple(converted)


def _ensure_topk_types(topk_weights, topk_ids):
    import torch

    if topk_weights.dtype is not torch.float32:
        topk_weights = topk_weights.float()
    if topk_ids.dtype is not torch.int32:
        topk_ids = topk_ids.int()
    return topk_weights, topk_ids


def _parse_int(value):
    try:
        return int(float(value))
    except Exception:
        return None


def _match_header(headers, candidates, allow_substring=True):
    lowered = {h.lower(): h for h in headers}
    for cand in candidates:
        if cand in lowered:
            return lowered[cand]
    if allow_substring:
        for h in headers:
            lh = h.lower()
            for cand in candidates:
                if cand in lh:
                    return h
    return None


def _load_csv_rows(path):
    import csv

    with open(path, "r", newline="") as f:
        reader = csv.reader(f)
        rows = list(reader)

    if not rows:
        return [], []

    header = rows[0]
    data = [row for row in rows[1:] if len(row) == len(header)]
    return header, data


def _write_csv_rows(path, header, rows):
    import csv
    import os

    tmp_path = f"{path}.tmp"
    with open(tmp_path, "w", newline="") as f:
        writer = csv.writer(f)
        writer.writerow(header)
        writer.writerows(rows)
    os.replace(tmp_path, path)


def _maybe_inject_fmoe_configs(config):
    global _INJECTED_FMOE_CONFIG
    if _INJECTED_FMOE_CONFIG:
        return
    _INJECTED_FMOE_CONFIG = True

    import os

    tuned_path = "/home/runner/aiter/aiter/configs/tuned_fmoe.csv"
    if not os.path.exists(tuned_path):
        return

    try:
        header, rows = _load_csv_rows(tuned_path)
    except Exception:
        return

    if not header:
        return

    # Sanitize malformed rows first to avoid pandas parser errors.
    raw_line_count = 0
    try:
        with open(tuned_path, "r", newline="") as f:
            raw_line_count = sum(1 for _ in f)
    except Exception:
        raw_line_count = 0

    expected_line_count = 1 + len(rows)
    if raw_line_count and raw_line_count != expected_line_count:
        try:
            _write_csv_rows(tuned_path, header, rows)
        except Exception:
            return

    col_m = _match_header(header, ["m", "bs", "batch", "tokens", "m_size"])
    if col_m is None:
        return

    col_n = _match_header(header, ["d_hidden", "hidden", "n", "hidden_size", "n_size"])
    col_k = _match_header(header, ["d_expert", "intermediate", "k", "k_size"])
    col_e = _match_header(header, ["num_experts", "experts", "n_experts", "e"])
    col_topk = _match_header(header, ["topk", "top_k", "n_experts_per_token", "k_top"])

    idx_m = header.index(col_m)
    idx_n = header.index(col_n) if col_n in header else None
    idx_k = header.index(col_k) if col_k in header else None
    idx_e = header.index(col_e) if col_e in header else None
    idx_topk = header.index(col_topk) if col_topk in header else None

    targets = [
        {"bs": 16, "d_expert": 256, "d_hidden": 7168, "experts": 257, "topk": 9},
        {"bs": 128, "d_expert": 256, "d_hidden": 7168, "experts": 257, "topk": 9},
        {"bs": 512, "d_expert": 256, "d_hidden": 7168, "experts": 257, "topk": 9},
        {"bs": 16, "d_expert": 512, "d_hidden": 7168, "experts": 33, "topk": 9},
        {"bs": 128, "d_expert": 512, "d_hidden": 7168, "experts": 33, "topk": 9},
        {"bs": 512, "d_expert": 512, "d_hidden": 7168, "experts": 33, "topk": 9},
        {"bs": 512, "d_expert": 2048, "d_hidden": 7168, "experts": 33, "topk": 9},
    ]

    existing_keys = set()
    for row in rows:
        key = (
            _parse_int(row[idx_m]) if idx_m is not None else None,
            _parse_int(row[idx_n]) if idx_n is not None else None,
            _parse_int(row[idx_k]) if idx_k is not None else None,
            _parse_int(row[idx_e]) if idx_e is not None else None,
            _parse_int(row[idx_topk]) if idx_topk is not None else None,
        )
        existing_keys.add(key)

    new_rows = []
    for target in targets:
        key = (
            target["bs"] if idx_m is not None else None,
            target["d_hidden"] if idx_n is not None else None,
            target["d_expert"] if idx_k is not None else None,
            target["experts"] if idx_e is not None else None,
            target["topk"] if idx_topk is not None else None,
        )
        if key in existing_keys:
            continue

        best_row = None
        best_score = -1
        for row in rows:
            score = 0
            if idx_n is not None and _parse_int(row[idx_n]) == target["d_hidden"]:
                score += 5
            if idx_k is not None and _parse_int(row[idx_k]) == target["d_expert"]:
                score += 5
            if idx_e is not None and _parse_int(row[idx_e]) == target["experts"]:
                score += 3
            if idx_topk is not None and _parse_int(row[idx_topk]) == target["topk"]:
                score += 2
            if score > best_score:
                best_score = score
                best_row = row

        if best_row is None:
            continue

        new_row = list(best_row)
        new_row[idx_m] = str(target["bs"])
        if idx_n is not None:
            new_row[idx_n] = str(target["d_hidden"])
        if idx_k is not None:
            new_row[idx_k] = str(target["d_expert"])
        if idx_e is not None:
            new_row[idx_e] = str(target["experts"])
        if idx_topk is not None:
            new_row[idx_topk] = str(target["topk"])
        new_rows.append(new_row)

    if new_rows:
        try:
            rows.extend(new_rows)
            _write_csv_rows(tuned_path, header, rows)
        except Exception:
            return


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

    (
        hidden_states,
        gate_up_weight,
        down_weight,
        gate_up_weight_scale,
        down_weight_scale,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data

    _maybe_inject_fmoe_configs(config)

    from aiter import ActivationType, QuantType
    from aiter.fused_moe import fused_moe

    with torch.inference_mode():
        bs = int(hidden_states.shape[0])
        d_expert = int(config["d_expert"])
        (
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
            topk_weights,
            topk_ids,
        ) = _maybe_contiguous(
            bs,
            d_expert,
            (
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
            ),
        )

        topk_weights, topk_ids = _ensure_topk_types(topk_weights, topk_ids)
        hidden_pad, intermediate_pad = _get_pads(config)

        kwargs = {
            "expert_mask": None,
            "activation": ActivationType.Silu,
            "quant_type": QuantType.per_1x32,
            "doweight_stage1": False,
            "w1_scale": gate_up_weight_scale_shuffled,
            "w2_scale": down_weight_scale_shuffled,
            "a1_scale": None,
            "a2_scale": None,
            "hidden_pad": hidden_pad,
            "intermediate_pad": intermediate_pad,
        }

        params = _get_fused_moe_params(fused_moe)
        if "out" in params:
            kwargs["out"] = _get_output_buffer(
                hidden_states.shape[0],
                config["d_hidden"],
                hidden_states.device,
                hidden_states.dtype,
            )

        return fused_moe(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            **kwargs,
        )
scrolls · 333 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