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
16 changes: 6 additions & 10 deletions megatron/training/models/dist_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from typing import Any, Callable

import torch
from megatron.core import tensor_parallel, mpu
from megatron.core import tensor_parallel
from megatron.core.distributed import (
DistributedDataParallel,
DistributedDataParallelConfig,
Expand Down Expand Up @@ -235,9 +235,7 @@ def _ddp_wrap(
# ring-reduce implementations are large enough to remain bandwidth-bound rather than
# latency-bound.
if ddp_config.bucket_size is None:
ddp_config.bucket_size = max(
40000000, 1000000 * mpu.get_data_parallel_world_size(with_context_parallel=True)
)
ddp_config.bucket_size = max(40000000, 1000000 * pg_collection.dp_cp.size())
# Set bucket_size to infinity if overlap_grad_reduce is False.
if not ddp_config.overlap_grad_reduce:
ddp_config.bucket_size = None
Expand All @@ -249,7 +247,7 @@ def _ddp_wrap(
ddp_stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(ddp_stream):
dp_init_kwargs = {}
if use_megatron_fsdp:
if not use_torch_fsdp2:
dp_init_kwargs["pg_collection"] = pg_collection

wrapped_model = []
Expand All @@ -266,7 +264,7 @@ def _ddp_wrap(
all_params = [
p for p in model_chunk.parameters() if p.requires_grad
]
pp_rank = mpu.get_pipeline_model_parallel_rank()
pp_rank = pg_collection.pp.rank()
effective_bucket_size = (
None
if disable_bucketing or pp_rank > 0
Expand All @@ -276,11 +274,9 @@ def _ddp_wrap(
DistributedOptimizer.compute_full_param_layout(
all_params,
effective_bucket_size,
mpu.get_data_parallel_world_size(with_context_parallel=True),
pg_collection.dp_cp.size(),
ddp_config,
expert_data_parallel_world_size=(
mpu.get_expert_data_parallel_world_size()
),
expert_data_parallel_world_size=pg_collection.expt_dp.size(),
)
)

Expand Down
39 changes: 10 additions & 29 deletions tests/unit_tests/training/models/test_dist_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,13 +22,15 @@


def _make_pg():
"""Mock ProcessGroupCollection with dp, cp, tp, pp sub-groups."""
"""Mock ProcessGroupCollection with dp, cp, tp, pp, dp_cp, expt_dp sub-groups."""
pg = Mock()
pg.dp.rank.return_value = 0
pg.cp.rank.return_value = 0
pg.tp.rank.return_value = 0
pg.pp.rank.return_value = 0
pg.pp.size.return_value = 1
pg.dp_cp.size.return_value = 1
pg.expt_dp.size.return_value = 1
return pg


Expand Down Expand Up @@ -314,15 +316,6 @@ def setup_method(self):
self.ddp_config.overlap_grad_reduce = True
self.ddp_config.use_distributed_optimizer = False
self.model = [_make_model_module(), _make_model_module()]
# Patch mpu so the default-bucket-size computation has integer return values.
self._mpu_patcher = patch("megatron.training.models.dist_utils.mpu")
mpu_mock = self._mpu_patcher.start()
mpu_mock.get_data_parallel_world_size.return_value = 1
mpu_mock.get_pipeline_model_parallel_rank.return_value = 0
mpu_mock.get_expert_data_parallel_world_size.return_value = 1

def teardown_method(self):
self._mpu_patcher.stop()

@patch("megatron.training.models.dist_utils.TorchFullyShardedDataParallel")
@patch("megatron.training.models.dist_utils.FullyShardedDataParallel")
Expand Down Expand Up @@ -538,14 +531,6 @@ class TestDdpWrapBucketSize:
def setup_method(self):
self.pg = _make_pg()
self.model = [_make_model_module()]
self._mpu_patcher = patch("megatron.training.models.dist_utils.mpu")
self._mpu = self._mpu_patcher.start()
self._mpu.get_data_parallel_world_size.return_value = 1
self._mpu.get_pipeline_model_parallel_rank.return_value = 0
self._mpu.get_expert_data_parallel_world_size.return_value = 1

def teardown_method(self):
self._mpu_patcher.stop()

def _ddp_config(self, **overrides):
cfg = Mock()
Expand Down Expand Up @@ -579,8 +564,8 @@ def test_num_buckets_divides_total_param_count(self, *_):
@patch("torch.cuda.Stream")
def test_bucket_size_default_uses_dp_world_size(self, *_):
ddp_config = self._ddp_config()
# dp world size 100 → 1_000_000 * 100 = 100M, bigger than 40M floor
self._mpu.get_data_parallel_world_size.return_value = 100
# dp_cp size 100 → 1_000_000 * 100 = 100M, bigger than 40M floor
self.pg.dp_cp.size.return_value = 100
_ddp_wrap(self.model, False, ddp_config, False, pg_collection=self.pg)
assert ddp_config.bucket_size == 100_000_000

Expand All @@ -591,8 +576,8 @@ def test_bucket_size_default_uses_dp_world_size(self, *_):
@patch("torch.cuda.Stream")
def test_bucket_size_default_minimum_floor(self, *_):
ddp_config = self._ddp_config()
# dp world size 1 → 1M, floor at 40M wins
self._mpu.get_data_parallel_world_size.return_value = 1
# dp_cp size 1 → 1M, floor at 40M wins
self.pg.dp_cp.size.return_value = 1
_ddp_wrap(self.model, False, ddp_config, False, pg_collection=self.pg)
assert ddp_config.bucket_size == 40_000_000

Expand Down Expand Up @@ -628,17 +613,13 @@ class TestDdpWrapFullParamLayout:

def setup_method(self):
self.pg = _make_pg()
self._mpu_patcher = patch("megatron.training.models.dist_utils.mpu")
self._mpu = self._mpu_patcher.start()
self._mpu.get_data_parallel_world_size.return_value = 4
self._mpu.get_pipeline_model_parallel_rank.return_value = 0
self._mpu.get_expert_data_parallel_world_size.return_value = 2
self.pg.dp_cp.size.return_value = 4
self.pg.expt_dp.size.return_value = 2
self._opt_patcher = patch("megatron.training.models.dist_utils.DistributedOptimizer")
self._opt = self._opt_patcher.start()
self._opt.compute_full_param_layout.return_value = "LAYOUT"

def teardown_method(self):
self._mpu_patcher.stop()
self._opt_patcher.stop()

def _ddp_config(self, **overrides):
Expand Down Expand Up @@ -777,7 +758,7 @@ def test_effective_bucket_size_none_when_pp_rank_nonzero(
):
mock_ctx.return_value.__enter__ = Mock(return_value=None)
mock_ctx.return_value.__exit__ = Mock(return_value=False)
self._mpu.get_pipeline_model_parallel_rank.return_value = 1
self.pg.pp.rank.return_value = 1
chunk, _ = self._make_chunk_with_params()
ddp_config = self._ddp_config()
_ddp_wrap([chunk], False, ddp_config, False, pg_collection=self.pg)
Expand Down
Loading