Skip to content
Merged
5 changes: 3 additions & 2 deletions megatron/core/distributed/param_and_grad_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -887,8 +887,9 @@ def group_params_for_buffers(
Each distinct buffer is identified by a BufferKey with three dimensions:
- param_dtype: storage dtype (torch.uint8 for FP8/NVFP4 parameters, else param.dtype).
- grad_dtype: gradient reduction dtype (torch.float if grad_reduce_in_fp32, else param.dtype).
- is_expert_parallel: whether the parameter is expert-parallel (param.allreduce == False),
which requires a separate buffer with a different data-parallel group.
- is_expert_parallel: whether the parameter uses the expert topology (param.allreduce == False),
which requires a separate buffer for the expert data-parallel group. This is true for experts
when expert-parallelism > 1 or expert-tensor-parallelism != tensor-parallelism.

The param_indices track each parameter's position among same-dtype params (using
the "fake" high-precision dtype for FP8/NVFP4 params), needed for loading non-native-fp8
Expand Down
86 changes: 64 additions & 22 deletions megatron/core/extensions/transformer_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import io
import os
import pickle
import re
import warnings
from contextlib import contextmanager, nullcontext
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple, cast
Expand Down Expand Up @@ -84,6 +85,42 @@
HAVE_TE = False

_TE_CONFIG_TYPE_KEY = "transformer_engine_config_type"
_EXPERT_PARAMETER_NAME_PATTERN = re.compile(r"(weight|bias)\d*")


def _set_expert_parameter_attributes(
module: torch.nn.Module, parallel_mode: Optional[str], use_expert_pgs: bool
) -> None:
"""Set process-group and tensor-partition metadata on an expert TE module.

``allreduce=False`` selects EDP for gradient reduction.

Weights and biases, including TEGroupedLinear's numbered parameters, are also marked as
TP-partitioned according to ``parallel_mode``; row-parallel biases remain replicated.

Any parameter which is partitioned along TP or ETP is marked with ``tensor_model_parallel``,
which ensures that all shards contribute to the gradient norm.

Args:
module: Transformer Engine module whose direct parameters should be marked.
parallel_mode: Tensor-parallel mode used by the module (``"column"``, ``"row"``, or None).
use_expert_pgs: Whether to use EP/ETP/EDP process groups instead of TP/CP/DP.
"""
for name, param in module.named_parameters(recurse=False):
Comment thread
philipcmonk marked this conversation as resolved.
param.allreduce = not use_expert_pgs

name_match = _EXPERT_PARAMETER_NAME_PATTERN.fullmatch(name)
parameter_kind = name_match.group(1) if name_match else None
is_weight = parameter_kind == "weight"
is_bias = parameter_kind == "bias"
is_partitioned = parallel_mode in ("column", "row") and (
is_weight or (parallel_mode == "column" and is_bias)
)
if is_weight or is_bias:
param.tensor_model_parallel = is_partitioned
if is_partitioned:
param.partition_dim = 1 if parallel_mode == "row" else 0
param.partition_stride = 1


class TransformerEngineConfigType(enum.Enum):
Expand Down Expand Up @@ -942,6 +979,10 @@ def __init__(
tp_size = get_pg_size(tp_group)

self.expert_parallel = self.config.expert_model_parallel_size > 1
use_expert_pgs = is_expert and (
self.expert_parallel
or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size
)
if is_expert:
rng_tracker_name = get_expert_parallel_rng_tracker_name()
else:
Expand Down Expand Up @@ -1007,11 +1048,10 @@ def __init__(
**extra_kwargs,
)

for param in self.parameters():
if is_expert:
# Reduce the gradient on the expert_data_parallel group for expert linear layers
setattr(param, "allreduce", not self.expert_parallel)
else:
if is_expert:
_set_expert_parameter_attributes(self, parallel_mode, use_expert_pgs)
else:
for param in self.parameters():
# Reduce the gradient on DP group
setattr(param, "allreduce", True)
if parallel_mode == "duplicated":
Expand Down Expand Up @@ -1431,6 +1471,13 @@ def __init__(
self.bias.zero_()
setattr(self.bias, "allreduce", True)

if is_expert:
use_expert_pgs = (
config.expert_model_parallel_size > 1
or config.expert_tensor_parallel_size != config.tensor_model_parallel_size
)
_set_expert_parameter_attributes(self, "column", use_expert_pgs)

def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None):
"""Sharding along axis 0, bias sharded"""
state_dict = self.state_dict(prefix="", keep_vars=True)
Expand Down Expand Up @@ -1680,6 +1727,13 @@ def __init__(
setattr(self.bias, "allreduce", True)
setattr(self.bias, "sequence_parallel", config.sequence_parallel)

if is_expert:
use_expert_pgs = (
config.expert_model_parallel_size > 1
or config.expert_tensor_parallel_size != config.tensor_model_parallel_size
)
_set_expert_parameter_attributes(self, "row", use_expert_pgs)

def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None):
"""Sharding along axis 1, bias not sharded"""
state_dict = self.state_dict(prefix="", keep_vars=True)
Expand Down Expand Up @@ -2106,6 +2160,10 @@ def __init__(
extra_kwargs["ub_name"] = tp_comm_buffer_name

self.expert_parallel = self.config.expert_model_parallel_size > 1
use_expert_pgs = is_expert and (
self.expert_parallel
or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size
)
if is_expert:
extra_kwargs["rng_tracker_name"] = get_expert_parallel_rng_tracker_name()

Expand Down Expand Up @@ -2182,23 +2240,7 @@ def __init__(
**extra_kwargs,
)

for param in self.parameters():
setattr(param, "allreduce", not (is_expert and self.expert_parallel))

# Explicitly stamp partition_dim and partition_stride on expert weight
# tensors when explicit_expert_comm cleared parallel_mode. TE ≤2.12
# set these internally; TE ≥2.13 no longer does (parallel_mode=None
# is passed due to explicit_expert_comm). The resharding/refit planner
# relies on partition_dim to correctly plan TP gather/scatter operations.
# NOTE: we intentionally do NOT stamp tensor_model_parallel here —
# doing so would change num-zeros gradient counting.
if self.explicit_expert_comm and original_parallel_mode in ("column", "row"):
part_dim = 0 if original_parallel_mode == "column" else 1
for i in range(num_gemms):
weight = getattr(self, f"weight{i}", None)
if weight is not None:
setattr(weight, "partition_dim", part_dim)
setattr(weight, "partition_stride", 1)
_set_expert_parameter_attributes(self, original_parallel_mode, use_expert_pgs)

self._register_load_state_dict_pre_hook(
type(self)._normalize_grouped_parameter_keys, with_module=True
Expand Down
17 changes: 8 additions & 9 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,6 @@
if HAVE_EMERGING_OPTIMIZERS:
from emerging_optimizers.scalar_optimizers import Lion

from megatron.core import parallel_state
from megatron.core.optimizer.cpu_offloading.hybrid_optimizer import HybridDeviceOptimizer
from megatron.core.optimizer_param_scheduler import (
ParamGroupOverride,
Expand Down Expand Up @@ -686,11 +685,12 @@ def init_state_fn(opt, config=None):
setattr(optimizer, 'grad_stats_parallel_group', model_parallel_group)

if pg_collection is None or not hasattr(pg_collection, 'tp'):
tp_group = parallel_state.get_tensor_model_parallel_group()
else:
tp_group = pg_collection.tp
# TODO(M4): plumb tp_group through optimizer constructors so this setattr disappears.
pg_collection = ProcessGroupCollection.use_mpu_process_groups()
Comment thread
deepakn94 marked this conversation as resolved.
tp_group = pg_collection.tp
expert_tp_group = getattr(pg_collection, 'expt_tp', tp_group)
# TODO(M4): plumb TP groups through optimizer constructors so these setattrs disappear.
setattr(optimizer, 'tp_group', tp_group)
setattr(optimizer, 'expert_tp_group', expert_tp_group)

return optimizer

Expand Down Expand Up @@ -898,11 +898,10 @@ def _get_megatron_emerging_optimizer(
else:
optimizer = FP32Optimizer(optimizer, config, init_state_fn)
setattr(optimizer, 'grad_stats_parallel_group', model_parallel_group)
if pg_collection is None or not hasattr(pg_collection, 'tp'):
tp_group = parallel_state.get_tensor_model_parallel_group()
else:
tp_group = pg_collection.tp
tp_group = pg_collection.tp
expert_tp_group = getattr(pg_collection, 'expt_tp', tp_group)
setattr(optimizer, 'tp_group', tp_group)
setattr(optimizer, 'expert_tp_group', expert_tp_group)
results.append(optimizer)
continue
else:
Expand Down
5 changes: 4 additions & 1 deletion megatron/core/optimizer/clip_grads.py
Original file line number Diff line number Diff line change
Expand Up @@ -197,6 +197,7 @@ def count_zeros_fp32(
grad_stats_parallel_group: torch.distributed.ProcessGroup,
use_decoupled_grad: bool = False,
tp_group: Optional[torch.distributed.ProcessGroup] = None,
expert_tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> float:
"""Counts the number of zero values in the gradients of the given parameters.

Expand Down Expand Up @@ -242,7 +243,9 @@ def count_zeros_fp32(
total_num_zeros += num_zeros
continue
is_not_shared = param_is_not_shared(param)
is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(param, tp_group=tp_group)
is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(
param, tp_group=tp_group, expert_tp_group=expert_tp_group
)
is_not_gtp_duplicate = param_is_not_gtp_duplicate(param)
if grad_not_none and is_not_shared and is_not_tp_duplicate and is_not_gtp_duplicate:
grad_obj = getattr(param, grad_attr)
Expand Down
9 changes: 8 additions & 1 deletion megatron/core/optimizer/optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,8 @@ def _is_separate_grad_norm_group(grad_norm_group: Optional[str]) -> bool:

def copy_optimizer_param_metadata(destination: torch.Tensor, source: torch.Tensor) -> None:
"""Copy optimizer-relevant metadata when creating param views/copies."""
if hasattr(source, 'allreduce'):
destination.allreduce = source.allreduce
if hasattr(source, 'shared'):
destination.shared = source.shared
if hasattr(source, GRAD_NORM_GROUP_ATTR):
Expand Down Expand Up @@ -266,7 +268,9 @@ def _filter_grads_for_norm(
grad_not_none = grad is not None
is_not_shared = param_is_not_shared(param)
is_not_tp_duplicate = tensor_parallel.param_is_not_tensor_parallel_duplicate(
param, getattr(self, 'tp_group', None)
param,
tp_group=getattr(self, 'tp_group', None),
expert_tp_group=getattr(self, 'expert_tp_group', None),
)
is_not_gtp_duplicate = tensor_parallel.param_is_not_gtp_duplicate(param)
if grad_not_none and is_not_shared and is_not_tp_duplicate and is_not_gtp_duplicate:
Expand Down Expand Up @@ -434,6 +438,7 @@ def count_zeros(self) -> float:
and getattr(params[0], "__fsdp_param__", False)
),
tp_group=getattr(self, 'tp_group', None),
expert_tp_group=getattr(self, 'expert_tp_group', None),
)

@abstractmethod
Expand Down Expand Up @@ -1785,6 +1790,8 @@ def count_zeros(self):
self.config.use_precision_aware_optimizer
and getattr(params[0], "__fsdp_param__", False)
),
tp_group=getattr(self.chained_optimizers[0], 'tp_group', None),
expert_tp_group=getattr(self.chained_optimizers[0], 'expert_tp_group', None),
)
else:
num_zeros_in_grad = 0
Expand Down
7 changes: 6 additions & 1 deletion megatron/core/post_training/modelopt/layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,12 @@ def __init__(
for param in self.parameters():
if is_expert:
# Reduce the gradient on the expert_data_parallel group for expert linear layers
setattr(param, "allreduce", self.config.expert_model_parallel_size == 1)
use_expert_groups = (
self.config.expert_model_parallel_size > 1
or self.config.expert_tensor_parallel_size
!= self.config.tensor_model_parallel_size
)
setattr(param, "allreduce", not use_expert_groups)
else:
# Reduce the gradient on DP group
setattr(param, "allreduce", True)
Expand Down
28 changes: 21 additions & 7 deletions megatron/core/tensor_parallel/layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,11 +92,17 @@
dist_reduce_scatter_func = torch.distributed._reduce_scatter_base


def param_is_not_tensor_parallel_duplicate(param, tp_group=None):
"""Returns true if the passed-in parameter is not a duplicate parameter
on another TP rank."""
def param_is_not_tensor_parallel_duplicate(param, tp_group=None, expert_tp_group=None):
"""Return whether a parameter contributes to a unique model-parallel shard.

Parameters reduced over expert data parallel groups use the expert tensor-parallel
group for duplicate filtering. Other parameters use the regular tensor-parallel group.
"""
if hasattr(param, "tensor_model_parallel") and param.tensor_model_parallel:
return True
# allreduce=False marks parameters reduced over expert DP, so filter their duplicates over ETP.
if not getattr(param, "allreduce", True) and expert_tp_group is not None:
Comment thread
philipcmonk marked this conversation as resolved.
tp_group = expert_tp_group
# Prefer provided tp_group when available (new explicit path).
if tp_group is not None:
return tp_group.rank() == 0
Expand Down Expand Up @@ -952,6 +958,10 @@ def __init__(
world_size = get_pg_size(self.tp_group)
rank = get_pg_rank(self.tp_group)
self.explicit_expert_comm = self.is_expert and (world_size > 1 or self.expert_parallel)
use_expert_pgs = self.is_expert and (
self.expert_parallel
or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size
)
self.output_size_per_partition = divide(output_size, world_size)

# Parameters.
Expand Down Expand Up @@ -1004,7 +1014,7 @@ def __init__(
tensor=self.weight, is_parallel=True, dim=0, stride=stride
)

setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel))
setattr(self.weight, "allreduce", not use_expert_pgs)
else:
self.weight = None

Expand Down Expand Up @@ -1037,7 +1047,7 @@ def __init__(
# Always initialize bias to zero.
with torch.no_grad():
self.bias.zero_()
setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel))
setattr(self.bias, "allreduce", not use_expert_pgs)
else:
self.register_parameter("bias", None)

Expand Down Expand Up @@ -1366,7 +1376,11 @@ def __init__(
set_tensor_model_parallel_attributes(
tensor=self.weight, is_parallel=True, dim=1, stride=stride
)
setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel))
use_expert_pgs = self.is_expert and (
self.expert_parallel
or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size
)
setattr(self.weight, "allreduce", not use_expert_pgs)

self.gtp_remat_size = 1
_pg = ProcessGroupCollection.use_mpu_process_groups(
Expand Down Expand Up @@ -1395,7 +1409,7 @@ def __init__(
# Always initialize bias to zero.
with torch.no_grad():
self.bias.zero_()
setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel))
setattr(self.bias, "allreduce", not use_expert_pgs)
setattr(self.bias, "sequence_parallel", self.sequence_parallel)
else:
self.register_parameter("bias", None)
Expand Down
6 changes: 5 additions & 1 deletion megatron/training/utils/common_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,14 +112,18 @@ def calc_params_l2_norm(model, force_create_fp32_copy=False):

gtp_rank = mpu.get_gtp_weight_remat_rank()
egtp_rank = mpu.get_expert_gtp_weight_remat_rank()
tp_group = mpu.get_tensor_model_parallel_group()
expert_tp_group = mpu.get_expert_tensor_parallel_group()

for model_chunk in model:
for param in model_chunk.parameters():
is_gtp = getattr(param, 'is_gtp_weight_remat', False)

# Filter TP duplicates. GTP_remat params are always unique across TP ranks
# so skip this check for them.
if not is_gtp and not param_is_not_tensor_parallel_duplicate(param):
if not is_gtp and not param_is_not_tensor_parallel_duplicate(
param, tp_group=tp_group, expert_tp_group=expert_tp_group
):
continue
is_expert = not getattr(param, 'allreduce', True)

Expand Down
Loading
Loading