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
21 changes: 20 additions & 1 deletion megatron/core/distributed/fsdp/mcore_fsdp_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

import logging
import random
from typing import Dict, List, Optional
from typing import Dict, List, Optional, Tuple, Type

try:
import einops
Expand Down Expand Up @@ -91,6 +91,22 @@ class FullyShardedDataParallel(_BaseDataParallel):
},
}

@staticmethod
def _fine_grained_recurse_module_types(
config: TransformerConfig, ddp_config: DistributedDataParallelConfig
) -> Tuple[Type[nn.Module], ...]:
"""Module classes needing ``parameters(recurse=True)`` for fine-grained hooks."""
if (
config.overlap_moe_expert_parallel_comm
and ddp_config.data_parallel_sharding_strategy == "optim_grads_params"
):
# Lazy import to avoid circular chain.
from megatron.core.transformer.moe.experts import TEGroupedMLP
from megatron.core.transformer.moe.shared_experts import SharedExpertMLP

return (TEGroupedMLP, SharedExpertMLP)
return ()

def __init__(
self,
config: TransformerConfig,
Expand Down Expand Up @@ -211,6 +227,9 @@ def __init__(
config.overlap_moe_expert_parallel_comm
and ddp_config.data_parallel_sharding_strategy == "optim_grads_params"
),
fine_grained_recurse_module_types=self._fine_grained_recurse_module_types(
config, ddp_config
),
),
)
self.param_and_grad_buffer = self.module.param_and_grad_buffer
Expand Down
69 changes: 54 additions & 15 deletions megatron/core/distributed/fsdp/src/megatron_fsdp/megatron_fsdp.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
from contextlib import contextmanager
from enum import Enum, auto
from functools import partial
from typing import Any, Dict, List, Optional, Tuple
from typing import Any, Dict, List, Literal, Optional, Tuple, Type

import torch
import torch.nn as nn
Expand Down Expand Up @@ -177,6 +177,12 @@ class MegatronFSDP(torch.nn.Module):
userbuffer registration when nccl_ub is set.
enable_fine_grained_param_gather (bool): Whether to enable "fine-grained" param all-gather,
which can improve performance when using MXFP8 parameters with activation recomputation.
enable_fine_grained_param_gather_backward_hook (bool): Register pre-backward unshard hooks

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Note to myself: this PR doesn't add this flag; it merely adds a missing comment.

on each submodule (used by 1F1B EP overlap and similar schedules).
fine_grained_recurse_module_types (Optional[Tuple[Type[nn.Module], ...]]):
Module classes for which fine-grained pre-forward / pre-backward unshard uses
``parameters(recurse=True)`` (container modules whose sharded weights live on
children). Checked with :func:`isinstance`. Defaults to empty (none).
report_nan_in_param_grad (bool): Whether to enable precise NaN-checking for parameter wgrad.
Can significantly degrade performance. Defaults to False.

Expand Down Expand Up @@ -217,6 +223,7 @@ def __init__(
disable_symmetric_registration: bool = False,
enable_fine_grained_param_gather_hook: bool = False,
enable_fine_grained_param_gather_backward_hook: bool = False,
fine_grained_recurse_module_types: Optional[Tuple[Type[nn.Module], ...]] = None,
report_nan_in_param_grad: bool = False,
):
super().__init__()
Expand Down Expand Up @@ -272,6 +279,8 @@ def __init__(
self.enable_fine_grained_param_gather_backward_hook = (
enable_fine_grained_param_gather_backward_hook
)
recurse_types = fine_grained_recurse_module_types or ()
self.fine_grained_recurse_module_types: Tuple[Type[nn.Module], ...] = recurse_types
self.report_nan_in_param_grad = report_nan_in_param_grad

# FSDPDistributedIndex stores the process groups and meshes used by Megatron-FSDP.
Expand Down Expand Up @@ -546,6 +555,48 @@ def _register_fsdp_hooks(self, root_module):
"""
fsdp_unit_modules = self.fsdp_unit_modules

def _param_list_for_submodule_unshard(
module: nn.Module, pass_direction: Literal["forward", "backward"]
) -> List[nn.Parameter]:
"""Build the parameter list for fine-grained or FSDP-unit unshard hooks.

Parameter buckets designated by this function are all-gathered and may
pre-fetch subsequent buckets in FSDP bucket order during runtime.
"""
# Fine-grained hooks are attached to all sub-modules; this function
# controls which parameters each hook should unshard.
fine_grained_enabled = (
self.enable_fine_grained_param_gather_backward_hook
if pass_direction == "backward"
else self.enable_fine_grained_param_gather_hook
)
if fine_grained_enabled:
# Fine-grained hooks run on every submodule: shallow params by
# default, including on FSDP units (e.g. TransformerLayer). Leaf
# child hooks gather their own nested weights. Container modules
# in fine_grained_recurse_module_types (e.g. TEGroupedMLP,
# SharedExpertMLP) need recurse=True because weights live on
# children and the container is the compute entry point.
if self.fine_grained_recurse_module_types and isinstance(
module, self.fine_grained_recurse_module_types
):
return list(module.parameters(recurse=True))
else:
# Only unshard direct parameters. Used when submodules are
# called in isolation of an FSDP-unit forward (e.g. mxfp8
# param gather, EP-overlap 1F1B schedule). Leaf modules
# (e.g. TELinear) still gather their own weights via
# separate hooks. Also limits unshard scope for activation
# recomputation on individual submodules.
return list(module.parameters(recurse=False))
else:
if isinstance(module, tuple(fsdp_unit_modules)):
# FSDP unit modules should be unsharded and communicated together.
return list(module.parameters())
else:
# Non-unit modules should only unshard the direct parameters they need.
return list(module.parameters(recurse=False))

def release_module_parameters(module, bwd, lazy=False, *unused):
"""
Release the parameters of a given module after completing the forward
Expand Down Expand Up @@ -736,16 +787,7 @@ def _pre_forward_param_unshard(module: nn.Module, *unused):
else:
module._training_state = TrainingState.FORWARD

if isinstance(module, tuple(fsdp_unit_modules)):
param_list = list(module.parameters())
else:
# All-gather the shallow parameters in every forward pass for modules
# that are not FSDP units. Do not recurse unless absolutely necessary,
# to allocate as little memory as possible for this forward pass.
param_list = list(module.parameters(recurse=False))

if self.enable_fine_grained_param_gather_hook:
param_list = list(module.parameters(recurse=False))
param_list = _param_list_for_submodule_unshard(module, "forward")

# All-gather the parameters before the forward pass.
self.all_gather_and_wait_parameters_ready(
Expand Down Expand Up @@ -861,10 +903,7 @@ def _pre_backward_param_unshard(module: nn.Module, *unused):
for sub_module in module.modules():
sub_module._training_state = TrainingState.PRE_BACKWARD

if isinstance(module, tuple(fsdp_unit_modules)):
param_list = list(module.parameters())
else:
param_list = list(module.parameters(recurse=False))
param_list = _param_list_for_submodule_unshard(module, "backward")

# All-gather / unshard the module parameters before the backward pass.
self.all_gather_and_wait_parameters_ready(
Expand Down
Loading