Skip to content
47 changes: 38 additions & 9 deletions megatron/core/distributed/fsdp/mcore_fsdp_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
from megatron.core.distributed.distributed_data_parallel_config import DistributedDataParallelConfig
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.ssm.mamba_layer import MambaLayer
from megatron.core.transformer.moe.moe_layer import MoELayer
from megatron.core.transformer.transformer_config import TransformerConfig
from megatron.core.transformer.transformer_layer import MoETransformerLayer, TransformerLayer
from megatron.core.utils import is_te_min_version, log_single_rank
Expand Down Expand Up @@ -560,7 +561,12 @@ def __init__(

dp_group = pg_collection.dp_cp
device_type = device.type if device is not None else "cuda"
mesh = DeviceMesh.from_group(dp_group, device_type=device_type, mesh_dim_names=("dp",))
dp_mesh = DeviceMesh.from_group(dp_group, device_type=device_type, mesh_dim_names=("dp",))
expert_dp_mesh = None
if config.expert_model_parallel_size > 1:
expert_dp_mesh = DeviceMesh.from_group(
pg_collection.expt_dp, device_type=device_type, mesh_dim_names=("expert_dp",)
)
placements = Placements(
dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]
)
Expand All @@ -569,6 +575,19 @@ def __init__(
# ncclCommWindowRegister:
# https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/usage/bufferreg.html#window-registration
with fully_shard_context(device=device, use_symmetric_memory=ddp_config.nccl_ub):
if expert_dp_mesh is not None:
# Expert parameters are replicated over expert-DP, not the full DP group.
# Their gradients need the EP divisor because the same expert receives
# contributions after dispatch from every EP rank.
for submodule in module.modules():
if isinstance(submodule, MoELayer):
fully_shard(
submodule.experts,
mesh=expert_dp_mesh,
placements=placements,
mixed_precision_policy=self.mp_policy,
grad_divisor=config.expert_model_parallel_size,
)
for submodule in reversed(list(module.modules())):
if submodule is module:
# The root is always sharded after selected child units so it is not
Expand All @@ -577,12 +596,12 @@ def __init__(
if any(isinstance(submodule, module_type) for module_type in fsdp_unit_modules):
fully_shard(
submodule,
mesh=mesh,
mesh=dp_mesh,
placements=placements,
mixed_precision_policy=self.mp_policy,
)
fully_shard(
module, mesh=mesh, placements=placements, mixed_precision_policy=self.mp_policy
module, mesh=dp_mesh, placements=placements, mixed_precision_policy=self.mp_policy
)
super().__init__(config=config, module=module)

Expand Down Expand Up @@ -617,7 +636,6 @@ def _validate_config(
"tensor_model_parallel_size",
"pipeline_model_parallel_size",
"context_parallel_size",
"expert_model_parallel_size",
]
if any(getattr(config, parallelism) != 1 for parallelism in unsupported_parallelisms):
raise ValueError(
Expand All @@ -630,18 +648,29 @@ def _validate_config(

# The config validates the requested topology, while these checks validate the
# materialized topology supplied by the caller's process-group collection.
for group_name in ("tp", "pp", "cp", "ep"):
for group_name in ("tp", "pp", "cp"):
group = getattr(pg_collection, group_name, None)
if group is not None and group.size() != 1:
raise ValueError(
f"MFSDP v2 currently requires {group_name.upper()} process-group size 1, "
f"got {group.size()}."
)

if getattr(config, "num_moe_experts", None) is not None or any(
not getattr(parameter, "allreduce", True) for parameter in module.parameters()
):
raise ValueError("MFSDP v2 does not currently support expert parameters.")
if config.expert_model_parallel_size > 1:
if (
pg_collection.ep is None
or pg_collection.ep.size() != config.expert_model_parallel_size
):
actual_ep_size = None if pg_collection.ep is None else pg_collection.ep.size()
raise ValueError(
"MFSDP v2 requires an EP process group matching "
f"expert_model_parallel_size={config.expert_model_parallel_size}, "
f"got {actual_ep_size}."
)
if pg_collection.expt_dp is None:
raise ValueError("MFSDP v2 with EP requires an explicit expert-DP process group.")
if not any(isinstance(submodule, MoELayer) for submodule in module.modules()):
raise ValueError("MFSDP v2 with EP requires MoE transformer layers.")
if ddp_config.data_parallel_sharding_strategy != "optim_grads_params":
raise ValueError(
"MFSDP v2 requires data_parallel_sharding_strategy='optim_grads_params'."
Expand Down
191 changes: 183 additions & 8 deletions tests/unit_tests/distributed/mfsdp_v2/test_mcore_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

"""MCore adapter and optimizer integration tests for experimental MFSDP v2."""

import logging
import os
from dataclasses import replace

import pytest
Expand All @@ -12,15 +14,20 @@
from megatron.core.distributed.fsdp.mcore_fsdp_adapter import FullyShardedDataParallel
from megatron.core.distributed.fsdp.src.megatron_fsdp.experimental.module import FsdpModule
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec
from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec
from megatron.core.models.hybrid.hybrid_model import HybridModel
from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer
from megatron.core.optimizer.fully_sharded_optimizer import FullyShardedOptimizer
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed
from megatron.core.transformer.enums import AttnBackend
from megatron.core.transformer.transformer_block import TransformerBlock
from megatron.core.transformer.transformer_config import TransformerConfig
from megatron.core.transformer.transformer_layer import TransformerLayer
from megatron.core.transformer.transformer_layer import MoETransformerLayer, TransformerLayer
from tests.unit_tests.test_utilities import Utils

logger = logging.getLogger(__name__)


def _build_layer(config: TransformerConfig) -> TransformerLayer:
return TransformerLayer(
Expand All @@ -37,13 +44,11 @@ def _build_block(config: TransformerConfig) -> TransformerBlock:
)


class TestMcoreAdapter:
class TestMcoreAdapterDense:
"""Exercise a dense MCore transformer block over two data-parallel ranks."""

def setup_method(self):
Utils.initialize_model_parallel(1, 1)
if torch.distributed.get_world_size() < 2:
pytest.skip("MFSDP v2 MCore integration test requires at least two ranks.")
self.pg_collection = ProcessGroupCollection.use_mpu_process_groups()
model_parallel_cuda_manual_seed(1234)

Expand Down Expand Up @@ -72,8 +77,6 @@ def test_wraps_fsdp_unit_modules_before_root(self):
megatron_fsdp_version=2,
use_distributed_optimizer=False,
data_parallel_sharding_strategy="optim_grads_params",
megatron_fsdp_main_params_dtype=torch.float32,
megatron_fsdp_main_grads_dtype=torch.float32,
),
module=model,
fsdp_unit_modules=[TransformerLayer],
Expand Down Expand Up @@ -160,8 +163,6 @@ def test_build_train_and_step(self):
megatron_fsdp_version=2,
use_distributed_optimizer=False,
data_parallel_sharding_strategy="optim_grads_params",
megatron_fsdp_main_params_dtype=torch.float32,
megatron_fsdp_main_grads_dtype=torch.bfloat16,
),
module=model,
pg_collection=self.pg_collection,
Expand Down Expand Up @@ -232,3 +233,177 @@ def test_build_train_and_step(self):
assert torch.isfinite(losses).all()
assert torch.isfinite(reference_losses).all()
torch.testing.assert_close(losses, reference_losses, rtol=1e-3, atol=0)


class TestMcoreAdapterExpertParallel:
"""Exercise the MFSDP v2 adapter over an MoE model with EP=2."""

def setup_method(self):
self.world_size = int(os.environ.get("WORLD_SIZE", "1"))
if self.world_size < 2 or self.world_size % 2:
pytest.skip("MFSDP v2 EP adapter test requires an even world size of at least two.")
Utils.initialize_model_parallel(1, 1, expert_model_parallel_size=2)
self.pg_collection = ProcessGroupCollection.use_mpu_process_groups()
assert self.pg_collection.ep.size() == 2
assert self.pg_collection.expt_dp.size() == self.world_size // 2
self.reference_group = torch.distributed.new_group(
[torch.distributed.get_rank()], use_local_synchronization=True
)
self.reference_pg_collection = ProcessGroupCollection(
tp=self.reference_group,
expt_tp=self.reference_group,
cp=self.reference_group,
pp=self.reference_group,
tp_cp=self.reference_group,
tp_dp_cp=self.reference_group,
ep=self.reference_group,
tp_ep=self.reference_group,
expt_dp=self.reference_group,
dp=self.reference_group,
dp_cp=self.reference_group,
embd=None,
pos_embd=None,
)
model_parallel_cuda_manual_seed(1234)

def teardown_method(self):
torch.distributed.destroy_process_group(self.reference_group)
Utils.destroy_model_parallel()

def test_build_train_and_step(self):
"""Shard experts over expert-DP and dense parameters over full DP."""
# The in-process EP=1 reference needs rank-invariant initialization. GPU expert
# initialization instead uses the globally configured EP=2 rank in its RNG seed.
config = TransformerConfig(
num_layers=2,
hidden_size=64,
num_attention_heads=4,
num_moe_experts=4,
expert_model_parallel_size=2,
moe_layer_freq=[0, 1],
moe_token_dispatcher_type="alltoall",
moe_router_topk=2,
moe_grouped_gemm=True,
moe_ffn_hidden_size=128,
add_bias_linear=False,
use_cpu_initialization=True,
params_dtype=torch.float32,
attention_dropout=0.0,
hidden_dropout=0.0,
gradient_accumulation_fusion=False,
attention_backend=AttnBackend.unfused,
)
# Pair CPU initialization with an explicit common seed for the reference and EP model.
torch.manual_seed(123)
reference_config = replace(config, expert_model_parallel_size=1)
reference_model = HybridModel(
config=reference_config,
hybrid_stack_spec=hybrid_stack_spec,
vocab_size=128,
max_sequence_length=8,
hybrid_layer_pattern="*E",
pg_collection=self.reference_pg_collection,
).cuda()
model = HybridModel(
config=config,
hybrid_stack_spec=hybrid_stack_spec,
vocab_size=128,
max_sequence_length=8,
hybrid_layer_pattern="*E",
pg_collection=self.pg_collection,
).cuda()
model.load_state_dict(reference_model.state_dict(), strict=False)
for model_layer, reference_layer in zip(
model.decoder.layers, reference_model.decoder.layers
):
if not isinstance(model_layer, MoETransformerLayer):
continue
for fc in ("linear_fc1", "linear_fc2"):
model_fc = getattr(model_layer.mlp.experts, fc)
reference_fc = getattr(reference_layer.mlp.experts, fc)
for local, global_ in enumerate(model_layer.mlp.local_expert_indices):
for parameter_name in ("weight", "bias"):
model_parameter = getattr(model_fc, f"{parameter_name}{local}", None)
reference_parameter = getattr(
reference_fc, f"{parameter_name}{global_}", None
)
if model_parameter is not None:
model_parameter.data.copy_(reference_parameter.data)
reference_model.ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False)
model = FullyShardedDataParallel(
config=config,
ddp_config=DistributedDataParallelConfig(
use_megatron_fsdp=True,
megatron_fsdp_version=2,
use_distributed_optimizer=False,
data_parallel_sharding_strategy="optim_grads_params",
fsdp_all_gather_in_start_param_sync=False,
),
module=model,
pg_collection=self.pg_collection,
)
assert isinstance(model.module, FsdpModule)
assert isinstance(model.module.decoder.layers[1].mlp.experts, FsdpModule)

optimizer_config = OptimizerConfig(
lr=1.0e-3, weight_decay=0.0, use_distributed_optimizer=False, clip_grad=0.0
)
reference_optimizer = get_megatron_optimizer(
optimizer_config, [reference_model], use_gloo_process_groups=False
)
optimizer = get_megatron_optimizer(
replace(optimizer_config), [model], use_gloo_process_groups=False
)
assert isinstance(optimizer, FullyShardedOptimizer)
optimizer.reload_model_params()

local_batch_size = 2
torch.manual_seed(4321)
input_ids = torch.randint(0, 128, (self.world_size * local_batch_size, 8), device="cuda")
position_ids = torch.arange(8, device="cuda").repeat(self.world_size * local_batch_size, 1)
targets = torch.randn(self.world_size * local_batch_size, 8, 128, device="cuda")
input_slice = slice(
torch.distributed.get_rank() * local_batch_size,
(torch.distributed.get_rank() + 1) * local_batch_size,
)
reference_losses = []
for _ in range(5):
reference_optimizer.zero_grad(set_to_none=True)
reference_loss = torch.nn.functional.mse_loss(
reference_model(
input_ids=input_ids, position_ids=position_ids, attention_mask=None
),
targets,
)
reference_loss.backward()
reference_success, _, _ = reference_optimizer.step()
assert reference_success
reference_losses.append(reference_loss.detach())

losses = []
for _ in range(5):
optimizer.zero_grad(set_to_none=True)
loss = torch.nn.functional.mse_loss(
model(
input_ids=input_ids[input_slice],
position_ids=position_ids[input_slice],
attention_mask=None,
),
targets[input_slice],
)
loss.backward()
success, _, _ = optimizer.step()
assert success
loss = loss.detach()
torch.distributed.all_reduce(loss, op=torch.distributed.ReduceOp.AVG)
losses.append(loss)

losses = torch.stack(losses)
reference_losses = torch.stack(reference_losses)
if torch.distributed.get_rank() == 0:
logger.info("MFSDP v2 EP loss curve: %s", losses.tolist())
logger.info("MFSDP v2 EP reference loss curve: %s", reference_losses.tolist())
assert torch.isfinite(losses).all()
assert torch.isfinite(reference_losses).all()
assert losses[-1] < losses[0]
torch.testing.assert_close(losses, reference_losses)
Loading