Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
42 commits
Select commit Hold shift + click to select a range
9d8420f
mega_moe: packaged fused MoE operator + autotune config path fix
GwilliamHu Jul 17, 2026
6e3bdd9
fix(moe gemm2): use host-int k_shift in scale prefetch (avoid arith.c…
yanboshao Jul 20, 2026
6707a1a
refactor(mega_moe): fold MegaMoeStage1/2 into MegaMoE, slim group-maj…
yanboshao Jul 20, 2026
bd0fbe7
refactor(mega_moe): move package to kernels/mega_moe (sibling of moe/…
yanboshao Jul 20, 2026
6fb5615
refactor: remove unused use_token_flag_sync token-flag-sync path (gem…
yanboshao Jul 20, 2026
7183e0a
add CI case
GwilliamHu Jul 20, 2026
93a6fe5
delete md files
GwilliamHu Jul 21, 2026
fbf34d0
style: black-format comm ops + dispatch-combine test (line-length 120)
GwilliamHu Jul 21, 2026
46321dd
fix import error
GwilliamHu Jul 21, 2026
dab1954
mega_moe: descriptive symbol name for stage1 gemm1 kernel (name=modul…
yanboshao Jul 22, 2026
5b0d911
refactor group gemm
GwilliamHu Jul 22, 2026
997d505
add refactor gemm
GwilliamHu Jul 23, 2026
605bdba
add refactor megamoe
GwilliamHu Jul 23, 2026
338eb6a
refactor megamoev2
GwilliamHu Jul 24, 2026
996c63e
mega_moe_exp: port aiter mxmoe gemm2 into fused stage2 + gemm2 autotune
yanboshao Jul 26, 2026
7adf931
mega_moe_exp: enable persistent gemm2 for fp8 in fused stage2 autotune
yanboshao Jul 26, 2026
12c90a3
megamoev2 released
GwilliamHu Jul 27, 2026
20997ee
fix megamoev2 percision err on bs 32768
GwilliamHu Jul 27, 2026
df11d07
code clean
GwilliamHu Jul 27, 2026
9f23660
tune for stage1
GwilliamHu Jul 27, 2026
8d1c9be
optimize megamoe stage2
Yaowu-Xiong Jul 27, 2026
af566aa
rebase megamoe stage2
Yaowu-Xiong Jul 27, 2026
e4f0ed4
optimize MegaMoE V2 FP8 stage2 transport
Yaowu-Xiong Jul 29, 2026
2a2704c
megamoev2
GwilliamHu Jul 29, 2026
1910320
Merge branch 'main' into mega_moe_v1
GwilliamHu Jul 29, 2026
52ee01d
tune megamoev2
GwilliamHu Jul 29, 2026
0183070
fix CI
GwilliamHu Jul 29, 2026
b954a4f
fix CI py check
GwilliamHu Jul 29, 2026
78b2536
reduce duplicate comments
GwilliamHu Jul 29, 2026
aaeecb5
autotune megamoev2
GwilliamHu Jul 29, 2026
f752c7d
optimize stage1
GwilliamHu Jul 29, 2026
c00df7d
reduce autotune file
GwilliamHu Jul 30, 2026
4821432
autotune bundle
GwilliamHu Jul 30, 2026
3dcd641
fix autotune deadlock and tuneconfig mismatch problem
GwilliamHu Jul 30, 2026
4c613f7
remove dup stage2 env
GwilliamHu Jul 30, 2026
e092035
Merge origin/main into mega_moe_v1
GwilliamHu Jul 30, 2026
cc020fc
remove unsafe autotune
GwilliamHu Jul 31, 2026
c5e1ef1
Merge branch 'main' into mega_moe_v1
GwilliamHu Jul 31, 2026
cb0b459
fix code check
GwilliamHu Jul 31, 2026
4ccb610
remove autotune change
GwilliamHu Jul 31, 2026
2f9da1f
add stage2 config
GwilliamHu Jul 31, 2026
5be20f5
apply review comments
GwilliamHu Jul 31, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions .github/workflows/flydsl.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -611,6 +611,18 @@ jobs:
--output-dir /tmp/flydsl_ci_sweep
"

# MegaMoEV2 A8W4 v4_pro accuracy against the torch f32 oracle. The test self-skips on gfx942.
- name: Run multi-GPU MegaMoEV2 A8W4 v4_pro accuracy tests
timeout-minutes: 45
run: |
docker exec flydsl_test bash -c "
cd /flydsl-test &&
python3 -c 'import torch; n = torch.cuda.device_count(); assert n >= 8, f\"requires 8 GPUs, found {n}\"' &&
python3 -m pytest \
tests/kernels/test_mega_moe_v2.py::test_mega_moe_8gpu_accuracy \
-v --no-header --tb=short
"

- name: Run multi-GPU allreduce tests
timeout-minutes: 30
run: |
Expand Down
72 changes: 68 additions & 4 deletions kernels/comm/communication_ops_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,16 +17,25 @@
from dataclasses import dataclass, field
from typing import Dict, Tuple

import flydsl.expr as fx
from flydsl._mlir import ir
from flydsl._mlir.dialects import llvm as _llvm_d
from flydsl.expr import arith

__all__ = [
"store_i32_system",
"store_i64_global_system",
"fence_acquire",
"fence_release",
"fence_system_acquire",
"fence_system_release",
"fence_agent_acquire",
"fence_agent_release",
"load_i64_global",
"atomic_add_global_at",
"atomic_add_agent",
"atomic_add_system",
"atomic_xchg_global_at",
"GeometryTuningTable",
]

Expand Down Expand Up @@ -57,9 +66,34 @@ def store_i64_global_system(addr_i64, val):
_llvm_d.StoreOp(arith.unwrap(val), gptr, alignment=8, ordering=_llvm_d.AtomicOrdering.release, syncscope="one-as")


def fence_acquire(syncscope):
"""Emit an acquire fence for the selected AMDGPU memory scope."""
_llvm_d.FenceOp(_llvm_d.AtomicOrdering.acquire, syncscope=syncscope)


def fence_release(syncscope):
"""Emit a release fence for the selected AMDGPU memory scope."""
_llvm_d.FenceOp(_llvm_d.AtomicOrdering.release, syncscope=syncscope)


def fence_system_acquire():
"""System-scope acquire fence."""
_llvm_d.FenceOp(_llvm_d.AtomicOrdering.acquire, syncscope="one-as")
fence_acquire(fx.rocdl.SyncScope.OneAs)


def fence_system_release():
"""System-scope release fence."""
fence_release(fx.rocdl.SyncScope.OneAs)


def fence_agent_acquire():
"""Agent-scope acquire fence."""
fence_acquire(fx.rocdl.SyncScope.AgentOneAs)


def fence_agent_release():
"""Agent-scope release fence."""
fence_release(fx.rocdl.SyncScope.AgentOneAs)


def load_i64_global(addr_i64):
Expand All @@ -69,10 +103,40 @@ def load_i64_global(addr_i64):
return _llvm_d.LoadOp(_i64, ptr, alignment=8).result


def atomic_add_global_at(addr_i64, val):
"""Monotonic global ``atomic fetch-and-add``; returns the old value."""
def atomic_add_global_at(addr_i64, val, syncscope="one-as"):
"""Monotonic global fetch-add with configurable agent/system visibility."""
ptr = _to_ptr_global(addr_i64)
kwargs = {} if syncscope is None else {"syncscope": syncscope}
return _llvm_d.AtomicRMWOp(
_llvm_d.AtomicBinOp.add,
ptr,
arith.unwrap(val),
_llvm_d.AtomicOrdering.monotonic,
**kwargs,
).res


def atomic_add_agent(addr_i64, val):
"""Agent-scope monotonic global fetch-and-add."""
return atomic_add_global_at(addr_i64, val, syncscope=fx.rocdl.SyncScope.Agent)


def atomic_add_system(addr_i64, val):
"""System-scope monotonic global fetch-and-add."""
return atomic_add_global_at(addr_i64, val)


def atomic_xchg_global_at(addr_i64, val, syncscope="agent"):
"""Monotonic global exchange with configurable agent/system visibility."""
ptr = _to_ptr_global(addr_i64)
return _llvm_d.AtomicRMWOp(_llvm_d.AtomicBinOp.add, ptr, arith.unwrap(val), _llvm_d.AtomicOrdering.monotonic).res
kwargs = {} if syncscope is None else {"syncscope": syncscope}
return _llvm_d.AtomicRMWOp(
_llvm_d.AtomicBinOp.xchg,
ptr,
arith.unwrap(val),
_llvm_d.AtomicOrdering.monotonic,
**kwargs,
).res


@dataclass
Expand Down
102 changes: 93 additions & 9 deletions kernels/comm/flydsl_dispatch_combine_intranode_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
cvt_pk_fp8_f32,
cvt_scalef32_pk_f32_fp4,
cvt_scalef32_pk_fp4_f32,
ds_bpermute,
readfirstlane,
readlane,
)
Expand All @@ -38,7 +39,7 @@
)

# Bump when generated kernel shape changes.
_DISPATCH_COMBINE_JIT_SCHEMA_VERSION = "v6-combine-batched-kload"
_DISPATCH_COMBINE_JIT_SCHEMA_VERSION = "v10-stage2-blockwise-fp8-scale-prefetch"

# Stage-3 switches from narrow step=64 to wide step=128/256 above this threshold.
_S3_WIDE_PATH_THRESHOLD_I32 = 895
Expand Down Expand Up @@ -374,6 +375,7 @@ def make_combine_kernel(
zero_copy: bool = False,
skip_stage1: bool = False,
fp8_direct_cast: bool = False,
blockwise_fp8_transport: bool = False,
max_recv: int = None,
):
"""Build the intranode combine ``@flyc.kernel``.
Expand All @@ -392,13 +394,20 @@ def make_combine_kernel(
"""
# Contract (op-layer _check_config): fp8_direct_cast => data_type==bf16 and
# not enable_std_moe. skip_stage1 and zero_copy are independent switches.
if blockwise_fp8_transport and (not skip_stage1 or data_type != torch.bfloat16):
raise ValueError("blockwise_fp8_transport requires skip_stage1=True and external bf16")
if blockwise_fp8_transport and fp8_direct_cast:
raise ValueError("blockwise_fp8_transport and fp8_direct_cast are mutually exclusive")
_xfer_bf16_to_fp8 = fp8_direct_cast
_transport_dtype = torch.float8_e4m3fn if _xfer_bf16_to_fp8 else data_type

if max_recv is None:
max_recv = npes * max_tok_per_rank
_is_fp4 = _transport_dtype == torch.float4_e2m1fn_x2
if _is_fp4:
if blockwise_fp8_transport:
n_i32 = hidden_dim // 4
nbytes = hidden_dim + hidden_dim // 32
elif _is_fp4:
n_i32 = hidden_dim // 8
nbytes = hidden_dim // 2
else:
Expand All @@ -407,13 +416,30 @@ def make_combine_kernel(

# Stage 1/3 strides diverge only under ``fp8_direct_cast``: external
# bf16 reads/writes vs fp8 staging. Other modes keep transport == external.
if _xfer_bf16_to_fp8:
if blockwise_fp8_transport:
inp_nbytes = hidden_dim * 2
out_n_i32 = (hidden_dim * 2) // 4
elif _xfer_bf16_to_fp8:
inp_nbytes = hidden_dim * 2
out_n_i32 = (hidden_dim * 2) // 4
else:
inp_nbytes = nbytes
out_n_i32 = n_i32
if _is_fp4:
if blockwise_fp8_transport:

def _to_accum(i32_val):
_v2f32_fp8 = T.VectorType.get([2], T.f32())
lo = cvt_pk_f32_fp8(res=_v2f32_fp8, src=i32_val, word_sel=False)
hi = cvt_pk_f32_fp8(res=_v2f32_fp8, src=i32_val, word_sel=True)
return lo.shuffle(hi, [0, 1, 2, 3])

def _from_accum(accum_val):
return accum_val.to(fx.BFloat16).bitcast(fx.Int32)

def _zero_accum():
return Vec.filled(4, 0.0, fx.Float32)

elif _is_fp4:

def _to_accum(i32_val):
_v2f32_fp4 = T.VectorType.get([2], T.f32())
Expand Down Expand Up @@ -784,7 +810,13 @@ def _maybe_load(rsrc, offset, vld_flag, **kwargs):
# Clamp denom to 1 when cur_rank_num_token == 0 (loop won't execute anyway).
safe_token_count = (cur_rank_num_token == 0).select(1, cur_rank_num_token)
warps_per_tok = (global_warp_num + safe_token_count - 1) // safe_token_count
hdim_per_warp = (n_elems + warps_per_tok - 1) // warps_per_tok
if const_expr(blockwise_fp8_transport):
# Align warp partitions to the 32-value blockwise FP8 scale.
scale_blocks = n_elems // 8
warps_per_tok = (warps_per_tok > scale_blocks).select(scale_blocks, warps_per_tok)
hdim_per_warp = ((scale_blocks + warps_per_tok - 1) // warps_per_tok) * 8
else:
hdim_per_warp = (n_elems + warps_per_tok - 1) // warps_per_tok
s3_total_work = cur_rank_num_token * warps_per_tok

for s3_work_idx in range(global_warp_id, s3_total_work, global_warp_num):
Expand All @@ -793,6 +825,7 @@ def _maybe_load(rsrc, offset, vld_flag, **kwargs):
hdim_off = part_id * hdim_per_warp

expert_rsrcs = []
expert_scale_rsrcs = []
expert_vlds = []

if const_expr(skip_stage1 and not zero_copy):
Expand All @@ -803,7 +836,12 @@ def _maybe_load(rsrc, offset, vld_flag, **kwargs):
expert_tok_off = fx.Int64(slot_idx) * nbytes
expert_tok_addr = as_ir_value(addr_shmem_tok + expert_tok_off)
# Warp-uniform base -> SGPR (avoids per-lane waterfall).
expert_rsrcs.append(create_buffer_resource_from_addr(_wave_uniform_i64(expert_tok_addr)))
expert_tok_addr = _wave_uniform_i64(expert_tok_addr)
expert_rsrcs.append(create_buffer_resource_from_addr(expert_tok_addr))
if const_expr(blockwise_fp8_transport):
expert_scale_rsrcs.append(
create_buffer_resource_from_addr(expert_tok_addr + fx.Int64(hidden_dim))
)
expert_vlds.append(fx.Boolean(1))
else:
# Baseline Stage 3: decode (peer_pe, dest_lid) from dest_tok_map and
Expand Down Expand Up @@ -832,6 +870,8 @@ def _maybe_load(rsrc, offset, vld_flag, **kwargs):

def _accum_step(ec_abs, U):
vals = [[] for _ in range(U)]
scales = [[] for _ in range(U)]
scale_raws = [[] for _ in range(U)]
for k_slot in range_constexpr(experts_per_token):
rsrc_k = expert_rsrcs[k_slot]
vld_k = expert_vlds[k_slot]
Expand All @@ -840,16 +880,56 @@ def _accum_step(ec_abs, U):
if u > 0:
kw["soffset_bytes"] = u * 256
vals[u].append(_maybe_load(rsrc_k, ec_abs, vld_k, **kw))

if const_expr(_xfer_bf16_to_fp8):
if const_expr(blockwise_fp8_transport):
scale_idx = (ec_abs + u * 64) // 8
sc_i32 = fx.Int32(0)
if (lane & 7) == 0:
sc_raw = _maybe_load(
expert_scale_rsrcs[k_slot],
scale_idx,
vld_k,
vec_width=1,
dtype=T.i8(),
cache_modifier=SLC_CACHE,
)
sc_i32 = fx.Uint8(sc_raw).to(fx.Uint32).bitcast(fx.Int32)
scale_raws[u].append(sc_i32)

if const_expr(blockwise_fp8_transport):
# Broadcast after issuing all VMEM loads to avoid serial vmcnt waits.
for u in range_constexpr(U):
for k_slot in range_constexpr(experts_per_token):
sc_i32 = fx.Int32(
ds_bpermute(
T.i32(),
(lane & fx.Int32(~7)) * fx.Int32(4),
scale_raws[u][k_slot],
)
)
scales[u].append(
fx.Float32(
arith.bitcast(
T.f32(),
arith.unwrap(sc_i32 << fx.Int32(23)),
)
)
)

if const_expr(_xfer_bf16_to_fp8 or blockwise_fp8_transport):
out_off = tok_id * out_n_i32 + ec_abs * 2
out_step = 512 # bf16 store on fp8-stride buffer: 2x
else:
out_off = tok_id * n_i32 + ec_abs
out_step = 256

for u in range_constexpr(U):
acc = _accum_experts(vals[u])
if const_expr(blockwise_fp8_transport):
fp32_acc = _zero_accum()
for k_slot in range_constexpr(experts_per_token):
fp32_acc = fp32_acc + _to_accum(vals[u][k_slot]) * scales[u][k_slot]
acc = _from_accum(fp32_acc)
else:
acc = _accum_experts(vals[u])
kw = dict(cache_modifier=SLC_CACHE)
if u > 0:
kw["soffset_bytes"] = u * out_step
Expand Down Expand Up @@ -1071,6 +1151,7 @@ def make_combine_jit(
zero_copy=False,
skip_stage1=False,
fp8_direct_cast: bool = False,
blockwise_fp8_transport: bool = False,
max_recv=None,
):
"""Build the JIT launcher for ``make_combine_kernel``. ``data_type`` is the
Expand All @@ -1094,6 +1175,7 @@ def make_combine_jit(
zero_copy=zero_copy,
skip_stage1=skip_stage1,
fp8_direct_cast=fp8_direct_cast,
blockwise_fp8_transport=blockwise_fp8_transport,
max_recv=max_recv,
)

Expand All @@ -1106,6 +1188,7 @@ def make_combine_jit(
_key_zero_copy = zero_copy
_key_skip_s1 = skip_stage1
_key_fp8_direct_cast = bool(fp8_direct_cast)
_key_blockwise_fp8_transport = bool(blockwise_fp8_transport)
_key_max_recv = max_recv if max_recv is not None else npes * max_tok_per_rank
# See dispatch launcher for the ``str(torch.dtype)`` rationale.
_key_data_type = str(data_type)
Expand Down Expand Up @@ -1145,6 +1228,7 @@ def combine_launch(
_key_zero_copy,
_key_skip_s1,
_key_fp8_direct_cast,
_key_blockwise_fp8_transport,
_key_max_recv,
_key_data_type,
_key_schema_version,
Expand Down
Loading
Loading