Skip to content
36 changes: 36 additions & 0 deletions src/megatron/bridge/training/comm_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

from megatron.core.distributed import DistributedDataParallelConfig
from megatron.core.optimizer import OptimizerConfig
from megatron.core.transformer.enums import CudaGraphScope
from megatron.core.utils import get_te_version, is_te_min_version, is_torch_min_version

from megatron.bridge.models import GPTModelProvider, T5ModelProvider
Expand Down Expand Up @@ -521,6 +522,41 @@ def _get_model_comm_overlap_cfgs(
"delay_wgrad_compute is not supported with legacy groupedgemm implementation"
)

# CUDA graph scope-specific validations for delayed wgrad.
cuda_graph_scope = getattr(model_cfg, "cuda_graph_scope", None)
if cuda_graph_scope is None or cuda_graph_scope == "full":
cuda_graph_scope = []
elif isinstance(cuda_graph_scope, (str, CudaGraphScope)):
cuda_graph_scope = [cuda_graph_scope]
attn_scope_enabled = (
CudaGraphScope.attn in cuda_graph_scope
or CudaGraphScope.attn.value in cuda_graph_scope
or f"CudaGraphScope.{CudaGraphScope.attn.value}" in cuda_graph_scope
)
moe_router_scope_enabled = (
CudaGraphScope.moe_router in cuda_graph_scope
or CudaGraphScope.moe_router.value in cuda_graph_scope
or f"CudaGraphScope.{CudaGraphScope.moe_router.value}" in cuda_graph_scope
)
wgrad_in_graph_scope = attn_scope_enabled or (
moe_router_scope_enabled
and getattr(model_cfg, "moe_shared_expert_intermediate_size", None) is not None
and not getattr(model_cfg, "moe_shared_expert_overlap", False)
)
if wgrad_in_graph_scope:
assert is_te_min_version("2.12.0"), (
"CUDA graph with delay_wgrad_compute requires TE version >= 2.12.0."
)
assert model_cfg.gradient_accumulation_fusion, (
"CUDA graph with delay_wgrad_compute requires gradient_accumulation_fusion "
"to be enabled. This is because default gradient accumulation does not use "
"static memory addresses, which breaks CUDA graph requirements."
)
if attn_scope_enabled:
assert not model_cfg.add_bias_linear and not model_cfg.add_qkv_bias, (
"CUDA graph with delay_wgrad_compute does not support attention bias for now."
)

comm_overlap_cfg = self._override_user_cfgs(comm_overlap_cfg)
return comm_overlap_cfg

Expand Down
41 changes: 37 additions & 4 deletions src/megatron/bridge/training/utils/flop_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import importlib

import torch.nn.functional as F

from megatron.bridge.training.config import ConfigContainer
Expand All @@ -28,9 +30,23 @@ def calculate_layer_counts():
"""Calculate the number of attention, Mamba, MLP, and MoE layers."""
if hasattr(cfg.model, "hybrid_override_pattern") and cfg.model.hybrid_override_pattern:
counts = {"M": 0, "*": 0, "-": 0, "E": 0}
for layer_type in cfg.model.hybrid_override_pattern:
if layer_type in counts:
counts[layer_type] += 1
try:
parse_hybrid_pattern = importlib.import_module(
"megatron.core.ssm.mamba_hybrid_layer_allocation"
).parse_hybrid_pattern
parsed = parse_hybrid_pattern(cfg.model.hybrid_override_pattern)
if parsed.main_pattern:
for layer_type in parsed.main_pattern:
if layer_type in counts:
counts[layer_type] += 1
if parsed.mtp_pattern and parsed.mtp_num_depths > 0:
for layer_type in parsed.mtp_pattern:
if layer_type in counts:
counts[layer_type] += parsed.mtp_num_depths
except (ImportError, ModuleNotFoundError):
for layer_type in cfg.model.hybrid_override_pattern:
if layer_type in counts:
counts[layer_type] += 1
return counts["*"], counts["M"], counts["-"], counts["E"]
else:
num_attn_layers = round(cfg.model.num_layers * getattr(cfg.model, "hybrid_attention_ratio", 0))
Expand Down Expand Up @@ -135,6 +151,7 @@ def hybrid_flops(
shared_expert_ffn_hidden_size=2048,
num_experts_routed_to=1,
vocab_size=256000,
mtp_num_layers=0,
):
"""Calculate total FLOPs for the hybrid model."""
flops_fwd = (
Expand Down Expand Up @@ -169,7 +186,7 @@ def hybrid_flops(
moe_latent_size,
swiglu,
)
+ (2 * batch_size * seq_len * hidden_size * vocab_size) # logits computation
+ (2 * batch_size * seq_len * hidden_size * vocab_size * (1 + mtp_num_layers)) # logits computation
)
return flops_fwd * 3

Expand Down Expand Up @@ -363,6 +380,21 @@ def transformer_flops():
if getattr(cfg.model, "is_hybrid_model", False):
# Calculate the number of each type of layer.
num_attn_layers, num_mamba_layers, num_mlp_layers, num_moe_layers = calculate_layer_counts()
mtp_num_layers = getattr(cfg.model, "mtp_num_layers", None)
if mtp_num_layers is None:
# When using unified hybrid patterns, infer MTP depth count from the pattern.
hybrid_pattern = getattr(cfg.model, "hybrid_override_pattern", None)
if hybrid_pattern:
try:
parse_hybrid_pattern = importlib.import_module(
"megatron.core.ssm.mamba_hybrid_layer_allocation"
).parse_hybrid_pattern
parsed = parse_hybrid_pattern(hybrid_pattern)
mtp_num_layers = parsed.mtp_num_depths if parsed.mtp_pattern else 0
except (ImportError, ModuleNotFoundError):
mtp_num_layers = 0
else:
mtp_num_layers = 0
padded_vocab_size = calculate_padded_vocab_size(
cfg.model.vocab_size,
cfg.model.make_vocab_size_divisible_by,
Expand Down Expand Up @@ -404,6 +436,7 @@ def transformer_flops():
),
num_experts_routed_to=getattr(cfg.model, "moe_router_topk", 1),
vocab_size=padded_vocab_size,
mtp_num_layers=mtp_num_layers,
)
else:
# Compute standard Transformer model FLOPs.
Expand Down
105 changes: 105 additions & 0 deletions tests/unit_tests/training/test_comm_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,9 @@
import os
from unittest.mock import MagicMock, patch

import pytest
from megatron.core.transformer.enums import CudaGraphScope

from megatron.bridge.models.gpt_provider import GPTModelProvider
from megatron.bridge.models.t5_provider import T5ModelProvider
from megatron.bridge.training.comm_overlap import (
Expand Down Expand Up @@ -608,3 +611,105 @@ def test_delay_wgrad_requires_ep_overlap(self):
assert False, "Expected AssertionError when EP overlap is not enabled"
except AssertionError:
pass

def test_delay_wgrad_cuda_graph_attn_requires_grad_accum_fusion(self):
"""CUDA graph attn scope with delay_wgrad_compute requires gradient_accumulation_fusion."""
comm_cfg = CommOverlapConfig(
tp_comm_overlap=False,
data_parallel_size=1,
delay_wgrad_compute=True,
overlap_moe_expert_parallel_comm=True,
)
comm_cfg.finalize()

model_cfg = create_gpt_config(
tensor_model_parallel_size=1,
pipeline_model_parallel_size=1,
virtual_pipeline_model_parallel_size=None,
sequence_parallel=False,
expert_model_parallel_size=2,
num_moe_experts=2,
moe_token_dispatcher_type="alltoall",
bf16=True,
moe_use_legacy_grouped_gemm=False,
add_bias_linear=False,
add_qkv_bias=False,
gradient_accumulation_fusion=False,
cuda_graph_scope=[CudaGraphScope.attn],
)
ddp_cfg = DistributedDataParallelConfig(use_distributed_optimizer=False)

with (
patch("megatron.bridge.training.comm_overlap.is_torch_min_version", return_value=True),
patch("megatron.bridge.training.comm_overlap.is_te_min_version", return_value=True),
pytest.raises(AssertionError, match="gradient_accumulation_fusion"),
):
comm_cfg._get_model_comm_overlap_cfgs(model_cfg, ddp_cfg)

def test_delay_wgrad_cuda_graph_attn_rejects_attention_bias(self):
"""CUDA graph attn scope with delay_wgrad_compute rejects attention bias."""
comm_cfg = CommOverlapConfig(
tp_comm_overlap=False,
data_parallel_size=1,
delay_wgrad_compute=True,
overlap_moe_expert_parallel_comm=True,
)
comm_cfg.finalize()

model_cfg = create_gpt_config(
tensor_model_parallel_size=1,
pipeline_model_parallel_size=1,
virtual_pipeline_model_parallel_size=None,
sequence_parallel=False,
expert_model_parallel_size=2,
num_moe_experts=2,
moe_token_dispatcher_type="alltoall",
bf16=True,
moe_use_legacy_grouped_gemm=False,
add_bias_linear=True,
add_qkv_bias=False,
gradient_accumulation_fusion=True,
cuda_graph_scope=[CudaGraphScope.attn],
)
ddp_cfg = DistributedDataParallelConfig(use_distributed_optimizer=False)

with (
patch("megatron.bridge.training.comm_overlap.is_torch_min_version", return_value=True),
patch("megatron.bridge.training.comm_overlap.is_te_min_version", return_value=True),
pytest.raises(AssertionError, match="attention bias"),
):
comm_cfg._get_model_comm_overlap_cfgs(model_cfg, ddp_cfg)

def test_delay_wgrad_cuda_graph_attn_validation_passes_with_supported_settings(self):
"""CUDA graph attn scope should pass delay_wgrad validation when all constraints are met."""
comm_cfg = CommOverlapConfig(
tp_comm_overlap=False,
data_parallel_size=1,
delay_wgrad_compute=True,
overlap_moe_expert_parallel_comm=True,
)
comm_cfg.finalize()

model_cfg = create_gpt_config(
tensor_model_parallel_size=1,
pipeline_model_parallel_size=1,
virtual_pipeline_model_parallel_size=None,
sequence_parallel=False,
expert_model_parallel_size=2,
num_moe_experts=2,
moe_token_dispatcher_type="alltoall",
bf16=True,
moe_use_legacy_grouped_gemm=False,
add_bias_linear=False,
add_qkv_bias=False,
gradient_accumulation_fusion=True,
cuda_graph_scope=[CudaGraphScope.attn],
)
ddp_cfg = DistributedDataParallelConfig(use_distributed_optimizer=False)

with (
patch("megatron.bridge.training.comm_overlap.is_torch_min_version", return_value=True),
patch("megatron.bridge.training.comm_overlap.is_te_min_version", return_value=True),
):
result = comm_cfg._get_model_comm_overlap_cfgs(model_cfg, ddp_cfg)
assert result.delay_wgrad_compute is True
47 changes: 47 additions & 0 deletions tests/unit_tests/training/utils/test_flop_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,8 @@
"""Unit tests for flop_utils module."""

from dataclasses import dataclass, field
from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import pytest

Expand Down Expand Up @@ -376,3 +378,48 @@ def test_swiglu_scaling_factor(self):
f"Non-SwiGLU: expected {expected_no_swiglu:.2e} but got {flops_no_swiglu:.2e}"
)
assert flops_swiglu == expected_swiglu, f"SwiGLU: expected {expected_swiglu:.2e} but got {flops_swiglu:.2e}"


class TestHybridMtpPatternParsing:
"""Tests for hybrid/MTP pattern parsing in FLOPs accounting."""

def test_inferred_mtp_depth_scales_hybrid_logit_flops(self):
"""When mtp_num_layers is inferred from parsed pattern, logits FLOPs should scale accordingly."""
batch_size = 1
seq_len = 256
hidden_size = 1024
vocab_size = 32000 # divisible by 128, so padded vocab is unchanged.

base_cfg = dict(
is_hybrid_model=True,
hybrid_override_pattern="M*/MM/MM",
num_layers=2,
hidden_size=hidden_size,
seq_length=seq_len,
ffn_hidden_size=4096,
num_attention_heads=8,
num_query_groups=8,
vocab_size=vocab_size,
moe_ffn_hidden_size=2048,
moe_shared_expert_intermediate_size=0,
moe_router_topk=1,
gated_linear_unit=False,
mtp_num_layers=0, # overridden below for inferred-vs-explicit comparison
)

cfg_explicit_zero = MockConfigContainer(model=MockModelConfig(**base_cfg))
cfg_inferred = MockConfigContainer(model=MockModelConfig(**(base_cfg | {"mtp_num_layers": None})))

parsed_pattern = SimpleNamespace(main_pattern="M*", mtp_pattern="MM", mtp_num_depths=2)
mock_module = MagicMock()
mock_module.parse_hybrid_pattern.return_value = parsed_pattern

with patch("megatron.bridge.training.utils.flop_utils.importlib.import_module", return_value=mock_module):
flops_explicit_zero = num_floating_point_operations(cfg_explicit_zero, batch_size=batch_size)
flops_inferred = num_floating_point_operations(cfg_inferred, batch_size=batch_size)

# Only the logits term should differ here:
# delta = 2 * B * S * H * vocab * inferred_mtp_num_layers, then *3 for fwd+bwd factor.
expected_delta = 2 * batch_size * seq_len * hidden_size * vocab_size * 2 * 3
actual_delta = flops_inferred - flops_explicit_zero
assert actual_delta == expected_delta, f"Expected logits delta {expected_delta:.2e} but got {actual_delta:.2e}"
Loading