diff --git a/deepspeed/checkpoint/affine.py b/deepspeed/checkpoint/affine.py index 8ca076ae6667..ac2a829dc611 100644 --- a/deepspeed/checkpoint/affine.py +++ b/deepspeed/checkpoint/affine.py @@ -234,6 +234,8 @@ def rebuild(self, shards, scale_power=1): written = {} for rank, pieces in self.pieces_by_rank.items(): + if not pieces: + continue flat_shard = _flat_buffer(shards[rank]) for piece in pieces: target = piece.source_view(full_param) diff --git a/deepspeed/checkpoint/affine_ir_spec.md b/deepspeed/checkpoint/affine_ir_spec.md index 4fd1ddf26bfc..993163d6cf15 100644 --- a/deepspeed/checkpoint/affine_ir_spec.md +++ b/deepspeed/checkpoint/affine_ir_spec.md @@ -568,11 +568,22 @@ and fp16 alike, because dividing by `2^k` only shifts the exponent — so a bias bit-exactly at every TP degree in normal use. It is lossy for non-power-of-two `N` (3, 6, 12), where a converted-and-restored bias may differ in the last bits from the original. -**8.4 ZeRO and offload placement.** delock's extension in #8230 — "a subset of a parameter -combined with a list of ranks holding this subset" — is what §2.1's `locations` implements. -ZeRO-1/3 partitions and offload replicas should fall out as pieces whose `locations` -describe the DP group rather than the TP group, but this spec does not yet work through -AutoEP's expert placement, where locations are per-expert rather than per-parameter. +**8.4 AutoEP, ZeRO, and offload placement.** Phase 1 additively introduces a versioned +AutoEP placement descriptor in EP-local rank coordinates and lowers one logical +`[num_experts, ...]` tensor through the existing `AffinePiece` / `ParamAffineMap` IR. The +descriptor records each rank's ordered global expert IDs, including uneven, +non-contiguous, replicated, and empty placements. It describes placement provenance, not +a scheduling policy. ZeRO and EDP fragments remain outside the map: callers first +normalize storage to one logical packed expert tensor per EP rank. + +This phase does not change the current runtime's uniform contiguous scheduling, choose an +arbitrary future expert schedule, or implement direct phase-2 shard-to-shard transfer. +Those remain follow-on work. Phase 2 may derive a target descriptor from runtime +scheduling and transfer directly between source and target maps; until then extraction +from the universal full tensor uses the target map. delock's extension in #8230 — "a +subset of a parameter combined with a list of ranks holding this subset" — is still what +§2.1's exact `locations` implements, and the same mechanism can later describe normalized +ZeRO/offload placement without introducing another geometry IR. --- diff --git a/deepspeed/checkpoint/autoep_affine.py b/deepspeed/checkpoint/autoep_affine.py new file mode 100644 index 000000000000..cd615c8c43de --- /dev/null +++ b/deepspeed/checkpoint/autoep_affine.py @@ -0,0 +1,198 @@ +# SPDX-License-Identifier: Apache-2.0 +# DeepSpeed Team +"""AutoEP expert placement descriptors and their affine lowering. + +The descriptor records placement in EP-local rank coordinates. ZeRO and EDP +fragments are deliberately outside this map: callers must first normalize +storage to one logical, packed expert tensor per EP rank. +""" + +from deepspeed.checkpoint.affine import AffinePiece, ParamAffineMap +from deepspeed.checkpoint.constants import (AUTOEP_PLACEMENT_EP_SIZE, AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_NUM_EXPERTS, AUTOEP_PLACEMENT_RANK, + AUTOEP_PLACEMENT_RANKS, AUTOEP_PLACEMENT_VERSION, + AUTOEP_PLACEMENT_VERSION_KEY) + +__all__ = [ + 'AUTOEP_PLACEMENT_VERSION', + 'make_autoep_placement_descriptor', + 'validate_autoep_placement_descriptor', + 'legacy_uniform_autoep_placement_descriptor', + 'autoep_placement_to_affine_map', + 'extract_autoep_rank_tensor', +] + + +def make_autoep_placement_descriptor(num_experts, experts_by_rank): + """Build a versioned descriptor from expert IDs in rank-local packed order.""" + ranks = [{ + AUTOEP_PLACEMENT_RANK: rank, + AUTOEP_PLACEMENT_EXPERTS: list(experts), + } for rank, experts in enumerate(experts_by_rank)] + descriptor = { + AUTOEP_PLACEMENT_VERSION_KEY: AUTOEP_PLACEMENT_VERSION, + AUTOEP_PLACEMENT_NUM_EXPERTS: num_experts, + AUTOEP_PLACEMENT_EP_SIZE: len(ranks), + AUTOEP_PLACEMENT_RANKS: ranks, + } + validate_autoep_placement_descriptor(descriptor) + return descriptor + + +def validate_autoep_placement_descriptor(descriptor): + """Validate a plain scalar/list/dict AutoEP placement descriptor.""" + if not isinstance(descriptor, dict): + raise TypeError('AutoEP placement descriptor must be a dict.') + + version = descriptor.get(AUTOEP_PLACEMENT_VERSION_KEY) + if version != AUTOEP_PLACEMENT_VERSION: + raise ValueError(f'Unsupported AutoEP placement descriptor version {version!r}; ' + f'expected {AUTOEP_PLACEMENT_VERSION}.') + + num_experts = _positive_int(descriptor.get(AUTOEP_PLACEMENT_NUM_EXPERTS), AUTOEP_PLACEMENT_NUM_EXPERTS) + ep_size = _positive_int(descriptor.get(AUTOEP_PLACEMENT_EP_SIZE), AUTOEP_PLACEMENT_EP_SIZE) + ranks = descriptor.get(AUTOEP_PLACEMENT_RANKS) + if not isinstance(ranks, list): + raise TypeError(f'{AUTOEP_PLACEMENT_RANKS} must be a list.') + if len(ranks) != ep_size: + raise ValueError(f'AutoEP placement descriptor has {len(ranks)} rank entries; expected exactly {ep_size}.') + + seen_ranks = set() + covered_experts = set() + for entry in ranks: + if not isinstance(entry, dict): + raise TypeError('Each AutoEP placement rank entry must be a dict.') + rank = entry.get(AUTOEP_PLACEMENT_RANK) + if not isinstance(rank, int) or isinstance(rank, bool) or not 0 <= rank < ep_size: + raise ValueError(f'AutoEP placement rank {rank!r} is outside [0, {ep_size}).') + if rank in seen_ranks: + raise ValueError(f'AutoEP placement rank {rank} appears more than once.') + seen_ranks.add(rank) + + experts = entry.get(AUTOEP_PLACEMENT_EXPERTS) + if not isinstance(experts, list): + raise TypeError(f'Experts for AutoEP placement rank {rank} must be a list defining local packed order.') + local_experts = set() + for expert_id in experts: + if not isinstance(expert_id, int) or isinstance(expert_id, bool) or not 0 <= expert_id < num_experts: + raise ValueError(f'AutoEP placement expert ID {expert_id!r} on rank {rank} is outside ' + f'[0, {num_experts}).') + if expert_id in local_experts: + raise ValueError(f'AutoEP placement rank {rank} contains duplicate expert ID {expert_id}.') + local_experts.add(expert_id) + covered_experts.add(expert_id) + + missing = sorted(set(range(num_experts)) - covered_experts) + if missing: + raise ValueError(f'AutoEP placement descriptor does not cover global expert IDs {missing}.') + + +def legacy_uniform_autoep_placement_descriptor(num_experts, num_local_experts, ep_size): + """Synthesize the legacy contiguous, uniform AutoEP placement.""" + num_experts = _positive_int(num_experts, 'num_experts') + num_local_experts = _positive_int(num_local_experts, 'num_local_experts') + ep_size = _positive_int(ep_size, 'ep_size') + if num_local_experts * ep_size != num_experts: + raise ValueError(f'Inconsistent legacy AutoEP metadata: num_local_experts ({num_local_experts}) * ' + f'ep_size ({ep_size}) != num_experts ({num_experts}).') + experts_by_rank = [] + for rank in range(ep_size): + start = rank * num_local_experts + experts_by_rank.append(list(range(start, start + num_local_experts))) + return make_autoep_placement_descriptor(num_experts, experts_by_rank) + + +def autoep_placement_to_affine_map(descriptor, logical_shape): + """Lower one ``[num_experts, ...]`` parameter to a :class:`ParamAffineMap`.""" + validate_autoep_placement_descriptor(descriptor) + logical_shape = tuple(int(dim) for dim in logical_shape) + num_experts = descriptor[AUTOEP_PLACEMENT_NUM_EXPERTS] + if not logical_shape or logical_shape[0] != num_experts: + raise ValueError(f'Expert parameter logical shape must start with num_experts ({num_experts}); ' + f'got {logical_shape}.') + if any(dim < 0 for dim in logical_shape): + raise ValueError(f'Expert parameter logical shape cannot contain negative dimensions: {logical_shape}.') + + entries_by_rank = {entry[AUTOEP_PLACEMENT_RANK]: entry for entry in descriptor[AUTOEP_PLACEMENT_RANKS]} + holders = _expert_holders(entries_by_rank, num_experts) + source_strides = _row_major_strides(logical_shape) + expert_shape = logical_shape[1:] + expert_numel = _product(expert_shape) + pieces_by_rank = {} + shard_shapes = {} + + for rank in range(descriptor[AUTOEP_PLACEMENT_EP_SIZE]): + experts = entries_by_rank[rank][AUTOEP_PLACEMENT_EXPERTS] + shard_shape = (len(experts), ) + expert_shape + shard_shapes[rank] = shard_shape + dest_strides = _row_major_strides(shard_shape) + pieces_by_rank[rank] = _pieces_for_rank(experts, holders, expert_shape, expert_numel, source_strides, + dest_strides) + + affine_map = ParamAffineMap(logical_shape=logical_shape, shard_shapes=shard_shapes, pieces_by_rank=pieces_by_rank) + # Descriptor validation already proves expert-level coverage. Expanding that + # check to every tensor element is prohibitively expensive for expert weights. + affine_map.validate() + return affine_map + + +def extract_autoep_rank_tensor(full_param, target_map, ep_rank): + """Extract one EP rank's packed local tensor from a universal full tensor.""" + if not isinstance(target_map, ParamAffineMap): + raise TypeError('target_map must be a ParamAffineMap.') + if ep_rank not in target_map.shard_shapes: + raise ValueError(f'EP rank {ep_rank} is not present in the target affine map.') + return target_map.extract(full_param, ep_rank) + + +def _pieces_for_rank(experts, holders, expert_shape, expert_numel, source_strides, dest_strides): + pieces = [] + run_start = 0 + for local_index in range(1, len(experts) + 1): + at_end = local_index == len(experts) + if not at_end: + previous_expert = experts[local_index - 1] + current_expert = experts[local_index] + homogeneous = (current_expert == previous_expert + 1 + and holders[current_expert] == holders[previous_expert] + and len(holders[current_expert]) == 1) + if at_end or not homogeneous: + first_expert = experts[run_start] + run_length = local_index - run_start + pieces.append( + AffinePiece(shape=(run_length, ) + expert_shape, + source_offset=first_expert * expert_numel, + source_strides=source_strides, + dest_offset=run_start * expert_numel, + dest_strides=dest_strides, + locations=holders[first_expert])) + run_start = local_index + return pieces + + +def _expert_holders(entries_by_rank, num_experts): + holders = {expert_id: set() for expert_id in range(num_experts)} + for rank, entry in entries_by_rank.items(): + for expert_id in entry[AUTOEP_PLACEMENT_EXPERTS]: + holders[expert_id].add(rank) + return holders + + +def _positive_int(value, name): + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f'{name} must be a positive integer; got {value!r}.') + return value + + +def _product(shape): + count = 1 + for dim in shape: + count *= dim + return count + + +def _row_major_strides(shape): + strides = [1] * len(shape) + for axis in range(len(shape) - 2, -1, -1): + strides[axis] = strides[axis + 1] * shape[axis + 1] + return tuple(strides) diff --git a/deepspeed/checkpoint/autoep_universal.py b/deepspeed/checkpoint/autoep_universal.py index 9ecfd11d35a1..c4d82b428f3a 100644 --- a/deepspeed/checkpoint/autoep_universal.py +++ b/deepspeed/checkpoint/autoep_universal.py @@ -16,9 +16,15 @@ from .constants import ( AUTOEP_EP_SIZE, + AUTOEP_EXPERT_PLACEMENT, AUTOEP_EXPERT_KEY_PREFIX, AUTOEP_NUM_EXPERTS, AUTOEP_NUM_LOCAL_EXPERTS, + AUTOEP_PLACEMENT_EP_SIZE, + AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_NUM_EXPERTS, + AUTOEP_PLACEMENT_RANK, + AUTOEP_PLACEMENT_RANKS, AUTOEP_ZERO12_REQUIRED_FIELDS, PARAM, CAT_DIM, @@ -40,6 +46,7 @@ FOLDING_FAMILY, FOLDING_PARAM_FAMILIES, ) +from .autoep_affine import autoep_placement_to_affine_map, validate_autoep_placement_descriptor def make_folding_metadata(*, @@ -229,7 +236,7 @@ def get_autoep_zero12_expert_param_info(autoep_layers_metadata): 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): + for field in (AUTOEP_NUM_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}.") @@ -237,7 +244,33 @@ def get_autoep_zero12_expert_param_info(autoep_layers_metadata): 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: + placement = layer_info.get(AUTOEP_EXPERT_PLACEMENT) + if isinstance(num_local_experts, bool) or not isinstance(num_local_experts, int): + raise RuntimeError(f"AutoEP {AUTOEP_NUM_LOCAL_EXPERTS} must be an integer, got " + f"{num_local_experts!r}.") + minimum_local_experts = 0 if placement is not None else 1 + if num_local_experts < minimum_local_experts: + qualifier = "non-negative" if placement is not None else "positive" + raise RuntimeError(f"AutoEP {AUTOEP_NUM_LOCAL_EXPERTS} must be a {qualifier} integer, got " + f"{num_local_experts!r}.") + if placement is not None: + try: + validate_autoep_placement_descriptor(placement) + except (TypeError, ValueError) as exc: + raise RuntimeError(f"Invalid AutoEP expert placement for {prefix}: {exc}") from exc + if (placement[AUTOEP_PLACEMENT_NUM_EXPERTS] != num_experts + or placement[AUTOEP_PLACEMENT_EP_SIZE] != ep_size): + raise RuntimeError(f"AutoEP expert placement disagrees with layer metadata for {prefix}.") + ep_rank = layer_info.get('ep_rank') + if ep_rank is not None: + rank_entries = {entry[AUTOEP_PLACEMENT_RANK]: entry for entry in placement[AUTOEP_PLACEMENT_RANKS]} + if ep_rank not in rank_entries: + raise RuntimeError(f"AutoEP expert placement does not contain metadata ep_rank {ep_rank}.") + expected_local_experts = len(rank_entries[ep_rank][AUTOEP_PLACEMENT_EXPERTS]) + if num_local_experts != expected_local_experts: + raise RuntimeError("AutoEP num_local_experts disagrees with the placement entry for " + f"EP rank {ep_rank}: {num_local_experts} != {expected_local_experts}.") + elif 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}.") @@ -245,6 +278,7 @@ def get_autoep_zero12_expert_param_info(autoep_layers_metadata): 'num_experts': num_experts, 'num_local_experts': num_local_experts, 'ep_size': ep_size, + 'expert_placement': placement, } for weight_name in ('w1', 'w2', 'w3'): param_name = f"{prefix}.{weight_name}" @@ -288,19 +322,27 @@ def consolidate_autoep_zero12_expert_states(temp_dir, output_dir, expert_param_i ep_size = metadata['ep_size'] num_experts = metadata['num_experts'] - num_local_experts = metadata['num_local_experts'] + placement = metadata.get('expert_placement') local_shape = tuple(slice_shapes[param_name]) - if not local_shape or local_shape[0] != num_local_experts: + if not local_shape: + raise RuntimeError(f"AutoEP local shape is empty for {param_name}.") + + affine_map = None + if placement is not None: + logical_shape = (num_experts, ) + local_shape[1:] + affine_map = autoep_placement_to_affine_map(placement, logical_shape) + elif local_shape[0] != metadata['num_local_experts']: raise RuntimeError(f"AutoEP local shape mismatch for {param_name}: shape={local_shape}, " - f"num_local_experts={num_local_experts}.") + f"num_local_experts={metadata['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 = [] + ep_tensors = {} for ep_rank in range(ep_size): + expected_shape = affine_map.shard_shapes[ep_rank] if affine_map is not None else local_shape fragments = [] dp_ranks = _autoep_zero12_dp_ranks(ep_rank, dp_degree, ep_size, use_data_before_expert_parallel) for dp_rank in dp_ranks: @@ -315,17 +357,24 @@ def consolidate_autoep_zero12_expert_states(temp_dir, output_dir, expert_param_i f"in {fragment_path}.") fragments.append(fragment.flatten()) + expected_numel = torch.Size(expected_shape).numel() + if not fragments and expected_numel == 0: + local_tensor = torch.empty(expected_shape, dtype=torch.float32) + ep_tensors[ep_rank] = local_tensor + continue 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)) + ep_tensors[ep_rank] = local_tensor.reshape(expected_shape) - full_tensor = torch.cat(ep_tensors, dim=0) + if affine_map is not None: + full_tensor = affine_map.rebuild(ep_tensors) + else: + full_tensor = torch.cat([ep_tensors[rank] for rank in range(ep_size)], 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}.") @@ -376,6 +425,11 @@ def consolidate_autoep_expert_files(checkpoint_dir, output_dir, autoep_layers_me moe_layer_id = layer_info['moe_layer_id'] num_experts = layer_info['num_experts'] prefix = layer_info['expert_key_prefix'] + placement = layer_info.get(AUTOEP_EXPERT_PLACEMENT) + if placement is not None: + validate_autoep_placement_descriptor(placement) + if placement[AUTOEP_NUM_EXPERTS] != num_experts: + raise RuntimeError(f"AutoEP expert placement disagrees with num_experts for {prefix}.") for wname in ('w1', 'w2', 'w3'): expert_tensors = [] diff --git a/deepspeed/checkpoint/autoep_zero3_metadata.py b/deepspeed/checkpoint/autoep_zero3_metadata.py index f854d17c2738..6b7766dd428f 100644 --- a/deepspeed/checkpoint/autoep_zero3_metadata.py +++ b/deepspeed/checkpoint/autoep_zero3_metadata.py @@ -4,7 +4,15 @@ # DeepSpeed Team """Shared validation for AutoEP ZeRO-3 checkpoint metadata.""" +from deepspeed.checkpoint.autoep_affine import (legacy_uniform_autoep_placement_descriptor, + validate_autoep_placement_descriptor) from deepspeed.checkpoint.constants import ( + AUTOEP_EXPERT_PLACEMENT, + AUTOEP_PLACEMENT_EP_SIZE, + AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_NUM_EXPERTS, + AUTOEP_PLACEMENT_RANK, + AUTOEP_PLACEMENT_RANKS, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY, @@ -39,6 +47,7 @@ def is_autoep_zero3_partitioned_entry(entry): def validate_autoep_zero3_partitioned_metadata(autoep_metadata, require_partitioned=True, expected_expert_prefixes=None, + expected_runtime_layers=None, version_context="This DeepSpeed build"): if not isinstance(autoep_metadata, list): raise RuntimeError(f"ds_autoep_layers metadata is malformed: expected list, got " @@ -47,6 +56,7 @@ def validate_autoep_zero3_partitioned_metadata(autoep_metadata, seen_layer_ids = set() seen_prefixes = set() partitioned_count = 0 + placements = [] for entry in autoep_metadata: if not isinstance(entry, dict): @@ -66,6 +76,10 @@ def validate_autoep_zero3_partitioned_metadata(autoep_metadata, raise RuntimeError(f"ds_autoep_layers metadata has duplicate expert_key_prefix: {prefix}") seen_prefixes.add(prefix) + placement = _validated_placement(entry) + placements.append(placement) + _validate_runtime_layer(entry, placement, expected_runtime_layers) + if not is_autoep_zero3_partitioned_entry(entry): continue @@ -78,20 +92,21 @@ def validate_autoep_zero3_partitioned_metadata(autoep_metadata, f"{version}. {version_context} supports version " f"{AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION}.") - num_experts = entry['num_experts'] - num_local_experts = entry['num_local_experts'] - ep_size = entry['ep_size'] - if num_local_experts * ep_size != num_experts: - raise RuntimeError("AutoEP ZeRO-3 checkpoint metadata is inconsistent: " - f"num_local_experts={num_local_experts}, ep_size={ep_size}, " - f"num_experts={num_experts}") - - expected_start = entry['ep_rank'] * num_local_experts - expected_end = expected_start + num_local_experts - if entry['global_expert_start'] != expected_start or entry['global_expert_end'] != expected_end: - raise RuntimeError("AutoEP ZeRO-3 checkpoint metadata has inconsistent global expert range: " - f"got [{entry['global_expert_start']}, {entry['global_expert_end']}), " - f"expected [{expected_start}, {expected_end})") + if AUTOEP_EXPERT_PLACEMENT not in entry: + num_experts = entry['num_experts'] + num_local_experts = entry['num_local_experts'] + ep_size = entry['ep_size'] + if num_local_experts * ep_size != num_experts: + raise RuntimeError("AutoEP ZeRO-3 checkpoint metadata is inconsistent: " + f"num_local_experts={num_local_experts}, ep_size={ep_size}, " + f"num_experts={num_experts}") + + expected_start = entry['ep_rank'] * num_local_experts + expected_end = expected_start + num_local_experts + if entry['global_expert_start'] != expected_start or entry['global_expert_end'] != expected_end: + raise RuntimeError("AutoEP ZeRO-3 checkpoint metadata has inconsistent global expert range: " + f"got [{entry['global_expert_start']}, {entry['global_expert_end']}), " + f"expected [{expected_start}, {expected_end})") if expected_expert_prefixes is not None: module_path = entry['module_path'] @@ -107,3 +122,50 @@ def validate_autoep_zero3_partitioned_metadata(autoep_metadata, if require_partitioned and partitioned_count == 0: raise RuntimeError("AutoEP ZeRO-3 partition-native checkpoint metadata was expected but no " "partitioned AutoEP layer entries were found") + return placements + + +def _validated_placement(entry): + try: + if AUTOEP_EXPERT_PLACEMENT in entry: + placement = entry[AUTOEP_EXPERT_PLACEMENT] + validate_autoep_placement_descriptor(placement) + else: + placement = legacy_uniform_autoep_placement_descriptor(entry['num_experts'], entry['num_local_experts'], + entry['ep_size']) + except (TypeError, ValueError) as exc: + raise RuntimeError(f"AutoEP placement metadata is invalid: {exc}") from exc + + if placement[AUTOEP_PLACEMENT_NUM_EXPERTS] != entry['num_experts']: + raise RuntimeError("AutoEP placement metadata num_experts does not match its enclosing layer entry: " + f"{placement[AUTOEP_PLACEMENT_NUM_EXPERTS]} != {entry['num_experts']}") + if placement[AUTOEP_PLACEMENT_EP_SIZE] != entry['ep_size']: + raise RuntimeError("AutoEP placement metadata ep_size does not match its enclosing layer entry: " + f"{placement[AUTOEP_PLACEMENT_EP_SIZE]} != {entry['ep_size']}") + return placement + + +def _validate_runtime_layer(entry, placement, expected_runtime_layers): + ep_rank = entry.get('ep_rank') + if ep_rank is None: + return + ranks = {rank_entry[AUTOEP_PLACEMENT_RANK]: rank_entry for rank_entry in placement[AUTOEP_PLACEMENT_RANKS]} + if ep_rank not in ranks: + raise RuntimeError(f"AutoEP placement metadata does not contain current ep_rank {ep_rank}") + listed_experts = ranks[ep_rank][AUTOEP_PLACEMENT_EXPERTS] + if len(listed_experts) != entry['num_local_experts']: + raise RuntimeError("AutoEP placement metadata current-rank expert count does not match num_local_experts: " + f"{len(listed_experts)} != {entry['num_local_experts']}") + + if expected_runtime_layers is None: + return + module_path = entry['module_path'] + if module_path not in expected_runtime_layers: + raise RuntimeError(f"AutoEP checkpoint metadata references missing module: {module_path}") + runtime = expected_runtime_layers[module_path] + if ep_rank != runtime['ep_rank']: + raise RuntimeError(f"AutoEP checkpoint metadata ep_rank {ep_rank} does not match runtime ep_rank " + f"{runtime['ep_rank']} for module {module_path}") + if listed_experts != runtime['local_experts']: + raise RuntimeError("AutoEP placement metadata current-rank expert order does not match runtime local packing: " + f"{listed_experts} != {runtime['local_experts']} for module {module_path}") diff --git a/deepspeed/checkpoint/constants.py b/deepspeed/checkpoint/constants.py index f84ef8718064..230f22fabc1e 100644 --- a/deepspeed/checkpoint/constants.py +++ b/deepspeed/checkpoint/constants.py @@ -113,6 +113,11 @@ AUTOEP_NUM_EXPERTS = 'num_experts' AUTOEP_NUM_LOCAL_EXPERTS = 'num_local_experts' AUTOEP_EP_SIZE = 'ep_size' +AUTOEP_EXPERT_PLACEMENT = 'expert_placement' +DS_AUTOEP_UC_META = 'ds_autoep_uc_meta' +AUTOEP_PARAM_LOGICAL_SHAPE = 'logical_shape' +AUTOEP_PARAM_EP_RANK = 'ep_rank' +AUTOEP_PARAM_LOCAL_EXPERTS = 'local_experts' AUTOEP_ZERO12_REQUIRED_FIELDS = ( AUTOEP_EXPERT_KEY_PREFIX, AUTOEP_NUM_EXPERTS, @@ -123,6 +128,13 @@ AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT = 'zero3_partitioned' AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY = 'checkpoint_format_version' AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION = 1 +AUTOEP_PLACEMENT_VERSION_KEY = 'version' +AUTOEP_PLACEMENT_VERSION = 1 +AUTOEP_PLACEMENT_NUM_EXPERTS = 'num_experts' +AUTOEP_PLACEMENT_EP_SIZE = 'ep_size' +AUTOEP_PLACEMENT_RANKS = 'ranks' +AUTOEP_PLACEMENT_RANK = 'rank' +AUTOEP_PLACEMENT_EXPERTS = 'experts' ######################################### # Universal Checkpoint EP keys diff --git a/deepspeed/checkpoint/ds_to_universal.py b/deepspeed/checkpoint/ds_to_universal.py index 032a33e7e151..08e90aa09700 100755 --- a/deepspeed/checkpoint/ds_to_universal.py +++ b/deepspeed/checkpoint/ds_to_universal.py @@ -50,7 +50,11 @@ AUTOTP_UNSUPPORTED_PARAMETER_PATTERNS, AUTOEP_LAYERS_KEY, AUTOEP_LAYERS_KEY_LEGACY, + AUTOEP_EXPERT_PLACEMENT, AUTOEP_EXPERT_KEY_PREFIX, + AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_RANK, + AUTOEP_PLACEMENT_RANKS, EP_IS_EXPERT_PARAM, EP_NUM_EXPERTS, EXPERT_PARAMETER_PATTERNS, @@ -61,6 +65,7 @@ validate_autoep_zero3_partitioned_metadata, ) from deepspeed.checkpoint.affine import ParamAffineMap, AFFINE_MAP_FORMAT_VERSION +from deepspeed.checkpoint.autoep_affine import autoep_placement_to_affine_map, validate_autoep_placement_descriptor def parse_arguments(): @@ -743,9 +748,20 @@ def _validate_autoep_expert_shapes(model_states_by_rank, metadata_by_rank): param_shapes = model_states_by_rank[rank][PARAM_SHAPES] zero_shape_names = {name for sub_group_shape in param_shapes for name in sub_group_shape} missing = set(expert_info) - zero_shape_names - if missing: + missing_nonempty = [] + for param_name in missing: + layer_info = expert_info[param_name] + placement = layer_info.get(AUTOEP_EXPERT_PLACEMENT) + if placement is None: + missing_nonempty.append(param_name) + continue + entries_by_rank = {entry[AUTOEP_PLACEMENT_RANK]: entry for entry in placement[AUTOEP_PLACEMENT_RANKS]} + rank_entry = entries_by_rank[layer_info['ep_rank']] + if rank_entry[AUTOEP_PLACEMENT_EXPERTS]: + missing_nonempty.append(param_name) + if missing_nonempty: raise RuntimeError(f"AutoEP expert parameters are missing from rank {rank} ZeRO param_shapes: " - f"{sorted(missing)}") + f"{sorted(missing_nonempty)}") frozen_shapes = model_states_by_rank[rank].get('frozen_param_shapes') or {} frozen_experts = set(expert_info).intersection(frozen_shapes) if frozen_experts: @@ -767,6 +783,34 @@ def _save_zero3_autoep_universal_tensor(output_dir, param_name, state_key, tenso ) +def _rebuild_zero3_autoep_rank_tensors(ep_tensors, placement, logical_shape, context): + if placement is None: + return torch.cat([ep_tensors[rank] for rank in sorted(ep_tensors)], dim=0) + + affine_map = autoep_placement_to_affine_map(placement, logical_shape) + expected_ranks = set(affine_map.shard_shapes) + actual_ranks = set(ep_tensors) + unexpected_ranks = actual_ranks - expected_ranks + missing_nonempty_ranks = { + rank + for rank in expected_ranks - actual_ranks if torch.Size(affine_map.shard_shapes[rank]).numel() != 0 + } + if unexpected_ranks or missing_nonempty_ranks: + raise RuntimeError( + f"Incomplete AutoEP EP-rank tensors for {context}: got {sorted(actual_ranks)}, " + f"expected nonempty ranks " + f"{sorted(rank for rank in expected_ranks if torch.Size(affine_map.shard_shapes[rank]).numel())}.") + for rank, tensor in ep_tensors.items(): + expected_shape = affine_map.shard_shapes[rank] + if tuple(tensor.shape) != expected_shape: + raise RuntimeError(f"AutoEP EP-rank tensor shape mismatch for {context}, rank {rank}: " + f"got {tuple(tensor.shape)}, expected {expected_shape}.") + try: + return affine_map.rebuild(ep_tensors) + except ValueError as exc: + raise RuntimeError(f"Failed to rebuild AutoEP universal tensor for {context}: {exc}") from exc + + def _consolidate_zero3_autoep_expert_states(output_dir, model_files, optim_files): model_rank_map, optim_rank_map = _validate_zero3_model_optim_rank_sets(model_files, optim_files) model_states_by_rank = { @@ -787,6 +831,8 @@ def _consolidate_zero3_autoep_expert_states(output_dir, model_files, optim_files num_experts_by_param = {} expected_dp_world_by_param_rank = {} expected_ep_ranks_by_param = {} + placement_by_param = {} + logical_shape_by_param = {} for rank, model_state in model_states_by_rank.items(): optim_state = optim_states_by_rank.get(rank) @@ -824,9 +870,29 @@ def _consolidate_zero3_autoep_expert_states(output_dir, model_files, optim_files if layer_info is not None: ep_rank = layer_info['ep_rank'] num_experts_by_param[param_name] = layer_info['num_experts'] + placement = layer_info.get(AUTOEP_EXPERT_PLACEMENT) + if placement is not None: + try: + validate_autoep_placement_descriptor(placement) + except (TypeError, ValueError) as exc: + raise RuntimeError(f"Invalid AutoEP expert placement for {param_name}: {exc}") from exc + existing_placement = placement_by_param.setdefault(param_name, placement) + if existing_placement != placement: + raise RuntimeError(f"AutoEP expert placement disagrees across ranks for {param_name}.") + logical_shape = (layer_info['num_experts'], ) + tuple(shape)[1:] + existing_logical_shape = logical_shape_by_param.setdefault(param_name, logical_shape) + if existing_logical_shape != logical_shape: + raise RuntimeError(f"AutoEP expert logical shape disagrees across ranks for {param_name}: " + f"{existing_logical_shape} != {logical_shape}.") + affine_map = autoep_placement_to_affine_map(placement, logical_shape) + expected_shape = affine_map.shard_shapes.get(ep_rank) + if expected_shape is None or tuple(shape) != expected_shape: + raise RuntimeError(f"AutoEP expert shard shape mismatch for {param_name}, EP rank " + f"{ep_rank}: got {tuple(shape)}, expected {expected_shape}.") expected_dp_world_by_param_rank[(param_name, ep_rank)] = layer_info['expert_data_parallel_world_size'] - expected_ep_ranks_by_param[param_name] = set(range(layer_info['ep_size'])) + expected_ep_ranks_by_param[param_name] = (set(affine_map.shard_shapes) if placement is not None + else set(range(layer_info['ep_size']))) for state_key, flat_tensor in flat_state.items(): if flat_tensor is None: raise RuntimeError(f"Missing optimizer state '{state_key}' for AutoEP expert " @@ -843,10 +909,18 @@ def _consolidate_zero3_autoep_expert_states(output_dir, model_files, optim_files for (param_name, state_key), ep_rank_fragments in grouped_by_param.items(): missing_ep_ranks = expected_ep_ranks_by_param[param_name] - set(ep_rank_fragments) + placement = placement_by_param.get(param_name) + logical_shape = logical_shape_by_param.get(param_name) + if placement is not None: + affine_map = autoep_placement_to_affine_map(placement, logical_shape) + missing_ep_ranks = { + rank + for rank in missing_ep_ranks if torch.Size(affine_map.shard_shapes[rank]).numel() != 0 + } if missing_ep_ranks: raise RuntimeError(f"Missing AutoEP universal fragments for {param_name}/{state_key} EP ranks: " f"{sorted(missing_ep_ranks)}") - ep_tensors = [] + ep_tensors = {} for ep_rank in sorted(ep_rank_fragments): fragments = sorted(ep_rank_fragments[ep_rank], key=lambda item: item[0]) expected_dp_world = expected_dp_world_by_param_rank[(param_name, ep_rank)] @@ -863,11 +937,12 @@ def _consolidate_zero3_autoep_expert_states(output_dir, model_files, optim_files raise RuntimeError(f"Inconsistent AutoEP expert fragment shapes for {param_name}/{state_key} " f"EP rank {ep_rank}") full_flat = torch.cat([fragment for _, fragment, _ in fragments], dim=0)[:_shape_numel(shape)] - ep_tensors.append(full_flat.view(shape)) + ep_tensors[ep_rank] = full_flat.view(shape) if not ep_tensors: continue - full_expert_tensor = torch.cat(ep_tensors, dim=0) + full_expert_tensor = _rebuild_zero3_autoep_rank_tensors(ep_tensors, placement, logical_shape, + f"{param_name}/{state_key}") if full_expert_tensor.shape[0] != num_experts_by_param[param_name]: raise RuntimeError(f"AutoEP universal tensor for {param_name}/{state_key} has wrong expert dimension: " f"got {full_expert_tensor.shape[0]}, expected {num_experts_by_param[param_name]}") diff --git a/deepspeed/checkpoint/universal_checkpoint.py b/deepspeed/checkpoint/universal_checkpoint.py index 33e0ee7e9d98..c07da4fc5c41 100644 --- a/deepspeed/checkpoint/universal_checkpoint.py +++ b/deepspeed/checkpoint/universal_checkpoint.py @@ -10,7 +10,11 @@ from typing import List, Tuple, Union from dataclasses import dataclass from .constants import (FP32_WEIGHT_KEY, PARAM, VOCAB_TENSOR, CAT_DIM, PARAM_N_SUB_PARAMS, SUB_PARAM_SHAPE, - EP_IS_EXPERT_PARAM, EP_NUM_EXPERTS, DS_AUTOTP_UC_META, UNIVERSAL_CHECKPOINT_VERSION_KEY) + EP_IS_EXPERT_PARAM, EP_NUM_EXPERTS, DS_AUTOEP_UC_META, DS_AUTOTP_UC_META, + AUTOEP_EXPERT_PLACEMENT, AUTOEP_PARAM_EP_RANK, AUTOEP_PARAM_LOCAL_EXPERTS, + AUTOEP_PARAM_LOGICAL_SHAPE, AUTOEP_PLACEMENT_EXPERTS, AUTOEP_PLACEMENT_RANK, + AUTOEP_PLACEMENT_RANKS, UNIVERSAL_CHECKPOINT_VERSION_KEY) +from .autoep_affine import autoep_placement_to_affine_map, extract_autoep_rank_tensor @dataclass @@ -31,6 +35,75 @@ def _get_param_uc_restore_meta(param): return getattr(param, DS_AUTOTP_UC_META, None) +def _resolve_autoep_partition(current_param, ckpt_dict, full_hp_param, ep_rank): + meta = getattr(current_param, DS_AUTOEP_UC_META, None) + if meta is None: + return None + if not isinstance(meta, dict): + raise RuntimeError(f"AutoEP universal checkpoint target metadata must be a dict, got {type(meta).__name__}.") + + required = { + AUTOEP_EXPERT_PLACEMENT, + AUTOEP_PARAM_LOGICAL_SHAPE, + AUTOEP_PARAM_EP_RANK, + AUTOEP_PARAM_LOCAL_EXPERTS, + } + missing = required - meta.keys() + if missing: + raise RuntimeError(f"AutoEP universal checkpoint target metadata is missing fields {sorted(missing)}.") + + etp_size = meta.get('expert_tensor_parallel_size', meta.get('etp_size', 1)) + if etp_size != 1: + raise NotImplementedError("Universal checkpoint restore for AutoEP expert tensor parallelism is not " + f"supported; got expert_tensor_parallel_size={etp_size}.") + + logical_shape_value = meta[AUTOEP_PARAM_LOGICAL_SHAPE] + if not isinstance(logical_shape_value, (tuple, list)): + raise RuntimeError("AutoEP universal checkpoint target logical_shape must be a tuple or list.") + logical_shape = tuple(logical_shape_value) + + checkpoint_num_experts = ckpt_dict.get(EP_NUM_EXPERTS) + if checkpoint_num_experts is None: + raise RuntimeError(f"AutoEP universal checkpoint is missing '{EP_NUM_EXPERTS}' metadata.") + placement = meta[AUTOEP_EXPERT_PLACEMENT] + placement_num_experts = placement.get('num_experts') if isinstance(placement, dict) else None + if logical_shape and logical_shape[0] != checkpoint_num_experts and placement_num_experts == logical_shape[0]: + ranks = placement.get(AUTOEP_PLACEMENT_RANKS, []) if isinstance(placement, dict) else [] + if ranks and checkpoint_num_experts % len(ranks) == 0: + checkpoint_local_experts = checkpoint_num_experts // len(ranks) + target_local_experts = len(meta[AUTOEP_PARAM_LOCAL_EXPERTS]) + raise ValueError("AutoEP expert shape mismatch: " + f"target_local_experts={target_local_experts}, " + f"checkpoint_local_experts={checkpoint_local_experts}.") + raise ValueError("AutoEP universal checkpoint expert count disagrees with target logical shape: " + f"checkpoint={checkpoint_num_experts}, target={logical_shape}.") + if tuple(full_hp_param.shape) != logical_shape: + raise RuntimeError("AutoEP universal checkpoint tensor shape disagrees with target logical shape: " + f"checkpoint={tuple(full_hp_param.shape)}, target={logical_shape}.") + if not logical_shape or logical_shape[0] != checkpoint_num_experts: + raise RuntimeError("AutoEP universal checkpoint expert count disagrees with target logical shape: " + f"checkpoint={checkpoint_num_experts}, target={logical_shape}.") + + try: + target_map = autoep_placement_to_affine_map(meta[AUTOEP_EXPERT_PLACEMENT], logical_shape) + if not isinstance(ep_rank, int) or isinstance(ep_rank, bool): + raise ValueError(f"ep_rank must be an integer, got {ep_rank!r}.") + metadata_ep_rank = meta[AUTOEP_PARAM_EP_RANK] + if metadata_ep_rank != ep_rank: + raise ValueError(f"target ep_rank {ep_rank} does not match parameter metadata ep_rank {metadata_ep_rank}.") + entries_by_rank = { + entry[AUTOEP_PLACEMENT_RANK]: entry + for entry in meta[AUTOEP_EXPERT_PLACEMENT][AUTOEP_PLACEMENT_RANKS] + } + expected_local_experts = entries_by_rank[ep_rank][AUTOEP_PLACEMENT_EXPERTS] + if meta[AUTOEP_PARAM_LOCAL_EXPERTS] != expected_local_experts: + raise ValueError("parameter metadata local_experts does not match the selected placement rank: " + f"{meta[AUTOEP_PARAM_LOCAL_EXPERTS]} != {expected_local_experts}.") + return extract_autoep_rank_tensor(full_hp_param, target_map, ep_rank) + except (AssertionError, KeyError, TypeError, ValueError) as exc: + raise RuntimeError(f"Invalid AutoEP universal checkpoint target metadata: {exc}") from exc + + def _narrow_sub_params(full_view, partition_dim, sub_dim_sizes, shard_widths, tp_rank, tp_world_size, uc_version): """Take this rank's piece of every sub-parameter and concatenate them back together. @@ -173,7 +246,11 @@ def load_hp_checkpoint_state(self, folder, tp_rank, tp_world_size, ep_rank=0, ep # Must happen BEFORE shape-match check so that after slicing, # full_hp_param.shape == self.shape triggers tp_rank=0, tp_world_size=1. is_expert_param = ckpt_dict.get(EP_IS_EXPERT_PARAM, False) - if is_expert_param and ep_size > 1: + autoep_hp_param = (_resolve_autoep_partition(self, ckpt_dict, full_hp_param, ep_rank) + if is_expert_param else None) + if autoep_hp_param is not None: + full_hp_param = autoep_hp_param + elif is_expert_param and ep_size > 1: ep_num_experts = ckpt_dict.get(EP_NUM_EXPERTS) assert ep_num_experts is not None, \ f"Expert param in {ckpt_file} missing '{EP_NUM_EXPERTS}' metadata" diff --git a/deepspeed/module_inject/auto_ep_layer.py b/deepspeed/module_inject/auto_ep_layer.py index 4b8fcec70b88..78ab72e55e44 100644 --- a/deepspeed/module_inject/auto_ep_layer.py +++ b/deepspeed/module_inject/auto_ep_layer.py @@ -19,6 +19,10 @@ import torch import torch.nn as nn import deepspeed.comm as dist +from deepspeed.checkpoint.autoep_affine import legacy_uniform_autoep_placement_descriptor +from deepspeed.checkpoint.constants import (AUTOEP_EXPERT_PLACEMENT, AUTOEP_PARAM_EP_RANK, AUTOEP_PARAM_LOCAL_EXPERTS, + AUTOEP_PARAM_LOGICAL_SHAPE, AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_RANKS, DS_AUTOEP_UC_META) from deepspeed.module_inject.auto_ep_config import AutoEPConfig, MoELayerSpec, resolve_autoep_config_defaults from deepspeed.module_inject.auto_ep_folding import mark_autoep_folding_router_parameter from deepspeed.ops.triton_ops import autoep_fused_token_ops as fused_token_ops @@ -415,6 +419,8 @@ def __init__( self.ep_rank = ep_rank self.num_experts = spec.num_experts self.num_local_experts = spec.num_experts // ep_size + self.expert_placement_descriptor = legacy_uniform_autoep_placement_descriptor( + self.num_experts, self.num_local_experts, self.ep_size) self.hidden_size = spec.hidden_size self.ep_group_name = f"ep_size_{ep_size}" self.ep_group = None # Set by set_deepspeed_parallelism() @@ -507,6 +513,21 @@ def __init__( self.experts.w1.requires_grad_(w1_requires_grad) self.experts.w2.requires_grad_(w2_requires_grad) self.experts.w3.requires_grad_(w3_requires_grad) + local_experts = self.expert_placement_descriptor[AUTOEP_PLACEMENT_RANKS][ + self.ep_rank][AUTOEP_PLACEMENT_EXPERTS] + for param in (self.experts.w1, self.experts.w2, self.experts.w3): + physical_shape = getattr(param, 'ds_shape', param.shape) + logical_shape = [self.num_experts, *physical_shape[1:]] + setattr( + param, + DS_AUTOEP_UC_META, + { + AUTOEP_EXPERT_PLACEMENT: self.expert_placement_descriptor, + AUTOEP_PARAM_LOGICAL_SHAPE: logical_shape, + AUTOEP_PARAM_EP_RANK: self.ep_rank, + AUTOEP_PARAM_LOCAL_EXPERTS: list(local_experts), + }, + ) self.shared_experts = getattr(source_module, spec.shared_experts_name, None) if spec.has_shared_experts else None diff --git a/deepspeed/runtime/engine.py b/deepspeed/runtime/engine.py index 06f302f5011f..b7d6dbfa267d 100644 --- a/deepspeed/runtime/engine.py +++ b/deepspeed/runtime/engine.py @@ -66,6 +66,9 @@ DATA_PARALLEL_GROUP, GLOBAL_RANK, DDP_BFLOAT16, GRADIENT_ALLREDUCE_OP_MEAN from deepspeed.runtime.zero.config import ZeroStageEnum from deepspeed.checkpoint.constants import ( + AUTOEP_EXPERT_PLACEMENT, + AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_RANKS, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY, @@ -4272,24 +4275,9 @@ def load_moe_state_dict(checkpoint_path, else: # Validate AutoEP metadata if present if autoep_layers is not None: - if not isinstance(autoep_layers, list): - raise RuntimeError( - f"ds_autoep_layers metadata is malformed: expected list, got {type(autoep_layers).__name__}") - seen_ids = set() - required_fields = { - 'moe_layer_id', 'module_path', 'num_experts', 'num_local_experts', 'ep_size', 'expert_key_prefix' - } - for entry in autoep_layers: - if not isinstance(entry, dict): - raise RuntimeError( - f"ds_autoep_layers entry is malformed: expected dict, got {type(entry).__name__}") - missing = required_fields - entry.keys() - if missing: - raise RuntimeError(f"ds_autoep_layers entry is invalid: missing fields {sorted(missing)}") - lid = entry['moe_layer_id'] - if lid in seen_ids: - raise RuntimeError(f"ds_autoep_layers metadata has duplicate moe_layer_id: {lid}") - seen_ids.add(lid) + DeepSpeedEngine._validate_autoep_zero3_partitioned_metadata(autoep_layers, + model=model, + require_partitioned=False) elif has_autoep_layers: logger.warning("Checkpoint does not contain ds_autoep_layers metadata. " "Loading AutoEP expert weights using best-effort module detection.") @@ -4601,13 +4589,17 @@ def _uses_autoep_zero3_partitioned_experts(autoep_layers): return any(is_autoep_zero3_partitioned_entry(entry) for entry in autoep_layers) @staticmethod - def _validate_autoep_zero3_partitioned_metadata(autoep_layers, model=None, require_partitioned=True): + def _validate_autoep_zero3_partitioned_metadata(autoep_layers, + model=None, + require_partitioned=True, + validate_runtime_placement=False): try: from deepspeed.module_inject.auto_ep_layer import AutoEPMoELayer as _AutoEPMoELayer except ImportError: _AutoEPMoELayer = None expected_expert_prefixes = None + expected_runtime_layers = None if _AutoEPMoELayer is not None and model is not None: expected_expert_prefixes = { module_name: f"{module_name}.experts" if module_name else "experts" @@ -4615,10 +4607,23 @@ def _validate_autoep_zero3_partitioned_metadata(autoep_layers, model=None, requi } if not expected_expert_prefixes: expected_expert_prefixes = None + if validate_runtime_placement: + expected_runtime_layers = {} + for module_name, module in model.named_modules(): + if not isinstance(module, _AutoEPMoELayer): + continue + rank_entry = module.expert_placement_descriptor[AUTOEP_PLACEMENT_RANKS][module.ep_rank] + expected_runtime_layers[module_name] = { + 'ep_rank': module.ep_rank, + 'local_experts': list(rank_entry[AUTOEP_PLACEMENT_EXPERTS]), + } + if not expected_runtime_layers: + expected_runtime_layers = None validate_autoep_zero3_partitioned_metadata(autoep_layers, require_partitioned=require_partitioned, expected_expert_prefixes=expected_expert_prefixes, + expected_runtime_layers=expected_runtime_layers, version_context="This DeepSpeed build") @staticmethod @@ -5290,6 +5295,8 @@ def autoep_expert_writer() -> bool: num_local_experts, 'ep_size': module.ep_size, + AUTOEP_EXPERT_PLACEMENT: + module.expert_placement_descriptor, 'expert_key_prefix': f"{module_prefix}experts", AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY: @@ -5354,6 +5361,12 @@ def autoep_expert_writer() -> bool: moe_layer_id += 1 + if autoep_layer_info: + DeepSpeedEngine._validate_autoep_zero3_partitioned_metadata(autoep_layer_info, + model=self.module, + require_partitioned=False, + validate_runtime_placement=True) + self._curr_ckpt_path = os.path.join(save_dir, tag) largest_group_name = groups._get_max_expert_size_name() diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 0afd8c1b9f89..3aa3d44ec166 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -3529,11 +3529,17 @@ def load_hp_checkpoint_state(self, folder, key, param=None): def _slice_autoep_universal_expert_param(self, checkpoint_state, param): full_expert_tensor = checkpoint_state[PARAM] - checkpoint_num_experts = checkpoint_state.get(EP_NUM_EXPERTS, full_expert_tensor.shape[0]) group_name = getattr(param, "ds_zero_partition_group_name", None) if group_name is None: raise ValueError("AutoEP universal expert checkpoint target parameter is missing its EP group name") ep_rank = groups._get_expert_parallel_rank(group_name) + + from deepspeed.checkpoint.universal_checkpoint import _resolve_autoep_partition + affine_partition = _resolve_autoep_partition(param, checkpoint_state, full_expert_tensor, ep_rank) + if affine_partition is not None: + return affine_partition + + checkpoint_num_experts = checkpoint_state.get(EP_NUM_EXPERTS, full_expert_tensor.shape[0]) ep_world_size = groups._get_expert_parallel_world_size(group_name) if checkpoint_num_experts % ep_world_size != 0: raise ValueError("AutoEP universal expert checkpoint tensor cannot be evenly split across the target " diff --git a/tests/unit/checkpoint/test_autoep_affine.py b/tests/unit/checkpoint/test_autoep_affine.py new file mode 100644 index 000000000000..57aa5a9251e3 --- /dev/null +++ b/tests/unit/checkpoint/test_autoep_affine.py @@ -0,0 +1,185 @@ +# SPDX-License-Identifier: Apache-2.0 +# DeepSpeed Team + +import json + +import pytest +import torch + +from deepspeed.checkpoint.affine import ParamAffineMap +from deepspeed.checkpoint.autoep_affine import (AUTOEP_PLACEMENT_VERSION, autoep_placement_to_affine_map, + extract_autoep_rank_tensor, legacy_uniform_autoep_placement_descriptor, + make_autoep_placement_descriptor, validate_autoep_placement_descriptor) + + +def _oracle_shards(full_tensor, experts_by_rank): + return { + rank: + torch.stack([full_tensor[expert_id] + for expert_id in experts]) if experts else full_tensor.new_empty((0, ) + + tuple(full_tensor.shape[1:])) + for rank, experts in enumerate(experts_by_rank) + } + + +@pytest.mark.parametrize('num_experts, ep_size', [(8, 2), (8, 4)]) +def test_uniform_placement_lowers_and_round_trips(num_experts, ep_size): + num_local_experts = num_experts // ep_size + descriptor = legacy_uniform_autoep_placement_descriptor(num_experts, num_local_experts, ep_size) + affine_map = autoep_placement_to_affine_map(descriptor, (num_experts, 3, 2)) + full_tensor = torch.arange(num_experts * 6, dtype=torch.float32).reshape(num_experts, 3, 2) + experts_by_rank = [ + list(range(rank * num_local_experts, (rank + 1) * num_local_experts)) for rank in range(ep_size) + ] + expected_shards = _oracle_shards(full_tensor, experts_by_rank) + + assert affine_map.shard_shapes == {rank: (num_local_experts, 3, 2) for rank in range(ep_size)} + assert all(len(pieces) == 1 for pieces in affine_map.pieces_by_rank.values()) + for rank in range(ep_size): + assert torch.equal(affine_map.extract(full_tensor, rank), expected_shards[rank]) + assert torch.equal(affine_map.rebuild(expected_shards), full_tensor) + + +def test_non_contiguous_non_uniform_placement_preserves_local_order(): + experts_by_rank = [[4, 1, 5], [0], [3, 2]] + descriptor = make_autoep_placement_descriptor(6, experts_by_rank) + affine_map = autoep_placement_to_affine_map(descriptor, (6, 2)) + full_tensor = torch.tensor([[10, 11], [20, 21], [30, 31], [40, 41], [50, 51], [60, 61]]) + expected_shards = _oracle_shards(full_tensor, experts_by_rank) + + assert affine_map.shard_shapes == {0: (3, 2), 1: (1, 2), 2: (2, 2)} + assert torch.equal(affine_map.extract(full_tensor, 0), torch.tensor([[50, 51], [20, 21], [60, 61]])) + assert torch.equal(affine_map.rebuild(expected_shards), full_tensor) + + +def test_replication_records_exact_holders_and_stops_piece_merging(): + experts_by_rank = [[0, 1, 2], [1, 2, 3]] + descriptor = make_autoep_placement_descriptor(4, experts_by_rank) + affine_map = autoep_placement_to_affine_map(descriptor, (4, 2)) + full_tensor = torch.arange(8).reshape(4, 2) + + assert [piece.locations + for piece in affine_map.pieces_by_rank[0]] == [frozenset({0}), + frozenset({0, 1}), + frozenset({0, 1})] + assert [piece.shape for piece in affine_map.pieces_by_rank[0]] == [(1, 2), (1, 2), (1, 2)] + assert [piece.locations + for piece in affine_map.pieces_by_rank[1]] == [frozenset({0, 1}), + frozenset({0, 1}), + frozenset({1})] + assert torch.equal(affine_map.rebuild(_oracle_shards(full_tensor, experts_by_rank)), full_tensor) + + +def test_replication_disagreement_with_different_local_order_fails(): + experts_by_rank = [[0, 1, 3, 4], [2, 4, 3]] + descriptor = make_autoep_placement_descriptor(5, experts_by_rank) + affine_map = autoep_placement_to_affine_map(descriptor, (5, 2)) + full_tensor = torch.arange(10, dtype=torch.float32).reshape(5, 2) + shards = _oracle_shards(full_tensor, experts_by_rank) + shards[1][1].add_(1) + + with pytest.raises(ValueError, match='different data'): + affine_map.rebuild(shards) + + +def test_empty_rank_shard_is_representable(): + descriptor = make_autoep_placement_descriptor(3, [[2, 0], [], [1]]) + affine_map = autoep_placement_to_affine_map(descriptor, (3, 4)) + full_tensor = torch.arange(12).reshape(3, 4) + + assert affine_map.shard_shapes[1] == (0, 4) + assert affine_map.pieces_by_rank[1] == [] + assert extract_autoep_rank_tensor(full_tensor, affine_map, 1).shape == (0, 4) + assert torch.equal(affine_map.rebuild(_oracle_shards(full_tensor, [[2, 0], [], [1]])), full_tensor) + shards_without_empty_rank = _oracle_shards(full_tensor, [[2, 0], [], [1]]) + del shards_without_empty_rank[1] + assert torch.equal(affine_map.rebuild(shards_without_empty_rank), full_tensor) + + +@pytest.mark.parametrize('descriptor, match', [ + ({ + 'version': AUTOEP_PLACEMENT_VERSION, + 'num_experts': 2, + 'ep_size': 2, + 'ranks': [{ + 'rank': 0, + 'experts': [0] + }], + }, 'exactly 2'), + ({ + 'version': AUTOEP_PLACEMENT_VERSION, + 'num_experts': 2, + 'ep_size': 2, + 'ranks': [{ + 'rank': 0, + 'experts': [0] + }, { + 'rank': 0, + 'experts': [1] + }], + }, 'rank 0 appears more than once'), + ({ + 'version': AUTOEP_PLACEMENT_VERSION, + 'num_experts': 2, + 'ep_size': 1, + 'ranks': [{ + 'rank': 0, + 'experts': [0, 0, 1] + }], + }, 'duplicate expert ID 0'), + ({ + 'version': AUTOEP_PLACEMENT_VERSION, + 'num_experts': 3, + 'ep_size': 2, + 'ranks': [{ + 'rank': 0, + 'experts': [0] + }, { + 'rank': 1, + 'experts': [2] + }], + }, 'does not cover global expert IDs \\[1\\]'), + ({ + 'version': AUTOEP_PLACEMENT_VERSION, + 'num_experts': 2, + 'ep_size': 1, + 'ranks': [{ + 'rank': 0, + 'experts': [0, 2] + }], + }, 'outside \\[0, 2\\)'), +]) +def test_malformed_placement_is_rejected(descriptor, match): + with pytest.raises(ValueError, match=match): + validate_autoep_placement_descriptor(descriptor) + + +@pytest.mark.parametrize('num_experts, num_local_experts, ep_size', [(8, 3, 2), (0, 1, 1), (4, 0, 2)]) +def test_inconsistent_legacy_metadata_is_rejected(num_experts, num_local_experts, ep_size): + with pytest.raises(ValueError): + legacy_uniform_autoep_placement_descriptor(num_experts, num_local_experts, ep_size) + + +def test_descriptor_and_affine_map_serialization_round_trip(): + descriptor = make_autoep_placement_descriptor(4, [[2, 0], [1, 2, 3]]) + serialized_descriptor = json.loads(json.dumps(descriptor)) + validate_autoep_placement_descriptor(serialized_descriptor) + affine_map = autoep_placement_to_affine_map(serialized_descriptor, (4, 3)) + restored_map = ParamAffineMap.from_dict(json.loads(json.dumps(affine_map.to_dict()))) + full_tensor = torch.randn(4, 3) + shards = _oracle_shards(full_tensor, [[2, 0], [1, 2, 3]]) + + assert serialized_descriptor == descriptor + assert restored_map.to_dict() == affine_map.to_dict() + assert torch.equal(restored_map.rebuild(shards), full_tensor) + + +def test_extract_target_rank_uses_target_packed_order(): + descriptor = make_autoep_placement_descriptor(5, [[0, 2], [4, 1, 3]]) + target_map = autoep_placement_to_affine_map(descriptor, (5, 2, 2)) + universal_tensor = torch.arange(20).reshape(5, 2, 2) + expected = torch.stack([universal_tensor[4], universal_tensor[1], universal_tensor[3]]) + + actual = extract_autoep_rank_tensor(universal_tensor, target_map, 1) + + assert torch.equal(actual, expected) diff --git a/tests/unit/checkpoint/test_autoep_affine_universal.py b/tests/unit/checkpoint/test_autoep_affine_universal.py new file mode 100644 index 000000000000..885998a9a030 --- /dev/null +++ b/tests/unit/checkpoint/test_autoep_affine_universal.py @@ -0,0 +1,494 @@ +# SPDX-License-Identifier: Apache-2.0 +# DeepSpeed Team + +import os +from types import SimpleNamespace + +import pytest +import torch + +from deepspeed.checkpoint.autoep_affine import make_autoep_placement_descriptor +from deepspeed.checkpoint.autoep_universal import (consolidate_autoep_zero12_expert_states, + get_autoep_zero12_expert_param_info) +from deepspeed.checkpoint.constants import ( + AUTOEP_EXPERT_PLACEMENT, AUTOEP_LAYERS_KEY, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY, + AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT, AUTOEP_PARAM_LOCAL_EXPERTS, AUTOEP_PLACEMENT_EXPERTS, + AUTOEP_PLACEMENT_RANK, AUTOEP_PLACEMENT_RANKS, DS_AUTOEP_UC_META, DS_AUTOTP_UC_META, EP_IS_EXPERT_PARAM, + EP_NUM_EXPERTS, OPTIMIZER_STATE_DICT, PARAM, PARAM_SHAPES) +from deepspeed.checkpoint.ds_to_universal import (_consolidate_zero3_autoep_expert_states, + _rebuild_zero3_autoep_rank_tensors) +from deepspeed.checkpoint.universal_checkpoint import _resolve_autoep_partition, load_hp_checkpoint_state +from deepspeed.runtime.zero import stage3 +from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3 + + +def _write_zero12_fragments(temp_dir, param_name, rank_tensors): + for state_name in ('fp32', 'exp_avg', 'exp_avg_sq'): + for dp_rank, tensor in rank_tensors.items(): + path = os.path.join(temp_dir, param_name, "0", f"{state_name}.{dp_rank:02d}") + os.makedirs(os.path.dirname(path), exist_ok=True) + torch.save(tensor.flatten().float(), path) + for dp_rank in rank_tensors: + torch.save(7, os.path.join(temp_dir, param_name, "0", f"step.{dp_rank:02d}")) + + +def _load_consolidated(output_dir, param_name, state_name='fp32'): + state = torch.load(os.path.join(output_dir, "zero", param_name, f"{state_name}.pt"), + map_location='cpu', + weights_only=False) + assert state[EP_IS_EXPERT_PARAM] + return state + + +def test_zero12_affine_consolidation_supports_nonuniform_noncontiguous_placement(tmp_path): + param_name = "layer.experts.w1" + placement = make_autoep_placement_descriptor(4, [[2, 0, 3], [1]]) + full = torch.arange(8, dtype=torch.float32).view(4, 2) + rank_tensors = { + 0: torch.stack([full[2], full[0], full[3]]), + 1: full[1:2], + } + _write_zero12_fragments(str(tmp_path / "fragments"), param_name, rank_tensors) + info = { + param_name: { + 'num_experts': 4, + 'num_local_experts': 2, + 'ep_size': 2, + 'expert_placement': placement, + } + } + + consolidate_autoep_zero12_expert_states(str(tmp_path / "fragments"), str(tmp_path / "output"), info, + {param_name: (2, 2)}, 2, 1, False) + + for state_name in ('fp32', 'exp_avg', 'exp_avg_sq'): + state = _load_consolidated(str(tmp_path / "output"), param_name, state_name) + assert state[EP_NUM_EXPERTS] == 4 + assert torch.equal(state[PARAM], full) + + +def test_zero12_metadata_accepts_explicit_empty_rank(): + placement = make_autoep_placement_descriptor(2, [[], [0, 1]]) + metadata = [{ + 'moe_layer_id': 0, + 'module_path': 'layer', + 'num_experts': 2, + 'num_local_experts': 0, + 'ep_size': 2, + 'ep_rank': 0, + 'expert_key_prefix': 'layer.experts', + AUTOEP_EXPERT_PLACEMENT: placement, + }] + + param_info = get_autoep_zero12_expert_param_info(metadata) + + assert param_info['layer.experts.w1']['num_local_experts'] == 0 + + +def test_zero12_uniform_affine_consolidation_matches_legacy_order(tmp_path): + param_name = "layer.experts.w1" + placement = make_autoep_placement_descriptor(4, [[0, 1], [2, 3]]) + full = torch.arange(8, dtype=torch.float32).view(4, 2) + _write_zero12_fragments(str(tmp_path / "fragments"), param_name, {0: full[:2], 1: full[2:]}) + info = { + param_name: { + 'num_experts': 4, + 'num_local_experts': 2, + 'ep_size': 2, + 'expert_placement': placement, + } + } + + consolidate_autoep_zero12_expert_states(str(tmp_path / "fragments"), str(tmp_path / "output"), info, + {param_name: (2, 2)}, 2, 1, False) + + assert torch.equal(_load_consolidated(str(tmp_path / "output"), param_name)[PARAM], full) + + +def test_zero12_affine_consolidation_validates_replicated_experts(tmp_path): + param_name = "layer.experts.w2" + placement = make_autoep_placement_descriptor(3, [[0, 1], [1, 2]]) + full = torch.arange(6, dtype=torch.float32).view(3, 2) + rank_tensors = { + 0: torch.stack([full[0], full[1]]), + 1: torch.stack([full[1] + 1, full[2]]), + } + _write_zero12_fragments(str(tmp_path / "fragments"), param_name, rank_tensors) + info = { + param_name: { + 'num_experts': 3, + 'num_local_experts': 2, + 'ep_size': 2, + 'expert_placement': placement, + } + } + + with pytest.raises(ValueError, match="different data"): + consolidate_autoep_zero12_expert_states(str(tmp_path / "fragments"), str(tmp_path / "output"), info, + {param_name: (2, 2)}, 2, 1, False) + + +def test_zero12_legacy_consolidation_still_concatenates_ep_ranks(tmp_path): + param_name = "layer.experts.w3" + full = torch.arange(8, dtype=torch.float32).view(4, 2) + _write_zero12_fragments(str(tmp_path / "fragments"), param_name, {0: full[:2], 1: full[2:]}) + info = { + param_name: { + 'num_experts': 4, + 'num_local_experts': 2, + 'ep_size': 2, + 'expert_placement': None, + } + } + + consolidate_autoep_zero12_expert_states(str(tmp_path / "fragments"), str(tmp_path / "output"), info, + {param_name: (2, 2)}, 2, 1, False) + + assert torch.equal(_load_consolidated(str(tmp_path / "output"), param_name)[PARAM], full) + + +def test_zero3_affine_rebuild_checks_ranks_shapes_and_replicas(): + placement = make_autoep_placement_descriptor(3, [[2, 0], [1, 2]]) + full = torch.arange(6, dtype=torch.float32).view(3, 2) + shards = {0: torch.stack([full[2], full[0]]), 1: torch.stack([full[1], full[2]])} + + assert torch.equal(_rebuild_zero3_autoep_rank_tensors(shards, placement, full.shape, "test"), full) + with pytest.raises(RuntimeError, match="expected nonempty ranks \\[0, 1\\]"): + _rebuild_zero3_autoep_rank_tensors({0: shards[0]}, placement, full.shape, "test") + with pytest.raises(RuntimeError, match="different data"): + _rebuild_zero3_autoep_rank_tensors({ + 0: shards[0], + 1: shards[1] + torch.tensor([[0, 0], [1, 0]]) + }, placement, full.shape, "test") + + +def test_zero3_affine_rebuild_allows_omitted_empty_rank(): + placement = make_autoep_placement_descriptor(3, [[2, 0, 1], []]) + full = torch.arange(6, dtype=torch.float32).view(3, 2) + shards = {0: torch.stack([full[2], full[0], full[1]])} + + assert torch.equal(_rebuild_zero3_autoep_rank_tensors(shards, placement, full.shape, "test"), full) + + +def test_zero3_partition_native_consolidation_uses_affine_placement(tmp_path): + param_names = [f"layer.experts.w{index}" for index in range(1, 4)] + placement = make_autoep_placement_descriptor(4, [[2, 0, 3], [1]]) + full = torch.arange(8, dtype=torch.float32).view(4, 2) + rank_tensors = { + 0: torch.stack([full[2], full[0], full[3]]), + 1: full[1:2], + } + model_files = [] + optim_files = [] + for rank, local_tensor in rank_tensors.items(): + layer_info = { + 'moe_layer_id': 0, + 'module_path': 'layer', + 'num_experts': 4, + 'num_local_experts': local_tensor.shape[0], + 'ep_size': 2, + 'expert_key_prefix': 'layer.experts', + AUTOEP_EXPERT_PLACEMENT: placement, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY: AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY: AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, + 'ep_group_name': 'ep_size_2', + 'ep_rank': rank, + 'expert_data_parallel_rank': 0, + 'expert_data_parallel_world_size': 1, + 'global_expert_start': 0, + 'global_expert_end': 0, + } + model_file = tmp_path / f"zero_pp_rank_{rank}_mp_rank_00_model_states.pt" + torch.save( + { + AUTOEP_LAYERS_KEY: [layer_info], + PARAM_SHAPES: [{ + param_name: local_tensor.shape + } for param_name in param_names], + }, model_file) + model_files.append(str(model_file)) + + flat = local_tensor.flatten() + optim_file = tmp_path / f"bf16_zero_pp_rank_{rank}_mp_rank_00_optim_states.pt" + torch.save( + { + OPTIMIZER_STATE_DICT: { + 'ds_zero_partition_groups': [{ + 'partition_count': 1, + 'partition_rank': 0, + } for _ in param_names], + 'optimizer_state_dict': { + 'state': { + index: { + 'exp_avg': flat.clone(), + 'exp_avg_sq': flat.clone(), + } + for index in range(len(param_names)) + } + }, + 'fp32_flat_groups': [flat.clone() for _ in param_names], + } + }, optim_file) + optim_files.append(str(optim_file)) + + output_dir = tmp_path / "universal" + _consolidate_zero3_autoep_expert_states(str(output_dir), model_files, optim_files) + + for param_name in param_names: + for state_name in ('fp32', 'exp_avg', 'exp_avg_sq'): + state = _load_consolidated(str(output_dir), param_name, state_name) + assert torch.equal(state[PARAM], full) + + +def test_zero3_partition_native_consolidation_allows_rank_without_fragments(tmp_path): + param_names = [f"layer.experts.w{index}" for index in range(1, 4)] + placement = make_autoep_placement_descriptor(3, [[2, 0, 1], []]) + full = torch.arange(6, dtype=torch.float32).view(3, 2) + model_files = [] + optim_files = [] + for rank in range(2): + local_tensor = torch.stack([full[2], full[0], full[1]]) if rank == 0 else None + layer_info = { + 'moe_layer_id': 0, + 'module_path': 'layer', + 'num_experts': 3, + 'num_local_experts': 3 if rank == 0 else 0, + 'ep_size': 2, + 'expert_key_prefix': 'layer.experts', + AUTOEP_EXPERT_PLACEMENT: placement, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY: AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY: AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, + 'ep_group_name': 'ep_size_2', + 'ep_rank': rank, + 'expert_data_parallel_rank': 0, + 'expert_data_parallel_world_size': 1, + 'global_expert_start': 0, + 'global_expert_end': 0, + } + model_file = tmp_path / f"zero_pp_rank_{rank}_mp_rank_00_model_states.pt" + param_shapes = [] if local_tensor is None else [{param_name: local_tensor.shape} for param_name in param_names] + torch.save({ + AUTOEP_LAYERS_KEY: [layer_info], + PARAM_SHAPES: param_shapes, + }, model_file) + model_files.append(str(model_file)) + + optim_file = tmp_path / f"bf16_zero_pp_rank_{rank}_mp_rank_00_optim_states.pt" + if local_tensor is None: + optimizer_state = { + 'ds_zero_partition_groups': [], + 'optimizer_state_dict': { + 'state': {} + }, + 'fp32_flat_groups': [], + } + else: + flat = local_tensor.flatten() + optimizer_state = { + 'ds_zero_partition_groups': [{ + 'partition_count': 1, + 'partition_rank': 0, + } for _ in param_names], + 'optimizer_state_dict': { + 'state': { + index: { + 'exp_avg': flat.clone(), + 'exp_avg_sq': flat.clone(), + } + for index in range(len(param_names)) + } + }, + 'fp32_flat_groups': [flat.clone() for _ in param_names], + } + torch.save({OPTIMIZER_STATE_DICT: optimizer_state}, optim_file) + optim_files.append(str(optim_file)) + + output_dir = tmp_path / "universal" + _consolidate_zero3_autoep_expert_states(str(output_dir), model_files, optim_files) + + for param_name in param_names: + for state_name in ('fp32', 'exp_avg', 'exp_avg_sq'): + state = _load_consolidated(str(output_dir), param_name, state_name) + assert torch.equal(state[PARAM], full) + + +def _target_param(placement, logical_shape, ep_rank): + entries_by_rank = {entry[AUTOEP_PLACEMENT_RANK]: entry for entry in placement[AUTOEP_PLACEMENT_RANKS]} + param = torch.nn.Parameter(torch.empty(1)) + setattr( + param, + DS_AUTOEP_UC_META, + { + AUTOEP_EXPERT_PLACEMENT: placement, + 'logical_shape': list(logical_shape), + 'ep_rank': ep_rank, + AUTOEP_PARAM_LOCAL_EXPERTS: list(entries_by_rank[ep_rank][AUTOEP_PLACEMENT_EXPERTS]), + }, + ) + return param + + +def test_autoep_target_extraction_precedes_tp_and_stage3_partitioning(monkeypatch): + placement = make_autoep_placement_descriptor(5, [[0, 3], [4, 1, 2]]) + full = torch.arange(20, dtype=torch.float32).view(5, 2, 2) + checkpoint_state = { + PARAM: full, + EP_IS_EXPERT_PARAM: True, + EP_NUM_EXPERTS: 5, + } + param = _target_param(placement, full.shape, 1) + expected = torch.stack([full[4], full[1], full[2]]) + + param.ds_zero_partition_group_name = "ep" + monkeypatch.setattr(stage3.groups, "_get_expert_parallel_rank", lambda _: 1) + + assert torch.equal(_resolve_autoep_partition(param, checkpoint_state, full, 1), expected) + assert torch.equal( + DeepSpeedZeroOptimizer_Stage3._slice_autoep_universal_expert_param(None, checkpoint_state, param), expected) + + +def test_zero12_restore_loads_target_packed_order(tmp_path): + placement = make_autoep_placement_descriptor(5, [[0, 3], [4, 1, 2]]) + full = torch.arange(10, dtype=torch.float32).view(5, 2) + expected = torch.stack([full[4], full[1], full[2]]) + checkpoint_dir = tmp_path / "checkpoint" + checkpoint_dir.mkdir() + torch.save({ + PARAM: full, + EP_IS_EXPERT_PARAM: True, + EP_NUM_EXPERTS: 5, + }, checkpoint_dir / "fp32.pt") + + param = torch.nn.Parameter(torch.empty_like(expected)) + setattr( + param, + DS_AUTOEP_UC_META, + { + AUTOEP_EXPERT_PLACEMENT: placement, + 'logical_shape': list(full.shape), + 'ep_rank': 1, + AUTOEP_PARAM_LOCAL_EXPERTS: [4, 1, 2], + }, + ) + destination = torch.empty(expected.numel(), dtype=torch.float32) + param._hp_mapping = SimpleNamespace( + optim_fragment={}, + lp_fragment_address=SimpleNamespace(start=0, numel=expected.numel()), + get_hp_fragment=lambda: destination, + ) + + load_hp_checkpoint_state(param, str(checkpoint_dir), tp_rank=0, tp_world_size=1, ep_rank=1, ep_size=2) + + assert torch.equal(destination.view_as(expected), expected) + + +def test_zero12_restore_applies_autoep_before_independent_autotp(tmp_path): + placement = make_autoep_placement_descriptor(4, [[2, 0], [3, 1]]) + full = torch.arange(16, dtype=torch.float32).view(4, 4) + ep_local = torch.stack([full[2], full[0]]) + checkpoint_dir = tmp_path / "checkpoint" + checkpoint_dir.mkdir() + torch.save({ + PARAM: full, + EP_IS_EXPERT_PARAM: True, + EP_NUM_EXPERTS: 4, + }, checkpoint_dir / "fp32.pt") + + param = _target_param(placement, full.shape, 0) + setattr( + param, + DS_AUTOTP_UC_META, + { + 'logical_shape': list(ep_local.shape), + 'partition_dim': 1, + 'partition_sizes': [2, 2], + }, + ) + destination = torch.empty(4, dtype=torch.float32) + param._hp_mapping = SimpleNamespace( + optim_fragment={}, + lp_fragment_address=SimpleNamespace(start=0, numel=4), + get_hp_fragment=lambda: destination, + ) + + load_hp_checkpoint_state(param, str(checkpoint_dir), tp_rank=1, tp_world_size=2, ep_rank=0, ep_size=2) + + assert torch.equal(destination.view(2, 2), ep_local[:, 2:]) + + +def test_zero3_restore_extracts_ep_rank_before_zero_partition(tmp_path, monkeypatch): + placement = make_autoep_placement_descriptor(5, [[0, 3], [4, 1, 2]]) + full = torch.arange(10, dtype=torch.float32).view(5, 2) + expected_local = torch.stack([full[4], full[1], full[2]]) + checkpoint_dir = tmp_path / "checkpoint" + checkpoint_dir.mkdir() + torch.save({ + PARAM: full, + EP_IS_EXPERT_PARAM: True, + EP_NUM_EXPERTS: 5, + }, checkpoint_dir / "fp32.pt") + param = _target_param(placement, full.shape, 1) + param.ds_zero_partition_group_name = "ep" + partition_group = object() + optimizer = SimpleNamespace( + _get_param_partition_group=lambda _: partition_group, + _slice_autoep_universal_expert_param=lambda checkpoint_state, target_param: DeepSpeedZeroOptimizer_Stage3. + _slice_autoep_universal_expert_param(None, checkpoint_state, target_param), + ) + monkeypatch.setattr(stage3.dist, "get_rank", lambda group=None: 1) + monkeypatch.setattr(stage3.dist, "get_world_size", lambda group=None: 2) + monkeypatch.setattr(stage3.groups, "_get_expert_parallel_rank", lambda _: 1) + + actual = DeepSpeedZeroOptimizer_Stage3.load_hp_checkpoint_state(optimizer, str(checkpoint_dir), "fp32", param) + + assert torch.equal(actual, expected_local.flatten()[3:]) + + +@pytest.mark.parametrize( + 'metadata, match', + [ + ({}, "missing fields"), + ({ + AUTOEP_EXPERT_PLACEMENT: make_autoep_placement_descriptor(2, [[0], [1]]), + 'logical_shape': [3, 2], + 'ep_rank': 0, + AUTOEP_PARAM_LOCAL_EXPERTS: [0], + }, "tensor shape disagrees"), + ({ + AUTOEP_EXPERT_PLACEMENT: make_autoep_placement_descriptor(2, [[0], [1]]), + 'logical_shape': [2, 2], + 'ep_rank': 2, + AUTOEP_PARAM_LOCAL_EXPERTS: [], + }, "Invalid AutoEP"), + ], +) +def test_autoep_target_metadata_fails_loudly(metadata, match): + full = torch.arange(4, dtype=torch.float32).view(2, 2) + checkpoint_state = {PARAM: full, EP_IS_EXPERT_PARAM: True, EP_NUM_EXPERTS: 2} + param = torch.nn.Parameter(torch.empty(1)) + setattr(param, DS_AUTOEP_UC_META, metadata) + + with pytest.raises(RuntimeError, match=match): + _resolve_autoep_partition(param, checkpoint_state, full, metadata.get('ep_rank', 0)) + + +def test_autoep_target_metadata_rejects_caller_ep_rank_mismatch(): + placement = make_autoep_placement_descriptor(2, [[0], [1]]) + full = torch.arange(4, dtype=torch.float32).view(2, 2) + checkpoint_state = {PARAM: full, EP_IS_EXPERT_PARAM: True, EP_NUM_EXPERTS: 2} + param = _target_param(placement, full.shape, 0) + + with pytest.raises(RuntimeError, match="target ep_rank 1 does not match parameter metadata ep_rank 0"): + _resolve_autoep_partition(param, checkpoint_state, full, 1) + + +def test_autoep_target_metadata_rejects_local_expert_list_mismatch(): + placement = make_autoep_placement_descriptor(3, [[2, 0], [1]]) + full = torch.arange(6, dtype=torch.float32).view(3, 2) + checkpoint_state = {PARAM: full, EP_IS_EXPERT_PARAM: True, EP_NUM_EXPERTS: 3} + param = _target_param(placement, full.shape, 0) + getattr(param, DS_AUTOEP_UC_META)[AUTOEP_PARAM_LOCAL_EXPERTS] = [0, 2] + + with pytest.raises(RuntimeError, match="local_experts does not match"): + _resolve_autoep_partition(param, checkpoint_state, full, 0) diff --git a/tests/unit/checkpoint/test_autoep_zero3_metadata.py b/tests/unit/checkpoint/test_autoep_zero3_metadata.py new file mode 100644 index 000000000000..38b4d1dab341 --- /dev/null +++ b/tests/unit/checkpoint/test_autoep_zero3_metadata.py @@ -0,0 +1,124 @@ +# SPDX-License-Identifier: Apache-2.0 +# DeepSpeed Team +"""Contract tests for AutoEP checkpoint placement metadata.""" + +import copy + +import pytest + +from deepspeed.checkpoint.autoep_affine import make_autoep_placement_descriptor +from deepspeed.checkpoint.autoep_zero3_metadata import validate_autoep_zero3_partitioned_metadata +from deepspeed.checkpoint.constants import (AUTOEP_EXPERT_PLACEMENT, AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY, + AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT) + + +def _metadata_entry(): + return { + "moe_layer_id": 0, + "module_path": "model.layers.0.mlp", + "num_experts": 4, + "num_local_experts": 2, + "ep_size": 2, + "expert_key_prefix": "model.layers.0.mlp.experts", + "ep_rank": 1, + } + + +def _partitioned_metadata_entry(): + entry = _metadata_entry() + entry.update({ + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY: AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY: AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, + "ep_group_name": "ep_size_2", + "expert_data_parallel_rank": 0, + "expert_data_parallel_world_size": 1, + "global_expert_start": 2, + "global_expert_end": 4, + }) + return entry + + +def test_legacy_metadata_synthesizes_uniform_placement(): + placements = validate_autoep_zero3_partitioned_metadata([_metadata_entry()], require_partitioned=False) + + assert placements == [{ + "version": 1, + "num_experts": 4, + "ep_size": 2, + "ranks": [{ + "rank": 0, + "experts": [0, 1] + }, { + "rank": 1, + "experts": [2, 3] + }], + }] + + +@pytest.mark.parametrize( + "placement", + [ + { + "version": 999, + "num_experts": 4, + "ep_size": 2, + "ranks": [], + }, + make_autoep_placement_descriptor(6, [[0, 1, 2], [3, 4, 5]]), + make_autoep_placement_descriptor(4, [[0, 1], [2, 3], []]), + ], +) +def test_new_invalid_or_mismatched_placement_is_rejected(placement): + entry = _metadata_entry() + entry[AUTOEP_EXPERT_PLACEMENT] = placement + + with pytest.raises(RuntimeError, match="placement metadata"): + validate_autoep_zero3_partitioned_metadata([entry], require_partitioned=False) + + +def test_rank_local_placement_must_match_runtime_packing(): + entry = _metadata_entry() + entry[AUTOEP_EXPERT_PLACEMENT] = make_autoep_placement_descriptor(4, [[0, 2], [1, 3]]) + expected_runtime_layers = { + "model.layers.0.mlp": { + "ep_rank": 1, + "local_experts": [2, 3], + } + } + + with pytest.raises(RuntimeError, match="expert order"): + validate_autoep_zero3_partitioned_metadata([entry], + require_partitioned=False, + expected_runtime_layers=expected_runtime_layers) + + +def test_rank_local_expert_count_must_match_layer_entry(): + entry = copy.deepcopy(_metadata_entry()) + entry["num_local_experts"] = 1 + entry[AUTOEP_EXPERT_PLACEMENT] = make_autoep_placement_descriptor(4, [[0, 1], [2, 3]]) + + with pytest.raises(RuntimeError, match="expert count"): + validate_autoep_zero3_partitioned_metadata([entry], require_partitioned=False) + + +def test_partitioned_metadata_with_explicit_uneven_noncontiguous_placement_is_accepted(): + entry = _partitioned_metadata_entry() + entry["num_experts"] = 5 + entry["num_local_experts"] = 3 + entry["global_expert_start"] = 99 + entry["global_expert_end"] = 102 + entry[AUTOEP_EXPERT_PLACEMENT] = make_autoep_placement_descriptor(5, [[1, 3], [4, 0, 2]]) + + placements = validate_autoep_zero3_partitioned_metadata([entry]) + + assert placements == [entry[AUTOEP_EXPERT_PLACEMENT]] + + +def test_partitioned_legacy_metadata_still_requires_uniform_contiguous_ranges(): + entry = _partitioned_metadata_entry() + entry["global_expert_start"] = 1 + + with pytest.raises(RuntimeError, match="global expert range"): + validate_autoep_zero3_partitioned_metadata([entry]) diff --git a/tests/unit/v1/moe/test_autoep_unit.py b/tests/unit/v1/moe/test_autoep_unit.py index e0e15c3db6ad..63f3b3f8db26 100644 --- a/tests/unit/v1/moe/test_autoep_unit.py +++ b/tests/unit/v1/moe/test_autoep_unit.py @@ -520,6 +520,92 @@ def test_autoep_layer_marks_zero3_param_placement_families(self): for param in autoep_layer.router.parameters(): assert param.ds_zero_placement_family == "replicated" + def test_autoep_layer_emits_uniform_placement_and_param_restore_metadata(self): + from deepspeed.checkpoint.constants import ( + AUTOEP_EXPERT_PLACEMENT, + AUTOEP_PARAM_EP_RANK, + AUTOEP_PARAM_LOCAL_EXPERTS, + AUTOEP_PARAM_LOGICAL_SHAPE, + DS_AUTOEP_UC_META, + ) + + source = MockMoEBlock(num_experts=4, ffn_hidden=128, hidden_size=64) + layer = AutoEPMoELayer(_make_spec(), + source, + ep_size=2, + ep_rank=1, + config=_runtime_config(enabled=True, autoep_size=2)) + peer_layer = AutoEPMoELayer(_make_spec(), + source, + ep_size=2, + ep_rank=0, + config=_runtime_config(enabled=True, autoep_size=2)) + expected_placement = { + "version": 1, + "num_experts": 4, + "ep_size": 2, + "ranks": [{ + "rank": 0, + "experts": [0, 1] + }, { + "rank": 1, + "experts": [2, 3] + }], + } + + assert layer.expert_placement_descriptor == expected_placement + assert peer_layer.expert_placement_descriptor == expected_placement + for param in layer.experts.parameters(): + restore_metadata = getattr(param, DS_AUTOEP_UC_META) + assert restore_metadata[AUTOEP_EXPERT_PLACEMENT] == expected_placement + assert restore_metadata[AUTOEP_PARAM_LOGICAL_SHAPE] == [4, *param.shape[1:]] + assert restore_metadata[AUTOEP_PARAM_EP_RANK] == 1 + assert restore_metadata[AUTOEP_PARAM_LOCAL_EXPERTS] == [2, 3] + + for param in layer.router.parameters(): + assert not hasattr(param, DS_AUTOEP_UC_META) + + def test_checkpoint_source_placement_does_not_need_to_match_target_topology(self): + from deepspeed.checkpoint.autoep_affine import make_autoep_placement_descriptor + from deepspeed.checkpoint.constants import ( + AUTOEP_EXPERT_PLACEMENT, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY, + AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT, + ) + + source = MockMoEBlock(num_experts=4, ffn_hidden=128, hidden_size=64) + target_layer = AutoEPMoELayer(_make_spec(), + source, + ep_size=4, + ep_rank=0, + config=_runtime_config(enabled=True, autoep_size=4)) + target_model = nn.Sequential(target_layer) + source_metadata = [{ + "moe_layer_id": 0, + "module_path": "0", + "num_experts": 4, + "num_local_experts": 2, + "ep_size": 2, + AUTOEP_EXPERT_PLACEMENT: make_autoep_placement_descriptor(4, [[0, 1], [2, 3]]), + "expert_key_prefix": "0.experts", + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_KEY: AUTOEP_ZERO3_PARTITIONED_EXPERT_STATE_FORMAT, + AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION_KEY: AUTOEP_ZERO3_EXPERT_STATE_FORMAT_VERSION, + "ep_group_name": "ep_size_2", + "ep_rank": 0, + "expert_data_parallel_rank": 0, + "expert_data_parallel_world_size": 1, + "global_expert_start": 0, + "global_expert_end": 2, + }] + + DeepSpeedEngine._validate_autoep_zero3_partitioned_metadata(source_metadata, model=target_model) + with pytest.raises(RuntimeError, match="expert order"): + DeepSpeedEngine._validate_autoep_zero3_partitioned_metadata(source_metadata, + model=target_model, + validate_runtime_placement=True) + def test_zero3_checkpoint_metadata_includes_partition_group_ranks(self): optimizer = object.__new__(DeepSpeedZeroOptimizer_Stage3) param = nn.Parameter(torch.empty(1))