chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from .module import PipelineModule, LayerSpec, TiedLayerSpec
|
||||
from .topology import ProcessTopology
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,698 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import os
|
||||
import glob
|
||||
|
||||
import re as regex
|
||||
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from deepspeed import comm as dist
|
||||
|
||||
from deepspeed.utils import logger
|
||||
from .. import utils as ds_utils
|
||||
from ..activation_checkpointing import checkpointing
|
||||
from .topology import PipeDataParallelTopology, PipelineParallelGrid
|
||||
from deepspeed.runtime.state_dict_factory import SDLoaderFactory
|
||||
from deepspeed.accelerator import get_accelerator
|
||||
from deepspeed.checkpoint.utils import clone_tensors_for_torch_save
|
||||
|
||||
|
||||
class PipelineError(Exception):
|
||||
"""Errors related to the use of deepspeed.PipelineModule """
|
||||
|
||||
|
||||
class LayerSpec:
|
||||
"""Building block for specifying pipeline-parallel modules.
|
||||
|
||||
LayerSpec stores the type information and parameters for each stage in a
|
||||
PipelineModule. For example:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
nn.Sequence(
|
||||
torch.nn.Linear(self.in_dim, self.hidden_dim, bias=False),
|
||||
torch.nn.Linear(self.hidden_hidden, self.out_dim)
|
||||
)
|
||||
|
||||
becomes
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
layer_specs = [
|
||||
LayerSpec(torch.nn.Linear, self.in_dim, self.hidden_dim, bias=False),
|
||||
LayerSpec(torch.nn.Linear, self.hidden_hidden, self.out_dim)]
|
||||
]
|
||||
"""
|
||||
|
||||
def __init__(self, typename, *module_args, **module_kwargs):
|
||||
self.typename = typename
|
||||
self.module_args = module_args
|
||||
self.module_kwargs = module_kwargs
|
||||
|
||||
if not issubclass(typename, nn.Module):
|
||||
raise RuntimeError('LayerSpec only supports torch.nn.Module types.')
|
||||
|
||||
if dist.is_initialized():
|
||||
self.global_rank = dist.get_rank()
|
||||
else:
|
||||
self.global_rank = -1
|
||||
|
||||
def __repr__(self):
|
||||
return ds_utils.call_to_str(self.typename.__name__, self.module_args, self.module_kwargs)
|
||||
|
||||
def build(self, log=False):
|
||||
"""Build the stored specification."""
|
||||
if log:
|
||||
logger.info(f'RANK={self.global_rank} building {repr(self)}')
|
||||
|
||||
return self.typename(*self.module_args, **self.module_kwargs)
|
||||
|
||||
|
||||
class TiedLayerSpec(LayerSpec):
|
||||
|
||||
def __init__(self, key, typename, *module_args, forward_fn=None, tied_weight_attr=['weight'], **module_kwargs):
|
||||
super().__init__(typename, *module_args, **module_kwargs)
|
||||
self.key = key
|
||||
self.forward_fn = forward_fn
|
||||
self.tied_weight_attr = [tied_weight_attr] if type(tied_weight_attr) == str else tied_weight_attr
|
||||
|
||||
|
||||
class PipelineModule(nn.Module):
|
||||
"""Modules to be parallelized with pipeline parallelism.
|
||||
|
||||
The key constraint that enables pipeline parallelism is the
|
||||
representation of the forward pass as a sequence of layers
|
||||
and the enforcement of a simple interface between them. The
|
||||
forward pass is implicitly defined by the module ``layers``. The key
|
||||
assumption is that the output of each layer can be directly fed as
|
||||
input to the next, like a ``torch.nn.Sequence``. The forward pass is
|
||||
implicitly:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def forward(self, inputs):
|
||||
x = inputs
|
||||
for layer in self.layers:
|
||||
x = layer(x)
|
||||
return x
|
||||
|
||||
.. note::
|
||||
Pipeline parallelism is not compatible with ZeRO-2 and ZeRO-3.
|
||||
|
||||
Args:
|
||||
layers (Iterable): A sequence of layers defining pipeline structure. Can be a ``torch.nn.Sequential`` module.
|
||||
num_stages (int, optional): The degree of pipeline parallelism. If not specified, ``topology`` must be provided.
|
||||
topology (``deepspeed.runtime.pipe.ProcessTopology``, optional): Defines the axes of parallelism axes for training. Must be provided if ``num_stages`` is ``None``.
|
||||
loss_fn (callable, optional): Loss is computed ``loss = loss_fn(outputs, label)``
|
||||
seed_layers(bool, optional): Use a different seed for each layer. Defaults to False.
|
||||
seed_fn(type, optional): The custom seed generating function. Defaults to random seed generator.
|
||||
base_seed (int, optional): The starting seed. Defaults to 1234.
|
||||
partition_method (str, optional): The method upon which the layers are partitioned. Defaults to 'parameters'.
|
||||
activation_checkpoint_interval (int, optional): The granularity activation checkpointing in terms of number of layers. 0 disables activation checkpointing.
|
||||
activation_checkpoint_func (callable, optional): The function to use for activation checkpointing. Defaults to ``deepspeed.checkpointing.checkpoint``.
|
||||
checkpointable_layers (list[str], optional): List of layer class names that are eligible for checkpointing. For GPT models,
|
||||
ParallelTransformerLayerPipe is always checkpointed regardless of this list. If None, all layers with parameters are
|
||||
considered checkpointable. Defaults to None.
|
||||
dynamic_shape: Allows dynamic shapes of inputs. This might have a performance impact.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
layers,
|
||||
num_stages=None,
|
||||
topology=None,
|
||||
loss_fn=None,
|
||||
seed_layers=False,
|
||||
seed_fn=None,
|
||||
base_seed=1234,
|
||||
partition_method='parameters',
|
||||
activation_checkpoint_interval=0,
|
||||
activation_checkpoint_func=checkpointing.checkpoint,
|
||||
checkpointable_layers=None,
|
||||
dynamic_shape=False):
|
||||
|
||||
super().__init__()
|
||||
|
||||
if num_stages is None and topology is None:
|
||||
raise RuntimeError('must provide num_stages or topology')
|
||||
|
||||
self.micro_offset = 0
|
||||
|
||||
self.loss_fn = loss_fn
|
||||
|
||||
self.checkpointable_layers = checkpointable_layers
|
||||
if checkpointable_layers is not None:
|
||||
assert isinstance(checkpointable_layers, list), "param `checkpointable_layers` must be type of list."
|
||||
|
||||
self.seed_layers = seed_layers
|
||||
self.seed_fn = seed_fn
|
||||
self.base_seed = base_seed
|
||||
if dist.get_rank() == 0:
|
||||
try:
|
||||
seed_str = self.seed_fn.__name__
|
||||
except AttributeError:
|
||||
seed_str = None
|
||||
print(f'SEED_LAYERS={self.seed_layers} BASE_SEED={self.base_seed} SEED_FN={seed_str}')
|
||||
|
||||
# Setup world info
|
||||
self.world_group = dist.new_group(ranks=range(dist.get_world_size()))
|
||||
self.global_rank = dist.get_rank(group=self.world_group)
|
||||
self.world_size = dist.get_world_size(group=self.world_group)
|
||||
self.local_rank = int(os.environ.get("LOCAL_RANK", None))
|
||||
assert self.local_rank is not None
|
||||
|
||||
if topology:
|
||||
self._topo = topology
|
||||
self.num_stages = self._topo.get_dim('pipe')
|
||||
else:
|
||||
self.num_stages = num_stages
|
||||
if topology is None:
|
||||
if self.world_size % self.num_stages != 0:
|
||||
raise RuntimeError(
|
||||
f'num_stages ({self.num_stages}) must divide distributed world size ({self.world_size})')
|
||||
dp = self.world_size // num_stages
|
||||
topology = PipeDataParallelTopology(num_pp=num_stages, num_dp=dp)
|
||||
self._topo = topology
|
||||
|
||||
# Construct communicators for pipeline topology
|
||||
self._grid = PipelineParallelGrid(process_group=self.world_group, topology=self._topo)
|
||||
|
||||
self.stage_id = self._topo.get_coord(self.global_rank).pipe
|
||||
|
||||
# Initialize partition information
|
||||
self._layer_specs = list(layers)
|
||||
self._num_layers = len(self._layer_specs)
|
||||
self._local_start = 0
|
||||
self._local_stop = None
|
||||
self._partition_layers(method=partition_method)
|
||||
|
||||
self.forward_funcs = []
|
||||
self.fwd_map = {}
|
||||
self.tied_modules = nn.ModuleDict()
|
||||
self.tied_weight_attrs = {}
|
||||
|
||||
# Offset the random seed by the stage ID.
|
||||
#newseed = get_accelerator().initial_seed() + self._grid.get_stage_id()
|
||||
#ds_utils.set_random_seed(newseed)
|
||||
|
||||
self.activation_checkpoint_interval = activation_checkpoint_interval
|
||||
|
||||
self.activation_checkpoint_func = activation_checkpoint_func
|
||||
|
||||
#storage for precomputed checkpointeble results
|
||||
self.is_checkpointable_results = []
|
||||
self.is_checkpointable_results_interval = None
|
||||
|
||||
# if configuration use_reentrant = False, self.activation_checkpoint_func will be set to ``checkpointing.non_reentrant_checkpoint``
|
||||
|
||||
#with torch.random.fork_rng(devices=[get_accelerator().current_device_name()]):
|
||||
self._build()
|
||||
self.to(get_accelerator().device_name(self.local_rank))
|
||||
|
||||
self.tied_comms = self._index_tied_modules()
|
||||
self._synchronize_tied_weights()
|
||||
|
||||
self.dynamic_shape = dynamic_shape
|
||||
|
||||
def _precompute_checkpointable_values(self):
|
||||
if self.activation_checkpoint_interval > 0 and self.is_checkpointable_results_interval != self.activation_checkpoint_interval:
|
||||
num_layers = len(self.forward_funcs)
|
||||
self.interval_was_zero = False
|
||||
for start_idx in range(0, num_layers, self.activation_checkpoint_interval):
|
||||
end_idx = min(start_idx + self.activation_checkpoint_interval, num_layers)
|
||||
funcs = self.forward_funcs[start_idx:end_idx]
|
||||
self.is_checkpointable_results.append(self._is_checkpointable(funcs))
|
||||
self.is_checkpointable_results_interval = self.activation_checkpoint_interval
|
||||
|
||||
def _build(self):
|
||||
specs = self._layer_specs
|
||||
|
||||
for local_idx, layer in enumerate(specs[self._local_start:self._local_stop]):
|
||||
layer_idx = local_idx + self._local_start
|
||||
if self.seed_layers:
|
||||
if self.seed_fn:
|
||||
self.seed_fn(self.base_seed + layer_idx)
|
||||
else:
|
||||
ds_utils.set_random_seed(self.base_seed + layer_idx)
|
||||
|
||||
# Recursively build PipelineModule objects
|
||||
if isinstance(layer, PipelineModule):
|
||||
raise NotImplementedError('RECURSIVE BUILD NOT YET IMPLEMENTED')
|
||||
|
||||
# LayerSpec objects contain an nn.Module that should be allocated now.
|
||||
elif isinstance(layer, nn.Module):
|
||||
name = str(layer_idx)
|
||||
self.forward_funcs.append(layer)
|
||||
self.fwd_map.update({name: len(self.forward_funcs) - 1})
|
||||
self.add_module(name, layer)
|
||||
|
||||
# TiedLayerSpec objects contain an nn.Module that should be allocated now.
|
||||
elif isinstance(layer, TiedLayerSpec):
|
||||
# Build and register the module if we haven't seen it before.
|
||||
if layer.key not in self.tied_modules:
|
||||
self.tied_modules[layer.key] = layer.build()
|
||||
self.tied_weight_attrs[layer.key] = layer.tied_weight_attr
|
||||
|
||||
if layer.forward_fn is None:
|
||||
# Just use forward()
|
||||
self.forward_funcs.append(self.tied_modules[layer.key])
|
||||
else:
|
||||
# User specified fn with args (module, input)
|
||||
self.forward_funcs.append(partial(layer.forward_fn, self.tied_modules[layer.key]))
|
||||
|
||||
# LayerSpec objects contain an nn.Module that should be allocated now.
|
||||
elif isinstance(layer, LayerSpec):
|
||||
module = layer.build()
|
||||
name = str(layer_idx)
|
||||
self.forward_funcs.append(module)
|
||||
self.fwd_map.update({name: len(self.forward_funcs) - 1})
|
||||
self.add_module(name, module)
|
||||
|
||||
# Last option: layer may be a functional (e.g., lambda). We do nothing in
|
||||
# that case and just use it in forward()
|
||||
else:
|
||||
self.forward_funcs.append(layer)
|
||||
|
||||
# All pipeline parameters should be considered as model parallel in the context
|
||||
# of our FP16 optimizer
|
||||
for p in self.parameters():
|
||||
p.ds_pipe_replicated = False
|
||||
|
||||
def _get_frozen_parameter_names(self, layer):
|
||||
""" Get names of frozen parameters in the layer.
|
||||
|
||||
Returns:
|
||||
A list of frozen parameter names
|
||||
"""
|
||||
if isinstance(layer, LayerSpec):
|
||||
l = layer.build()
|
||||
return [n for n, p in l.named_parameters() if not p.requires_grad]
|
||||
elif isinstance(layer, nn.Module):
|
||||
return [n for n, p in layer.named_parameters() if not p.requires_grad]
|
||||
|
||||
return []
|
||||
|
||||
def _count_layer_params(self):
|
||||
"""Count the trainable parameters in individual layers.
|
||||
|
||||
This routine will only build one layer at a time.
|
||||
|
||||
Returns:
|
||||
A list of the number of parameters in each layer.
|
||||
"""
|
||||
param_counts = [0] * len(self._layer_specs)
|
||||
for idx, layer in enumerate(self._layer_specs):
|
||||
if isinstance(layer, LayerSpec):
|
||||
l = layer.build()
|
||||
params = filter(lambda p: p.requires_grad, l.parameters())
|
||||
param_counts[idx] = sum(p.numel() for p in params)
|
||||
elif isinstance(layer, nn.Module):
|
||||
params = filter(lambda p: p.requires_grad, layer.parameters())
|
||||
param_counts[idx] = sum(p.numel() for p in params)
|
||||
return param_counts
|
||||
|
||||
def _find_layer_type(self, layername):
|
||||
idxs = []
|
||||
typeregex = regex.compile(layername, regex.IGNORECASE)
|
||||
for idx, layer in enumerate(self._layer_specs):
|
||||
name = None
|
||||
if isinstance(layer, LayerSpec):
|
||||
name = layer.typename.__name__
|
||||
elif isinstance(layer, nn.Module):
|
||||
name = layer.__class__.__name__
|
||||
else:
|
||||
try:
|
||||
name = layer.__name__
|
||||
except AttributeError:
|
||||
continue
|
||||
if typeregex.search(name):
|
||||
idxs.append(idx)
|
||||
|
||||
if len(idxs) == 0:
|
||||
raise RuntimeError(f"Partitioning '{layername}' found no valid layers to partition.")
|
||||
return idxs
|
||||
|
||||
def forward(self, forward_input):
|
||||
# We need to offset the seed by the microbatch ID. Save it in a local var to
|
||||
# ensure it is preserved in the closure. Otherwise checkpointed forward funcs
|
||||
# will see a different offset.
|
||||
self.micro_offset += 1
|
||||
|
||||
def exec_range_func(start, end):
|
||||
''' Helper function to be used with checkpoint()
|
||||
Adapted from torch.utils.checkpoint:checkpoint_sequential()
|
||||
'''
|
||||
local_micro_offset = self.micro_offset + 1
|
||||
|
||||
def exec_func(*inputs):
|
||||
# Single tensor inputs need to be unwrapped
|
||||
if len(inputs) == 1:
|
||||
inputs = inputs[0]
|
||||
for idx, layer in enumerate(self.forward_funcs[start:end]):
|
||||
self.curr_layer = idx + self._local_start
|
||||
if self.seed_layers:
|
||||
new_seed = (self.base_seed * local_micro_offset) + self.curr_layer
|
||||
if self.seed_fn:
|
||||
self.seed_fn(new_seed)
|
||||
else:
|
||||
ds_utils.set_random_seed(new_seed)
|
||||
|
||||
inputs = layer(inputs)
|
||||
return inputs
|
||||
|
||||
return exec_func
|
||||
|
||||
if self.activation_checkpoint_interval == 0:
|
||||
func = exec_range_func(0, len(self.forward_funcs))
|
||||
x = func(forward_input)
|
||||
else:
|
||||
num_layers = len(self.forward_funcs)
|
||||
x = forward_input
|
||||
for start_idx, is_checkpointable_result in \
|
||||
zip(range(0, num_layers, self.activation_checkpoint_interval), self.is_checkpointable_results):
|
||||
|
||||
end_idx = min(start_idx + self.activation_checkpoint_interval, num_layers)
|
||||
|
||||
funcs = self.forward_funcs[start_idx:end_idx]
|
||||
# Since we either pass tensors or tuples of tensors without unpacking, we
|
||||
# need to be careful not to double-wrap tensors with tuple.
|
||||
if not isinstance(x, tuple):
|
||||
x = (x, )
|
||||
|
||||
if is_checkpointable_result:
|
||||
x = self.activation_checkpoint_func(exec_range_func(start_idx, end_idx), *x)
|
||||
else:
|
||||
x = exec_range_func(start_idx, end_idx)(*x)
|
||||
return x
|
||||
|
||||
def _partition_layers(self, method='uniform'):
|
||||
num_stages = self._topo.get_dim('pipe')
|
||||
stage_id = self._topo.get_coord(self.global_rank).pipe
|
||||
|
||||
if self.global_rank == 0:
|
||||
logger.info(f'Partitioning pipeline stages with method {method}')
|
||||
|
||||
method = method.lower()
|
||||
|
||||
# Each stage gets a simple uniform number of layers.
|
||||
if method == 'uniform':
|
||||
num_layers = len(self._layer_specs)
|
||||
self.parts = ds_utils.partition_uniform(num_items=num_layers, num_parts=num_stages)
|
||||
elif method == 'parameters':
|
||||
param_counts = self._count_layer_params()
|
||||
self.parts = ds_utils.partition_balanced(weights=param_counts, num_parts=num_stages)
|
||||
elif method.startswith('type:'):
|
||||
layertype = method.split(':')[1]
|
||||
binary_weights = [0] * len(self._layer_specs)
|
||||
for idx in self._find_layer_type(layertype):
|
||||
binary_weights[idx] = 1
|
||||
self.parts = ds_utils.partition_balanced(weights=binary_weights, num_parts=num_stages)
|
||||
elif method == 'profile':
|
||||
raise NotImplementedError(f'Partitioning method {method} not implemented.')
|
||||
else:
|
||||
raise NotImplementedError(f'Partitioning method {method} not implemented.')
|
||||
|
||||
# Print some information on the partitioning.
|
||||
if self.global_rank == 0:
|
||||
for stage in range(num_stages):
|
||||
start = self.parts[stage]
|
||||
stop = self.parts[stage + 1]
|
||||
print(f'stage={stage} layers={stop - start}')
|
||||
for idx, layer in enumerate(self._layer_specs[start:stop]):
|
||||
name = str(layer)
|
||||
if isinstance(layer, LayerSpec):
|
||||
name = layer.typename.__name__
|
||||
if isinstance(layer, nn.Module):
|
||||
name = layer.__class__.__name__
|
||||
else:
|
||||
try:
|
||||
name = layer.__name__
|
||||
except AttributeError:
|
||||
pass
|
||||
print(f' {idx+start:2d}: {name}')
|
||||
if self.loss_fn:
|
||||
try:
|
||||
print(f' loss: {self.loss_fn.__name__}')
|
||||
except AttributeError:
|
||||
print(f' loss: {self.loss_fn.__class__.__name__}')
|
||||
|
||||
self._set_bounds(start=self.parts[stage_id], stop=self.parts[stage_id + 1])
|
||||
|
||||
@staticmethod
|
||||
def _recursive_getattr(module: torch.nn.Module, attr_name: str) -> torch.Tensor:
|
||||
'''Allow getting an attribute like "linear.weight"'''
|
||||
weight = module
|
||||
for item in attr_name.split("."):
|
||||
weight = getattr(weight, item)
|
||||
return weight
|
||||
|
||||
def allreduce_tied_weight_gradients(self):
|
||||
'''All reduce the gradients of the tied weights between tied stages'''
|
||||
for key, comm in self.tied_comms.items():
|
||||
for attr_name in comm['weight_attr']:
|
||||
weight = self._recursive_getattr(self.tied_modules[key], attr_name)
|
||||
dist.all_reduce(weight.grad, group=comm['group'])
|
||||
|
||||
def get_tied_weights_and_groups(self):
|
||||
weight_group_list = []
|
||||
for key, comm in self.tied_comms.items():
|
||||
for attr_name in comm['weight_attr']:
|
||||
weight = self._recursive_getattr(self.tied_modules[key], attr_name)
|
||||
weight_group_list.append((weight, comm['group']))
|
||||
return weight_group_list
|
||||
|
||||
def _synchronize_tied_weights(self):
|
||||
for key, comm in self.tied_comms.items():
|
||||
for attr_name in comm['weight_attr']:
|
||||
dist.broadcast(
|
||||
self._recursive_getattr(comm['module'], attr_name),
|
||||
src=min(comm['ranks']),
|
||||
group=comm['group'],
|
||||
)
|
||||
|
||||
def _index_tied_modules(self):
|
||||
''' Build communication structures for tied modules. '''
|
||||
tied_comms = {}
|
||||
if self._topo.get_dim('pipe') == 1:
|
||||
return tied_comms
|
||||
|
||||
specs = self._layer_specs
|
||||
tie_keys = set(s.key for s in specs if isinstance(s, TiedLayerSpec))
|
||||
# Since Python 3.7, "Dictionary order is guaranteed to be insertion order."
|
||||
# Sort tie_keys here so that orders of self.tied_comms.items() are consistent
|
||||
# among ranks.
|
||||
for key in sorted(tie_keys):
|
||||
# Find the layers that the tied module appears in
|
||||
tied_layers = []
|
||||
for idx, layer in enumerate(specs):
|
||||
if isinstance(layer, TiedLayerSpec) and layer.key == key:
|
||||
tied_layers.append(idx)
|
||||
# Find all stages with this tied module
|
||||
# TODO: Would be nice to remove the nested data/model parallelism loops and
|
||||
# TODO: instead generalize in some way, since we really just care about the
|
||||
# TODO: stage that owns the tied layer. Then loop over each (dp, mp, ...)
|
||||
# TODO: fiber to generate process groups.
|
||||
tied_stages = set(self.stage_owner(idx) for idx in tied_layers)
|
||||
for dp in range(self._grid.data_parallel_size):
|
||||
for mp in range(self._grid.get_slice_parallel_world_size()):
|
||||
tied_ranks = []
|
||||
for s in sorted(tied_stages):
|
||||
if self._grid.get_slice_parallel_world_size() > 1:
|
||||
tied_ranks.append(self._grid.stage_to_global(stage_id=s, data=dp, model=mp))
|
||||
else:
|
||||
tied_ranks.append(self._grid.stage_to_global(stage_id=s, data=dp))
|
||||
group = dist.new_group(ranks=tied_ranks)
|
||||
|
||||
# Record this tied module if we own a local copy of it.
|
||||
if self.global_rank in tied_ranks:
|
||||
assert key in self.tied_modules
|
||||
if key in self.tied_modules:
|
||||
tied_comms[key] = {
|
||||
'ranks': tied_ranks,
|
||||
'group': group,
|
||||
'weight_attr': self.tied_weight_attrs[key],
|
||||
'module': self.tied_modules[key],
|
||||
}
|
||||
# Only count the tied module once in the eyes of the FP16 optimizer
|
||||
if self.global_rank != tied_ranks[0]:
|
||||
for p in self.tied_modules[key].parameters():
|
||||
p.ds_pipe_replicated = True
|
||||
'''
|
||||
if len(tied_comms) > 0:
|
||||
print(f'RANK={self.global_rank} tied_comms={tied_comms}')
|
||||
'''
|
||||
|
||||
return tied_comms
|
||||
|
||||
def partitions(self):
|
||||
return self.parts
|
||||
|
||||
def stage_owner(self, layer_idx):
|
||||
assert 0 <= layer_idx < self._num_layers
|
||||
for stage in range(self._topo.get_dim('pipe')):
|
||||
if self.parts[stage] <= layer_idx < self.parts[stage + 1]:
|
||||
return stage
|
||||
raise RuntimeError(f'Layer {layer_idx} not owned? parts={self.parts}')
|
||||
|
||||
def _set_bounds(self, start=None, stop=None):
|
||||
"""Manually define the range of layers that will be built on this process.
|
||||
|
||||
These boundaries are treated as list slices and so start is inclusive and stop is
|
||||
exclusive. The default of None for both results in all layers being built
|
||||
locally.
|
||||
"""
|
||||
self._local_start = start
|
||||
self._local_stop = stop
|
||||
|
||||
def set_checkpoint_interval(self, interval):
|
||||
assert interval >= 0
|
||||
self.checkpoint_interval = interval
|
||||
|
||||
def topology(self):
|
||||
""" ProcessTopology object to query process mappings. """
|
||||
return self._topo
|
||||
|
||||
def mpu(self):
|
||||
return self._grid
|
||||
|
||||
def num_pipeline_stages(self):
|
||||
return self._topo.get_dim('pipe')
|
||||
|
||||
def ckpt_prefix(self, checkpoints_path, tag):
|
||||
"""Build a prefix for all checkpoint files written by this module. """
|
||||
# All checkpoint files start with this
|
||||
rank_name = 'module'
|
||||
|
||||
# Data parallelism is omitted from the naming convention because we are agnostic
|
||||
# to this in the checkpoint.
|
||||
omit_dims = frozenset(['data'])
|
||||
axes = [a for a in self._grid._topo.get_axis_names() if a not in omit_dims]
|
||||
for dim in axes:
|
||||
rank = getattr(self._grid._topo.get_coord(rank=self.global_rank), dim)
|
||||
rank_name += f'-{dim}_{rank:02d}'
|
||||
|
||||
ckpt_name = os.path.join(checkpoints_path, str(tag), rank_name)
|
||||
return ckpt_name
|
||||
|
||||
def ckpt_layer_path(self, ckpt_dir, local_layer_idx):
|
||||
"""Customize a prefix for a specific pipeline module layer. """
|
||||
idx = local_layer_idx + self._local_start
|
||||
layer_ckpt_path = os.path.join(ckpt_dir, f'layer_{idx:02d}')
|
||||
rank_repr = self._grid._topo.get_rank_repr(rank=self.global_rank)
|
||||
if rank_repr != '':
|
||||
layer_ckpt_path += f'-{rank_repr}'
|
||||
layer_ckpt_path += '-model_states.pt'
|
||||
return layer_ckpt_path
|
||||
|
||||
def ckpt_layer_path_list(self, ckpt_dir, local_layer_idx):
|
||||
"""Get all ckpt file list for a specific pipeline module layer. """
|
||||
idx = local_layer_idx + self._local_start
|
||||
layer_ckpt_path = os.path.join(ckpt_dir, f'layer_{idx:02d}-')
|
||||
layer_ckpt_path += "*model_states.pt"
|
||||
ckpt_files = glob.glob(layer_ckpt_path)
|
||||
ckpt_files.sort()
|
||||
return ckpt_files
|
||||
|
||||
def save_state_dict(self, save_dir, checkpoint_engine, exclude_frozen_params=False):
|
||||
# TODO: Need to validate interaction of checkpoint_parallel_write_pipeline and fastwriter
|
||||
|
||||
# Processes having the same model parallel rank on different data parallel instances
|
||||
# have identical layer weights. We can distribute the task of saving the layer weights
|
||||
# among the data parallel ranks. For example, if a pipeline stage has 9 layers and
|
||||
# if there are 2 data parallel instances, rank 0 will save the first 5 layers and
|
||||
# rank 1 will save the last 4.
|
||||
dp_rank = self._grid.data_parallel_id
|
||||
dp_size = self._grid.data_parallel_size
|
||||
num_layers = len(self.forward_funcs)
|
||||
if self.checkpoint_parallel_write_pipeline:
|
||||
# spread layers evenly across data parallel ranks
|
||||
offsets = ds_utils.partition_uniform(num_layers, dp_size)
|
||||
start, end = offsets[dp_rank], offsets[dp_rank + 1]
|
||||
else:
|
||||
# data parallel rank 0 writes all layers
|
||||
if dp_rank != 0:
|
||||
return
|
||||
start, end = 0, num_layers
|
||||
layer_list = self.forward_funcs[start:end]
|
||||
|
||||
checkpoint_engine.makedirs(save_dir, exist_ok=True)
|
||||
should_clone = checkpoint_engine.preserves_storage_sharing()
|
||||
for idx, layer in enumerate(layer_list):
|
||||
model_ckpt_path = self.ckpt_layer_path(save_dir, start + idx)
|
||||
if not hasattr(layer, 'state_dict'):
|
||||
continue
|
||||
|
||||
orig_state_dict = layer.state_dict()
|
||||
if exclude_frozen_params:
|
||||
for n in self._get_frozen_parameter_names(layer):
|
||||
del orig_state_dict[n]
|
||||
final_state_dict = orig_state_dict
|
||||
if should_clone:
|
||||
final_state_dict = clone_tensors_for_torch_save(orig_state_dict)
|
||||
checkpoint_engine.save(state_dict=final_state_dict, path=model_ckpt_path)
|
||||
|
||||
def load_state_dir(self, load_dir, checkpoint_engine, strict=True):
|
||||
for idx, layer in enumerate(self.forward_funcs):
|
||||
# Functions, etc. will not have state_dicts
|
||||
if not hasattr(layer, 'load_state_dict'):
|
||||
continue
|
||||
|
||||
# get all checkpoint files for the layer.
|
||||
model_ckpt_list = self.ckpt_layer_path_list(load_dir, idx)
|
||||
mp_rank = self._grid.get_slice_parallel_rank()
|
||||
mp_world_size = self._grid.get_slice_parallel_world_size()
|
||||
|
||||
sd_loader = SDLoaderFactory.get_sd_loader(model_ckpt_list,
|
||||
version=2.0,
|
||||
checkpoint_engine=checkpoint_engine)
|
||||
load_path, checkpoint, _ = sd_loader.load(mp_world_size, mp_rank, module_key=None, is_pipe_parallel=True)
|
||||
|
||||
layer.load_state_dict(checkpoint, strict=strict)
|
||||
|
||||
# if self._grid.data_parallel_id == 0:
|
||||
# logger.info(
|
||||
# f'RANK={self.global_rank} Loaded layer={idx+self._local_start} file={load_path}'
|
||||
# )
|
||||
|
||||
self._synchronize_tied_weights()
|
||||
|
||||
def _is_checkpointable(self, funcs):
|
||||
|
||||
if self.activation_checkpoint_func is not checkpointing.non_reentrant_checkpoint:
|
||||
# This hook excludes the embedding layer
|
||||
# because only non_reentrant_checkpoint can accept inputs with requires_grad=False
|
||||
# otherwise, the backward of the embedding layer won't receive gradients.
|
||||
if self.__class__.__name__ in ('GPTModelPipe', 'GPT2ModelPipe'):
|
||||
# For GPT models, checkpoint both transformer layers and any additional
|
||||
# layers specified in checkpointable_layers (if provided)
|
||||
return all('ParallelTransformerLayerPipe' in f.__class__.__name__ or (
|
||||
self.checkpointable_layers is not None and f.__class__.__name__ in self.checkpointable_layers)
|
||||
for f in funcs)
|
||||
|
||||
if self.checkpointable_layers is not None:
|
||||
# For non-GPT models, only checkpoint layers specified in checkpointable_layers
|
||||
return all(f.__class__.__name__ in self.checkpointable_layers for f in funcs)
|
||||
|
||||
# Default behavior: checkpoint any layer that has parameters
|
||||
params = [f.parameters() for f in funcs if isinstance(f, torch.nn.Module)]
|
||||
return any(len(list(p)) > 0 for p in params)
|
||||
|
||||
def get_additional_losses(self):
|
||||
""" Returns model specific additional losses for reporting
|
||||
|
||||
Return a dictionary of {"loss name": loss_value} or None if no additional losses.
|
||||
"""
|
||||
return None
|
||||
|
||||
def compile(self, *args, **kwargs):
|
||||
for idx, layer in enumerate(self.forward_funcs):
|
||||
if isinstance(layer, nn.Module):
|
||||
layer.compile(*args, **kwargs)
|
||||
else:
|
||||
new_layer = torch.compile(layer, *args, **kwargs)
|
||||
self.forward_funcs[idx] = new_layer
|
||||
@@ -0,0 +1,182 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import msgpack
|
||||
import typing
|
||||
|
||||
import torch
|
||||
from deepspeed import comm as dist
|
||||
|
||||
from deepspeed.utils.torch import required_torch_version
|
||||
from deepspeed.accelerator import get_accelerator
|
||||
|
||||
_groups = None
|
||||
_grid = None
|
||||
|
||||
_async = []
|
||||
|
||||
|
||||
def can_send_recv() -> bool:
|
||||
return required_torch_version(min_version=1.8)
|
||||
|
||||
|
||||
#initializes adjacent process groups
|
||||
#run this only after deepspeed.init_distributed() has been called
|
||||
def init_process_groups(grid):
|
||||
global _groups, _grid
|
||||
_grid = grid
|
||||
|
||||
assert _grid.pipe_parallel_size > 1, "There is no pipeline parallelism"
|
||||
|
||||
if not can_send_recv():
|
||||
_groups = [dist.new_group(ranks=group) for group in _grid.p2p_groups]
|
||||
|
||||
|
||||
def _is_valid_send_recv(src_stage, dest_stage):
|
||||
first_stage = 0
|
||||
last_stage = _grid.pipe_parallel_size - 1
|
||||
assert abs(src_stage-dest_stage) == 1 or \
|
||||
(src_stage == first_stage and dest_stage == last_stage) or \
|
||||
(src_stage == last_stage and dest_stage == first_stage), \
|
||||
"Functionality currently limited to send and receive between adjacent ranks only"
|
||||
|
||||
|
||||
def send(tensor, dest_stage, async_op=False):
|
||||
global _groups
|
||||
assert async_op == False, "Doesn't support async_op true"
|
||||
src_stage = _grid.get_stage_id()
|
||||
_is_valid_send_recv(src_stage, dest_stage)
|
||||
|
||||
dest_rank = _grid.stage_to_global(stage_id=dest_stage)
|
||||
if async_op:
|
||||
global _async
|
||||
op = dist.isend(tensor, dest_rank)
|
||||
_async.append(op)
|
||||
else:
|
||||
|
||||
if can_send_recv():
|
||||
return dist.send(tensor, dest_rank)
|
||||
else:
|
||||
group = _get_send_recv_group(src_stage, dest_stage)
|
||||
src_rank = _grid.stage_to_global(stage_id=src_stage)
|
||||
return dist.broadcast(tensor, src_rank, group=group, async_op=async_op)
|
||||
|
||||
|
||||
def recv(tensor, src_stage, async_op=False):
|
||||
global _groups
|
||||
assert async_op == False, "Doesn't support async_op true"
|
||||
dest_stage = _grid.get_stage_id()
|
||||
_is_valid_send_recv(src_stage, dest_stage)
|
||||
|
||||
src_rank = _grid.stage_to_global(stage_id=src_stage)
|
||||
|
||||
if async_op:
|
||||
global _async
|
||||
op = dist.irecv(tensor, src_rank)
|
||||
_async.append(op)
|
||||
else:
|
||||
if can_send_recv():
|
||||
return dist.recv(tensor, src_rank)
|
||||
else:
|
||||
group = _get_send_recv_group(src_stage, dest_stage)
|
||||
return dist.broadcast(tensor, src_rank, group=group, async_op=async_op)
|
||||
|
||||
|
||||
def wait():
|
||||
global _async
|
||||
for op in _async:
|
||||
op.wait()
|
||||
_async = []
|
||||
|
||||
get_accelerator().synchronize()
|
||||
|
||||
|
||||
def send_obj(msg: typing.Any, dest: int):
|
||||
"""Send an arbitrary python object to ``dest``.
|
||||
|
||||
Note: ``msg`` must be serializable by msgpack.
|
||||
|
||||
WARN: This incurs a CPU -> GPU transfer and should be used sparingly
|
||||
for performance reasons.
|
||||
|
||||
Args:
|
||||
msg (typing.Any): The object to send.
|
||||
dest (int): Destination rank.
|
||||
"""
|
||||
# serialize the message
|
||||
msg = msgpack.packb(msg)
|
||||
# construct a tensor to send
|
||||
msg = torch.ByteTensor(torch.ByteStorage.from_buffer(msg)).to(get_accelerator().device_name())
|
||||
|
||||
# Send meta and message
|
||||
length_tensor = torch.tensor([len(msg)], dtype=torch.long).to(get_accelerator().device_name())
|
||||
dist.send(length_tensor, dst=dest)
|
||||
dist.send(msg, dst=dest)
|
||||
|
||||
|
||||
def recv_obj(sender: int) -> typing.Any:
|
||||
"""Receive an arbitrary python object from ``sender``.
|
||||
|
||||
WARN: This incur a CPU <-> GPU transfers and should be used sparingly
|
||||
for performance reasons.
|
||||
|
||||
Args:
|
||||
sender (int): The rank sending the message.
|
||||
"""
|
||||
# Get message meta
|
||||
length = torch.tensor([0], dtype=torch.long).to(get_accelerator().device_name())
|
||||
dist.recv(length, src=sender)
|
||||
|
||||
# Receive and deserialize
|
||||
msg = torch.empty(length.item(), dtype=torch.uint8).to(get_accelerator().device_name())
|
||||
dist.recv(msg, src=sender)
|
||||
|
||||
msg = msgpack.unpackb(msg.cpu().numpy().tobytes())
|
||||
|
||||
def _to(x):
|
||||
"""Recursively move to the current device."""
|
||||
if torch.is_tensor(x):
|
||||
return x.to(get_accelerator().device_name())
|
||||
if isinstance(x, (tuple, list)):
|
||||
ret = [_to(x_) for x_ in x]
|
||||
if isinstance(x, tuple):
|
||||
ret = tuple(ret)
|
||||
return ret
|
||||
# handle kwargs
|
||||
if isinstance(x, dict):
|
||||
ret = dict()
|
||||
for key, val in x.items():
|
||||
ret[_to(key)] = _to(val)
|
||||
return ret
|
||||
|
||||
# Anything else is a no-op
|
||||
return x
|
||||
|
||||
msg = _to(msg)
|
||||
return msg
|
||||
|
||||
|
||||
def _get_send_recv_group(src_stage, dest_stage):
|
||||
'''the group id is always the smaller rank unless its a wrap around'''
|
||||
|
||||
stage_id = None
|
||||
|
||||
first_stage = 0
|
||||
last_stage = _grid.pipe_parallel_size - 1
|
||||
|
||||
if (src_stage == first_stage and dest_stage == last_stage
|
||||
or dest_stage == first_stage and src_stage == last_stage):
|
||||
stage_id = last_stage
|
||||
elif src_stage > dest_stage:
|
||||
stage_id = dest_stage
|
||||
else:
|
||||
stage_id = src_stage
|
||||
'''group_id corresponds to group of [group_id, group_id+1]
|
||||
unless group_id is the rank of the last stage
|
||||
in which case group_id corresponds to group[group_id-num_stages+1, group_id]
|
||||
'''
|
||||
group_id = _grid.stage_to_global(stage_id=stage_id)
|
||||
|
||||
return _groups[group_id]
|
||||
@@ -0,0 +1,494 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from ..utils import call_to_str
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class PipeSchedule(ABC):
|
||||
"""Directs the execution of a pipeline engine by generating sequences of
|
||||
:class:`PipeInstruction`.
|
||||
|
||||
Schedules are generators that yield sequences of
|
||||
:class:`PipeInstruction` to process the micro-batches in one batch.
|
||||
Each yielded step is atomic in the sense that a barrier
|
||||
synchronization can be placed between successive steps without
|
||||
deadlock.
|
||||
|
||||
Below is an example schedule that implements data parallelism with gradient accumulation:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
class DataParallelSchedule(PipeSchedule):
|
||||
def steps(self):
|
||||
for step_id in range(self.micro_batches):
|
||||
cmds = [
|
||||
LoadMicroBatch(buffer_id=0),
|
||||
ForwardPass(buffer_id=0),
|
||||
BackwardPass(buffer_id=0),
|
||||
]
|
||||
if step_id == self.micro_batches - 1:
|
||||
cmds.extend([
|
||||
ReduceGrads(),
|
||||
OptimizerStep(),
|
||||
])
|
||||
yield cmds
|
||||
|
||||
def num_pipe_buffers(self):
|
||||
return 1
|
||||
|
||||
Args:
|
||||
micro_batches (int): The number of micro-batches that comprise a batch.
|
||||
stages (int): The number of pipeline stages.
|
||||
stage_id (int): The pipe stage that will execute the generated schedule.
|
||||
"""
|
||||
|
||||
def __init__(self, micro_batches, stages, stage_id):
|
||||
super().__init__()
|
||||
self.micro_batches = micro_batches
|
||||
self.stages = stages
|
||||
self.stage_id = stage_id
|
||||
self.prev_stage = self.stage_id - 1
|
||||
self.next_stage = self.stage_id + 1
|
||||
|
||||
@abstractmethod
|
||||
def steps(self):
|
||||
"""Yield a list of :class:`PipeInstruction` for each step in the schedule.
|
||||
|
||||
.. note::
|
||||
Schedules must implement ``steps()`` to define the schedule.
|
||||
|
||||
Returns:
|
||||
Instructions to be executed as one step of the pipeline
|
||||
"""
|
||||
pass
|
||||
|
||||
def num_pipe_buffers(self):
|
||||
"""The number of pipeline buffers that will be used by this stage.
|
||||
|
||||
.. note::
|
||||
Schedules should specialize ``num_pipe_buffers()`` for memory savings at scale.
|
||||
|
||||
Returns:
|
||||
The number of buffers for the engine to allocate.
|
||||
"""
|
||||
return self.micro_batches
|
||||
|
||||
def _valid_micro_batch(self, micro_batch_id):
|
||||
return 0 <= micro_batch_id < self.micro_batches
|
||||
|
||||
def _valid_stage(self, stage_id):
|
||||
return 0 <= stage_id < self.stages
|
||||
|
||||
@property
|
||||
def stage(self):
|
||||
"""Stage index used to configure this schedule."""
|
||||
return self.stage_id
|
||||
|
||||
@property
|
||||
def num_stages(self):
|
||||
"""The number of total pipeline stages used to configure this schedule."""
|
||||
return self.stages
|
||||
|
||||
@property
|
||||
def num_micro_batches(self):
|
||||
"""The number of total micro_batches used to configure this schedule."""
|
||||
return self.micro_batches
|
||||
|
||||
@property
|
||||
def is_first_stage(self):
|
||||
"""True if the configured ``stage_id`` is the first stage in the pipeline."""
|
||||
return self.stage_id == 0
|
||||
|
||||
@property
|
||||
def is_last_stage(self):
|
||||
"""True if the configured ``stage_id`` is the last stage in the pipeline."""
|
||||
return self.stage_id == self.stages - 1
|
||||
|
||||
def _buffer_idx(self, micro_batch_id):
|
||||
"""Map a micro-batch index to a pipeline buffer index.
|
||||
|
||||
This method uses a cyclic allocation strategy.
|
||||
|
||||
Args:
|
||||
micro_batch_id (int): The micro-batch index relative to the beginning of the schedule.
|
||||
|
||||
Returns:
|
||||
int: The index of the buffer that should store data.
|
||||
"""
|
||||
assert self._valid_micro_batch(micro_batch_id)
|
||||
return micro_batch_id % self.num_pipe_buffers()
|
||||
|
||||
def __iter__(self):
|
||||
self.it = None
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
if self.it is None:
|
||||
self.it = self.steps()
|
||||
return next(self.it)
|
||||
|
||||
|
||||
class InferenceSchedule(PipeSchedule):
|
||||
"""A schedule for inferencing batches using pipeline parallelism.
|
||||
"""
|
||||
|
||||
def steps(self):
|
||||
""""""
|
||||
prev_micro_batch_id = -1
|
||||
total_steps = self.micro_batches + self.stages - 1
|
||||
for step_id in range(total_steps):
|
||||
cmds = []
|
||||
micro_batch_id = step_id - self.stage_id
|
||||
|
||||
# Alternate send/recv buffers
|
||||
if _is_even(self.stage_id):
|
||||
recv_buf = step_id % 2
|
||||
send_buf = (step_id + 1) % 2
|
||||
else:
|
||||
recv_buf = (step_id + 1) % 2
|
||||
send_buf = step_id % 2
|
||||
|
||||
if self.is_first_stage or self.is_last_stage:
|
||||
if self._valid_micro_batch(micro_batch_id):
|
||||
cmds.append(LoadMicroBatch(recv_buf))
|
||||
|
||||
if _is_even(self.stage_id):
|
||||
if self._valid_stage(self.next_stage):
|
||||
if self._valid_micro_batch(micro_batch_id - 1):
|
||||
cmds.append(SendActivation(send_buf))
|
||||
if self._valid_stage(self.prev_stage):
|
||||
if self._valid_micro_batch(micro_batch_id):
|
||||
cmds.append(RecvActivation(recv_buf))
|
||||
else:
|
||||
if self._valid_stage(self.prev_stage):
|
||||
if self._valid_micro_batch(micro_batch_id):
|
||||
cmds.append(RecvActivation(recv_buf))
|
||||
|
||||
if self._valid_stage(self.next_stage):
|
||||
if self._valid_micro_batch(micro_batch_id - 1):
|
||||
cmds.append(SendActivation(send_buf))
|
||||
|
||||
if self._valid_micro_batch(micro_batch_id):
|
||||
cmds.append(ForwardPass(recv_buf))
|
||||
|
||||
yield cmds
|
||||
|
||||
def num_pipe_buffers(self):
|
||||
"""Only two pipeline buffers are required for inferencing.
|
||||
|
||||
Returns:
|
||||
``2``
|
||||
"""
|
||||
return 2
|
||||
|
||||
|
||||
class TrainSchedule(PipeSchedule):
|
||||
"""A schedule for training a batch using hybrid parallelism.
|
||||
|
||||
Pipeline parallelism is extracted through gradient accumulation and thus
|
||||
convergence follows that of a data parallel approach with the same batch
|
||||
size.
|
||||
"""
|
||||
|
||||
def steps(self):
|
||||
""""""
|
||||
prev_micro_batch_id = -1
|
||||
total_steps = 2 * (self.micro_batches + self.stages - 1)
|
||||
for step_id in range(total_steps):
|
||||
# Map the step of the pipeline to the micro-batch id and also whether it is a
|
||||
# forward or backward pass step.
|
||||
micro_batch_id, is_forward = self._step_to_micro_batch(step_id)
|
||||
|
||||
if self._valid_micro_batch(prev_micro_batch_id):
|
||||
prev_buffer = self._buffer_idx(prev_micro_batch_id)
|
||||
if self._valid_micro_batch(micro_batch_id):
|
||||
curr_buffer = self._buffer_idx(micro_batch_id)
|
||||
|
||||
cmds = []
|
||||
|
||||
# Exchange activations
|
||||
if is_forward:
|
||||
if self._valid_micro_batch(prev_micro_batch_id) and self._valid_stage(self.prev_stage):
|
||||
cmds.append(SendGrad(prev_buffer))
|
||||
if self._valid_micro_batch(micro_batch_id) and self._valid_stage(self.prev_stage):
|
||||
cmds.append(RecvActivation(curr_buffer))
|
||||
else:
|
||||
if self._valid_micro_batch(micro_batch_id) and self._valid_stage(self.next_stage):
|
||||
cmds.append(RecvGrad(curr_buffer))
|
||||
if self._valid_micro_batch(prev_micro_batch_id) and self._valid_stage(self.next_stage):
|
||||
cmds.append(SendActivation(prev_buffer))
|
||||
|
||||
# First/last stage loads
|
||||
if self.stage_id == 0 or self.stage_id == self.stages - 1:
|
||||
if is_forward and self._valid_micro_batch(micro_batch_id):
|
||||
cmds.append(LoadMicroBatch(curr_buffer))
|
||||
|
||||
# Computation
|
||||
if self._valid_micro_batch(micro_batch_id):
|
||||
if is_forward:
|
||||
cmds.append(ForwardPass(curr_buffer))
|
||||
else:
|
||||
cmds.append(BackwardPass(curr_buffer))
|
||||
|
||||
# Model step at the end of the batch
|
||||
if step_id == total_steps - 1:
|
||||
cmds.append(ReduceTiedGrads())
|
||||
cmds.append(ReduceGrads())
|
||||
cmds.append(OptimizerStep())
|
||||
|
||||
# Prepare state for next time
|
||||
prev_micro_batch_id = micro_batch_id
|
||||
yield cmds
|
||||
|
||||
def num_pipe_buffers(self):
|
||||
"""Return the number of pipeline buffers required for this stage.
|
||||
|
||||
This is equivalent to the maximum number of in-flight forward passes,
|
||||
since we need to remember the activations of forward passes in order
|
||||
to run backpropagation. For synchronous 1F1B, this is equivalent to
|
||||
the index difference between this stage and the last stage.
|
||||
"""
|
||||
buffers = min(self.stages - self.stage_id, self.micro_batches)
|
||||
return max(2, buffers)
|
||||
|
||||
def _step_to_micro_batch(self, step_id):
|
||||
if _is_even(step_id) and _is_even(self.stage_id):
|
||||
micro_batch_id = self._even_step_forward_id(step_id)
|
||||
is_forward = True
|
||||
|
||||
elif _is_odd(step_id) and _is_odd(self.stage_id):
|
||||
micro_batch_id = self._odd_step_forward_id(step_id)
|
||||
is_forward = True
|
||||
|
||||
elif _is_even(step_id) and _is_odd(self.stage_id):
|
||||
micro_batch_id = self._even_step_backward_id(step_id)
|
||||
is_forward = False
|
||||
|
||||
elif _is_odd(step_id) and _is_even(self.stage_id):
|
||||
micro_batch_id = self._odd_step_backward_id(step_id)
|
||||
is_forward = False
|
||||
|
||||
else:
|
||||
assert False
|
||||
|
||||
return micro_batch_id, is_forward
|
||||
|
||||
def _even_step_forward_id(self, step_id):
|
||||
base = step_id // 2
|
||||
micro_batch_id = int(base - self.stage_id // 2)
|
||||
return micro_batch_id
|
||||
|
||||
def _odd_step_forward_id(self, step_id):
|
||||
base = (step_id - 1) // 2
|
||||
micro_batch_id = int(base - self.stage_id // 2)
|
||||
return micro_batch_id
|
||||
|
||||
def _even_step_backward_id(self, step_id):
|
||||
base = step_id // 2
|
||||
micro_batch_id = int(base - self.stages + (self.stage_id + 1) // 2)
|
||||
return micro_batch_id
|
||||
|
||||
def _odd_step_backward_id(self, step_id):
|
||||
base = ((step_id - 1) // 2) - self.stages + 1
|
||||
micro_batch_id = int(base + self.stage_id // 2)
|
||||
return micro_batch_id
|
||||
|
||||
|
||||
class DataParallelSchedule(PipeSchedule):
|
||||
"""An example schedule that trains using traditional data parallelism with gradient
|
||||
accumulation.
|
||||
"""
|
||||
|
||||
def steps(self):
|
||||
""""""
|
||||
for step_id in range(self.micro_batches):
|
||||
cmds = [
|
||||
LoadMicroBatch(buffer_id=0),
|
||||
ForwardPass(buffer_id=0),
|
||||
BackwardPass(buffer_id=0),
|
||||
]
|
||||
if step_id == self.micro_batches - 1:
|
||||
cmds.extend([
|
||||
ReduceGrads(),
|
||||
OptimizerStep(),
|
||||
])
|
||||
yield cmds
|
||||
|
||||
def num_pipe_buffers(self):
|
||||
"""Only one pipeline buffer needed.
|
||||
"""
|
||||
return 1
|
||||
|
||||
|
||||
class PipeInstruction:
|
||||
"""Base class for all instructions to be executed by the pipeline engine.
|
||||
|
||||
All keyword arguments are stored as members similar to a ``namedtuple``. These are
|
||||
then accessible to the :class:`PipeEngine` during execution.
|
||||
|
||||
Args:
|
||||
kwargs (optional): keyword arguments to store as members
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.name = self.__class__.__name__
|
||||
self.kwargs = kwargs
|
||||
for key, val in kwargs.items():
|
||||
setattr(self, key, val)
|
||||
|
||||
def __repr__(self):
|
||||
return call_to_str(self.name, **self.kwargs)
|
||||
|
||||
|
||||
class OptimizerStep(PipeInstruction):
|
||||
"""Performs one step with the optimizer and zeros gradients.
|
||||
|
||||
.. note:: Should be issued after :class:`ReduceGrads` and :class:`ReduceTiedGrads`.
|
||||
|
||||
.. note:: Can be a synchronization point among data-parallel ranks.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class ReduceGrads(PipeInstruction):
|
||||
"""Reduce the computed gradients among data-parallel processes within the stage.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class ReduceTiedGrads(PipeInstruction):
|
||||
"""Reduce the computed gradients of tied modules within a pipeline-parallel group.
|
||||
|
||||
.. warning::
|
||||
The stages included in this synchronization point are not known until
|
||||
the model is partitioned among pipeline stages. In the worst case, it
|
||||
includes all pipeline stages. This instruction should be scheduled
|
||||
carefully to avoid deadlocks.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class BufferOpInstruction(PipeInstruction):
|
||||
"""A pipeline instruction that operates on pipeline buffer(s).
|
||||
|
||||
Args:
|
||||
buffer_id (int): the index of the pipeline buffer() to modify.
|
||||
"""
|
||||
|
||||
def __init__(self, buffer_id, **kwargs):
|
||||
super().__init__(buffer_id=buffer_id, **kwargs)
|
||||
|
||||
|
||||
# IO
|
||||
class LoadMicroBatch(BufferOpInstruction):
|
||||
"""Load a micro-batch into a buffer.
|
||||
|
||||
Roughly:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
buffers['inputs'][buffer_id] = next(data_iter)
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
# Compute
|
||||
class ForwardPass(BufferOpInstruction):
|
||||
"""Compute a forward pass.
|
||||
|
||||
Roughly:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
buffers['outputs'][buffer_id] = forward(buffers['inputs'][buffer_id])
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class BackwardPass(BufferOpInstruction):
|
||||
"""Compute a backward pass and accumulate gradients.
|
||||
|
||||
Roughly:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
outputs = buffers['outputs'][buffer_id]
|
||||
gradients = buffers['gradients'][buffer_id]
|
||||
torch.autograd.backward(tensors=outputs,
|
||||
grad_tensors=gradients)
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
# Communication
|
||||
class SendActivation(BufferOpInstruction):
|
||||
"""Send activations to the next stage in the pipeline.
|
||||
|
||||
Roughly:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
send(buffers['outputs'][buffer_id])
|
||||
|
||||
.. note::
|
||||
The communication is blocking and must be paired with a :class:`RecvActivation`
|
||||
on the next pipeline stage to avoid deadlock.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class RecvActivation(BufferOpInstruction):
|
||||
"""Receive activations from the previous stage in the pipeline.
|
||||
|
||||
Roughly:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
buffers['inputs'][buffer_id] = recv()
|
||||
|
||||
.. note::
|
||||
The communication is blocking and must be paired with a :class:`SendActivation`
|
||||
on the previous pipeline stage to avoid deadlock.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class SendGrad(BufferOpInstruction):
|
||||
"""Send computed gradients to the previous pipeline stage.
|
||||
with respect to the received activations
|
||||
|
||||
.. note::
|
||||
Only received tensors with ``requires_grad==True`` will produce gradients.
|
||||
Missing gradients will be replaced with ``None`` on the receiving stage.
|
||||
|
||||
.. note::
|
||||
The communication is blocking and must be paired with a :class:`RecvGrad`
|
||||
on the previous pipeline stage to avoid deadlock.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class RecvGrad(BufferOpInstruction):
|
||||
"""Receive computed gradients the next pipeline stage.
|
||||
|
||||
.. note::
|
||||
Only activations with ``requires_grad==True`` will produce gradients.
|
||||
Missing gradients will be replaced with ``None``.
|
||||
|
||||
.. note::
|
||||
The communication is blocking and must be paired with a :class:`SendGrad`
|
||||
on the next pipeline stage to avoid deadlock.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
def _is_even(x):
|
||||
return x % 2 == 0
|
||||
|
||||
|
||||
def _is_odd(x):
|
||||
return x % 2 != 0
|
||||
@@ -0,0 +1,472 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from deepspeed import comm as dist
|
||||
|
||||
from collections import namedtuple
|
||||
from itertools import product as cartesian_product
|
||||
|
||||
|
||||
class ProcessTopology:
|
||||
""" Manages the mapping of n-dimensional Cartesian coordinates to linear
|
||||
indices. This mapping is used to map the rank of processes to the grid
|
||||
for various forms of parallelism.
|
||||
|
||||
Each axis of the tensor is accessed by its name. The provided ordering
|
||||
of the axes defines the layout of the topology. ProcessTopology uses a "row-major"
|
||||
layout of the tensor axes, and so axes=['x', 'y'] would map coordinates (x,y) and
|
||||
(x,y+1) to adjacent linear indices. If instead axes=['y', 'x'] was used, coordinates
|
||||
(x,y) and (x+1,y) would be adjacent.
|
||||
|
||||
Some methods return ProcessCoord namedtuples.
|
||||
"""
|
||||
|
||||
def __init__(self, axes, dims):
|
||||
"""Create a mapping of n-dimensional tensor coordinates to linear indices.
|
||||
|
||||
Arguments:
|
||||
axes (list): the names of the tensor axes
|
||||
dims (list): the dimension (length) of each axis of the topology tensor
|
||||
"""
|
||||
|
||||
self.axes = axes # names of each topology axis
|
||||
self.dims = dims # length of each topology axis
|
||||
|
||||
# This is actually a class that lets us hash {'row':3, 'col':2} mappings
|
||||
self.ProcessCoord = namedtuple('ProcessCoord', axes)
|
||||
|
||||
self.mapping = {}
|
||||
ranges = [range(d) for d in dims]
|
||||
# example: 1, (0,0,1)
|
||||
for global_rank, coord in enumerate(cartesian_product(*ranges)):
|
||||
key = {axis: coord[self.axes.index(axis)] for axis in self.axes}
|
||||
key = self.ProcessCoord(**key)
|
||||
# for example, {ProcessCoord(row=0, col=1) : 1}
|
||||
self.mapping[key] = global_rank
|
||||
|
||||
def get_rank(self, **coord_kwargs):
|
||||
"""Return the global rank of a process via its coordinates.
|
||||
|
||||
Coordinates are specified as kwargs. For example:
|
||||
|
||||
>>> X = ProcessTopology(axes=['x', 'y'], dims=[2,3])
|
||||
>>> X.get_rank(x=0, y=1)
|
||||
1
|
||||
"""
|
||||
if len(coord_kwargs) != len(self.axes):
|
||||
raise ValueError('get_rank() does not support slices. Use filter_match())')
|
||||
|
||||
key = self.ProcessCoord(**coord_kwargs)
|
||||
assert key in self.mapping, f'key {coord_kwargs} invalid'
|
||||
return self.mapping[key]
|
||||
|
||||
def get_axis_names(self):
|
||||
"""Return a list of the axis names in the ordering of the topology. """
|
||||
return self.axes
|
||||
|
||||
def get_rank_repr(self, rank, omit_axes=['data', 'pipe'], inner_sep='_', outer_sep='-'):
|
||||
"""Return a string representation of a rank.
|
||||
|
||||
This method is primarily used for checkpointing model data.
|
||||
|
||||
For example:
|
||||
>>> topo = Topo(axes=['a', 'b'], dims=[2, 2])
|
||||
>>> topo.get_rank_repr(rank=3)
|
||||
'a_01-b_01'
|
||||
>>> topo.get_rank_repr(rank=3, omit_axes=['a'])
|
||||
'b_01'
|
||||
|
||||
Args:
|
||||
rank (int): A rank in the topology.
|
||||
omit_axes (list, optional): Axes that should not be in the representation. Defaults to ['data', 'pipe'].
|
||||
inner_sep (str, optional): [description]. Defaults to '_'.
|
||||
outer_sep (str, optional): [description]. Defaults to '-'.
|
||||
|
||||
Returns:
|
||||
str: A string representation of the coordinate owned by ``rank``.
|
||||
"""
|
||||
omit_axes = frozenset(omit_axes)
|
||||
axes = [a for a in self.get_axis_names() if a not in omit_axes]
|
||||
names = []
|
||||
for ax in axes:
|
||||
ax_rank = getattr(self.get_coord(rank=rank), ax)
|
||||
names.append(f'{ax}{inner_sep}{ax_rank:02d}')
|
||||
return outer_sep.join(names)
|
||||
|
||||
def get_dim(self, axis):
|
||||
"""Return the number of processes along the given axis.
|
||||
|
||||
For example:
|
||||
>>> X = ProcessTopology(axes=['x', 'y'], dims=[2,3])
|
||||
>>> X.get_dim('y')
|
||||
3
|
||||
"""
|
||||
if axis not in self.axes:
|
||||
return 0
|
||||
return self.dims[self.axes.index(axis)]
|
||||
|
||||
def get_coord(self, rank):
|
||||
"""Return the coordinate owned by a process rank.
|
||||
|
||||
The axes of the returned namedtuple can be directly accessed as members. For
|
||||
example:
|
||||
>>> X = ProcessTopology(axes=['x', 'y'], dims=[2,3])
|
||||
>>> coord = X.get_coord(rank=1)
|
||||
>>> coord.x
|
||||
0
|
||||
>>> coord.y
|
||||
1
|
||||
"""
|
||||
for coord, idx in self.mapping.items():
|
||||
if idx == rank:
|
||||
return coord
|
||||
raise ValueError(f'rank {rank} not found in topology.')
|
||||
|
||||
def get_axis_comm_lists(self, axis):
|
||||
""" Construct lists suitable for a communicator group along axis ``axis``.
|
||||
|
||||
Example:
|
||||
>>> topo = Topo(axes=['pipe', 'data', 'model'], dims=[2, 2, 2])
|
||||
>>> topo.get_axis_comm_lists('pipe')
|
||||
[
|
||||
[0, 4], # data=0, model=0
|
||||
[1, 5], # data=0, model=1
|
||||
[2, 6], # data=1, model=0
|
||||
[3, 7], # data=1, model=1
|
||||
]
|
||||
|
||||
Returns:
|
||||
A list of lists whose coordinates match in all axes *except* ``axis``.
|
||||
"""
|
||||
|
||||
# We don't want to RuntimeError because it allows us to write more generalized
|
||||
# code for hybrid parallelisms.
|
||||
if axis not in self.axes:
|
||||
return []
|
||||
|
||||
# Grab all axes but `axis`
|
||||
other_axes = [a for a in self.axes if a != axis]
|
||||
|
||||
lists = []
|
||||
|
||||
# Construct all combinations of coords with other_axes
|
||||
ranges = [range(self.get_dim(a)) for a in other_axes]
|
||||
for coord in cartesian_product(*ranges):
|
||||
other_keys = {a: coord[other_axes.index(a)] for a in other_axes}
|
||||
# now go over all ranks in `axis`.
|
||||
sub_list = []
|
||||
for axis_key in range(self.get_dim(axis)):
|
||||
key = self.ProcessCoord(**other_keys, **{axis: axis_key})
|
||||
sub_list.append(self.mapping[key])
|
||||
lists.append(sub_list)
|
||||
|
||||
return lists
|
||||
|
||||
def filter_match(self, **filter_kwargs):
|
||||
"""Return the list of ranks whose coordinates match the provided criteria.
|
||||
|
||||
Example:
|
||||
>>> X = ProcessTopology(axes=['pipe', 'data', 'model'], dims=[2, 2, 2])
|
||||
>>> X.filter_match(pipe=0, data=1)
|
||||
[2, 3]
|
||||
>>> [X.get_coord(rank) for rank in X.filter_match(pipe=0, data=1)]
|
||||
[ProcessCoord(pipe=0, data=1, model=0), ProcessCoord(pipe=0, data=1, model=1)]
|
||||
|
||||
Arguments:
|
||||
**filter_kwargs (dict): criteria used to select coordinates.
|
||||
|
||||
Returns:
|
||||
The list of ranks whose coordinates match filter_kwargs.
|
||||
"""
|
||||
|
||||
def _filter_helper(x):
|
||||
for key, val in filter_kwargs.items():
|
||||
if getattr(x, key) != val:
|
||||
return False
|
||||
return True
|
||||
|
||||
coords = filter(_filter_helper, self.mapping.keys())
|
||||
return [self.mapping[coord] for coord in coords]
|
||||
|
||||
def get_axis_list(self, axis, idx):
|
||||
"""Returns the list of global ranks whose coordinate in an axis is idx.
|
||||
|
||||
For example:
|
||||
>>> X = ProcessTopology(axes=['x', 'y'], dims=[2,3])
|
||||
>>> X.get_axis_list(axis='x', idx=0)
|
||||
[0, 1, 2]
|
||||
>>> X.get_axis_list(axis='y', idx=0)
|
||||
[0, 3]
|
||||
"""
|
||||
|
||||
# This could be faster by generating the desired keys directly instead of
|
||||
# filtering.
|
||||
axis_num = self.axes.index(axis)
|
||||
ranks = [self.mapping[k] for k in self.mapping.keys() if k[axis_num] == idx]
|
||||
return ranks
|
||||
|
||||
def world_size(self):
|
||||
return len(self.mapping)
|
||||
|
||||
def __str__(self):
|
||||
return str(self.mapping)
|
||||
|
||||
|
||||
def _prime_factors(N):
|
||||
""" Returns the prime factorization of positive integer N. """
|
||||
if N <= 0:
|
||||
raise ValueError("Values must be strictly positive.")
|
||||
|
||||
primes = []
|
||||
while N != 1:
|
||||
for candidate in range(2, N + 1):
|
||||
if N % candidate == 0:
|
||||
primes.append(candidate)
|
||||
N //= candidate
|
||||
break
|
||||
return primes
|
||||
|
||||
|
||||
class PipeDataParallelTopology(ProcessTopology):
|
||||
""" A topology specialization for hybrid data and pipeline parallelism.
|
||||
|
||||
Uses data parallelism on the last dimension to encourage gradient
|
||||
reductions to use high-bandwidth intra-node links and lower-volume
|
||||
pipeline communications to use low-bandwidth inter-node links.
|
||||
"""
|
||||
|
||||
def __init__(self, num_pp, num_dp):
|
||||
super().__init__(axes=['pipe', 'data'], dims=[num_pp, num_dp])
|
||||
|
||||
|
||||
class PipeModelDataParallelTopology(ProcessTopology):
|
||||
""" A topology for hybrid pipeline, model, and data parallelism. """
|
||||
|
||||
def __init__(self, num_pp, num_mp, num_dp):
|
||||
super().__init__(axes=['pipe', 'data', 'model'], dims=[num_pp, num_dp, num_mp])
|
||||
|
||||
|
||||
class PipelineParallelGrid:
|
||||
"""Implements a grid object that stores the data parallel ranks
|
||||
corresponding to each of the model parallel stages
|
||||
|
||||
The grid object organizes the processes in a distributed pytorch job
|
||||
into a 2D grid, of stage_id and data_parallel_id.
|
||||
|
||||
self.stage_id and self.data_parallel_id stores the stage id
|
||||
and the data parallel id of current process.
|
||||
|
||||
self.dp_group groups the processes by stage_id.
|
||||
self.dp_group[i], is a list containing all process ranks whose
|
||||
stage_id is i.
|
||||
|
||||
self.p2p_groups stores a list of tuple, where each tuple
|
||||
stores process ranks of adjacent stages for a given data_parallel_id.
|
||||
For example if num_stage is 5 then a tuple [7,8] represents stages [3, 4],
|
||||
with data_parallel id = 1. A stage wrap around will appear as non-adjacent ranks,
|
||||
for example tuple [4,0] with representing wrap-around stage 4 and 0, for
|
||||
data_parallel_id = 0, or similarly [9,5] represents wrapped around stages [4,0]
|
||||
for data_parallel_id = 1.
|
||||
"""
|
||||
|
||||
def __init__(self, topology=None, process_group=None):
|
||||
# TODO use process_group if provided
|
||||
self.global_rank = dist.get_rank()
|
||||
self.world_size = dist.get_world_size()
|
||||
if topology is not None:
|
||||
if self.global_rank == 0:
|
||||
print('Using topology:', topology)
|
||||
self._topo = topology
|
||||
else:
|
||||
num_pp = 1
|
||||
num_dp = 1
|
||||
for idx, prime in enumerate(_prime_factors(self.world_size)):
|
||||
if idx % 2 == 0:
|
||||
num_pp *= prime
|
||||
else:
|
||||
num_dp *= prime
|
||||
self._topo = PipeDataParallelTopology(num_dp=num_dp, num_pp=num_pp)
|
||||
self.data_parallel_size = max(self._topo.get_dim('data'), 1)
|
||||
self.pipe_parallel_size = max(self._topo.get_dim('pipe'), 1)
|
||||
self.model_parallel_size = max(self._topo.get_dim('model'), 1)
|
||||
self.slice_parallel_size = self.model_parallel_size
|
||||
assert self._is_grid_valid(), "Invalid Grid"
|
||||
|
||||
self.stage_id = self.get_stage_id()
|
||||
self.data_parallel_id = self.get_data_parallel_id()
|
||||
|
||||
# Create new ProcessGroups for all model parallelism. DeepSpeedLight uses these
|
||||
# to detect overflow, etc.
|
||||
self.ds_model_proc_group = None
|
||||
self.ds_model_rank = -1
|
||||
for dp in range(self.data_parallel_size):
|
||||
ranks = sorted(self._topo.get_axis_list(axis='data', idx=dp))
|
||||
if self.global_rank == 0:
|
||||
#print(f'RANK={self.global_rank} building DeepSpeed model group: {ranks}')
|
||||
pass
|
||||
proc_group = dist.new_group(ranks=ranks)
|
||||
if self.global_rank in ranks:
|
||||
self.ds_model_proc_group = proc_group
|
||||
self.ds_model_world_size = len(ranks)
|
||||
self.ds_model_rank = ranks.index(self.global_rank)
|
||||
assert self.ds_model_rank > -1
|
||||
assert self.ds_model_proc_group is not None
|
||||
|
||||
# Create new ProcessGroup for gradient all-reduces - these are the data parallel groups
|
||||
self.dp_group = []
|
||||
self.dp_groups = self._topo.get_axis_comm_lists('data')
|
||||
for g in self.dp_groups:
|
||||
proc_group = dist.new_group(ranks=g)
|
||||
if self.global_rank in g:
|
||||
self.dp_group = g
|
||||
self.dp_proc_group = proc_group
|
||||
|
||||
self.is_first_stage = (self.stage_id == 0)
|
||||
self.is_last_stage = (self.stage_id == (self.pipe_parallel_size - 1))
|
||||
|
||||
self.p2p_groups = self._build_p2p_groups()
|
||||
|
||||
# Create new ProcessGroup for pipeline collectives - these are pipe parallel groups
|
||||
self.pp_group = []
|
||||
self.pp_proc_group = None
|
||||
self.pipe_groups = self._topo.get_axis_comm_lists('pipe')
|
||||
for ranks in self.pipe_groups:
|
||||
if self.global_rank == 0:
|
||||
#print(f'RANK={self.global_rank} building pipeline group: {ranks}')
|
||||
pass
|
||||
proc_group = dist.new_group(ranks=ranks)
|
||||
if self.global_rank in ranks:
|
||||
self.pp_group = ranks
|
||||
self.pp_proc_group = proc_group
|
||||
assert self.pp_proc_group is not None
|
||||
|
||||
# Create new ProcessGroup for model (tensor-slicing) collectives
|
||||
|
||||
# Short circuit case without model parallelism.
|
||||
# TODO: it would be nice if topology had bcast semantics to avoid this branching
|
||||
# case?
|
||||
if self.model_parallel_size == 1:
|
||||
for group_rank in range(self.world_size):
|
||||
group_rank = [group_rank]
|
||||
group = dist.new_group(ranks=group_rank)
|
||||
if group_rank[0] == self.global_rank:
|
||||
self.slice_group = group_rank
|
||||
self.slice_proc_group = group
|
||||
return
|
||||
else:
|
||||
self.mp_group = []
|
||||
self.model_groups = self._topo.get_axis_comm_lists('model')
|
||||
for g in self.model_groups:
|
||||
proc_group = dist.new_group(ranks=g)
|
||||
if self.global_rank in g:
|
||||
self.slice_group = g
|
||||
self.slice_proc_group = proc_group
|
||||
|
||||
def get_stage_id(self):
|
||||
return self._topo.get_coord(rank=self.global_rank).pipe
|
||||
|
||||
def get_data_parallel_id(self):
|
||||
return self._topo.get_coord(rank=self.global_rank).data
|
||||
|
||||
def _build_p2p_groups(self):
|
||||
"""Groups for sending and receiving activations and gradients across model
|
||||
parallel stages.
|
||||
"""
|
||||
comm_lists = self._topo.get_axis_comm_lists('pipe')
|
||||
p2p_lists = []
|
||||
for rank in range(self.world_size):
|
||||
for l in comm_lists:
|
||||
assert len(l) == self.pipe_parallel_size
|
||||
if rank in l:
|
||||
idx = l.index(rank)
|
||||
buddy_rank = l[(idx + 1) % self.pipe_parallel_size]
|
||||
p2p_lists.append([rank, buddy_rank])
|
||||
break # next global rank
|
||||
assert len(p2p_lists) == self.world_size
|
||||
return p2p_lists
|
||||
|
||||
def _is_grid_valid(self):
|
||||
ranks = 1
|
||||
for ax in self._topo.get_axis_names():
|
||||
ranks *= self._topo.get_dim(ax)
|
||||
return ranks == dist.get_world_size()
|
||||
|
||||
#returns the global rank of the process with the provided stage id
|
||||
#which has the same data_parallel_id as caller process
|
||||
def stage_to_global(self, stage_id, **kwargs):
|
||||
me = self._topo.get_coord(self.global_rank)
|
||||
transform = me._replace(pipe=stage_id, **kwargs)._asdict()
|
||||
return self._topo.get_rank(**transform)
|
||||
|
||||
def topology(self):
|
||||
return self._topo
|
||||
|
||||
# MPU functions for DeepSpeed integration
|
||||
def get_global_rank(self):
|
||||
return self.global_rank
|
||||
|
||||
def get_pipe_parallel_rank(self):
|
||||
""" The stage of the pipeline this rank resides in. """
|
||||
return self.get_stage_id()
|
||||
|
||||
def get_pipeline_model_parallel_rank(self):
|
||||
return self.get_pipe_parallel_rank()
|
||||
|
||||
def get_pipe_parallel_world_size(self):
|
||||
""" The number of stages in the pipeline. """
|
||||
return self.pipe_parallel_size
|
||||
|
||||
def get_pipeline_model_parallel_world_size(self):
|
||||
return self.get_pipe_parallel_world_size()
|
||||
|
||||
def get_pipe_parallel_group(self):
|
||||
""" The group of ranks within the same pipeline. """
|
||||
return self.pp_proc_group
|
||||
|
||||
def get_data_parallel_rank(self):
|
||||
""" Which pipeline this rank resides in. """
|
||||
return self.data_parallel_id
|
||||
|
||||
def get_data_parallel_world_size(self):
|
||||
""" The number of pipelines. """
|
||||
return self.data_parallel_size
|
||||
|
||||
def get_data_parallel_group(self):
|
||||
""" The group of ranks within the same stage of all pipelines. """
|
||||
return self.dp_proc_group
|
||||
|
||||
def get_data_parallel_group_ranks(self):
|
||||
""" List of ranks in the data parallel group. """
|
||||
return self.dp_group
|
||||
|
||||
# These are model parallel groups across all types of model parallelism.
|
||||
# Deepspeed uses them to detect overflow, etc.
|
||||
def get_model_parallel_rank(self):
|
||||
return self.ds_model_rank
|
||||
|
||||
def get_model_parallel_world_size(self):
|
||||
return self.ds_model_world_size
|
||||
|
||||
def get_model_parallel_group(self):
|
||||
return self.ds_model_proc_group
|
||||
|
||||
# For Megatron-style tensor slicing
|
||||
def get_slice_parallel_rank(self):
|
||||
if 'model' in self._topo.get_axis_names():
|
||||
return self._topo.get_coord(rank=self.global_rank).model
|
||||
else:
|
||||
return 0
|
||||
|
||||
def get_tensor_model_parallel_rank(self):
|
||||
return self.get_slice_parallel_rank()
|
||||
|
||||
def get_slice_parallel_world_size(self):
|
||||
return self.slice_parallel_size
|
||||
|
||||
def get_tensor_model_parallel_world_size(self):
|
||||
return self.get_slice_parallel_world_size()
|
||||
|
||||
def get_slice_parallel_group(self):
|
||||
return self.slice_proc_group
|
||||
Reference in New Issue
Block a user