Skip to content
KernelIndex
Search⌘K

submission 670884

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

moe_v64_copy.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-670884?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
120.5µs
#39 of 782
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c7ea6b09062b85751c8791dd7f2fcb6481c070b328021554c50849805b79290c
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Techniques

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

fp4"stage":2,"a_dtype":"fp4","b_dtype":"fp4","out_dtype":"bf16",

Kernel source

moe_v64_copy.py335 lines
# HIP MoE v64 -- Direct dispatch + NT load fix
#
# v64 bypasses fused_moe's Python dispatch chain (fused_moe → fused_moe_ →
# fused_moe_2stages, 3+ nested functions, enum conversions, conditional logic)
# and calls the GPU kernels directly. Saves ~8-12µs of Python overhead per call.
#
# Also fixes TWO bugs in NT load patching that existed since v19:
#   1. Wrong key construction: a[:13] doesn't include cu_num, so _CC.get(k)
#      never matched. Fixed: reconstruct keys matching cfg_2stages format.
#   2. Wrong keyword name: 'non_temporal_load' vs 'use_non_temporal_load'.
#      Fixed: use correct keyword name.
#
# Added NT loads for bs=128 E=33 (heuristic: tokens_per_expert=35 < 64).
#
# Flow:
#   1st call per shape → fused_moe (populates metadata cache, verified correct)
#   2nd+ calls → direct dispatch (sorting → quant → stage1 → re-quant → stage2)
#   Any error → falls back to fused_moe permanently for that shape

from task import input_t, output_t
import torch
import os
import sys
import functools

from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
from aiter.ops.moe_sorting import moe_sorting_fwd
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
import aiter.fused_moe as _fm

# =========================================================================== #
# FlyDSL tile registration (needed for stage2 kernel names)
# =========================================================================== #
try:
    import aiter.ops.flydsl.moe_kernels as _flydsl
    for tm, tn in [(32,128),(32,256),(16,256),(16,128),(64,128),(64,256)]:
        _flydsl._KERNEL_PARAMS[f"flydsl_moe2_afp4_wfp4_bf16_t{tm}x{tn}x128_atomic"] = {
            "stage":2,"a_dtype":"fp4","b_dtype":"fp4","out_dtype":"bf16",
            "tile_m":tm,"tile_n":tn,"tile_k":128,"mode":"atomic","MPerBlock":tm,
        }
except ImportError:
    pass

# =========================================================================== #
# Per-shape configs (injected into AITER's cfg_2stages)
# =========================================================================== #
def _key(t, i, e):
    return (256, t, 7168, i, e, 9, "ActivationType.Silu", "torch.bfloat16",
            "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
            "QuantType.per_1x32", True, False)

_4WG128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_4WG64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_F16 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"

_CC = {}
# E=33, TP=4 shapes
_CC[_key(16, 512, 33)] = {
    "block_m": 16, "ksplit": 2,
    "kernelName1": "", "kernelName2": "",
    "run_1stage": False, "use_non_temporal_load": True,
}
_CC[_key(128, 512, 33)] = {
    "block_m": 32, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False, "use_non_temporal_load": True,  # tokens/expert=35<64
}
_CC[_key(512, 512, 33)] = {
    "block_m": 64, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False,
}
_CC[_key(512, 2048, 33)] = {
    "block_m": 64, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False,
}
# E=257, TP=8 shapes
_CC[_key(16, 256, 257)] = {
    "block_m": 16, "ksplit": 2,
    "kernelName1": "", "kernelName2": "",
    "run_1stage": False, "use_non_temporal_load": True,
}
_CC[_key(128, 256, 257)] = {
    "block_m": 16, "ksplit": 2,
    "kernelName1": "", "kernelName2": "",
    "run_1stage": False, "use_non_temporal_load": True,
}
_CC[_key(512, 256, 257)] = {
    "block_m": 32, "ksplit": 0,
    "kernelName1": _4WG64, "kernelName2": _F16,
    "run_1stage": False, "use_non_temporal_load": True,  # tokens/expert=18<64
}

# =========================================================================== #
# Injection + metadata capture with FIXED NT load patching
# =========================================================================== #
_injected = False
_captured_meta = {}  # (padded_M, model_dim, inter_dim, E, topk) → MOEMetadata
_cu = None


def _inject():
    global _injected
    if _injected:
        return
    _injected = True

    if _fm.cfg_2stages is None:
        _fm.cfg_2stages = {}
    for k, v in _CC.items():
        _fm.cfg_2stages[k] = v

    orig = _fm.get_2stage_cfgs

    @functools.lru_cache(maxsize=2048)
    def _p(*a):
        global _cu
        m = orig(*a)

        # FIX #1: Reconstruct keys matching cfg_2stages format
        # get_2stage_cfgs builds keys = (cu_num, token, model_dim, inter_dim,
        #   expert, topk, str(activation), str(dtype), str(q_dtype_a),
        #   str(q_dtype_w), str(q_type), use_g1u1, doweight_stage1)
        # But function args a = (token[0], model_dim[1], inter_dim[2],
        #   expert[3], topk[4], dtype[5], q_dtype_a[6], q_dtype_w[7],
        #   q_type[8], use_g1u1[9], activation[10], doweight_stage1[11], ...)
        if _cu is None:
            try:
                from aiter.jit.utils.chip_info import get_cu_num
                _cu = get_cu_num()
            except Exception:
                _cu = 256

        keys = (_cu, a[0], a[1], a[2], a[3], a[4],
                str(a[10]), str(a[5]), str(a[6]), str(a[7]),
                str(a[8]), a[9], a[11])

        c = _CC.get(keys)
        if c and c.get("use_non_temporal_load"):
            s = m.stage1
            kw = getattr(s, 'keywords', None) or {}
            # FIX #2: correct keyword is 'use_non_temporal_load', not 'non_temporal_load'
            if hasattr(s, 'func') and 'use_non_temporal_load' in kw:
                nk = dict(kw)
                nk['use_non_temporal_load'] = True
                m = _fm.MOEMetadata(
                    functools.partial(s.func, **nk),
                    m.stage2, m.block_m, m.ksplit,
                    m.run_1stage, m.has_bias, True)

        # Cache metadata for direct dispatch
        _captured_meta[(a[0], a[1], a[2], a[3], a[4])] = m
        return m

    _fm.get_2stage_cfgs = _p


# =========================================================================== #
# Per-shape cache for direct dispatch
# =========================================================================== #
_shape_cache = {}  # shape_key → cache dict or None (fallback)


def _init_shape(data, shape_key):
    """First call: run fused_moe to populate metadata, build dispatch cache."""
    hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
    guws_sh = data[7]; dws_sh = data[8]
    tw = data[9]; ti = data[10]; cfg = data[11]

    M = cfg['bs']
    E = cfg['n_routed_experts'] + cfg['n_shared_experts']
    topk = ti.shape[1]
    hp = cfg['d_hidden_pad'] - cfg['d_hidden']
    ip = cfg['d_expert_pad'] - cfg['d_expert']
    device = hs.device

    # Run fused_moe once to populate metadata cache and return correct output
    result = fused_moe(
        hs, guw_sh, dw_sh, tw, ti,
        expert_mask=None, activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32, doweight_stage1=False,
        w1_scale=guws_sh, w2_scale=dws_sh,
        a1_scale=None, a2_scale=None,
        hidden_pad=hp, intermediate_pad=ip,
    )

    # Find captured metadata
    _, model_dim, inter_dim = _fm.get_inter_dim(guw_sh.shape, dw_sh.shape)
    padded_M = _fm.get_padded_M(M)
    meta_key = (padded_M, model_dim, inter_dim, E, topk)

    metadata = _captured_meta.get(meta_key)
    if metadata is None:
        print(f"[V64] No metadata for {shape_key}, using fallback",
              file=sys.stderr, flush=True)
        _shape_cache[shape_key] = None
        return result

    block_m = int(metadata.block_m)
    max_tok = M * topk + E * block_m - topk
    max_blk = (max_tok + block_m - 1) // block_m

    _shape_cache[shape_key] = {
        'meta': metadata,
        'block_m': block_m,
        'topk': topk,
        'E': E,
        'M': M,
        'model_dim': model_dim,
        'inter_dim': inter_dim,
        'hp': hp,
        'ip': ip,
        # Pre-allocated sorting buffers (reused across calls)
        'sid': torch.empty(max_tok, dtype=torch.int32, device=device),
        'sw': torch.empty(max_tok, dtype=torch.float32, device=device),
        'seid': torch.empty(max_blk, dtype=torch.int32, device=device),
        'nvi': torch.empty(2, dtype=torch.int32, device=device),
    }

    # Log metadata details
    s1 = metadata.stage1
    s1k = getattr(s1, 'keywords', {}) if hasattr(s1, 'func') else {}
    nt_val = s1k.get('use_non_temporal_load', 'N/A')
    s1name = s1.func.__name__ if hasattr(s1, 'func') else str(s1)[:40]
    print(f"[V64] Init {shape_key}: block_m={block_m} inter={inter_dim} "
          f"model={model_dim} nt={nt_val} stage1={s1name}",
          file=sys.stderr, flush=True)

    return result


def _direct_dispatch(data, c):
    """Direct dispatch: 5 GPU kernel calls with minimal Python overhead."""
    hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
    guws_sh = data[7]; dws_sh = data[8]
    tw = data[9]; ti = data[10]

    meta = c['meta']
    block_m = c['block_m']
    topk = c['topk']
    E = c['E']
    M = c['M']
    model_dim = c['model_dim']
    inter_dim = c['inter_dim']

    sid = c['sid']; sw = c['sw']; seid = c['seid']; nvi = c['nvi']
    device = hs.device

    # 1. Sorting (pre-allocated output buffers, fresh moe_buf for atomicAdd)
    moe_buf = torch.empty(M, model_dim, dtype=torch.bfloat16, device=device)
    moe_sorting_fwd(ti, tw, sid, sw, seid, nvi, moe_buf,
                    E, block_m, None, None, 0)

    # 2. Activation quantization (fused quant + scale sorting, Triton kernel)
    a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
        hs, sorted_ids=sid, num_valid_ids=nvi,
        token_num=M, topk=1, block_size=block_m)

    # 3. Stage 1: gate_up GEMM + SwiGLU (CK or CKtile kernel)
    a2 = torch.empty(M, topk, inter_dim, dtype=torch.bfloat16, device=device)
    a2 = meta.stage1(
        a1, guw_sh, dw_sh, sid, seid, nvi, a2, topk,
        block_m=block_m,
        a1_scale=a1_scale,
        w1_scale=guws_sh.view(dtypes.fp8_e8m0),
        sorted_weights=None)

    # 4. Intermediate re-quantization (bf16 → fp4x2, Triton kernel)
    a2_flat = a2.view(-1, inter_dim)
    a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        a2_flat, sorted_ids=sid, num_valid_ids=nvi,
        token_num=M, topk=topk, block_size=block_m)
    a2_q = a2_q.view(M, topk, -1)

    # 5. Stage 2: down GEMM + weighted scatter-reduce (CK/FlyDSL kernel)
    meta.stage2(
        a2_q, guw_sh, dw_sh, sid, seid, nvi, moe_buf, topk,
        w2_scale=dws_sh.view(dtypes.fp8_e8m0),
        a2_scale=a2_scale,
        block_m=block_m,
        sorted_weights=sw)

    return moe_buf


def _fallback(data):
    """Safe fallback: full fused_moe dispatch."""
    hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
    guws_sh = data[7]; dws_sh = data[8]
    tw = data[9]; ti = data[10]; cfg = data[11]
    hp = cfg['d_hidden_pad'] - cfg['d_hidden']
    ip = cfg['d_expert_pad'] - cfg['d_expert']
    return fused_moe(
        hs, guw_sh, dw_sh, tw, ti,
        expert_mask=None, activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32, doweight_stage1=False,
        w1_scale=guws_sh, w2_scale=dws_sh,
        a1_scale=None, a2_scale=None,
        hidden_pad=hp, intermediate_pad=ip,
    )


# =========================================================================== #
# Main entry point
# =========================================================================== #
def custom_kernel(data: input_t) -> output_t:
    _inject()

    cfg = data[11]
    shape_key = (cfg['bs'], cfg['d_expert'],
                 cfg['n_routed_experts'] + cfg['n_shared_experts'])

    # First call per shape: use fused_moe to populate metadata cache
    if shape_key not in _shape_cache:
        return _init_shape(data, shape_key)

    c = _shape_cache[shape_key]

    # Fallback if metadata capture failed
    if c is None:
        return _fallback(data)

    # Direct dispatch (bypasses ~8-12µs Python overhead)
    try:
        return _direct_dispatch(data, c)
    except Exception as e:
        print(f"[V64] Dispatch err {shape_key}: {str(e)[:200]}",
              file=sys.stderr, flush=True)
        import traceback
        traceback.print_exc(file=sys.stderr)
        _shape_cache[shape_key] = None
        return _fallback(data)
scrolls · 335 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 643890.

- # MoE v36 — Try 1WG_M32 stage1 for E=33 bs=16 (sparse)
+ # HIP MoE v64 -- Direct dispatch + NT load fix
#
- # v30 uses default CKTile for E=33 bs=16 (ksplit=2).
- # v34 tries 1WG_M16 for E=257. v36 tries 1WG_M32 for E=33 bs=16.
+ # v64 bypasses fused_moe's Python dispatch chain (fused_moe → fused_moe_ →
+ # fused_moe_2stages, 3+ nested functions, enum conversions, conditional logic)
+ # and calls the GPU kernels directly. Saves ~8-12µs of Python overhead per call.
#
- # E=33 bs=16: 16 tokens / ~9 experts = ~4.4 tokens per expert
- # With block_m=32, most tiles have 1-2 valid tokens (wasteful).
- # A 1WG stage1 kernel with M32 has lower launch overhead than 4WG,
- # which could help for this very sparse case.
+ # Also fixes TWO bugs in NT load patching that existed since v19:
+ # 1. Wrong key construction: a[:13] doesn't include cu_num, so _CC.get(k)
+ # never matched. Fixed: reconstruct keys matching cfg_2stages format.
+ # 2. Wrong keyword name: 'non_temporal_load' vs 'use_non_temporal_load'.
+ # Fixed: use correct keyword name.
#
- # Also try: block_m=16 for E=33 bs=16 (instead of 32)
- # With 4.4 tokens/expert, block_m=16 wastes less than block_m=32.
+ # Added NT loads for bs=128 E=33 (heuristic: tokens_per_expert=35 < 64).
#
- # Test: popcorn submit --gpu MI355X --leaderboard moe-mxfp4 --mode test moe_v36.py
- # Benchmark: popcorn submit --gpu MI355X --leaderboard moe-mxfp4 --mode benchmark moe_v36.py
+ # Flow:
+ # 1st call per shape → fused_moe (populates metadata cache, verified correct)
+ # 2nd+ calls → direct dispatch (sorting → quant → stage1 → re-quant → stage2)
+ # Any error → falls back to fused_moe permanently for that shape
+ from task import input_t, output_t
+ import torch
import os
import sys
import functools
- import torch
- from task import input_t, output_t
- from aiter import ActivationType, QuantType
+
+ from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
- import aiter.fused_moe as _fused_moe_module
+ from aiter.ops.moe_sorting import moe_sorting_fwd
+ from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
+ import aiter.fused_moe as _fm
- # ── FlyDSL registration (identical to v19) ──
+ # =========================================================================== #
+ # FlyDSL tile registration (needed for stage2 kernel names)
+ # =========================================================================== #
try:
- import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
- for tile_m, tile_n in [(32, 128), (32, 256), (16, 256), (16, 128), (64, 128), (64, 256)]:
- name = f"flydsl_moe2_afp4_wfp4_bf16_t{tile_m}x{tile_n}x128_atomic"
- _flydsl_moe_kernels._KERNEL_PARAMS[name] = {
- "stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16",
- "tile_m": tile_m, "tile_n": tile_n, "tile_k": 128,
- "mode": "atomic", "MPerBlock": tile_m,
+ import aiter.ops.flydsl.moe_kernels as _flydsl
+ for tm, tn in [(32,128),(32,256),(16,256),(16,128),(64,128),(64,256)]:
+ _flydsl._KERNEL_PARAMS[f"flydsl_moe2_afp4_wfp4_bf16_t{tm}x{tn}x128_atomic"] = {
+ "stage":2,"a_dtype":"fp4","b_dtype":"fp4","out_dtype":"bf16",
+ "tile_m":tm,"tile_n":tn,"tile_k":128,"mode":"atomic","MPerBlock":tm,
}
except ImportError:
pass
- def _key(token, inter_dim, expert):
- return (
- 256, token, 7168, inter_dim, expert, 9,
- "ActivationType.Silu", "torch.bfloat16",
- "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
- "QuantType.per_1x32", True, False,
- )
+ # =========================================================================== #
+ # Per-shape configs (injected into AITER's cfg_2stages)
+ # =========================================================================== #
+ def _key(t, i, e):
+ return (256, t, 7168, i, e, 9, "ActivationType.Silu", "torch.bfloat16",
+ "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
+ "QuantType.per_1x32", True, False)
- # ── Kernel names ──
- _4WG_M128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _4WG_M64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _4WG_M32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _1WG_M32 = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _1WG_M16 = "moe_ck2stages_gemm1_256x16x128x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _FLY_16x128 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
+ _4WG128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ _4WG64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
+ _F16 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
- _CUSTOM_CONFIGS = {}
-
- # ═══════════════════════════════════════════════════════════════
- # E=33 shapes
- # ═══════════════════════════════════════════════════════════════
-
- # bs=16: CHANGED — block_m=16 + ksplit=2 (smaller tiles for 4.4 tokens/expert)
- _CUSTOM_CONFIGS[_key(16, 512, 33)] = {
- "block_m": 16, "ksplit": 2, # CHANGED: block_m 32→16
+ _CC = {}
+ # E=33, TP=4 shapes
+ _CC[_key(16, 512, 33)] = {
+ "block_m": 16, "ksplit": 2,
"kernelName1": "", "kernelName2": "",
- "run_1stage": False,
- "use_non_temporal_load": True, # sparse — bypass cache
+ "run_1stage": False, "use_non_temporal_load": True,
}
-
- # bs=128: CHANGED from v19 — block_m=32 instead of 64 (v29: 91.4→89.6, 2% win)
- _CUSTOM_CONFIGS[_key(128, 512, 33)] = {
- "block_m": 32, "ksplit": 0, # CHANGED: block_m 64→32
- "kernelName1": _4WG_M128, "kernelName2": _FLY_16x128,
- "run_1stage": False,
+ _CC[_key(128, 512, 33)] = {
+ "block_m": 32, "ksplit": 0,
+ "kernelName1": _4WG128, "kernelName2": _F16,
+ "run_1stage": False, "use_non_temporal_load": True, # tokens/expert=35<64
}
-
- # bs=512 d=512: EXACT v19
- _CUSTOM_CONFIGS[_key(512, 512, 33)] = {
+ _CC[_key(512, 512, 33)] = {
"block_m": 64, "ksplit": 0,
- "kernelName1": _4WG_M128, "kernelName2": _FLY_16x128,
+ "kernelName1": _4WG128, "kernelName2": _F16,
"run_1stage": False,
}
-
- # bs=512 d=2048: EXACT v19 (FLY_64x128 was 24% worse in v29)
- _CUSTOM_CONFIGS[_key(512, 2048, 33)] = {
+ _CC[_key(512, 2048, 33)] = {
"block_m": 64, "ksplit": 0,
- "kernelName1": _4WG_M128, "kernelName2": _FLY_16x128,
+ "kernelName1": _4WG128, "kernelName2": _F16,
"run_1stage": False,
}
-
- # ═══════════════════════════════════════════════════════════════
- # E=257 shapes
- # ═══════════════════════════════════════════════════════════════
-
- # bs=16: EXACT v19
- _CUSTOM_CONFIGS[_key(16, 256, 257)] = {
+ # E=257, TP=8 shapes
+ _CC[_key(16, 256, 257)] = {
"block_m": 16, "ksplit": 2,
"kernelName1": "", "kernelName2": "",
- "run_1stage": False,
- "use_non_temporal_load": True,
+ "run_1stage": False, "use_non_temporal_load": True,
}
-
- # bs=128: EXACT v19
- _CUSTOM_CONFIGS[_key(128, 256, 257)] = {
+ _CC[_key(128, 256, 257)] = {
"block_m": 16, "ksplit": 2,
"kernelName1": "", "kernelName2": "",
- "run_1stage": False,
- "use_non_temporal_load": True,
+ "run_1stage": False, "use_non_temporal_load": True,
}
-
- # bs=512: CHANGED from v19 — 4WG_M64 instead of 4WG_M32 (v29: 208→180, 13% win!)
- _CUSTOM_CONFIGS[_key(512, 256, 257)] = {
+ _CC[_key(512, 256, 257)] = {
"block_m": 32, "ksplit": 0,
- "kernelName1": _4WG_M64, # CHANGED: M32→M64
- "kernelName2": _FLY_16x128,
- "run_1stage": False,
- "use_non_temporal_load": True,
+ "kernelName1": _4WG64, "kernelName2": _F16,
+ "run_1stage": False, "use_non_temporal_load": True, # tokens/expert=18<64
}
- # ── Config injection (identical to v19) ──
+ # =========================================================================== #
+ # Injection + metadata capture with FIXED NT load patching
+ # =========================================================================== #
_injected = False
- def _inject_configs():
+ _captured_meta = {} # (padded_M, model_dim, inter_dim, E, topk) → MOEMetadata
+ _cu = None
+
+
+ def _inject():
global _injected
if _injected:
return
_injected = True
- if _fused_moe_module.cfg_2stages is None:
- import pandas as pd
- from aiter.jit.core import AITER_CONFIGS
- tune_file = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
- if os.path.exists(tune_file):
- _INDEX_COLS = ["cu_num", "token", "model_dim", "inter_dim", "expert", "topk",
- "act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type", "use_g1u1", "doweight_stage1"]
- df = pd.read_csv(tune_file)
- if "_tag" in df.columns:
- df = df[df["_tag"].fillna("") == ""]
- _fused_moe_module.cfg_2stages = df.set_index(_INDEX_COLS).to_dict("index")
- else:
- _fused_moe_module.cfg_2stages = {}
- _fused_moe_module.cfg_2stages.update(_CUSTOM_CONFIGS)
- _original = _fused_moe_module.get_2stage_cfgs
+ if _fm.cfg_2stages is None:
+ _fm.cfg_2stages = {}
+ for k, v in _CC.items():
+ _fm.cfg_2stages[k] = v
+
+ orig = _fm.get_2stage_cfgs
+
@functools.lru_cache(maxsize=2048)
- def _patched(token, model_dim, inter_dim, expert, topk,
- dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
- activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True):
- metadata = _original(token, model_dim, inter_dim, expert, topk,
- dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
- activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled)
- from aiter.jit.utils.chip_info import get_cu_num
- cu_num = get_cu_num()
- keys = (cu_num, token, model_dim, inter_dim, expert, topk,
- str(activation), str(dtype), str(q_dtype_a), str(q_dtype_w),
- str(q_type), use_g1u1, doweight_stage1)
- cfg = _fused_moe_module.cfg_2stages.get(keys)
- if cfg and cfg.get("use_non_temporal_load") is not None:
- nt = cfg["use_non_temporal_load"]
- old_s1 = metadata.stage1
- kw = getattr(old_s1, 'keywords', None) or {}
- if hasattr(old_s1, 'func') and 'non_temporal_load' in kw:
- new_kw = dict(kw)
- new_kw['non_temporal_load'] = nt
- metadata = _fused_moe_module.MOEMetadata(
- functools.partial(old_s1.func, **new_kw),
- metadata.stage2, metadata.block_m, metadata.ksplit,
- metadata.run_1stage, metadata.has_bias, nt)
- return metadata
- _fused_moe_module.get_2stage_cfgs = _patched
+ def _p(*a):
+ global _cu
+ m = orig(*a)
+ # FIX #1: Reconstruct keys matching cfg_2stages format
+ # get_2stage_cfgs builds keys = (cu_num, token, model_dim, inter_dim,
+ # expert, topk, str(activation), str(dtype), str(q_dtype_a),
+ # str(q_dtype_w), str(q_type), use_g1u1, doweight_stage1)
+ # But function args a = (token[0], model_dim[1], inter_dim[2],
+ # expert[3], topk[4], dtype[5], q_dtype_a[6], q_dtype_w[7],
+ # q_type[8], use_g1u1[9], activation[10], doweight_stage1[11], ...)
+ if _cu is None:
+ try:
+ from aiter.jit.utils.chip_info import get_cu_num
+ _cu = get_cu_num()
+ except Exception:
+ _cu = 256
- def custom_kernel(data: input_t) -> output_t:
- (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
+ keys = (_cu, a[0], a[1], a[2], a[3], a[4],
+ str(a[10]), str(a[5]), str(a[6]), str(a[7]),
+ str(a[8]), a[9], a[11])
- _inject_configs()
+ c = _CC.get(keys)
+ if c and c.get("use_non_temporal_load"):
+ s = m.stage1
+ kw = getattr(s, 'keywords', None) or {}
+ # FIX #2: correct keyword is 'use_non_temporal_load', not 'non_temporal_load'
+ if hasattr(s, 'func') and 'use_non_temporal_load' in kw:
+ nk = dict(kw)
+ nk['use_non_temporal_load'] = True
+ m = _fm.MOEMetadata(
+ functools.partial(s.func, **nk),
+ m.stage2, m.block_m, m.ksplit,
+ m.run_1stage, m.has_bias, True)
- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
- intermediate_pad = config["d_expert_pad"] - config["d_expert"]
+ # Cache metadata for direct dispatch
+ _captured_meta[(a[0], a[1], a[2], a[3], a[4])] = m
+ return m
- output = fused_moe(
- hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
- topk_weights, topk_ids,
+ _fm.get_2stage_cfgs = _p
+
+
+ # =========================================================================== #
+ # Per-shape cache for direct dispatch
+ # =========================================================================== #
+ _shape_cache = {} # shape_key → cache dict or None (fallback)
+
+
+ def _init_shape(data, shape_key):
+ """First call: run fused_moe to populate metadata, build dispatch cache."""
+ hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
+ guws_sh = data[7]; dws_sh = data[8]
+ tw = data[9]; ti = data[10]; cfg = data[11]
+
+ M = cfg['bs']
+ E = cfg['n_routed_experts'] + cfg['n_shared_experts']
+ topk = ti.shape[1]
+ hp = cfg['d_hidden_pad'] - cfg['d_hidden']
+ ip = cfg['d_expert_pad'] - cfg['d_expert']
+ device = hs.device
+
+ # Run fused_moe once to populate metadata cache and return correct output
+ result = fused_moe(
+ hs, guw_sh, dw_sh, tw, ti,
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,
+ w1_scale=guws_sh, w2_scale=dws_sh,
a1_scale=None, a2_scale=None,
- hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
+ hidden_pad=hp, intermediate_pad=ip,
)
- return output
+ # Find captured metadata
+ _, model_dim, inter_dim = _fm.get_inter_dim(guw_sh.shape, dw_sh.shape)
+ padded_M = _fm.get_padded_M(M)
+ meta_key = (padded_M, model_dim, inter_dim, E, topk)
+
+ metadata = _captured_meta.get(meta_key)
+ if metadata is None:
+ print(f"[V64] No metadata for {shape_key}, using fallback",
+ file=sys.stderr, flush=True)
+ _shape_cache[shape_key] = None
+ return result
+
+ block_m = int(metadata.block_m)
+ max_tok = M * topk + E * block_m - topk
+ max_blk = (max_tok + block_m - 1) // block_m
+
+ _shape_cache[shape_key] = {
+ 'meta': metadata,
+ 'block_m': block_m,
+ 'topk': topk,
+ 'E': E,
+ 'M': M,
+ 'model_dim': model_dim,
+ 'inter_dim': inter_dim,
+ 'hp': hp,
+ 'ip': ip,
+ # Pre-allocated sorting buffers (reused across calls)
+ 'sid': torch.empty(max_tok, dtype=torch.int32, device=device),
+ 'sw': torch.empty(max_tok, dtype=torch.float32, device=device),
+ 'seid': torch.empty(max_blk, dtype=torch.int32, device=device),
+ 'nvi': torch.empty(2, dtype=torch.int32, device=device),
+ }
+
+ # Log metadata details
+ s1 = metadata.stage1
+ s1k = getattr(s1, 'keywords', {}) if hasattr(s1, 'func') else {}
+ nt_val = s1k.get('use_non_temporal_load', 'N/A')
+ s1name = s1.func.__name__ if hasattr(s1, 'func') else str(s1)[:40]
+ print(f"[V64] Init {shape_key}: block_m={block_m} inter={inter_dim} "
+ f"model={model_dim} nt={nt_val} stage1={s1name}",
+ file=sys.stderr, flush=True)
+
+ return result
+
+
+ def _direct_dispatch(data, c):
+ """Direct dispatch: 5 GPU kernel calls with minimal Python overhead."""
+ hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
+ guws_sh = data[7]; dws_sh = data[8]
+ tw = data[9]; ti = data[10]
+
+ meta = c['meta']
+ block_m = c['block_m']
+ topk = c['topk']
+ E = c['E']
+ M = c['M']
+ model_dim = c['model_dim']
+ inter_dim = c['inter_dim']
+
+ sid = c['sid']; sw = c['sw']; seid = c['seid']; nvi = c['nvi']
+ device = hs.device
+
+ # 1. Sorting (pre-allocated output buffers, fresh moe_buf for atomicAdd)
+ moe_buf = torch.empty(M, model_dim, dtype=torch.bfloat16, device=device)
+ moe_sorting_fwd(ti, tw, sid, sw, seid, nvi, moe_buf,
+ E, block_m, None, None, 0)
+
+ # 2. Activation quantization (fused quant + scale sorting, Triton kernel)
+ a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ hs, sorted_ids=sid, num_valid_ids=nvi,
+ token_num=M, topk=1, block_size=block_m)
+
+ # 3. Stage 1: gate_up GEMM + SwiGLU (CK or CKtile kernel)
+ a2 = torch.empty(M, topk, inter_dim, dtype=torch.bfloat16, device=device)
+ a2 = meta.stage1(
+ a1, guw_sh, dw_sh, sid, seid, nvi, a2, topk,
+ block_m=block_m,
+ a1_scale=a1_scale,
+ w1_scale=guws_sh.view(dtypes.fp8_e8m0),
+ sorted_weights=None)
+
+ # 4. Intermediate re-quantization (bf16 → fp4x2, Triton kernel)
+ a2_flat = a2.view(-1, inter_dim)
+ a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ a2_flat, sorted_ids=sid, num_valid_ids=nvi,
+ token_num=M, topk=topk, block_size=block_m)
+ a2_q = a2_q.view(M, topk, -1)
+
+ # 5. Stage 2: down GEMM + weighted scatter-reduce (CK/FlyDSL kernel)
+ meta.stage2(
+ a2_q, guw_sh, dw_sh, sid, seid, nvi, moe_buf, topk,
+ w2_scale=dws_sh.view(dtypes.fp8_e8m0),
+ a2_scale=a2_scale,
+ block_m=block_m,
+ sorted_weights=sw)
+
+ return moe_buf
+
+
+ def _fallback(data):
+ """Safe fallback: full fused_moe dispatch."""
+ hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
+ guws_sh = data[7]; dws_sh = data[8]
+ tw = data[9]; ti = data[10]; cfg = data[11]
+ hp = cfg['d_hidden_pad'] - cfg['d_hidden']
+ ip = cfg['d_expert_pad'] - cfg['d_expert']
+ return fused_moe(
+ hs, guw_sh, dw_sh, tw, ti,
+ expert_mask=None, activation=ActivationType.Silu,
+ quant_type=QuantType.per_1x32, doweight_stage1=False,
+ w1_scale=guws_sh, w2_scale=dws_sh,
+ a1_scale=None, a2_scale=None,
+ hidden_pad=hp, intermediate_pad=ip,
+ )
+
+
+ # =========================================================================== #
+ # Main entry point
+ # =========================================================================== #
+ def custom_kernel(data: input_t) -> output_t:
+ _inject()
+
+ cfg = data[11]
+ shape_key = (cfg['bs'], cfg['d_expert'],
+ cfg['n_routed_experts'] + cfg['n_shared_experts'])
+
+ # First call per shape: use fused_moe to populate metadata cache
+ if shape_key not in _shape_cache:
+ return _init_shape(data, shape_key)
+
+ c = _shape_cache[shape_key]
+
+ # Fallback if metadata capture failed
+ if c is None:
+ return _fallback(data)
+
+ # Direct dispatch (bypasses ~8-12µs Python overhead)
+ try:
+ return _direct_dispatch(data, c)
+ except Exception as e:
+ print(f"[V64] Dispatch err {shape_key}: {str(e)[:200]}",
+ file=sys.stderr, flush=True)
+ import traceback
+ traceback.print_exc(file=sys.stderr)
+ _shape_cache[shape_key] = None
+ return _fallback(data)
scrolls · 475 diff lines total

Best evidence level for this revision: reported

JSON