From 081635aa71c7bf006911dfdcc2a073917bc06603 Mon Sep 17 00:00:00 2001 From: Shijie Wang Date: Thu, 4 Jun 2026 10:19:19 +0800 Subject: [PATCH] Fix LatentMoE theoretical memory estimate Signed-off-by: Shijie Wang --- megatron/training/theoretical_memory_usage.py | 36 ++++++++++++++++--- .../test_weight_and_optimizer_memory.py | 19 ++++++++++ 2 files changed, 51 insertions(+), 4 deletions(-) diff --git a/megatron/training/theoretical_memory_usage.py b/megatron/training/theoretical_memory_usage.py index fc208640c3a..ee398d3bf66 100644 --- a/megatron/training/theoretical_memory_usage.py +++ b/megatron/training/theoretical_memory_usage.py @@ -102,14 +102,28 @@ def compute_weight_and_optimizer_memory(args, verbose=False): attention_params = self_attn_term dense_mlp_params = 2 * args.hidden_size * args.ffn_hidden_size * gated_linear_multiplier shared_expert_params = 2 * args.hidden_size * shared_expert_ffn_hidden_size * gated_linear_multiplier + # Latent MoE projects tokens down before routed experts and projects them back after + # combine. Shared experts still operate on the full hidden dimension. + routed_expert_hidden_size = ( + args.moe_latent_size if args.moe_latent_size is not None else args.hidden_size + ) routed_expert_params = ( - 2 * args.hidden_size * moe_ffn_hidden_size * num_experts * gated_linear_multiplier + 2 * routed_expert_hidden_size * moe_ffn_hidden_size * num_experts * gated_linear_multiplier ) active_routed_expert_params = ( - 2 * args.hidden_size * moe_ffn_hidden_size * args.moe_router_topk * gated_linear_multiplier + 2 + * routed_expert_hidden_size + * moe_ffn_hidden_size + * args.moe_router_topk + * gated_linear_multiplier if args.num_experts is not None else 0 ) + latent_projection_params = ( + 2 * args.hidden_size * args.moe_latent_size + if args.num_experts is not None and args.moe_latent_size is not None + else 0 + ) layernorm_params = 2 * args.hidden_size * norm_size router_params = ( (args.hidden_size * num_experts) + (num_experts if args.add_bias_linear else 0) @@ -129,6 +143,7 @@ def compute_weight_and_optimizer_memory(args, verbose=False): attention_params + shared_expert_params + routed_expert_params + + latent_projection_params + layernorm_params + router_params + shared_expert_gate_params @@ -137,6 +152,7 @@ def compute_weight_and_optimizer_memory(args, verbose=False): attention_params + shared_expert_params + active_routed_expert_params + + latent_projection_params + layernorm_params + router_params + shared_expert_gate_params @@ -202,7 +218,13 @@ def compute_weight_and_optimizer_memory(args, verbose=False): ) replicated_params_in_transformer_block = ( layernorm_params * num_dense_layers - + (layernorm_params + router_params + shared_expert_gate_params) * num_moe_layers + + ( + layernorm_params + + latent_projection_params + + router_params + + shared_expert_gate_params + ) + * num_moe_layers + final_layernorm ) expert_sharded_params_in_transformer_block = routed_expert_params * num_moe_layers @@ -212,7 +234,13 @@ def compute_weight_and_optimizer_memory(args, verbose=False): ) replicated_params_in_mtp_block = ( layernorm_params * mtp_num_dense_layers - + (layernorm_params + router_params + shared_expert_gate_params) * mtp_num_moe_layers + + ( + layernorm_params + + latent_projection_params + + router_params + + shared_expert_gate_params + ) + * mtp_num_moe_layers ) expert_sharded_params_in_mtp_block = routed_expert_params * mtp_num_moe_layers diff --git a/tests/unit_tests/training/test_weight_and_optimizer_memory.py b/tests/unit_tests/training/test_weight_and_optimizer_memory.py index c172fde39ae..19a532c35c2 100644 --- a/tests/unit_tests/training/test_weight_and_optimizer_memory.py +++ b/tests/unit_tests/training/test_weight_and_optimizer_memory.py @@ -17,6 +17,7 @@ def _make_args(**overrides): hidden_size=8, kv_channels=4, moe_ffn_hidden_size=16, + moe_latent_size=None, moe_layer_freq=[0, 1], moe_router_topk=1, moe_shared_expert_gate=False, @@ -88,3 +89,21 @@ def test_weight_and_optimizer_memory_decreases_with_expert_parallelism(): ] assert memories[0] > memories[1] > memories[2] > memories[3] + + +def test_weight_and_optimizer_memory_accounts_for_latent_moe_experts(): + args = _make_args(moe_latent_size=4) + + # Latent MoE routes experts through moe_latent_size instead of hidden_size. + # The hidden<->latent projections are duplicated non-expert params. + tp_sharded_params_on_rank = ((256 + 256) + (256 + 128) + 256) / 2 + replicated_params_on_rank = 16 + (16 + 64 + 32) + 8 + expert_sharded_params_on_rank = 512 / (4 * 2) + + # DP = 32 // 2(TP) = 16 + # EDP = 32 // 4(ETP) // 2(EP) = 4 + expected_memory = (tp_sharded_params_on_rank + replicated_params_on_rank) * ( + 6 + 12 / 16 + ) + expert_sharded_params_on_rank * (6 + 12 / 4) + + assert math.isclose(compute_weight_and_optimizer_memory(args), expected_memory)