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
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 torchimport osimport sysimport functools- import torch- from task import input_t, output_t- from aiter import ActivationType, QuantType++ from aiter import ActivationType, QuantType, dtypesfrom 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 _injectedif _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