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
62 changes: 56 additions & 6 deletions megatron/core/extensions/transformer_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import io
import os
import pickle
import re
import warnings
from contextlib import nullcontext
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple
Expand Down Expand Up @@ -77,6 +78,31 @@
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,
partition_stride: int = 1,
) -> None:
"""Route expert gradients and restore TP metadata hidden from TE."""
for name, param in module.named_parameters(recurse=False):
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 = partition_stride


class TransformerEngineConfigType(enum.Enum):
Expand Down Expand Up @@ -588,6 +614,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 @@ -640,10 +670,10 @@ def __init__(

for param in self.parameters():
setattr(param, "parallel_mode", parallel_mode)
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 @@ -1014,6 +1044,15 @@ 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, partition_stride=stride
)

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 @@ -1114,6 +1153,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 @@ -1560,6 +1606,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 All @@ -1575,6 +1625,7 @@ def __init__(
tp_group_for_te = tp_group

self.explicit_expert_comm = is_expert and (tp_size > 1 or self.expert_parallel)
original_parallel_mode = parallel_mode

if self.explicit_expert_comm:
if parallel_mode == "column":
Expand Down Expand Up @@ -1603,8 +1654,7 @@ def __init__(
**extra_kwargs,
)
self.te_quant_params: Optional[TEQuantizationParams] = None
for param in self.parameters():
setattr(param, "allreduce", not (is_expert and self.expert_parallel))
_set_expert_parameter_attributes(self, original_parallel_mode, use_expert_pgs)

def merge_extra_states(
self,
Expand Down
7 changes: 6 additions & 1 deletion megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -467,10 +467,15 @@ def init_state_fn(opt, config=None):

if pg_collection is None or not hasattr(pg_collection, 'tp'):
tp_group = parallel_state.get_tensor_model_parallel_group()
expert_tp_group = (
parallel_state.get_expert_tensor_parallel_group(check_initialized=False) or tp_group
)
else:
tp_group = pg_collection.tp
# TODO(M4): plumb tp_group through optimizer constructors so this setattr disappears.
expert_tp_group = getattr(pg_collection, 'expt_tp', None) or 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
5 changes: 4 additions & 1 deletion megatron/core/optimizer/clip_grads.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,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 zeros in gradients associated with the passed-in list of
parameters.
Expand Down Expand Up @@ -219,7 +220,9 @@ def count_zeros_fp32(
grad_attr = "decoupled_grad" if use_decoupled_grad else "grad"
grad_not_none = hasattr(param, grad_attr) and getattr(param, grad_attr) is not None
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
)
if grad_not_none and is_not_shared and is_not_tp_duplicate:
grad_obj = getattr(param, grad_attr)
data_parallel_group = get_data_parallel_group_if_dtensor(grad_obj, data_parallel_group)
Expand Down
66 changes: 41 additions & 25 deletions megatron/core/optimizer/layer_wise_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,14 @@ def __init__(
opt, config, None, init_state_fn_list[i] if init_state_fn_list else None
)

self.tp_group = self.pg_collection.tp
self.expert_tp_group = getattr(self.pg_collection, 'expt_tp', None) or self.tp_group
for optimizer in optimizers:
# Child optimizers filter TP replicas when collecting gradients for
# the world-reduced LayerWise gradient statistics.
optimizer.tp_group = self.tp_group
optimizer.expert_tp_group = self.expert_tp_group

super().__init__(optimizers)

# TODO(kunlun, deyuf): potential future perf optimization
Expand All @@ -95,15 +103,17 @@ def shard_params(self, optimizers):
# example of 4 ranks and 10 parameters p0-p9 after sorting, then dp_cp_params_list will be
# [[p0, p7, p8], [p1, p6, p9], [p2, p5], [p3, p4]]

# simplify when dp_cp group size is 1
if get_pg_size(self.pg_collection.dp_cp) == 1:
dp_cp_size = get_pg_size(self.pg_collection.dp_cp)
expt_dp_size = get_pg_size(self.pg_collection.expt_dp)

# Dense and expert parameters use independent data-parallel ownership
# domains. Only skip sharding when neither domain has replicas.
if dp_cp_size == 1 and expt_dp_size == 1:
self.dp_cp_params_list = None
self.expt_dp_params_list = None
return

dp_cp_idx, expt_dp_idx = 0, 0
dp_cp_size = get_pg_size(self.pg_collection.dp_cp)
expt_dp_size = get_pg_size(self.pg_collection.expt_dp)
# create ping-pong style loop so memory is more balanced
dp_cp_loop = list(range(dp_cp_size)) + list(range(dp_cp_size))[::-1]
expt_dp_loop = list(range(expt_dp_size)) + list(range(expt_dp_size))[::-1]
Expand Down Expand Up @@ -139,8 +149,11 @@ def shard_params(self, optimizers):
for groups, params in zip(param_groups, param_groups_this_rank):
groups["params"] = params

# simplify when expt_dp group size is 1 or expert parallel is off
if expt_dp_size == 1 or len(self.expt_dp_params_list[0]) == 0:
# Dense and expert all-gathers are independent. A singleton ownership
# domain or a domain with no parameters needs no synchronization.
if dp_cp_size == 1 or not any(self.dp_cp_params_list):
self.dp_cp_params_list = None
if expt_dp_size == 1 or not any(self.expt_dp_params_list):
self.expt_dp_params_list = None

@torch.no_grad()
Expand All @@ -149,20 +162,23 @@ def allgather_params(self) -> None:

# helper function to flatten local params, allgather, unflatten and copy to model params
def _allgather_helper(params_list, group):
# flatten this rank's params and create empty tensor output list
device = params_list[0][0].device
dtype = params_list[0][0].dtype
flat_sizes = [sum(p.numel() for p in params) for params in params_list]
if not any(flat_sizes):
return

# The first ownership shard may legitimately be empty (for example,
# a pure-expert optimizer with no dense parameters).
prototype = next(p for params in params_list for p in params)
device = prototype.device
dtype = prototype.dtype
rank = get_pg_rank(group)
# for rank without params create empty tensor and participate in allgather
src = (
_flatten_dense_tensors(params_list[rank])
if len(params_list[rank]) > 0
else torch.empty(0, device=device, dtype=dtype)
)
output_list = [
torch.empty(sum([p.numel() for p in params]), device=device, dtype=dtype)
for params in params_list
]
output_list = [torch.empty(size, device=device, dtype=dtype) for size in flat_sizes]
# single all_gather_v to collect all updated params
torch.distributed.all_gather(output_list, src, group=group)
# unflatten and copy gathered params for each rank i
Expand All @@ -185,18 +201,16 @@ def _allgather_helper(params_list, group):
def broadcast_params(self):
"""All rank broadcast updated local params."""
# Broadcast linear layer weights to all other ranks. Kept as reference test.
if self.dp_cp_params_list is None:
return
for i, params in enumerate(self.dp_cp_params_list):
src_global_rank = torch.distributed.get_global_rank(self.pg_collection.dp_cp, i)
for p in params:
torch.distributed.broadcast(p, src_global_rank, self.pg_collection.dp_cp)
if self.expt_dp_params_list is None:
return
for i, params in enumerate(self.expt_dp_params_list):
src_global_rank = torch.distributed.get_global_rank(self.pg_collection.expt_dp, i)
for p in params:
torch.distributed.broadcast(p, src_global_rank, self.pg_collection.expt_dp)
if self.dp_cp_params_list is not None:
for i, params in enumerate(self.dp_cp_params_list):
src_global_rank = torch.distributed.get_global_rank(self.pg_collection.dp_cp, i)
for p in params:
torch.distributed.broadcast(p, src_global_rank, self.pg_collection.dp_cp)
if self.expt_dp_params_list is not None:
for i, params in enumerate(self.expt_dp_params_list):
src_global_rank = torch.distributed.get_global_rank(self.pg_collection.expt_dp, i)
for p in params:
torch.distributed.broadcast(p, src_global_rank, self.pg_collection.expt_dp)

@torch.no_grad()
def get_grad_norm(self):
Expand All @@ -216,6 +230,8 @@ def count_zeros(self):
params,
grad_stats_parallel_group=None,
use_decoupled_grad=self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8,
tp_group=self.tp_group,
expert_tp_group=self.expert_tp_group,
)

@torch.no_grad()
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 @@ -154,7 +154,9 @@ def get_main_grads_for_grad_norm(self) -> List[torch.Tensor]:
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_witness = not getattr(param, "_is_witness_param", False)
if grad_not_none and is_not_shared and is_not_tp_duplicate and is_not_witness:
Expand Down Expand Up @@ -229,6 +231,7 @@ def count_zeros(self) -> float:
grad_stats_parallel_group=self.get_grad_stats_parallel_group(),
use_decoupled_grad=self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8,
tp_group=getattr(self, 'tp_group', None),
expert_tp_group=getattr(self, 'expert_tp_group', None),
)

@abstractmethod
Expand Down Expand Up @@ -675,6 +678,8 @@ def __init__(
main_param = param.detach().clone().float()
# Copy tensor model parallel attributes.
tensor_parallel.copy_tensor_model_parallel_attributes(main_param, param)
if hasattr(param, 'allreduce'):
main_param.allreduce = param.allreduce
if hasattr(param, 'shared'):
main_param.shared = param.shared
# Replace the optimizer params with the new fp32 copy.
Expand Down Expand Up @@ -1296,6 +1301,8 @@ def count_zeros(self):
params,
grad_stats_parallel_group=self.get_grad_stats_parallel_group(),
use_decoupled_grad=self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8,
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 @@ -145,7 +145,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
Loading