Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
40 commits
Select commit Hold shift + click to select a range
51f16cd
sign off
HaochenYuan Nov 6, 2025
ccee48d
fix UT error
HaochenYuan Nov 13, 2025
19786d9
refactor code
HaochenYuan Dec 12, 2025
873fd7f
Merge branch 'main' into main
HaochenYuan Dec 12, 2025
dda7ba6
fix linting
HaochenYuan Dec 12, 2025
9e8023b
fix linting
HaochenYuan Dec 12, 2025
9811067
fix linting
HaochenYuan Dec 12, 2025
c9a13c9
fix linting
HaochenYuan Dec 12, 2025
991b19c
fix linting
HaochenYuan Dec 12, 2025
0cbba67
slice padding_mask for SP
HaochenYuan Dec 18, 2025
53457c1
fix bug in 1f1b & recompute_mlp
HaochenYuan Jan 5, 2026
8c96f90
Merge branch 'main' into main
HaochenYuan Jan 5, 2026
4903c76
fix linting
HaochenYuan Jan 5, 2026
46923ea
Merge branch 'main' into main
HaochenYuan Jan 13, 2026
052fa75
add mamba support
HaochenYuan Jan 13, 2026
c4bd7f5
move padding_mask position
HaochenYuan Jan 13, 2026
321b950
Move padding_mask parameter to keyword-only arguments
HaochenYuan Jan 13, 2026
ef71a8e
Merge branch 'main' into main
HaochenYuan Jan 14, 2026
c0aecc9
Merge branch 'main' into main
HaochenYuan Jan 15, 2026
767638c
Merge branch 'main' into main
HaochenYuan Jan 16, 2026
ea1bb4f
Merge branch 'main' into main
HaochenYuan Jan 16, 2026
2525c62
Merge branch 'main' into main
Phlip79 Jan 21, 2026
be75ed4
Merge branch 'main' into main
HaochenYuan Jan 22, 2026
c766b90
fix linting
HaochenYuan Jan 22, 2026
bfdce92
fix linting
HaochenYuan Jan 22, 2026
111e48e
add padding_mask in MoETransformerLayer
HaochenYuan Jan 22, 2026
2e55f2f
fix linting
HaochenYuan Jan 22, 2026
1564d08
Merge branch 'main' into main
HaochenYuan Jan 23, 2026
867a0b8
fix linting
HaochenYuan Jan 23, 2026
b59d39c
fix linting
HaochenYuan Jan 23, 2026
26e5e93
fix API bug in MoELayer.forward.custom_forward
HaochenYuan Jan 27, 2026
fd04810
Merge branch 'main' into main
HaochenYuan Jan 27, 2026
1d898d4
fix linting
HaochenYuan Jan 27, 2026
c680e24
fix linting
HaochenYuan Jan 27, 2026
af811e2
fix linting
HaochenYuan Jan 27, 2026
ef5aaaf
Merge branch 'main' into main
ko3n1g Jan 27, 2026
ab23b47
fix bug in mbridge
HaochenYuan Jan 28, 2026
9be9463
add moe recompute test with padding_mask
HaochenYuan Jan 30, 2026
4ef2ef2
Merge branch 'main' into main
HaochenYuan Jan 30, 2026
80039b3
fix linting
HaochenYuan Jan 30, 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
3 changes: 2 additions & 1 deletion megatron/core/transformer/moe/moe_layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -437,11 +437,12 @@ def custom_forward(hidden_states, intermediate_tensors, padding_mask=None):
tensor_parallel.random.get_cuda_rng_tracker,
parallel_state.get_tensor_model_parallel_group(),
hidden_states,
intermediate_tensors,
padding_mask,
)
else:
outputs = tensor_parallel.checkpoint(
custom_forward, False, hidden_states, padding_mask
custom_forward, False, hidden_states, intermediate_tensors, padding_mask
)
else:
outputs = custom_forward(hidden_states, intermediate_tensors, padding_mask)
Expand Down
120 changes: 120 additions & 0 deletions tests/unit_tests/transformer/moe/test_moe_layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -276,3 +276,123 @@ def test_moe_layer_fp16_forward_backward(

def teardown_method(self, method):
Utils.destroy_model_parallel()


class TestMoELayerRecompute:
"""Test MoE layer with recompute enabled (activation checkpointing).

Tests both code paths:
- fp8=False: uses tensor_parallel.checkpoint
- fp8=True: uses te_checkpoint (requires TE >= 1.7.0)
"""

def setup_method(self, method):
pass

@pytest.mark.parametrize("moe_token_dispatcher_type", ["allgather", "alltoall"])
@pytest.mark.parametrize("num_moe_experts", [2, 4])
@pytest.mark.parametrize("with_padding_mask", [True, False])
@pytest.mark.parametrize("tp_size,ep_size", [(1, 1), (4, 2)])
@pytest.mark.parametrize("fp8", [False, True])
def test_moe_layer_recompute_forward_backward(
self, num_moe_experts, moe_token_dispatcher_type, with_padding_mask, tp_size, ep_size, fp8
):
"""Test MoE layer forward and backward pass with recompute enabled.

When fp8=False, uses tensor_parallel.checkpoint.
When fp8=True, uses te_checkpoint (requires TE >= 1.7.0).
"""
# Skip fp8 tests if TE version is not sufficient
if fp8 and not is_te_min_version("1.7.0.dev0"):
pytest.skip("FP8 MoE recompute requires TE 1.7.0 and later.")

Utils.initialize_model_parallel(
tensor_model_parallel_size=tp_size, expert_model_parallel_size=ep_size
)
_set_random_seed(seed_=123, data_parallel_random_init=False)

hidden_size = 64
sequence_length = 32
micro_batch_size = 2

transformer_config = TransformerConfig(
num_layers=1,
hidden_size=hidden_size,
num_attention_heads=4,
num_moe_experts=num_moe_experts,
use_cpu_initialization=False,
moe_token_dispatcher_type=moe_token_dispatcher_type,
moe_router_load_balancing_type="aux_loss",
moe_router_topk=2,
moe_aux_loss_coeff=0.01,
moe_grouped_gemm=False,
moe_ffn_hidden_size=256,
add_bias_linear=False,
# Enable recompute for MoE layer
recompute_granularity="selective",
recompute_modules=["moe"],
tensor_model_parallel_size=tp_size,
expert_model_parallel_size=ep_size,
sequence_parallel=tp_size > 1,
fp8=fp8,
bf16=True,
params_dtype=torch.bfloat16,
)

# Use TE spec for fp8, local spec otherwise
if fp8:
transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(
num_experts=num_moe_experts, moe_grouped_gemm=False
)
else:
transformer_layer_spec = get_gpt_layer_local_spec(
num_experts=num_moe_experts, moe_grouped_gemm=False
)

moe_layer = MoELayer(
transformer_config, transformer_layer_spec.submodules.mlp.submodules
).cuda()

hidden_states = torch.randn(
sequence_length,
micro_batch_size,
hidden_size,
device=torch.cuda.current_device(),
dtype=torch.bfloat16,
requires_grad=True,
)

# Create padding mask if needed: shape [batch_size, sequence_length]
padding_mask = None
if with_padding_mask:
padding_mask = torch.ones(
micro_batch_size,
sequence_length,
device=torch.cuda.current_device(),
dtype=torch.bool,
)
# Mark last 4 tokens as padding for each batch
padding_mask[:, -4:] = False

output, _ = moe_layer(hidden_states, padding_mask=padding_mask)

assert output.dtype == torch.bfloat16, f"Expected bf16 output, got {output.dtype}"
assert output.shape == hidden_states.shape, f"Output shape mismatch"

# Backward pass - this is where recompute/checkpoint is actually used
loss = output.sum()
loss.backward()

assert hidden_states.grad is not None, "Input gradients should exist"
assert (
hidden_states.grad.dtype == torch.bfloat16
), f"Expected bf16 gradients, got {hidden_states.grad.dtype}"

for name, param in moe_layer.named_parameters():
if param.requires_grad:
assert param.grad is not None, f"Gradient for {name} should exist"

Utils.destroy_model_parallel()

def teardown_method(self, method):
Utils.destroy_model_parallel()
Loading