Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
136 changes: 136 additions & 0 deletions deepspeed/checkpoint/autoep_universal.py
Original file line number Diff line number Diff line change
Expand Up @@ -204,6 +204,142 @@ 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.")

required_fields = ('expert_key_prefix', 'num_experts', 'num_local_experts', 'ep_size')
Comment thread
sfc-gh-truwase marked this conversation as resolved.
Outdated
missing_fields = [field for field in 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['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 ('num_experts', 'num_local_experts', '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['num_experts']
num_local_experts = layer_info['num_local_experts']
ep_size = layer_info['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 _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']
if dp_degree % ep_size != 0:
raise RuntimeError(f"ZeRO DP degree {dp_degree} is not divisible by AutoEP size {ep_size} "
f"for {param_name}.")

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 = []
if use_data_before_expert_parallel:
edp_size = dp_degree // ep_size
dp_ranks = range(ep_rank * edp_size, (ep_rank + 1) * edp_size)
else:
dp_ranks = range(ep_rank, dp_degree, ep_size)
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
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
48 changes: 32 additions & 16 deletions deepspeed/checkpoint/ds_to_universal.py
Original file line number Diff line number Diff line change
Expand Up @@ -397,10 +397,14 @@ def _extract_zero_shard_files_stage3(args, optim_files, param_shapes, dp_degree,
_do_parallel_work(do_work, list(range(dp_degree)), args.num_extract_workers)


def _merge_tp_slice_files(args, ds_checkpoint, slice_shapes, temp_dir):
def _merge_tp_slice_files(args, ds_checkpoint, 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, ds_checkpoint, zero_output_folder, temp_dir, ds_checkpoint.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 @@ -788,32 +792,44 @@ def main(args):
slice_shapes = dict((k, v) for d in slice_shapes for k, v in d.items())
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, 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
# Load AutoEP metadata before the generic merge so expert parameters can
# be reconstructed by EP rank from the authoritative ZeRO fragments.
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 {})
source_ds_config = main_sd.get('ds_config', {})
if not isinstance(source_ds_config, dict):
raise RuntimeError("DeepSpeed checkpoint ds_config must be a dictionary.")
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.")

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,
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')
consolidate_autoep_zero12_expert_states(temp_dir, args.output_folder, autoep_expert_param_info,
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 Down
Loading
Loading