diff --git a/miles/backends/experimental/fsdp_utils/actor.py b/miles/backends/experimental/fsdp_utils/actor.py index 47c2540f98d..c5832e50f99 100644 --- a/miles/backends/experimental/fsdp_utils/actor.py +++ b/miles/backends/experimental/fsdp_utils/actor.py @@ -20,6 +20,7 @@ from miles.utils.timer import Timer, inverse_timer, timer from miles.utils.tracking_utils import init_tracking +from ....utils.data import get_rollout_data_ref_fingerprint from ....utils.profile_utils import TrainProfiler from ...training_utils.ci_utils import check_grad_norm from ...training_utils.data import DataIterator, get_batch, get_data_iterator, get_rollout_data @@ -391,24 +392,61 @@ def _compute_log_prob( self.model.cuda() dist.barrier(group=get_gloo_group()) - def train(self, rollout_id: int, rollout_data_ref: Box) -> None: - """Run one training update over a rollout batch. + def _get_parallel_config(self): + parallel_state = get_parallel_state() + return { + "world_rank": dist.get_rank(), + "dp_rank": parallel_state.intra_dp.rank, + "dp_size": parallel_state.intra_dp.size, + "pp_rank": parallel_state.pp.rank, + "pp_size": parallel_state.pp.size, + "cp_rank": parallel_state.cp.rank, + "cp_size": parallel_state.cp.size, + "tp_size": parallel_state.tp.size, + "routing_replay_layer_indices": None, + } + + def preload_rollout_data(self, rollout_id: int, rollout_data_ref: Box) -> dict: + parallel_state = get_parallel_state() + object_fingerprint = get_rollout_data_ref_fingerprint( + rollout_data_ref, + parallel_state.intra_dp.rank, + pp_rank=parallel_state.pp.rank, + cp_rank=parallel_state.cp.rank, + include_routed_experts=False, + ) + cached = self._get_cached_rollout(rollout_id, object_fingerprint) + if cached is not None: + return { + "rank": dist.get_rank(), + "rollout_id": rollout_id, + "num_samples": len(cached["tokens"]), + "cached": True, + } - Parameters: - rollout_id: Monotonic id for logging. - rollout_data_ref: A Box handle wrapping a Ray object reference to a - dictionary with rollout tensors and metadata (e.g., `tokens`, - `loss_masks`, `rewards`, `response_lengths`, optional - `rollout_log_probs`, etc.). It will be fetched and partitioned - by `process_rollout_data` based on data-parallel rank/size. - """ if self.args.offload_train: self.wake_up() + with timer("data_preprocess"): + rollout_data = get_rollout_data( + self.args, + rollout_data_ref, + include_routed_experts=False, + ) + self._store_preloaded_rollout(rollout_id, object_fingerprint, rollout_data) + return { + "rank": dist.get_rank(), + "rollout_id": rollout_id, + "num_samples": len(rollout_data["tokens"]), + "cached": False, + } + + def train_preloaded(self, rollout_id: int) -> None: + rollout_data = self._take_preloaded_rollout(rollout_id) + if self.args.debug_rollout_only: + return + with inverse_timer("train_wait"), timer("train"): - rollout_data = get_rollout_data(self.args, rollout_data_ref) - if self.args.debug_rollout_only: - return self._train_core(rollout_id=rollout_id, rollout_data=rollout_data) train_metric_utils.log_perf_data_raw( @@ -418,6 +456,10 @@ def train(self, rollout_id: int, rollout_data_ref: Box) -> None: compute_total_fwd_flops=None, ) + def train(self, rollout_id: int, rollout_data_ref: Box) -> None: + self.preload_rollout_data(rollout_id, rollout_data_ref) + return self.train_preloaded(rollout_id) + def _train_core(self, rollout_id: int, rollout_data) -> None: data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) data_iterator = data_iterator[0] diff --git a/miles/backends/experimental/fsdp_utils/parallel.py b/miles/backends/experimental/fsdp_utils/parallel.py index fa444975bd4..c96a3b85aa9 100644 --- a/miles/backends/experimental/fsdp_utils/parallel.py +++ b/miles/backends/experimental/fsdp_utils/parallel.py @@ -58,6 +58,11 @@ def create_fsdp_parallel_state(args: Namespace) -> ParallelState: size=1, group=dist.new_group([rank]), ), + pp=GroupInfo( + rank=0, + size=1, + group=None, + ), ) parallel_state.dp_mesh = mesh["dp"] diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 5bb4600c57a..6bccf1fca13 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -22,11 +22,13 @@ from miles.utils.ray_utils import Box from miles.utils.reloadable_process_group import destroy_process_groups, monkey_patch_torch_dist, reload_process_groups from miles.utils.replay_base import all_replay_managers +from miles.utils.rollout_sharding import ROUTED_EXPERTS_SHARD_META_KEY from miles.utils.timer import Timer, inverse_timer, timer from miles.utils.tracking_utils import init_tracking from miles.utils.types import RolloutBatch from ...utils.profile_utils import TrainProfiler +from ...utils.data import get_rollout_data_ref_fingerprint from ...utils.tensor_backper import TensorBackuper from ..training_utils.cp_utils import slice_with_cp from ..training_utils.data import DataIterator, get_data_iterator, get_rollout_data, sync_actor_critic_data @@ -38,7 +40,7 @@ from .lora_utils import is_lora_enabled from .model import forward_only, initialize_model_and_optimizer, save, train from .parallel import verify_megatron_parallel_state -from .replay_utils import get_register_replay_list_func +from .replay_utils import get_register_replay_list_func, get_replay_layer_indices from .update_weight.common import named_params_and_buffers from .update_weight.update_weight_from_distributed.broadcast import UpdateWeightFromDistributed from .update_weight.update_weight_from_distributed.p2p import UpdateWeightP2P @@ -600,6 +602,11 @@ def init( ) verify_megatron_parallel_state(self.model) + self._routing_replay_layer_indices = ( + get_replay_layer_indices(self.model) + if role == "actor" and self.args.use_rollout_routing_replay + else None + ) if role == "critic": if self.args.offload_train: @@ -724,6 +731,29 @@ def _fill_replay_data( tp_rank = parallel_state.tp.rank tp_size = parallel_state.tp.size qkv_format = self.args.qkv_format + shard_metadata = rollout_data.get(ROUTED_EXPERTS_SHARD_META_KEY) + is_destination_sharded = shard_metadata is not None + if is_destination_sharded: + if shard_metadata["pp_size"] != parallel_state.pp.size: + raise ValueError( + f"routing replay PP size mismatch: shard={shard_metadata['pp_size']}, " + f"actor={parallel_state.pp.size}" + ) + if shard_metadata["cp_rank"] != parallel_state.cp.rank: + raise ValueError( + f"routing replay CP rank mismatch: shard={shard_metadata['cp_rank']}, " + f"actor={parallel_state.cp.rank}" + ) + if shard_metadata["cp_size"] != parallel_state.cp.size: + raise ValueError( + f"routing replay CP size mismatch: shard={shard_metadata['cp_size']}, " + f"actor={parallel_state.cp.size}" + ) + if shard_metadata["qkv_format"] != qkv_format: + raise ValueError( + f"routing replay qkv_format mismatch: shard={shard_metadata['qkv_format']}, " + f"actor={qkv_format}" + ) def pad_func(data, pad): _, num_layers, topk = data.shape @@ -740,23 +770,37 @@ def pad_func(data, pad): replay_data = batch[data_key] tokens = batch["tokens"] assert len(replay_data) == len(tokens) - for a, b in zip(replay_data, tokens, strict=False): - assert a.shape[0] == b.shape[0] - 1, f"{a.shape}, {b.shape}" - - # We need to pad the experts to the last token. We won't calculate loss on this token so this should be fine. - # TODO: fuse this padding with the following slice_with_cp to reduce memory copy. - replay_data = [pad_func(r, 1) for r in replay_data] - # TODO: maybe extract a common process function for here and get_batch? - - if qkv_format == "bshd": - max_seqlen = batch["max_seq_lens"][0] - replay_data = [slice_with_cp(r, pad_func, qkv_format, max_seqlen) for r in replay_data] - replay_data = torch.stack(replay_data, dim=0) - batch_size, seqlen, num_layers, topk = replay_data.shape - replay_data = replay_data.reshape(batch_size * seqlen, num_layers, topk) + if is_destination_sharded: + expected_layers = len(shard_metadata["layer_indices"]) + if any(r.shape[1] != expected_layers for r in replay_data): + raise ValueError( + f"routing replay shard has unexpected layer dimension; expected {expected_layers}, " + f"got {[r.shape for r in replay_data]}" + ) + if qkv_format == "bshd": + replay_data = torch.stack(replay_data, dim=0) + batch_size, seqlen, num_layers, topk = replay_data.shape + replay_data = replay_data.reshape(batch_size * seqlen, num_layers, topk) + else: + replay_data = torch.cat(replay_data, dim=0) else: - replay_data = [slice_with_cp(r, pad_func, qkv_format) for r in replay_data] - replay_data = torch.cat(replay_data, dim=0) + for a, b in zip(replay_data, tokens, strict=False): + assert a.shape[0] == b.shape[0] - 1, f"{a.shape}, {b.shape}" + + # Pad the omitted final token before legacy CP slicing. + replay_data = [pad_func(r, 1) for r in replay_data] + + if qkv_format == "bshd": + max_seqlen = batch["max_seq_lens"][0] + replay_data = [slice_with_cp(r, pad_func, qkv_format, max_seqlen) for r in replay_data] + replay_data = torch.stack(replay_data, dim=0) + batch_size, seqlen, num_layers, topk = replay_data.shape + replay_data = replay_data.reshape(batch_size * seqlen, num_layers, topk) + else: + replay_data = [slice_with_cp(r, pad_func, qkv_format) for r in replay_data] + replay_data = torch.cat(replay_data, dim=0) + + if qkv_format == "thd": pad_size = parallel_state.tp.size * self.args.data_pad_size_multiplier pad = (pad_size - replay_data.size(0) % pad_size) % pad_size if pad != 0: @@ -768,9 +812,15 @@ def pad_func(data, pad): start, end = seqlen // tp_size * tp_rank, seqlen // tp_size * (tp_rank + 1) replay_data = replay_data[start:end] - register_replay_list_func(replay_list, replay_data, self.model) + register_replay_list_func( + replay_list, + replay_data, + self.model, + source_layer_indices=shard_metadata["layer_indices"] if is_destination_sharded else None, + ) del rollout_data[data_key] + rollout_data.pop(ROUTED_EXPERTS_SHARD_META_KEY, None) for iterator in data_iterator: iterator.reset() @@ -792,22 +842,83 @@ def compute_log_prob( store_prefix=store_prefix, ) - def train(self, rollout_id: int, rollout_data_ref: Box) -> None: + def _include_rollout_routed_experts(self) -> bool: + return self.role == "actor" and self.args.use_rollout_routing_replay + + def _get_parallel_config(self): + parallel_state = get_parallel_state() + return { + "world_rank": dist.get_rank(), + "dp_rank": parallel_state.intra_dp.rank, + "dp_size": parallel_state.intra_dp.size, + "pp_rank": parallel_state.pp.rank, + "pp_size": parallel_state.pp.size, + "cp_rank": parallel_state.cp.rank, + "cp_size": parallel_state.cp.size, + "tp_size": parallel_state.tp.size, + "routing_replay_layer_indices": getattr(self, "_routing_replay_layer_indices", None), + } + + def preload_rollout_data(self, rollout_id: int, rollout_data_ref: Box) -> dict: self._last_rollout_id = rollout_id + parallel_state = get_parallel_state() + include_routed_experts = self._include_rollout_routed_experts() + object_fingerprint = get_rollout_data_ref_fingerprint( + rollout_data_ref, + parallel_state.intra_dp.rank, + pp_rank=parallel_state.pp.rank, + cp_rank=parallel_state.cp.rank, + include_routed_experts=include_routed_experts, + ) + cached = self._get_cached_rollout(rollout_id, object_fingerprint) + if cached is not None: + return { + "rank": dist.get_rank(), + "rollout_id": rollout_id, + "num_samples": len(cached["tokens"]), + "cached": True, + } + if self.args.offload_train: self.wake_up() with timer("data_preprocess"): - rollout_data = get_rollout_data(self.args, rollout_data_ref) - if self.args.debug_rollout_only: - log_rollout_data(rollout_id, self.args, rollout_data) - return + rollout_data = get_rollout_data( + self.args, + rollout_data_ref, + include_routed_experts=include_routed_experts, + ) + if include_routed_experts and ROUTED_EXPERTS_SHARD_META_KEY in rollout_data: + expected_layer_indices = self._routing_replay_layer_indices + actual_layer_indices = rollout_data[ROUTED_EXPERTS_SHARD_META_KEY]["layer_indices"] + if actual_layer_indices != expected_layer_indices: + raise ValueError( + f"routing replay layer shard {actual_layer_indices} does not match " + f"local model layers {expected_layer_indices}" + ) + self._store_preloaded_rollout(rollout_id, object_fingerprint, rollout_data) + return { + "rank": dist.get_rank(), + "rollout_id": rollout_id, + "num_samples": len(rollout_data["tokens"]), + "cached": False, + } + + def train_preloaded(self, rollout_id: int) -> None: + rollout_data = self._take_preloaded_rollout(rollout_id) + if self.args.debug_rollout_only: + log_rollout_data(rollout_id, self.args, rollout_data) + return if self.role == "critic": return self.train_critic(rollout_id, rollout_data) else: return self.train_actor(rollout_id, rollout_data) + def train(self, rollout_id: int, rollout_data_ref: Box) -> None: + self.preload_rollout_data(rollout_id, rollout_data_ref) + return self.train_preloaded(rollout_id) + def train_critic(self, rollout_id: int, rollout_data: RolloutBatch) -> None: # Create data iterator for log_probs and train. data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) diff --git a/miles/backends/megatron_utils/parallel.py b/miles/backends/megatron_utils/parallel.py index fc8333f42c3..a2516f77b3f 100644 --- a/miles/backends/megatron_utils/parallel.py +++ b/miles/backends/megatron_utils/parallel.py @@ -37,6 +37,11 @@ def _create_intra_dp(with_context_parallel: bool): size=mpu.get_tensor_model_parallel_world_size(), group=mpu.get_tensor_model_parallel_group(), ), + pp=GroupInfo( + rank=mpu.get_pipeline_model_parallel_rank(), + size=mpu.get_pipeline_model_parallel_world_size(), + group=mpu.get_pipeline_model_parallel_group(), + ), is_pp_last_stage=mpu.is_pipeline_last_stage(), vpp_size=vpp_size, microbatch_group_size_per_vp_stage=microbatch_group_size_per_vp_stage, diff --git a/miles/backends/megatron_utils/replay_utils.py b/miles/backends/megatron_utils/replay_utils.py index 2429964e236..caeaad3dc46 100644 --- a/miles/backends/megatron_utils/replay_utils.py +++ b/miles/backends/megatron_utils/replay_utils.py @@ -4,9 +4,8 @@ from miles.utils.replay_base import BaseReplayManager, RoutingReplayManager -def _register_replay_list_moe(replay_list, replay_data, models): +def get_replay_layer_indices(models) -> list[int]: layer_indices = [] - replay_idx = 0 for vp_stage, model in enumerate(models): config = model.module.config num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage) @@ -20,9 +19,29 @@ def _register_replay_list_moe(replay_list, replay_data, models): if config.moe_layer_freq[layer_id] == 0: continue layer_indices.append(layer_id) + return layer_indices - for replay_idx, layer_idx in enumerate(layer_indices): - layer_data = replay_data[:, layer_idx] + +def _register_replay_list_moe( + replay_list, + replay_data, + models, + *, + source_layer_indices=None, +): + layer_indices = get_replay_layer_indices(models) + if source_layer_indices is None: + replay_columns = layer_indices + else: + source_layer_indices = [int(layer_idx) for layer_idx in source_layer_indices] + if source_layer_indices != layer_indices: + raise ValueError( + f"routing replay layer shard {source_layer_indices} does not match local model layers {layer_indices}" + ) + replay_columns = range(len(layer_indices)) + + for replay_idx, replay_column in enumerate(replay_columns): + layer_data = replay_data[:, replay_column] replay_list[replay_idx].record(layer_data) diff --git a/miles/backends/training_utils/data.py b/miles/backends/training_utils/data.py index 5109676285c..1c561080b39 100644 --- a/miles/backends/training_utils/data.py +++ b/miles/backends/training_utils/data.py @@ -19,7 +19,12 @@ logger = logging.getLogger(__name__) -def get_rollout_data(args: Namespace, rollout_data_ref: Box) -> RolloutBatch: +def get_rollout_data( + args: Namespace, + rollout_data_ref: Box, + *, + include_routed_experts: bool = True, +) -> RolloutBatch: parallel_state = get_parallel_state() # Fetch data through ray on CPU, not sure if this will be performance bottleneck. # Both first pp stage and the last pp stage will receive the data. @@ -28,6 +33,11 @@ def get_rollout_data(args: Namespace, rollout_data_ref: Box) -> RolloutBatch: rollout_data_ref, parallel_state.intra_dp.rank, parallel_state.intra_dp.size, + pp_rank=parallel_state.pp.rank if parallel_state.pp is not None else 0, + pp_size=parallel_state.pp.size if parallel_state.pp is not None else 1, + cp_rank=parallel_state.cp.rank, + cp_size=parallel_state.cp.size, + include_routed_experts=include_routed_experts, ) # move tokens to GPU in advance rollout_data["tokens"] = [ diff --git a/miles/backends/training_utils/parallel.py b/miles/backends/training_utils/parallel.py index a989b759923..db1b3946292 100644 --- a/miles/backends/training_utils/parallel.py +++ b/miles/backends/training_utils/parallel.py @@ -53,6 +53,7 @@ class ParallelState: intra_dp_cp: GroupInfo cp: GroupInfo tp: GroupInfo + pp: GroupInfo | None = None is_pp_last_stage: bool = True vpp_size: int | None = 1 microbatch_group_size_per_vp_stage: int | None = None diff --git a/miles/ray/actor_group.py b/miles/ray/actor_group.py index 669988b52c3..0e2aeb319b1 100644 --- a/miles/ray/actor_group.py +++ b/miles/ray/actor_group.py @@ -8,6 +8,60 @@ from miles.ray.utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST +def merge_train_parallel_configs(configs: list[dict]) -> dict: + if not configs: + raise ValueError("at least one training-rank parallel config is required") + + size_keys = ("dp_size", "pp_size", "cp_size", "tp_size") + merged = {key: configs[0][key] for key in size_keys} + for config in configs[1:]: + for key in size_keys: + if config[key] != merged[key]: + raise ValueError( + f"inconsistent {key}: rank 0 reported {merged[key]}, " + f"rank {config['world_rank']} reported {config[key]}" + ) + + routing_specs = {} + routing_enabled = any(config["routing_replay_layer_indices"] is not None for config in configs) + if routing_enabled: + layers_by_pp_rank = {} + for config in configs: + layer_indices = config["routing_replay_layer_indices"] + if layer_indices is None: + raise ValueError( + f"rank {config['world_rank']} did not report routing-replay layers while other ranks did" + ) + layer_indices = list(layer_indices) + previous_layers = layers_by_pp_rank.setdefault(config["pp_rank"], layer_indices) + if previous_layers != layer_indices: + raise ValueError( + f"inconsistent routing-replay layers for PP rank {config['pp_rank']}: " + f"{previous_layers} != {layer_indices} on world rank {config['world_rank']}" + ) + destination = (config["dp_rank"], config["pp_rank"], config["cp_rank"]) + spec = { + "dp_rank": config["dp_rank"], + "pp_rank": config["pp_rank"], + "cp_rank": config["cp_rank"], + "layer_indices": layer_indices, + } + previous = routing_specs.setdefault(destination, spec) + if previous != spec: + raise ValueError(f"inconsistent routing-replay shard spec for destination {destination}") + + expected_destinations = merged["dp_size"] * merged["pp_size"] * merged["cp_size"] + if len(routing_specs) != expected_destinations: + raise ValueError( + f"expected {expected_destinations} routing-replay destinations, got {len(routing_specs)}" + ) + + merged["routing_replay_shard_specs"] = [ + routing_specs[key] for key in sorted(routing_specs) + ] + return merged + + class RayTrainGroup: """ A group of ray actors @@ -113,7 +167,31 @@ async def init(self): async def train(self, rollout_id, rollout_data_ref): """Do one rollout training""" - await self._broadcast("train", rollout_id, rollout_data_ref) + await self.preload_rollout_data(rollout_id, rollout_data_ref) + await self.train_preloaded(rollout_id) + + async def preload_rollout_data(self, rollout_id, rollout_data_ref): + """Materialize rollout data on every rank without entering collectives.""" + refs = [ + actor.preload_rollout_data.remote(rollout_id, rollout_data_ref) + for actor in self._actor_handles + ] + results = await asyncio.gather(*refs, return_exceptions=True) + errors = [(rank, result) for rank, result in enumerate(results) if isinstance(result, BaseException)] + if errors: + # Avoid distributed cleanup after a partial preload failure. + cleanup_refs = [ + actor.discard_preloaded_rollout.remote(rollout_id) + for actor in self._actor_handles + ] + await asyncio.gather(*cleanup_refs, return_exceptions=True) + details = "; ".join(f"rank {rank}: {error!r}" for rank, error in errors) + raise RuntimeError(f"rollout {rollout_id} preload failed on {len(errors)} rank(s): {details}") + return results + + async def train_preloaded(self, rollout_id): + """Start training only after all ranks have acknowledged preload.""" + return await self._broadcast("train_preloaded", rollout_id) async def save_model(self, rollout_id, force_sync=False): """Save actor model""" @@ -141,6 +219,9 @@ async def connect(self, critic_group): async def set_rollout_manager(self, rollout_manager): await self._broadcast("set_rollout_manager", rollout_manager) + configs = await self._broadcast("get_parallel_config") + merged_config = merge_train_parallel_configs(configs) + await rollout_manager.set_train_parallel_config.remote(merged_config) async def _broadcast(self, method_name: str, *args, **kwargs) -> list: refs = [getattr(actor, method_name).remote(*args, **kwargs) for actor in self._actor_handles] diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index 0d462d46cdf..7ce10a91cd5 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -49,6 +49,13 @@ ) from miles.utils.misc import load_function from miles.utils.ray_utils import Box +from miles.utils.rollout_sharding import ( + ROLLOUT_DATA_REF_FORMAT, + ROUTED_EXPERTS_SHARD_META_KEY, + rollout_destination_key, + shard_routed_experts_for_destination, + validate_routed_experts, +) from miles.utils.seqlen_balancing import get_seqlen_balanced_partitions from miles.utils.tracking_utils import init_tracking from miles.utils.types import Sample @@ -935,7 +942,10 @@ def _stat(xs): else: partitions = [range(i, len(total_lengths), dp_size) for i in range(dp_size)] + routing_shard_specs = self.train_parallel_config.get("routing_replay_shard_specs", []) + destination_shard_routing = "rollout_routed_experts" in data and bool(routing_shard_specs) rollout_data_refs = [] + routed_expert_refs = {} dp_summaries = [] for i in range(dp_size): @@ -943,7 +953,7 @@ def _stat(xs): partition = list(partitions[i]) rollout_data["partition"] = partition - for key in [ + per_sample_keys = [ "tokens", "multimodal_train_inputs", "response_lengths", @@ -953,12 +963,15 @@ def _stat(xs): "round_number", "sample_indices", "rollout_log_probs", - "rollout_routed_experts", "prompt", "teacher_log_probs", "weight_versions", "domains", - ]: + ] + if not destination_shard_routing: + per_sample_keys.append("rollout_routed_experts") + + for key in per_sample_keys: if key not in data: continue rollout_data[key] = [data[key][j] for j in partition] @@ -977,9 +990,78 @@ def _stat(xs): response_lens = [data["response_lengths"][j] for j in partition] if "response_lengths" in data else [] loss_mask_lens = [_safe_len(data["loss_masks"][j]) for j in partition] if "loss_masks" in data else [] - payload_bytes = _estimate_payload_bytes(rollout_data) + routed_experts = ( + [data["rollout_routed_experts"][j] for j in partition] if destination_shard_routing else None + ) + if routed_experts is not None: + validate_routed_experts(routed_experts, token_lens) + base_payload_bytes = _estimate_payload_bytes(rollout_data) + routed_payload_bytes = _estimate_payload_bytes(routed_experts) if routed_experts is not None else 0 + payload_bytes = base_payload_bytes + routed_payload_bytes ref = ray.put(rollout_data) + routing_shard_payload_bytes = [] + if routed_experts is not None: + max_seq_len = None + if self.args.qkv_format == "bshd": + pad_size = ( + self.train_parallel_config["tp_size"] * self.args.data_pad_size_multiplier + ) + max_seq_len = (max(token_lens) + pad_size - 1) // pad_size * pad_size + + specs_for_dp = [spec for spec in routing_shard_specs if spec["dp_rank"] == i] + expected_specs = ( + self.train_parallel_config["pp_size"] * self.train_parallel_config["cp_size"] + ) + if len(specs_for_dp) != expected_specs: + raise ValueError( + f"DP rank {i} expected {expected_specs} routing destinations, " + f"got {len(specs_for_dp)}" + ) + + for spec in specs_for_dp: + local_routed_experts = shard_routed_experts_for_destination( + routed_experts, + layer_indices=spec["layer_indices"], + cp_rank=spec["cp_rank"], + cp_size=self.train_parallel_config["cp_size"], + qkv_format=self.args.qkv_format, + max_seq_len=max_seq_len, + ) + shard_metadata = { + "version": 1, + "dp_rank": i, + "pp_rank": spec["pp_rank"], + "pp_size": self.train_parallel_config["pp_size"], + "cp_rank": spec["cp_rank"], + "cp_size": self.train_parallel_config["cp_size"], + "qkv_format": self.args.qkv_format, + "max_seq_len": max_seq_len, + "layer_indices": list(spec["layer_indices"]), + } + routing_payload = { + "rollout_routed_experts": local_routed_experts, + ROUTED_EXPERTS_SHARD_META_KEY: shard_metadata, + } + routing_payload_bytes = _estimate_payload_bytes(routing_payload) + routing_ref = ray.put(routing_payload) + destination = rollout_destination_key( + i, + spec["pp_rank"], + spec["cp_rank"], + ) + routed_expert_refs[destination] = Box(routing_ref) + routing_shard_payload_bytes.append(routing_payload_bytes) + logger.warning( + "ROLLOUT_ROUTING_SHARD destination=%s samples=%s layers=%s " + "payload_mb_est=%.2f object_ref=%s", + destination, + len(partition), + spec["layer_indices"], + routing_payload_bytes / 1024 / 1024, + routing_ref.hex(), + ) + summary = { "dp_rank": i, "num_samples": len(partition), @@ -988,12 +1070,17 @@ def _stat(xs): "responses": _stat(response_lens), "loss_masks": _stat(loss_mask_lens), "payload_mb_est": round(payload_bytes / 1024 / 1024, 2), + "base_payload_mb_est": round(base_payload_bytes / 1024 / 1024, 2), + "max_routing_shard_mb_est": round( + max(routing_shard_payload_bytes, default=0) / 1024 / 1024, + 2, + ), "object_ref": ref.hex(), } dp_summaries.append(summary) logger.warning( - "ROLLOUT_DP_SHARD dp=%s samples=%s token_sum=%s token_min=%s token_max=%s token_avg=%s response_sum=%s response_min=%s response_max=%s response_avg=%s loss_mask_sum=%s payload_mb_est=%.2f object_ref=%s partition=%s", + "ROLLOUT_DP_SHARD dp=%s samples=%s token_sum=%s token_min=%s token_max=%s token_avg=%s response_sum=%s response_min=%s response_max=%s response_avg=%s loss_mask_sum=%s payload_mb_est=%.2f base_payload_mb_est=%.2f max_routing_shard_mb_est=%.2f object_ref=%s partition=%s", summary["dp_rank"], summary["num_samples"], summary["tokens"]["sum"], @@ -1006,6 +1093,8 @@ def _stat(xs): summary["responses"]["avg"], summary["loss_masks"]["sum"], summary["payload_mb_est"], + summary["base_payload_mb_est"], + summary["max_routing_shard_mb_est"], summary["object_ref"], summary["partition"], ) @@ -1037,6 +1126,21 @@ def _ratio(xs): _ratio(payload_mbs), ) + if destination_shard_routing: + expected_routing_refs = ( + dp_size + * self.train_parallel_config["pp_size"] + * self.train_parallel_config["cp_size"] + ) + if len(routed_expert_refs) != expected_routing_refs: + raise ValueError( + f"expected {expected_routing_refs} routing-replay refs, got {len(routed_expert_refs)}" + ) + return { + "format": ROLLOUT_DATA_REF_FORMAT, + "base": rollout_data_refs, + "rollout_routed_experts": routed_expert_refs, + } return rollout_data_refs diff --git a/miles/ray/train_actor.py b/miles/ray/train_actor.py index a4145ca270d..c8a3a1657d6 100644 --- a/miles/ray/train_actor.py +++ b/miles/ray/train_actor.py @@ -47,6 +47,8 @@ def __init__(self, world_size, rank, master_addr, master_port): # os.environ.pop("CUDA_VISIBLE_DEVICES", None) # os.environ["LOCAL_RANK"] = str(ray.get_gpu_ids()[0]) os.environ["LOCAL_RANK"] = str(get_local_gpu_id()) + self._preloaded_rollout_key = None + self._preloaded_rollout_data = None def init(self, args, role, with_ref=False): self.args = args @@ -120,6 +122,14 @@ def wake_up(self, tags): def train(self, rollout_id, rollout_data_ref): raise NotImplementedError + @abc.abstractmethod + def preload_rollout_data(self, rollout_id, rollout_data_ref): + raise NotImplementedError + + @abc.abstractmethod + def train_preloaded(self, rollout_id): + raise NotImplementedError + @abc.abstractmethod def save_model(self, rollout_id, force_sync=False): raise NotImplementedError @@ -136,7 +146,50 @@ def connect_actor_critic(self, critic_group): def _get_parallel_config(self): raise NotImplementedError + def get_parallel_config(self): + return self._get_parallel_config() + def set_rollout_manager(self, rollout_manager): self.rollout_manager = rollout_manager - if self.args.rank == 0: - ray.get(self.rollout_manager.set_train_parallel_config.remote(self.train_parallel_config)) + + def _get_cached_rollout(self, rollout_id, object_fingerprint): + requested_key = (rollout_id, object_fingerprint) + if self._preloaded_rollout_key is None: + return None + if self._preloaded_rollout_key != requested_key: + raise RuntimeError( + f"cannot preload rollout {requested_key}: unconsumed rollout " + f"{self._preloaded_rollout_key} is already cached" + ) + return self._preloaded_rollout_data + + def _store_preloaded_rollout(self, rollout_id, object_fingerprint, rollout_data): + requested_key = (rollout_id, object_fingerprint) + if self._preloaded_rollout_key is not None: + raise RuntimeError( + f"cannot store rollout {requested_key}: unconsumed rollout " + f"{self._preloaded_rollout_key} is already cached" + ) + self._preloaded_rollout_key = requested_key + self._preloaded_rollout_data = rollout_data + + def _take_preloaded_rollout(self, rollout_id): + if self._preloaded_rollout_key is None: + raise RuntimeError(f"rollout {rollout_id} was not preloaded") + cached_rollout_id, _ = self._preloaded_rollout_key + if cached_rollout_id != rollout_id: + raise RuntimeError(f"requested rollout {rollout_id}, but cached rollout is {cached_rollout_id}") + rollout_data = self._preloaded_rollout_data + self._preloaded_rollout_key = None + self._preloaded_rollout_data = None + return rollout_data + + def discard_preloaded_rollout(self, rollout_id): + if self._preloaded_rollout_key is None: + return False + cached_rollout_id, _ = self._preloaded_rollout_key + if cached_rollout_id != rollout_id: + return False + self._preloaded_rollout_key = None + self._preloaded_rollout_data = None + return True diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 9d696bdb3a0..64d61d4c83a 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -2091,9 +2091,6 @@ def miles_validate_args(args): if args.enable_mtp_training: assert args.mtp_num_layers, "mtp_num_layers must be set when enable_mtp_training is set" - if args.use_rollout_routing_replay: - args.use_routing_replay = True - if args.custom_config_path: with open(args.custom_config_path) as f: data = yaml.safe_load(f) or {} @@ -2102,6 +2099,14 @@ def miles_validate_args(args): logger.info(f"Warning: Argument {k} is already set to {getattr(args, k)}, will override with {v}.") setattr(args, k, v) + if args.use_rollout_routing_replay: + if args.allgather_cp: + raise ValueError( + "--use-rollout-routing-replay is incompatible with --allgather-cp: " + "routing replay uses per-sample zig-zag CP ordering" + ) + args.use_routing_replay = True + if args.eval_max_context_len is None: logger.info( f"args.eval_max_context_len is not set. Use args.rollout_max_context_len {args.rollout_max_context_len} as default value." diff --git a/miles/utils/data.py b/miles/utils/data.py index 6e64ef678de..c12a7d6b80a 100644 --- a/miles/utils/data.py +++ b/miles/utils/data.py @@ -15,6 +15,11 @@ from miles.utils.types import MultimodalTypes, Sample +from .rollout_sharding import ( + ROLLOUT_DATA_REF_FORMAT, + ROUTED_EXPERTS_SHARD_META_KEY, + rollout_destination_key, +) from .timer import Timer __all__ = ["Dataset"] @@ -269,9 +274,138 @@ def get_minimum_num_micro_batch_size(total_lengths, max_tokens_per_gpu): return len(batches) -def process_rollout_data(args, rollout_data_ref, dp_rank, dp_size): - assert len(rollout_data_ref) == dp_size - rollout_data = ray.get(rollout_data_ref[dp_rank].inner) +def get_rollout_data_ref_fingerprint( + rollout_data_ref, + dp_rank, + *, + pp_rank=0, + cp_rank=0, + include_routed_experts=True, +): + """Identify exactly the local Ray objects backing one actor's preload.""" + if isinstance(rollout_data_ref, dict) and rollout_data_ref.get("format") == ROLLOUT_DATA_REF_FORMAT: + refs = [rollout_data_ref["base"][dp_rank].inner] + if include_routed_experts: + destination = rollout_destination_key(dp_rank, pp_rank, cp_rank) + routed_experts_ref = rollout_data_ref.get("rollout_routed_experts", {}).get(destination) + if routed_experts_ref is None: + raise KeyError(f"missing rollout routing-replay shard for destination {destination}") + refs.append(routed_experts_ref.inner) + else: + refs = [rollout_data_ref[dp_rank].inner] + + return ":".join(ref.hex() for ref in refs) + + +def process_rollout_data( + args, + rollout_data_ref, + dp_rank, + dp_size, + *, + pp_rank=0, + pp_size=1, + cp_rank=0, + cp_size=1, + include_routed_experts=True, +): + if isinstance(rollout_data_ref, dict) and rollout_data_ref.get("format") == ROLLOUT_DATA_REF_FORMAT: + base_refs = rollout_data_ref["base"] + assert len(base_refs) == dp_size + refs = [base_refs[dp_rank].inner] + + routed_experts_ref = None + if include_routed_experts: + destination = rollout_destination_key(dp_rank, pp_rank, cp_rank) + routed_experts_ref = rollout_data_ref.get("rollout_routed_experts", {}).get(destination) + if routed_experts_ref is None: + raise KeyError( + f"missing rollout routing-replay shard for destination {destination}; " + f"available={sorted(rollout_data_ref.get('rollout_routed_experts', {}))}" + ) + refs.append(routed_experts_ref.inner) + + fetched = ray.get(refs) + rollout_data = fetched[0] + if routed_experts_ref is not None: + routed_experts = fetched[1] + shard_metadata = routed_experts[ROUTED_EXPERTS_SHARD_META_KEY] + if shard_metadata.get("version") != 1: + raise ValueError( + f"unsupported routing-replay shard version {shard_metadata.get('version')}" + ) + expected_destination = (dp_rank, pp_rank, cp_rank) + actual_destination = tuple( + shard_metadata[key] for key in ("dp_rank", "pp_rank", "cp_rank") + ) + if actual_destination != expected_destination: + raise ValueError( + f"routing-replay shard destination mismatch: expected {expected_destination}, " + f"got {actual_destination}" + ) + if shard_metadata["pp_size"] != pp_size or shard_metadata["cp_size"] != cp_size: + raise ValueError( + "routing-replay shard parallel-size mismatch: " + f"shard pp/cp=({shard_metadata['pp_size']}, {shard_metadata['cp_size']}), " + f"actor pp/cp=({pp_size}, {cp_size})" + ) + if shard_metadata["qkv_format"] != args.qkv_format: + raise ValueError( + f"routing-replay qkv_format mismatch: shard={shard_metadata['qkv_format']}, " + f"actor={args.qkv_format}" + ) + if len(routed_experts["rollout_routed_experts"]) != len(rollout_data["tokens"]): + raise ValueError( + "routing-replay shard sample count does not match base payload: " + f"{len(routed_experts['rollout_routed_experts'])} != {len(rollout_data['tokens'])}" + ) + num_local_layers = len(shard_metadata["layer_indices"]) + topk = None + for sample_idx, (sample_routing, tokens) in enumerate( + zip( + routed_experts["rollout_routed_experts"], + rollout_data["tokens"], + strict=True, + ) + ): + num_tokens = len(tokens) + if shard_metadata["qkv_format"] == "thd": + target_length = num_tokens + else: + target_length = shard_metadata["max_seq_len"] + if target_length is None or target_length < num_tokens: + raise ValueError( + f"invalid BSHD routing-replay max_seq_len {target_length} " + f"for sample {sample_idx} with {num_tokens} tokens" + ) + if cp_size == 1: + expected_rows = target_length + else: + chunk_size = (target_length + 2 * cp_size - 1) // (2 * cp_size) + expected_rows = 2 * chunk_size + shape = getattr(sample_routing, "shape", None) + if shape is None or len(shape) != 3: + raise ValueError( + f"routing-replay sample {sample_idx} must be rank 3, got shape={shape}" + ) + if shape[0] != expected_rows or shape[1] != num_local_layers: + raise ValueError( + f"routing-replay sample {sample_idx} shape mismatch: got {shape}, " + f"expected ({expected_rows}, {num_local_layers}, topk)" + ) + if topk is None: + topk = shape[2] + elif shape[2] != topk: + raise ValueError( + f"routing-replay topk mismatch in sample {sample_idx}: {shape[2]} != {topk}" + ) + rollout_data["rollout_routed_experts"] = routed_experts["rollout_routed_experts"] + rollout_data[ROUTED_EXPERTS_SHARD_META_KEY] = shard_metadata + else: + assert len(rollout_data_ref) == dp_size + rollout_data = ray.get(rollout_data_ref[dp_rank].inner) + if not include_routed_experts: + rollout_data.pop("rollout_routed_experts", None) partition = rollout_data.pop("partition") total_lengths = rollout_data["total_lengths"] diff --git a/miles/utils/rollout_sharding.py b/miles/utils/rollout_sharding.py new file mode 100644 index 00000000000..2adecb7e573 --- /dev/null +++ b/miles/utils/rollout_sharding.py @@ -0,0 +1,142 @@ +"""Helpers for routing rollout payloads to their actual training consumers.""" + +from collections.abc import Sequence + +import numpy as np + + +ROLLOUT_DATA_REF_FORMAT = "destination_sharded_v1" +ROUTED_EXPERTS_SHARD_META_KEY = "_rollout_routed_experts_shard" + + +def rollout_destination_key(dp_rank: int, pp_rank: int, cp_rank: int) -> str: + """Return a stable, Ray-serializable key for a training destination.""" + return f"{dp_rank}:{pp_rank}:{cp_rank}" + + +def validate_routed_experts( + routed_experts: Sequence[np.ndarray], + token_lengths: Sequence[int], +) -> None: + """Validate the source routing tensors before any lossy PP/CP compaction.""" + if len(routed_experts) != len(token_lengths): + raise ValueError( + f"routing-replay sample count {len(routed_experts)} does not match " + f"token-length count {len(token_lengths)}" + ) + + expected_tail = None + for sample_idx, (sample, token_length) in enumerate( + zip(routed_experts, token_lengths, strict=True) + ): + shape = getattr(sample, "shape", None) + if not isinstance(sample, np.ndarray) or shape is None or len(shape) != 3: + raise ValueError( + f"routing-replay sample {sample_idx} must be a rank-3 numpy array, got shape={shape}" + ) + expected_rows = token_length - 1 + if shape[0] != expected_rows: + raise ValueError( + f"routing-replay sample {sample_idx} has {shape[0]} token rows; " + f"expected {expected_rows} for {token_length} tokens" + ) + if expected_tail is None: + expected_tail = shape[1:] + elif shape[1:] != expected_tail: + raise ValueError( + f"routing-replay sample {sample_idx} layer/topk shape {shape[1:]} " + f"does not match {expected_tail}" + ) + + +def shard_routed_experts_for_destination( + routed_experts: Sequence[np.ndarray], + *, + layer_indices: Sequence[int], + cp_rank: int, + cp_size: int, + qkv_format: str, + max_seq_len: int | None = None, +) -> list[np.ndarray]: + """Shard routing data for one PP/CP destination.""" + if cp_size < 1: + raise ValueError(f"cp_size must be positive, got {cp_size}") + if not 0 <= cp_rank < cp_size: + raise ValueError(f"cp_rank must be in [0, {cp_size}), got {cp_rank}") + if qkv_format not in {"thd", "bshd"}: + raise ValueError(f"unsupported qkv_format {qkv_format!r}") + if qkv_format == "bshd" and max_seq_len is None: + raise ValueError("max_seq_len is required for qkv_format='bshd'") + + layers = tuple(int(layer_idx) for layer_idx in layer_indices) + return [ + _shard_one_routed_experts( + sample, + layer_indices=layers, + cp_rank=cp_rank, + cp_size=cp_size, + qkv_format=qkv_format, + max_seq_len=max_seq_len, + ) + for sample in routed_experts + ] + + +def _shard_one_routed_experts( + routed_experts: np.ndarray, + *, + layer_indices: tuple[int, ...], + cp_rank: int, + cp_size: int, + qkv_format: str, + max_seq_len: int | None, +) -> np.ndarray: + if not isinstance(routed_experts, np.ndarray) or routed_experts.ndim != 3: + shape = getattr(routed_experts, "shape", None) + raise ValueError(f"routed experts must be a rank-3 numpy array, got shape={shape}") + + num_routed_tokens, num_layers, topk = routed_experts.shape + if any(layer_idx < 0 or layer_idx >= num_layers for layer_idx in layer_indices): + raise ValueError( + f"layer indices {layer_indices} are outside routed-expert layer dimension {num_layers}" + ) + + # The replay path appends one row for the final token before CP slicing. + num_tokens = num_routed_tokens + 1 + target_length = num_tokens if qkv_format == "thd" else int(max_seq_len) + if target_length < num_tokens: + raise ValueError( + f"max_seq_len {target_length} is shorter than sample token length {num_tokens}" + ) + + if cp_size == 1: + ranges = ((0, target_length),) + else: + chunk_size = (target_length + 2 * cp_size - 1) // (2 * cp_size) + ranges = ( + (chunk_size * cp_rank, chunk_size * (cp_rank + 1)), + ( + chunk_size * (2 * cp_size - cp_rank - 1), + chunk_size * (2 * cp_size - cp_rank), + ), + ) + + output_length = sum(end - start for start, end in ranges) + output = np.full( + (output_length, len(layer_indices), topk), + fill_value=-1, + dtype=routed_experts.dtype, + ) + + output_start = 0 + for start, end in ranges: + chunk_length = end - start + valid_end = min(end, num_routed_tokens) + if start < valid_end: + valid_length = valid_end - start + output[output_start : output_start + valid_length] = routed_experts[ + start:valid_end, layer_indices, : + ] + output_start += chunk_length + + return output diff --git a/train.py b/train.py index 9de6c30e0dd..6b881e0a5a3 100644 --- a/train.py +++ b/train.py @@ -80,13 +80,23 @@ async def save(rollout_id): offload_tags.append(GPU_MEMORY_TYPE_WEIGHTS) await rollout_manager.offload.remote(tags=offload_tags) + train_actor_this_step = not args.use_critic or rollout_id >= args.num_critic_only_steps + preload_groups = [] if args.use_critic: - critic_task = await eager_create_task(critic_model.train(rollout_id, rollout_data_ref)) - if rollout_id >= args.num_critic_only_steps: - await actor_model.train(rollout_id, rollout_data_ref) + preload_groups.append(critic_model) + if train_actor_this_step: + preload_groups.append(actor_model) + await asyncio.gather( + *(group.preload_rollout_data(rollout_id, rollout_data_ref) for group in preload_groups) + ) + + if args.use_critic: + critic_task = await eager_create_task(critic_model.train_preloaded(rollout_id)) + if train_actor_this_step: + await actor_model.train_preloaded(rollout_id) await critic_task else: - await actor_model.train(rollout_id, rollout_data_ref) + await actor_model.train_preloaded(rollout_id) if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout): await save(rollout_id) diff --git a/train_async.py b/train_async.py index cc46612b8bf..9c4908140a0 100644 --- a/train_async.py +++ b/train_async.py @@ -42,17 +42,27 @@ async def train(args): if rollout_data_next_future is not None: rollout_data_curr_ref = await rollout_data_next_future - # Start the next rollout early. + train_actor_this_step = not args.use_critic or rollout_id >= args.num_critic_only_steps + preload_groups = [] + if args.use_critic: + preload_groups.append(critic_model) + if train_actor_this_step: + preload_groups.append(actor_model) + await asyncio.gather( + *(group.preload_rollout_data(rollout_id, rollout_data_curr_ref) for group in preload_groups) + ) + + # Preload before the next rollout competes for transfer bandwidth. if rollout_id + 1 < args.num_rollout: rollout_data_next_future = rollout_manager.generate.remote(rollout_id + 1) if args.use_critic: - critic_task = await eager_create_task(critic_model.train(rollout_id, rollout_data_curr_ref)) - if rollout_id >= args.num_critic_only_steps: - await actor_model.train(rollout_id, rollout_data_curr_ref) + critic_task = await eager_create_task(critic_model.train_preloaded(rollout_id)) + if train_actor_this_step: + await actor_model.train_preloaded(rollout_id) await critic_task else: - await actor_model.train(rollout_id, rollout_data_curr_ref) + await actor_model.train_preloaded(rollout_id) if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout): await actor_model.save_model(