Skip to content
Merged
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
36 changes: 32 additions & 4 deletions megatron/training/theoretical_memory_usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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

Expand Down
19 changes: 19 additions & 0 deletions tests/unit_tests/training/test_weight_and_optimizer_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Loading