Skip to content
Merged
Show file tree
Hide file tree
Changes from 6 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
147 changes: 147 additions & 0 deletions deepspeed/checkpoint/autoep_universal.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,14 @@
import glob
import torch

from deepspeed.utils import logger

from .constants import (
AUTOEP_EP_SIZE,
AUTOEP_EXPERT_KEY_PREFIX,
AUTOEP_NUM_EXPERTS,
AUTOEP_NUM_LOCAL_EXPERTS,
AUTOEP_ZERO12_REQUIRED_FIELDS,
PARAM,
CAT_DIM,
EP_IS_EXPERT_PARAM,
Expand Down Expand Up @@ -204,6 +211,144 @@ def resolve_expert_ckpt_path(checkpoint_dir, moe_layer_id, global_expert_id):
return matches[0]


def get_autoep_zero12_expert_param_info(autoep_layers_metadata):
"""Validate AutoEP metadata and map fused expert parameter names to layer metadata."""
if not isinstance(autoep_layers_metadata, list) or not autoep_layers_metadata:
raise RuntimeError("AutoEP metadata must be a non-empty list for ZeRO-1/2 universal conversion.")

param_info = {}
for layer_info in autoep_layers_metadata:
if not isinstance(layer_info, dict):
raise RuntimeError("AutoEP layer metadata must contain dictionaries.")

missing_fields = [field for field in AUTOEP_ZERO12_REQUIRED_FIELDS if field not in layer_info]
if missing_fields:
raise RuntimeError(f"AutoEP layer metadata is missing fields: {missing_fields}")

prefix = layer_info[AUTOEP_EXPERT_KEY_PREFIX]
if not isinstance(prefix, str) or not prefix:
raise RuntimeError("AutoEP expert_key_prefix must be a non-empty string.")

for field in (AUTOEP_NUM_EXPERTS, AUTOEP_NUM_LOCAL_EXPERTS, AUTOEP_EP_SIZE):
value = layer_info[field]
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
raise RuntimeError(f"AutoEP {field} must be a positive integer, got {value!r}.")

num_experts = layer_info[AUTOEP_NUM_EXPERTS]
num_local_experts = layer_info[AUTOEP_NUM_LOCAL_EXPERTS]
ep_size = layer_info[AUTOEP_EP_SIZE]
if num_experts != num_local_experts * ep_size:
raise RuntimeError(f"AutoEP expert count mismatch for {prefix}: num_experts={num_experts}, "
f"num_local_experts={num_local_experts}, ep_size={ep_size}.")

normalized = {
'num_experts': num_experts,
'num_local_experts': num_local_experts,
'ep_size': ep_size,
}
for weight_name in ('w1', 'w2', 'w3'):
param_name = f"{prefix}.{weight_name}"
if param_name in param_info:
raise RuntimeError(f"Duplicate AutoEP expert parameter metadata for {param_name}.")
param_info[param_name] = normalized

return param_info


def _zero12_fragment_path(temp_dir, param_name, state_name, dp_rank):
return os.path.join(temp_dir, param_name, "0", f"{state_name}.{dp_rank:0>2d}")


def _autoep_zero12_dp_ranks(ep_rank, dp_degree, ep_size, use_data_before_expert_parallel):
"""Return stage-local ZeRO DP ranks that own one EP rank's fragments."""
if dp_degree % ep_size != 0:
raise RuntimeError(f"ZeRO DP degree {dp_degree} is not divisible by AutoEP size {ep_size}.")
if use_data_before_expert_parallel:
edp_size = dp_degree // ep_size
return range(ep_rank * edp_size, (ep_rank + 1) * edp_size)
return range(ep_rank, dp_degree, ep_size)


def _zero12_values_equal(left, right):
if torch.is_tensor(left) and torch.is_tensor(right):
return torch.equal(left, right)
return left == right


def consolidate_autoep_zero12_expert_states(temp_dir, output_dir, expert_param_info, slice_shapes, dp_degree,
tp_degree, use_data_before_expert_parallel):
"""Consolidate AutoEP expert FP32 and Adam states from ZeRO-1/2 fragments."""
if tp_degree != 1:
raise NotImplementedError("ZeRO-1/2 Universal Checkpoint conversion for AutoEP with tensor parallelism "
"is not supported.")

for param_name, metadata in expert_param_info.items():
if param_name not in slice_shapes:
raise RuntimeError(f"AutoEP expert parameter {param_name} is missing from checkpoint parameter shapes.")

ep_size = metadata['ep_size']
num_experts = metadata['num_experts']
num_local_experts = metadata['num_local_experts']

local_shape = tuple(slice_shapes[param_name])
if not local_shape or local_shape[0] != num_local_experts:
raise RuntimeError(f"AutoEP local shape mismatch for {param_name}: shape={local_shape}, "
f"num_local_experts={num_local_experts}.")

param_dir = os.path.join(output_dir, "zero", param_name)
os.makedirs(param_dir, exist_ok=True)

for state_name in ('fp32', 'exp_avg', 'exp_avg_sq'):
ep_tensors = []
for ep_rank in range(ep_size):
fragments = []
dp_ranks = _autoep_zero12_dp_ranks(ep_rank, dp_degree, ep_size, use_data_before_expert_parallel)
for dp_rank in dp_ranks:
fragment_path = _zero12_fragment_path(temp_dir, param_name, state_name, dp_rank)
if not os.path.isfile(fragment_path):
continue
fragment = torch.load(fragment_path, map_location='cpu', weights_only=False)
if not torch.is_tensor(fragment):
raise RuntimeError(f"AutoEP {state_name} fragment is not a tensor: {fragment_path}")
if fragment.dtype != torch.float32:
raise RuntimeError(f"AutoEP {state_name} fragment must be FP32, got {fragment.dtype} "
f"in {fragment_path}.")
fragments.append(fragment.flatten())

if not fragments:
raise RuntimeError(f"Missing AutoEP {state_name} fragments for {param_name}, EP rank {ep_rank}.")

local_tensor = torch.cat(fragments, dim=0)
expected_numel = torch.Size(local_shape).numel()
if local_tensor.numel() != expected_numel:
raise RuntimeError(f"AutoEP {state_name} fragment size mismatch for {param_name}, "
f"EP rank {ep_rank}: got {local_tensor.numel()}, expected {expected_numel}.")
ep_tensors.append(local_tensor.reshape(local_shape))

full_tensor = torch.cat(ep_tensors, dim=0)
if full_tensor.shape[0] != num_experts:
raise RuntimeError(f"AutoEP consolidated expert count mismatch for {param_name}: "
f"got {full_tensor.shape[0]}, expected {num_experts}.")

torch.save({
PARAM: full_tensor,
CAT_DIM: 0,
EP_IS_EXPERT_PARAM: True,
EP_NUM_EXPERTS: num_experts,
}, os.path.join(param_dir, f"{state_name}.pt"))

step_values = []
for dp_rank in range(dp_degree):
step_path = _zero12_fragment_path(temp_dir, param_name, "step", dp_rank)
if os.path.isfile(step_path):
step_values.append(torch.load(step_path, map_location='cpu', weights_only=False))
if not step_values:
raise RuntimeError(f"Missing AutoEP optimizer step for {param_name}.")
if not all(_zero12_values_equal(step_values[0], value) for value in step_values[1:]):
raise RuntimeError(f"Inconsistent AutoEP optimizer steps for {param_name}.")
torch.save(step_values[0], os.path.join(param_dir, "step.pt"))


def consolidate_autoep_expert_files(checkpoint_dir, output_dir, autoep_layers_metadata):
"""Consolidate per-expert checkpoint files into full-expert universal format.

Expand Down Expand Up @@ -303,6 +448,8 @@ def consolidate_autoep_optimizer_states(checkpoint_dir, output_dir, autoep_layer
# Extract optimizer state dict
optim_sd = optim_states[0].get('optimizer')
if optim_sd is None:
logger.warning("AutoEP per-expert optimizer checkpoint has no optimizer payload; ZeRO-1/2 checkpoints "
"must be consolidated from their ZeRO optimizer shards instead.")
return

state = optim_sd.get('state', {})
Expand Down
13 changes: 13 additions & 0 deletions deepspeed/checkpoint/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,9 @@
LAYER_FILE_PREFIX = 'layer_'
BF16_ZERO_FILE_PREFIX = 'bf16_' + ZERO_FILE_PREFIX
FP16_ZERO_FILE_PREFIX = 'fp16_' + ZERO_FILE_PREFIX
CHECKPOINT_PARALLEL_DIMS = 'checkpoint_parallel_dimensions'
CHECKPOINT_PP_DEGREE = 'pp_degree'
CHECKPOINT_TP_DEGREE = 'tp_degree'

#########################################
# Checkpoint utility keys
Expand Down Expand Up @@ -93,6 +96,16 @@
#########################################
AUTOEP_LAYERS_KEY = 'ds_autoep_layers'
AUTOEP_LAYERS_KEY_LEGACY = 'autoep_layers'
AUTOEP_EXPERT_KEY_PREFIX = 'expert_key_prefix'
AUTOEP_NUM_EXPERTS = 'num_experts'
AUTOEP_NUM_LOCAL_EXPERTS = 'num_local_experts'
AUTOEP_EP_SIZE = 'ep_size'
AUTOEP_ZERO12_REQUIRED_FIELDS = (
AUTOEP_EXPERT_KEY_PREFIX,
AUTOEP_NUM_EXPERTS,
AUTOEP_NUM_LOCAL_EXPERTS,
AUTOEP_EP_SIZE,
)
AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY = 'checkpoint_format'
AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT = 'zero3_partitioned'
AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY = 'checkpoint_format_version'
Expand Down
11 changes: 8 additions & 3 deletions deepspeed/checkpoint/deepspeed_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,11 @@
LAYER_CONCAT_DIM = {'self_attention.dense.weight': 1, 'mlp.dense_4h_to_h.weight': 1}


def _get_pipeline_layer_files(file_list):
return sorted(
[file_path for file_path in file_list if re.fullmatch(LAYER_FILE_PREFIX_PATTERN, os.path.basename(file_path))])


class DeepSpeedCheckpoint(object):

def __init__(self,
Expand All @@ -43,14 +48,14 @@ def __init__(self,
self.final_layer_norm_idx = final_layer_norm_idx
self.dir = dir

pipeline_parallel = len(get_files_with_prefix(get_files(dir), LAYER_FILE_PREFIX)) > 0
self.file_list = get_files(dir)
self.layer_files = _get_pipeline_layer_files(self.file_list)
pipeline_parallel = len(self.layer_files) > 0

self._validate_folder(dir, pipeline_parallel)

self.zero_checkpoint = ZeROCheckpoint(dir)

self.file_list = get_files(dir)
self.layer_files = get_files_with_prefix(self.file_list, LAYER_FILE_PREFIX)
self.mp_rank_files = get_files_with_prefix(self.file_list, MODEL_FILE_PREFIX)

self.layer_keys = self._get_layer_keys()
Expand Down
105 changes: 78 additions & 27 deletions deepspeed/checkpoint/ds_to_universal.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
PARAMETER_WITH_SUB_PARAMS,
AUTOEP_LAYERS_KEY,
AUTOEP_LAYERS_KEY_LEGACY,
AUTOEP_EXPERT_KEY_PREFIX,
EP_IS_EXPERT_PARAM,
EP_NUM_EXPERTS,
EXPERT_PARAMETER_PATTERNS,
Expand Down Expand Up @@ -423,10 +424,14 @@ def _extract_zero_shard_files_stage3(args,
_do_parallel_work(do_work, work_items, args.num_extract_workers)


def _merge_tp_slice_files(args, uc_info, tp_degree, slice_shapes, temp_dir):
def _merge_tp_slice_files(args, uc_info, tp_degree, slice_shapes, temp_dir, exclude_param_names=None):
exclude_param_names = exclude_param_names or set()
zero_output_folder = os.path.join(args.output_folder, "zero")
do_work = partial(merge_tp_slices, uc_info, zero_output_folder, temp_dir, tp_degree)
unmatched_patterns_lists = _do_parallel_work(do_work, list(slice_shapes.items()), args.num_merge_workers)
merge_shapes = [(name, shape) for name, shape in slice_shapes.items() if name not in exclude_param_names]
unmatched_patterns_lists = _do_parallel_work(do_work, merge_shapes, args.num_merge_workers)
if not unmatched_patterns_lists:
return

# verify that all patterns were used
# if a pattern was not used by any of the workers, then it was not used at all -> assert/alert
Expand Down Expand Up @@ -838,6 +843,54 @@ def _classify_autoep_expert_file_consolidation(autoep_metadata, expert_files):
return 'none'


def _aggregate_autoep_zero12_metadata(model_files):
slice_shapes_by_tp = []
metadata_by_prefix = {}
metadata_source_by_prefix = {}
use_data_before_expert_parallel_by_file = {}

for model_file in model_files:
model_state = torch.load(model_file, map_location=torch.device('cpu'), weights_only=False)
slice_shapes_by_tp.append(dict((k, v) for group in model_state[PARAM_SHAPES] for k, v in group.items()))

autoep_metadata = _get_autoep_metadata(model_state)
if autoep_metadata is not None:
if not isinstance(autoep_metadata, list):
raise RuntimeError(f"AutoEP metadata must be a list in {model_file}.")
for layer_info in autoep_metadata:
if not isinstance(layer_info, dict):
raise RuntimeError(f"AutoEP layer metadata must contain dictionaries in {model_file}.")
prefix = layer_info.get(AUTOEP_EXPERT_KEY_PREFIX)
if not isinstance(prefix, str) or not prefix:
raise RuntimeError(f"AutoEP expert_key_prefix must be a non-empty string in {model_file}.")
existing = metadata_by_prefix.get(prefix)
if existing is not None and existing != layer_info:
raise RuntimeError(f"Conflicting AutoEP metadata for expert_key_prefix {prefix!r} in "
f"{metadata_source_by_prefix[prefix]} and {model_file}.")
if existing is None:
metadata_by_prefix[prefix] = layer_info
metadata_source_by_prefix[prefix] = model_file

source_ds_config = model_state.get('ds_config', {})
if not isinstance(source_ds_config, dict):
raise RuntimeError(f"DeepSpeed checkpoint ds_config must be a dictionary in {model_file}.")
use_data_before_expert_parallel = source_ds_config.get('use_data_before_expert_parallelism', False)
if not isinstance(use_data_before_expert_parallel, bool):
raise RuntimeError("DeepSpeed checkpoint use_data_before_expert_parallelism must be a boolean "
f"in {model_file}.")
use_data_before_expert_parallel_by_file[model_file] = use_data_before_expert_parallel
del model_state

use_data_before_expert_parallel_values = set(use_data_before_expert_parallel_by_file.values())
if len(use_data_before_expert_parallel_values) != 1:
raise RuntimeError("DeepSpeed checkpoint use_data_before_expert_parallelism disagrees across model-state "
f"files: {use_data_before_expert_parallel_by_file}")

autoep_metadata = list(metadata_by_prefix.values()) or None
use_data_before_expert_parallel = next(iter(use_data_before_expert_parallel_values))
return slice_shapes_by_tp, autoep_metadata, use_data_before_expert_parallel


def main(args):
print('Convert DeepSpeed Checkpoint to Universal Checkpoint')

Expand All @@ -861,6 +914,8 @@ def main(args):
checkpoint_paths = _create_checkpoint_paths(args.output_folder, iteration, ds_checkpoint.tp_degree,
ds_checkpoint.pp_degree)

slice_shapes_by_tp, autoep_metadata, use_data_before_expert_parallel = _aggregate_autoep_zero12_metadata(
ds_checkpoint.mp_rank_files)
# Each mp_rank file stores one TP rank's PARAM_SHAPES for one PP stage.
# mp_rank_files are ordered pp-major (pp0_tp0, pp0_tp1, ..., pp1_tp0, ...),
# matching meg_2d_parallel_map.simple_init() (i -> pp=i // tp_degree,
Expand All @@ -869,40 +924,36 @@ def main(args):
# and returns one per-TP shape list per parameter name. The per-TP shapes
# let _merge_zero_shards reshape each TP rank's slice to its own shape,
# which matters for uneven TP splits.
slice_shapes_by_tp = []
for mp_rank_file in ds_checkpoint.mp_rank_files:
mp_sd = torch.load(mp_rank_file, map_location=torch.device('cpu'), weights_only=False)
slice_shapes_by_tp.append(dict((k, v) for d in mp_sd[PARAM_SHAPES] for k, v in d.items()))
slice_shapes = _group_per_tp_shapes(slice_shapes_by_tp, ds_checkpoint.pp_degree, ds_checkpoint.tp_degree)
temp_dir = os.path.join(args.output_folder, 'tmp')

print('*** 1. Extracting ZeRO fragments')
_extract_zero_shard_files(args, ds_checkpoint, temp_dir)

print('*** 2. Merging slices .....')
_merge_tp_slice_files(args, ds_checkpoint.get_checkpoint_info(UNIVERSAL_CHECKPOINT_INFO),
ds_checkpoint.tp_degree, slice_shapes, temp_dir)

print('*** 2.5. Consolidating AutoEP expert files')
from deepspeed.checkpoint.autoep_universal import (
consolidate_autoep_expert_files,
consolidate_autoep_optimizer_states,
consolidate_autoep_zero12_expert_states,
get_autoep_zero12_expert_param_info,
)

# Load AutoEP metadata from main checkpoint
main_sd = torch.load(ds_checkpoint.mp_rank_files[0], map_location=torch.device('cpu'), weights_only=False)
autoep_metadata = main_sd.get(AUTOEP_LAYERS_KEY)
if autoep_metadata is None:
autoep_metadata = main_sd.get(AUTOEP_LAYERS_KEY_LEGACY)

# Check for expert files in checkpoint directory
expert_files = glob.glob(os.path.join(args.input_folder, 'layer_*_expert_*_model_states.pt'))
autoep_expert_file_type = _classify_autoep_expert_file_consolidation(autoep_metadata, expert_files)
autoep_expert_param_info = (get_autoep_zero12_expert_param_info(autoep_metadata)
if autoep_expert_file_type == 'autoep' else {})

print('*** 1. Extracting ZeRO fragments')
_extract_zero_shard_files(args, ds_checkpoint, temp_dir)

print('*** 2. Merging slices .....')
_merge_tp_slice_files(args,
ds_checkpoint.get_checkpoint_info(UNIVERSAL_CHECKPOINT_INFO),
ds_checkpoint.tp_degree,
slice_shapes,
temp_dir,
exclude_param_names=set(autoep_expert_param_info))

if autoep_expert_file_type == 'autoep':
consolidate_autoep_expert_files(args.input_folder, args.output_folder, autoep_metadata)
ep_size = autoep_metadata[0]['ep_size'] if autoep_metadata else 1
consolidate_autoep_optimizer_states(args.input_folder, args.output_folder, autoep_metadata, ep_size)
print('*** 2.5. Consolidating AutoEP ZeRO-1/2 expert states')
autoep_slice_shapes = {name: per_tp_shapes[0] for name, per_tp_shapes in slice_shapes.items()}
consolidate_autoep_zero12_expert_states(temp_dir, args.output_folder, autoep_expert_param_info,
autoep_slice_shapes, ds_checkpoint.dp_degree,
ds_checkpoint.tp_degree, use_data_before_expert_parallel)
print(f' Consolidated {len(autoep_metadata)} AutoEP layer(s)')
elif autoep_expert_file_type == 'native_moe':
print(f' Found {len(expert_files)} expert checkpoint file(s) but no AutoEP metadata; '
Expand All @@ -917,7 +968,7 @@ def main(args):
shutil.rmtree(temp_dir, ignore_errors=True)

# Copy mp* files into output folder, injecting AutoEP metadata into UNIVERSAL_CHECKPOINT_INFO
for f in glob.glob(os.path.join(args.input_folder, 'mp*')):
for f in ds_checkpoint.mp_rank_files:
if autoep_metadata is not None:
# Load -> update with AutoEP metadata -> save
mp_sd = torch.load(f, map_location=torch.device('cpu'), weights_only=False)
Expand Down
Loading
Loading