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 .constants import *
|
||||
from .writer_factory import CheckpointWriterFactory
|
||||
@@ -0,0 +1,77 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from deepspeed.runtime.config_utils import get_scalar_param
|
||||
from .constants import *
|
||||
|
||||
VALID_VALUES = {
|
||||
CHECKPOINT_TAG_VALIDATION: CHECKPOINT_TAG_VALIDATION_MODES,
|
||||
CHECKPOINT_WRITER_TYPE: CHECKPOINT_WRITER_TYPES,
|
||||
CHECKPOINT_DATA_PARALLEL: CHECKPOINT_DATA_PARALLEL_UNITS
|
||||
}
|
||||
|
||||
CHECKPOINT_DEFAULT_DICT = {
|
||||
CHECKPOINT_TAG_VALIDATION: CHECKPOINT_TAG_VALIDATION_DEFAULT,
|
||||
CHECKPOINT_SERIALIZATION: CHECKPOINT_SERIALIZATION_DEFAULT,
|
||||
CHECKPOINT_WRITER: CHECKPOINT_WRITER_DEFAULT
|
||||
}
|
||||
|
||||
|
||||
def _validate_config_values(config_name, config_dict, valid_values):
|
||||
for key, value in config_dict.items():
|
||||
if value is None:
|
||||
continue
|
||||
if key in valid_values.keys():
|
||||
assert value in valid_values[key], \
|
||||
f"{config_name} contains invalid value {value} for {key}, expecting one of {valid_values[key]}"
|
||||
|
||||
|
||||
def _make_upper_case(value):
|
||||
return value if value is None else value.upper()
|
||||
|
||||
|
||||
def get_checkpoint_writer_config(param_dict):
|
||||
writer_dict = param_dict.get(CHECKPOINT_WRITER, None)
|
||||
if writer_dict is None:
|
||||
return CHECKPOINT_WRITER_DEFAULT
|
||||
|
||||
writer_config = {
|
||||
CHECKPOINT_WRITER_TYPE:
|
||||
_make_upper_case(get_scalar_param(writer_dict, CHECKPOINT_WRITER_TYPE, CHECKPOINT_WRITER_TYPE_DEFAULT)),
|
||||
CHECKPOINT_IO_BUFFER_SIZE:
|
||||
get_scalar_param(writer_dict, CHECKPOINT_IO_BUFFER_SIZE, CHECKPOINT_IO_BUFFER_SIZE_DEFAULT),
|
||||
CHECKPOINT_IO_BUFFER_DOUBLE:
|
||||
get_scalar_param(writer_dict, CHECKPOINT_IO_BUFFER_DOUBLE, CHECKPOINT_IO_BUFFER_DOUBLE_DEFAULT),
|
||||
CHECKPOINT_IO_STATISTICS:
|
||||
get_scalar_param(writer_dict, CHECKPOINT_IO_STATISTICS, CHECKPOINT_IO_STATISTICS_DEFAULT),
|
||||
CHECKPOINT_DATA_PARALLEL:
|
||||
_make_upper_case(get_scalar_param(writer_dict, CHECKPOINT_DATA_PARALLEL, CHECKPOINT_DATA_PARALLEL_DEFAULT)),
|
||||
CHECKPOINT_WRITER_DECOUPLED:
|
||||
get_scalar_param(writer_dict, CHECKPOINT_WRITER_DECOUPLED, CHECKPOINT_WRITER_DECOUPLED_DEFAULT),
|
||||
CHECKPOINT_IO_MULTIPLIER:
|
||||
get_scalar_param(writer_dict, CHECKPOINT_IO_MULTIPLIER, CHECKPOINT_IO_MULTIPLIER_DEFAULT),
|
||||
}
|
||||
_validate_config_values(CHECKPOINT_WRITER, writer_config, VALID_VALUES)
|
||||
|
||||
return writer_config
|
||||
|
||||
|
||||
def get_checkpoint_config(param_dict):
|
||||
checkpoint_dict = param_dict.get(CHECKPOINT, None)
|
||||
if checkpoint_dict is None:
|
||||
return CHECKPOINT_DEFAULT_DICT
|
||||
|
||||
checkpoint_config = {
|
||||
CHECKPOINT_TAG_VALIDATION:
|
||||
get_scalar_param(checkpoint_dict, CHECKPOINT_TAG_VALIDATION, CHECKPOINT_TAG_VALIDATION_DEFAULT).upper(),
|
||||
CHECKPOINT_SERIALIZATION:
|
||||
get_scalar_param(checkpoint_dict, CHECKPOINT_SERIALIZATION, CHECKPOINT_SERIALIZATION_DEFAULT),
|
||||
CHECKPOINT_WRITER:
|
||||
get_checkpoint_writer_config(checkpoint_dict)
|
||||
}
|
||||
|
||||
_validate_config_values(CHECKPOINT, checkpoint_config, VALID_VALUES)
|
||||
|
||||
return checkpoint_config
|
||||
@@ -0,0 +1,85 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
|
||||
#########################################
|
||||
# Validation modes
|
||||
#########################################
|
||||
class ValidationMode:
|
||||
WARN = "WARN"
|
||||
IGNORE = "IGNORE"
|
||||
FAIL = "FAIL"
|
||||
|
||||
|
||||
#########################################
|
||||
# Checkpoint config params
|
||||
#########################################
|
||||
# "checkpoint": {tag_validation=["Ignore"|"Warn"|"Fail"]}
|
||||
CHECKPOINT_FORMAT = '''
|
||||
"checkpoint": {
|
||||
"tag_validation": [Ignore|Warn|Fail],
|
||||
"checkpoint_serialization": False,
|
||||
"writer": {
|
||||
"type": [mock|python|fast],
|
||||
"decoupled": [True|False]
|
||||
"io_buffer_size": 64e6,
|
||||
"io_buffer_double": True,
|
||||
"show_statistics": False,
|
||||
"data_parallel": [replica|socket|machine],
|
||||
"io_multiplier": 1,
|
||||
}
|
||||
}
|
||||
'''
|
||||
CHECKPOINT = "checkpoint"
|
||||
CHECKPOINT_TAG_VALIDATION = "tag_validation"
|
||||
CHECKPOINT_TAG_VALIDATION_DEFAULT = ValidationMode.WARN
|
||||
CHECKPOINT_TAG_VALIDATION_MODES = [ValidationMode.WARN, ValidationMode.IGNORE, ValidationMode.FAIL]
|
||||
|
||||
CHECKPOINT_SERIALIZATION = "checkpoint_serialization"
|
||||
CHECKPOINT_SERIALIZATION_DEFAULT = True
|
||||
|
||||
CHECKPOINT_WRITER = "writer"
|
||||
CHECKPOINT_WRITER_DEFAULT = None
|
||||
|
||||
CHECKPOINT_WRITER_TYPE = "type"
|
||||
|
||||
|
||||
class CheckpointWriterType:
|
||||
MOCK = "MOCK"
|
||||
PYTHON = "PYTHON"
|
||||
FAST = "FAST"
|
||||
|
||||
|
||||
CHECKPOINT_WRITER_TYPE_DEFAULT = CheckpointWriterType.FAST
|
||||
CHECKPOINT_WRITER_TYPES = [CheckpointWriterType.MOCK, CheckpointWriterType.PYTHON, CheckpointWriterType.FAST]
|
||||
|
||||
CHECKPOINT_IO_BUFFER_SIZE = "io_buffer_size"
|
||||
CHECKPOINT_IO_BUFFER_SIZE_DEFAULT = 64 * (1024**2)
|
||||
|
||||
CHECKPOINT_IO_BUFFER_DOUBLE = "io_buffer_double"
|
||||
CHECKPOINT_IO_BUFFER_DOUBLE_DEFAULT = True
|
||||
|
||||
CHECKPOINT_IO_MULTIPLIER = "io_multiplier"
|
||||
CHECKPOINT_IO_MULTIPLIER_DEFAULT = 1
|
||||
|
||||
CHECKPOINT_IO_STATISTICS = "show_statistics"
|
||||
CHECKPOINT_IO_STATISTICS_DEFAULT = False
|
||||
|
||||
CHECKPOINT_DATA_PARALLEL = "data_parallel"
|
||||
CHECKPOINT_DATA_PARALLEL_DEFAULT = None
|
||||
|
||||
|
||||
class CheckpointDataParallel:
|
||||
REPLICA = "REPLICA"
|
||||
SOCKET = "SOCKET"
|
||||
MACHINE = "MACHINE"
|
||||
|
||||
|
||||
CHECKPOINT_DATA_PARALLEL_UNITS = [
|
||||
CheckpointDataParallel.REPLICA, CheckpointDataParallel.SOCKET, CheckpointDataParallel.MACHINE
|
||||
]
|
||||
|
||||
CHECKPOINT_WRITER_DECOUPLED = "decoupled"
|
||||
CHECKPOINT_WRITER_DECOUPLED_DEFAULT = False
|
||||
@@ -0,0 +1,216 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from dataclasses import dataclass
|
||||
from deepspeed.checkpoint.reshape_utils import partition_data
|
||||
from deepspeed.runtime.zero.config import ZeroStageEnum
|
||||
from .constants import *
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataParallelWriterConfig(object):
|
||||
world_size: int
|
||||
rank: int
|
||||
global_rank: int
|
||||
local_rank: int
|
||||
pure_dp: bool
|
||||
|
||||
|
||||
class DataParallelWriterFactory(object):
|
||||
|
||||
def __init__(self, uni_parallel_info, parallel_unit):
|
||||
self._uni_parallel_info = uni_parallel_info
|
||||
self._parallel_unit = parallel_unit
|
||||
if parallel_unit == CheckpointDataParallel.SOCKET:
|
||||
self._num_resources = uni_parallel_info.num_sockets
|
||||
else:
|
||||
self._num_resources = uni_parallel_info.num_machines
|
||||
self._ranks_per_resource = max(1, self._uni_parallel_info.global_world_size // self._num_resources)
|
||||
|
||||
def create_config(self, zero_stage, has_moe_layers):
|
||||
if zero_stage == ZeroStageEnum.weights:
|
||||
return self._create_config(1, 0)
|
||||
|
||||
if has_moe_layers:
|
||||
writer_config = self._get_expert_data_parallel_config()
|
||||
else:
|
||||
writer_config = self._get_data_parallel_config()
|
||||
|
||||
if writer_config is None and zero_stage >= ZeroStageEnum.optimizer_states:
|
||||
return self._create_config(1, 0)
|
||||
|
||||
return writer_config
|
||||
|
||||
def _create_config(self, world_size, rank):
|
||||
return DataParallelWriterConfig(world_size=world_size,
|
||||
rank=rank,
|
||||
global_rank=self._uni_parallel_info.global_rank,
|
||||
local_rank=self._uni_parallel_info.local_rank,
|
||||
pure_dp=self._uni_parallel_info.pure_dp)
|
||||
|
||||
def _get_expert_data_parallel_config(self):
|
||||
ep_info = self._uni_parallel_info.ep_info
|
||||
if self._parallel_unit is None:
|
||||
dp_rank = ep_info.dp_rank
|
||||
return self._create_config(1, 0) if dp_rank == 0 else None
|
||||
|
||||
assert self._uni_parallel_info.pure_dp, \
|
||||
'3D parallelism is not yet supported for data parallel checkpointing.'
|
||||
|
||||
if self._parallel_unit == CheckpointDataParallel.REPLICA or ep_info.ep_world_size == 1:
|
||||
return self._get_parallel_write_for_ddp(ep_info.dp_world_size, ep_info.dp_rank)
|
||||
|
||||
return self._get_expert_parallel_write_for_2d()
|
||||
|
||||
def _get_expert_parallel_write_for_2d(self):
|
||||
ep_info = self._uni_parallel_info.ep_info
|
||||
|
||||
def _get_expert_slice_resources(expert_resources, resource_name):
|
||||
ep_world_size = ep_info.ep_world_size
|
||||
slices_per_resource = min(self._ranks_per_resource, ep_world_size)
|
||||
assert slices_per_resource <= len(expert_resources)
|
||||
|
||||
ep_num_resources = len(expert_resources)
|
||||
assert ep_num_resources % slices_per_resource == 0, f'{resource_name}: Expected ep_num_resources={ep_num_resources} to multiple of slices_per_resource={slices_per_resource} for ep_world_size={ep_world_size}'
|
||||
|
||||
slice_partitions = partition_data(expert_resources, slices_per_resource)
|
||||
# print(
|
||||
# f'edp_resource_partition: self._uni_parallel_info.global_rank={self._uni_parallel_info.global_rank} expert_resources={expert_resources} slices_per_resource={slices_per_resource} ep_world_size={ep_world_size} slice_partitions={slice_partitions}'
|
||||
# )
|
||||
resource_index = ep_info.ep_rank % slice_resources
|
||||
return slice_partitions[resource_index]
|
||||
|
||||
dp_ranks = ep_info.dp_peer_ranks
|
||||
expert_resources = [r // self._ranks_per_resource for r in dp_ranks]
|
||||
slice_resources = _get_expert_slice_resources(expert_resources, self._parallel_unit)
|
||||
assert all([idx < self._num_resources for idx in expert_resources]), \
|
||||
f'Detected invalid resource index in expert_resources={expert_resources}, self._num_resources={self._num_resources}'
|
||||
return self._assign_resources_to_tensor_slice(slice_resources, ep_info.ep_rank, dp_ranks)
|
||||
|
||||
def _get_data_parallel_config(self):
|
||||
mpu_info = self._uni_parallel_info.mpu_info
|
||||
if self._parallel_unit is None:
|
||||
dp_rank = self._uni_parallel_info.dp_rank if mpu_info is None else mpu_info.dp_rank
|
||||
return self._create_config(1, 0) if dp_rank == 0 else None
|
||||
|
||||
if self._uni_parallel_info.pure_dp:
|
||||
return self._get_parallel_write_for_ddp(self._uni_parallel_info.global_world_size,
|
||||
self._uni_parallel_info.global_rank)
|
||||
|
||||
if self._parallel_unit == CheckpointDataParallel.REPLICA:
|
||||
return self._create_config(mpu_info.dp_world_size, mpu_info.dp_rank)
|
||||
|
||||
return self._get_parallel_write_for_3d()
|
||||
|
||||
def _get_parallel_write_for_3d(self):
|
||||
mpu_info = self._uni_parallel_info.mpu_info
|
||||
my_global_rank = self._uni_parallel_info.global_rank
|
||||
|
||||
def _expand_resources(resource_list, new_size):
|
||||
old_size = len(resource_list)
|
||||
if old_size >= new_size:
|
||||
return resource_list
|
||||
|
||||
assert new_size % old_size == 0, f'Expect new_size={new_size} to be multiple of old_size={old_size}'
|
||||
multiplier = new_size // old_size
|
||||
new_resource_list = []
|
||||
for r in resource_list:
|
||||
new_resource_list += [r] * multiplier
|
||||
# print(f'expand_resources: {my_global_rank=} {old_size=} {new_size=} {resource_list=} {new_resource_list=}')
|
||||
return new_resource_list
|
||||
|
||||
# Getting resource partition for a tensor slice is a 2-step process
|
||||
# 1. Get resource partitions for all pipeline stages. A pipeline stage is a 2D grid of size TP x DP
|
||||
def _get_pipeline_stage_resources(resource_indices):
|
||||
num_resources = len(resource_indices)
|
||||
pp_world_size = mpu_info.pp_world_size
|
||||
if num_resources < pp_world_size:
|
||||
resource_indices = _expand_resources(resource_indices, pp_world_size)
|
||||
num_resources = pp_world_size
|
||||
global_resource_partitions = partition_data(resource_indices, pp_world_size)
|
||||
pp_rank = mpu_info.pp_rank
|
||||
return global_resource_partitions[pp_rank]
|
||||
|
||||
# 2. Get resource partition for tensor slice. A tensor slice is a 1D vector of size DP
|
||||
def _get_tensor_slice_resources(resource_indices, resource_name):
|
||||
pipe_stage_resources = _get_pipeline_stage_resources(resource_indices)
|
||||
tp_world_size = mpu_info.tp_world_size
|
||||
if len(pipe_stage_resources) < tp_world_size:
|
||||
pipe_stage_resources = _expand_resources(pipe_stage_resources, tp_world_size)
|
||||
tp_num_resources = len(pipe_stage_resources)
|
||||
assert tp_num_resources % tp_world_size == 0, \
|
||||
f'{resource_name}: Expected tp_num_resources={tp_num_resources} to multiple of tp_world_size={tp_world_size}'
|
||||
|
||||
pipe_stage_resource_partitions = partition_data(pipe_stage_resources, tp_world_size)
|
||||
tp_rank = mpu_info.tp_rank
|
||||
return pipe_stage_resource_partitions[tp_rank]
|
||||
|
||||
def _get_model_parallel_slice_resources():
|
||||
# Get resources of my dp peer ranks
|
||||
resources = [(r // self._ranks_per_resource) for r in mpu_info.dp_peer_ranks]
|
||||
if len(resources) < self._ranks_per_resource:
|
||||
resources = _expand_resources(resources, self._ranks_per_resource)
|
||||
|
||||
resource_partitions = partition_data(resources, self._ranks_per_resource)
|
||||
mp_rank = (mpu_info.pp_rank * mpu_info.tp_world_size) + mpu_info.tp_rank
|
||||
slice_rank = mp_rank % self._ranks_per_resource
|
||||
return resource_partitions[slice_rank]
|
||||
|
||||
num_slices = mpu_info.tp_world_size * mpu_info.pp_world_size
|
||||
if num_slices > self._ranks_per_resource:
|
||||
slice_resources = _get_model_parallel_slice_resources()
|
||||
else:
|
||||
all_resources = list(range(self._num_resources))
|
||||
slice_resources = _get_tensor_slice_resources(all_resources, self._parallel_unit)
|
||||
|
||||
return self._assign_resources_to_tensor_slice(slice_resources, mpu_info.tp_rank, mpu_info.dp_peer_ranks)
|
||||
|
||||
def _get_slice_writers(self, slice_resources, my_dp_ranks):
|
||||
resource_map = {}
|
||||
for res in slice_resources:
|
||||
resource_map[res] = [r for r in my_dp_ranks if (r // self._ranks_per_resource) == res]
|
||||
|
||||
# Only one writer per resource, and we conventionally pick the first rank as writer.
|
||||
return [ranks[0] for ranks in resource_map.values()]
|
||||
|
||||
def _assign_resources_to_tensor_slice(self, slice_resources, my_slice_index, my_dp_ranks):
|
||||
my_global_rank = self._uni_parallel_info.global_rank
|
||||
slice_writer_ranks = self._get_slice_writers(slice_resources, my_dp_ranks)
|
||||
my_resource_index = my_global_rank // self._ranks_per_resource
|
||||
print(
|
||||
f'resource_assign: my_global_rank={my_global_rank} my_slice_index={my_slice_index} my_dp_ranks={my_dp_ranks} slice_resources={slice_resources} slice_writer_ranks={slice_writer_ranks}'
|
||||
)
|
||||
if my_resource_index in slice_resources and my_global_rank in slice_writer_ranks:
|
||||
my_writer_index = (my_global_rank - slice_writer_ranks[0]) // self._ranks_per_resource
|
||||
num_slice_writers = len(slice_writer_ranks)
|
||||
print(
|
||||
f'slice_writer: my_global_rank={my_global_rank} my_writer_index={my_writer_index} num_slice_writers={num_slice_writers}'
|
||||
)
|
||||
return self._create_config(num_slice_writers, my_writer_index)
|
||||
|
||||
return None
|
||||
|
||||
def _get_parallel_write_for_ddp(self, dp_world_size, dp_rank):
|
||||
if self._parallel_unit == CheckpointDataParallel.REPLICA:
|
||||
return self._create_config(dp_world_size, dp_rank)
|
||||
|
||||
num_machines = self._uni_parallel_info.num_machines
|
||||
if self._parallel_unit == CheckpointDataParallel.SOCKET:
|
||||
if dp_world_size == num_machines:
|
||||
# There is one rank per machine
|
||||
return self._create_config(num_machines, dp_rank)
|
||||
|
||||
num_sockets = self._uni_parallel_info.num_sockets
|
||||
ranks_per_socket = dp_world_size // num_sockets
|
||||
if dp_rank % ranks_per_socket == 0:
|
||||
return self._create_config(num_sockets, dp_rank // ranks_per_socket)
|
||||
else:
|
||||
return None
|
||||
|
||||
ranks_per_machine = dp_world_size // num_machines
|
||||
if dp_rank % ranks_per_machine == 0:
|
||||
return self._create_config(num_machines, self._uni_parallel_info.machine_rank)
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,84 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from deepspeed import comm as dist
|
||||
from deepspeed.constants import CROSS_RANK, CROSS_SIZE, LOCAL_RANK
|
||||
from .data_parallel_writer_factory import DataParallelWriterFactory
|
||||
|
||||
# TODO: parse socket number from env.
|
||||
SOCKETS_PER_MACHINE = 2
|
||||
|
||||
|
||||
@dataclass
|
||||
class MPUInfo(object):
|
||||
pp_world_size: int
|
||||
pp_rank: int
|
||||
tp_world_size: int
|
||||
tp_rank: int
|
||||
dp_world_size: int
|
||||
dp_peer_ranks: list
|
||||
dp_rank: int
|
||||
|
||||
|
||||
def _create_model_parallel_info(mpu):
|
||||
return MPUInfo(pp_world_size=mpu.get_pipeline_model_parallel_world_size(),
|
||||
pp_rank=mpu.get_pipeline_model_parallel_rank(),
|
||||
tp_world_size=mpu.get_tensor_model_parallel_world_size(),
|
||||
tp_rank=mpu.get_tensor_model_parallel_rank(),
|
||||
dp_world_size=mpu.get_data_parallel_world_size(),
|
||||
dp_peer_ranks=mpu.get_data_parallel_group_ranks(),
|
||||
dp_rank=mpu.get_data_parallel_rank())
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExpertParallelInfo(object):
|
||||
ep_world_size: int
|
||||
ep_rank: int
|
||||
dp_world_size: int
|
||||
dp_peer_ranks: list
|
||||
dp_rank: int
|
||||
|
||||
|
||||
def _create_expert_parallel_info(groups):
|
||||
group_name = groups._get_max_expert_size_name()
|
||||
return ExpertParallelInfo(ep_world_size=groups._get_expert_parallel_world_size(group_name),
|
||||
ep_rank=groups._get_expert_parallel_rank(group_name),
|
||||
dp_world_size=groups._get_expert_data_parallel_world_size(group_name),
|
||||
dp_peer_ranks=groups._get_expert_data_parallel_group_ranks(group_name),
|
||||
dp_rank=groups._get_expert_data_parallel_rank(group_name))
|
||||
|
||||
|
||||
@dataclass
|
||||
class UniversalParallelInfo(object):
|
||||
global_world_size: int
|
||||
global_rank: int
|
||||
local_rank: int
|
||||
mpu_info: MPUInfo
|
||||
ep_info: ExpertParallelInfo
|
||||
pure_dp: bool
|
||||
num_machines: int
|
||||
machine_rank: int
|
||||
num_sockets: int
|
||||
|
||||
|
||||
def create_universal_parallel_info(groups, has_moe_layers):
|
||||
return UniversalParallelInfo(global_world_size=dist.get_world_size(),
|
||||
global_rank=dist.get_rank(),
|
||||
local_rank=int(os.environ[LOCAL_RANK]),
|
||||
mpu_info=None if groups.mpu is None else _create_model_parallel_info(groups.mpu),
|
||||
ep_info=_create_expert_parallel_info(groups) if has_moe_layers else None,
|
||||
pure_dp=groups.mpu is None
|
||||
or groups.mpu.get_data_parallel_world_size() == dist.get_world_size(),
|
||||
num_machines=int(os.environ[CROSS_SIZE]),
|
||||
machine_rank=int(os.environ[CROSS_RANK]),
|
||||
num_sockets=int(os.environ[CROSS_SIZE]) * SOCKETS_PER_MACHINE)
|
||||
|
||||
|
||||
def create_data_parallel_writer_config(groups, parallel_unit, zero_stage, has_moe_layers):
|
||||
uni_parallel_info = create_universal_parallel_info(groups, has_moe_layers)
|
||||
writer_factory = DataParallelWriterFactory(uni_parallel_info, parallel_unit)
|
||||
return writer_factory.create_config(zero_stage, has_moe_layers)
|
||||
@@ -0,0 +1,95 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import torch
|
||||
from deepspeed.ops.op_builder import AsyncIOBuilder, GDSBuilder
|
||||
from deepspeed.io import MockFileWriter, PyFileWriter, FastFileWriter, FastFileWriterConfig
|
||||
from deepspeed.runtime.swap_tensor.constants import *
|
||||
from .constants import *
|
||||
from deepspeed.accelerator import get_accelerator
|
||||
|
||||
|
||||
class CheckpointWriterFactory(object):
|
||||
|
||||
def __init__(self, writer_config, aio_config, dp_writer_config):
|
||||
self._type = writer_config[CHECKPOINT_WRITER_TYPE]
|
||||
self._io_buffer_size = writer_config[CHECKPOINT_IO_BUFFER_SIZE]
|
||||
self._io_buffer_double = writer_config[CHECKPOINT_IO_BUFFER_DOUBLE]
|
||||
self._data_parallel_writer = dp_writer_config
|
||||
self._io_multiplier = writer_config[CHECKPOINT_IO_MULTIPLIER]
|
||||
if self._data_parallel_writer.pure_dp:
|
||||
self._show_statistics = writer_config[CHECKPOINT_IO_STATISTICS] and self._data_parallel_writer is not None
|
||||
else:
|
||||
self._show_statistics = writer_config[CHECKPOINT_IO_STATISTICS] and self._data_parallel_writer is not None
|
||||
self._io_buffer = None
|
||||
self._dnvme_handle = None
|
||||
self._writer = None
|
||||
self._use_gds = False
|
||||
|
||||
if self._type == CheckpointWriterType.FAST:
|
||||
self._use_gds = aio_config[AIO_USE_GDS]
|
||||
if self._use_gds:
|
||||
self._setup_for_gds(aio_config)
|
||||
else:
|
||||
self._setup_for_aio(aio_config)
|
||||
print(
|
||||
f'WriterFactory: self._data_parallel_writer={self._data_parallel_writer} self._show_statistics={self._show_statistics}'
|
||||
)
|
||||
|
||||
def create_writer(self, file_path, optimize_dp_state):
|
||||
assert self._writer is None, \
|
||||
f'Cannot create checkpoint writer for {file_path} because writer is currently used for {self._writer.file_path()}.\
|
||||
Must call writer.release() before reusing to avoid this error.'
|
||||
|
||||
if self._type == CheckpointWriterType.MOCK:
|
||||
self._writer = MockFileWriter(file_path)
|
||||
elif self._type == CheckpointWriterType.PYTHON:
|
||||
self._writer = PyFileWriter(file_path)
|
||||
else:
|
||||
if optimize_dp_state:
|
||||
num_parallel_writers = self._data_parallel_writer.world_size * self._io_multiplier
|
||||
writer_rank = self._data_parallel_writer.rank
|
||||
file_path = f'{file_path}-{writer_rank}.{num_parallel_writers}'
|
||||
# print(f'create_dp_writer: {self._data_parallel_writer.global_rank=} {writer_rank=} {num_parallel_writers=} {file_path=}')
|
||||
else:
|
||||
num_parallel_writers = 1
|
||||
writer_rank = 0
|
||||
# print(f'create_rank0_writer: {self._data_parallel_writer.global_rank=} {writer_rank=} {num_parallel_writers=} {file_path=}')
|
||||
|
||||
config = FastFileWriterConfig(dnvme_handle=self._dnvme_handle,
|
||||
pinned_tensor=self._io_buffer,
|
||||
double_buffer=self._io_buffer_double,
|
||||
num_parallel_writers=num_parallel_writers,
|
||||
writer_rank=writer_rank,
|
||||
global_rank=self._data_parallel_writer.global_rank)
|
||||
self._writer = FastFileWriter(file_path=file_path, config=config)
|
||||
|
||||
return self._writer
|
||||
|
||||
def release_writer(self):
|
||||
self._writer.close()
|
||||
if self._show_statistics:
|
||||
self._writer._dump_state()
|
||||
self._writer = None
|
||||
|
||||
def _setup_for_aio(self, aio_config):
|
||||
self._io_buffer = torch.zeros(self._io_buffer_size, dtype=torch.uint8, device='cpu').pin_memory()
|
||||
self._dnvme_handle = AsyncIOBuilder().load().aio_handle(
|
||||
block_size=aio_config[AIO_BLOCK_SIZE],
|
||||
queue_depth=aio_config[AIO_QUEUE_DEPTH],
|
||||
single_submit=aio_config[AIO_SINGLE_SUBMIT],
|
||||
overlap_events=aio_config[AIO_OVERLAP_EVENTS],
|
||||
intra_op_parallelism=aio_config[AIO_INTRA_OP_PARALLELISM])
|
||||
|
||||
def _setup_for_gds(self, aio_config):
|
||||
self._io_buffer = torch.zeros(self._io_buffer_size,
|
||||
dtype=torch.uint8,
|
||||
device=get_accelerator().current_device_name())
|
||||
self._dnvme_handle = GDSBuilder().load().gds_handle(block_size=aio_config[AIO_BLOCK_SIZE],
|
||||
queue_depth=aio_config[AIO_QUEUE_DEPTH],
|
||||
single_submit=aio_config[AIO_SINGLE_SUBMIT],
|
||||
overlap_events=aio_config[AIO_OVERLAP_EVENTS],
|
||||
intra_op_parallelism=aio_config[AIO_INTRA_OP_PARALLELISM])
|
||||
self._dnvme_handle.pin_device_tensor(self._io_buffer)
|
||||
Reference in New Issue
Block a user