Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
8 changes: 5 additions & 3 deletions megatron/core/transformer/transformer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -2658,9 +2658,11 @@ def _scope_to_str(s):
assert (
not self.moe_shared_expert_overlap
), 'disable moe_shared_expert_overlap when enabling overlap_moe_expert_parallel_comm'
assert (
self.mtp_num_layers is None or self.mtp_num_layers == 1
), 'MTP layernum only supports 1 when enabling overlap_moe_expert_parallel_comm.'
assert self.mtp_num_layers in (
None,
0,
1,
), 'MTP supports at most one layer when enabling overlap_moe_expert_parallel_comm.'

# NCCL EP (ncclep flex backend) mirrors hybridep's comm/compute overlap, but a few
# configs are not yet safe under the 1F1B split and are gated here.
Expand Down
32 changes: 32 additions & 0 deletions tests/unit_tests/transformer/test_transformer_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.

import pytest

from megatron.core.transformer.transformer_config import TransformerConfig


def _make_overlap_config(mtp_num_layers: int | None) -> TransformerConfig:
return TransformerConfig(
num_layers=1,
hidden_size=128,
num_attention_heads=4,
num_moe_experts=2,
expert_model_parallel_size=2,
moe_token_dispatcher_type="alltoall",
overlap_moe_expert_parallel_comm=True,
bf16=True,
mtp_num_layers=mtp_num_layers,
)


@pytest.mark.parametrize("mtp_num_layers", [None, 0, 1])
def test_ep_a2a_overlap_accepts_supported_mtp_layer_counts(mtp_num_layers: int | None):
config = _make_overlap_config(mtp_num_layers)

assert config.mtp_num_layers == mtp_num_layers


@pytest.mark.parametrize("mtp_num_layers", [-1, 2])
def test_ep_a2a_overlap_rejects_unsupported_mtp_layer_counts(mtp_num_layers: int):
with pytest.raises(AssertionError, match="MTP supports at most one layer"):
_make_overlap_config(mtp_num_layers)
Loading