From 8d35808c0925b232e89002f598da47cd0b57b03e Mon Sep 17 00:00:00 2001 From: dyurk-lila Date: Thu, 16 Jul 2026 21:27:44 +0000 Subject: [PATCH 1/2] fix(r3): align routed-expert metadata and padding Fix routed-expert replay (R3) correctness for RL training and introduce a shared token-metadata layout that later routed-expert and sampler-support work builds on. - Scope global RouterReplay state to one Megatron pipeline schedule so a forward-only logprob pass can no longer leak backward replay state into the next training schedule (clear before the schedule and in `finally`). - Keep `rollout_expert_indices` ragged and treat its length as the captured-prefix length. Derive a `router_padding_mask` after left padding that marks alignment padding and the uncaptured trajectory suffix, and carry it through the training data, replay experiences, microbatch padding, and the Megatron model call. - Build one `TokenMetadataLayout` per microbatch and apply it to both routes and the padding mask. Generic construction, alignment, next-token shifting, and packed-output restoration live in `skyrl/utils/token_metadata.py`. - Pass Megatron's `padding_mask` through the model and apply a narrow compatibility shim so `[tokens]` masks broadcast over experts in expert-bias accounting. - Slice every per-trajectory generator field generically during dynamic-sampling replacement and filtering so route metadata stays attached to its trajectory. Synthetic padding rows use distinct dummy experts `[0, ..., topk - 1]`; the mask excludes them from expert-bias accounting while preserving Megatron's dropless `tokens * topk` dispatcher invariant. --- skyrl/backends/skyrl_train/training_batch.py | 17 +- .../skyrl_train/utils/replay_utils.py | 264 +++++++----------- .../megatron/megatron_model_wrapper.py | 104 ++++--- .../workers/megatron/megatron_worker.py | 14 + .../skyrl_train/workers/worker_utils.py | 10 +- skyrl/train/dataset/preprocess.py | 98 +++++-- skyrl/train/dataset/replay_buffer.py | 7 +- skyrl/train/generators/skyrl_gym_generator.py | 27 +- skyrl/train/generators/utils.py | 17 +- skyrl/train/trainer.py | 10 + skyrl/train/utils/trainer_utils.py | 35 +-- skyrl/utils/routed_experts.py | 14 + skyrl/utils/token_metadata.py | 204 ++++++++++++++ .../gpu/gpu_ci/megatron/test_router_replay.py | 19 +- .../test_token_based_batching_utils.py | 12 + .../backends/skyrl_train/test_train_batch.py | 16 +- .../skyrl_train/utils/test_replay_utils.py | 206 ++++++++++++++ tests/train/dataset/test_preprocess.py | 46 ++- .../generators/test_skyrl_gym_generator.py | 21 +- tests/train/test_trainer_utils.py | 9 + tests/utils/test_token_metadata.py | 88 ++++++ 21 files changed, 951 insertions(+), 287 deletions(-) create mode 100644 skyrl/utils/routed_experts.py create mode 100644 skyrl/utils/token_metadata.py create mode 100644 tests/backends/skyrl_train/utils/test_replay_utils.py create mode 100644 tests/utils/test_token_metadata.py diff --git a/skyrl/backends/skyrl_train/training_batch.py b/skyrl/backends/skyrl_train/training_batch.py index f2bac3ab5a..901ae0ce1d 100644 --- a/skyrl/backends/skyrl_train/training_batch.py +++ b/skyrl/backends/skyrl_train/training_batch.py @@ -7,7 +7,9 @@ import numpy as np import torch -from jaxtyping import Float, Integer +from jaxtyping import Bool, Float, Integer + +from skyrl.utils.routed_experts import make_replay_padding_indices DictType = TypeVar("DictType") @@ -476,6 +478,7 @@ class TrainingInput(TypedDict, total=False): rewards: Optional[Float[torch.Tensor, "batch_size seq_len"]] rollout_logprobs: Optional[Float[torch.Tensor, "batch_size seq_len"]] rollout_expert_indices: Optional[Integer[torch.Tensor, "batch_size seq_len layer_num topk"]] + router_padding_mask: Optional[Bool[torch.Tensor, "batch_size seq_len"]] pixel_values: Optional[TensorList] # list of `batch_size` [num_patches_i, dim] tensors image_grid_thw: Optional[TensorList] # list of `batch_size` [num_images_i, 3] tensors @@ -524,6 +527,18 @@ def pad_training_input_batch(unpadded_batch: TrainingInputBatch, pad_size: int) additional_dims = tensor.shape[1:] padding_tensor = torch.zeros(pad_size, *additional_dims, dtype=tensor.dtype, device=tensor.device) new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) + elif key == "rollout_expert_indices": + additional_dims = tensor.shape[1:] + padding_tensor = make_replay_padding_indices( + (pad_size, *additional_dims), + dtype=tensor.dtype, + device=tensor.device, + ) + new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) + elif key == "router_padding_mask": + additional_dims = tensor.shape[1:] + padding_tensor = torch.ones(pad_size, *additional_dims, dtype=torch.bool, device=tensor.device) + new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) else: # Copy row 0 `pad_size` times. Loss masked so values don't affect the loss. Just need valid shape/dtype. assert tensor.shape[0] > 0, f"Cannot pad empty tensor field {key!r}" diff --git a/skyrl/backends/skyrl_train/utils/replay_utils.py b/skyrl/backends/skyrl_train/utils/replay_utils.py index 3d8ed7b56a..0953ea7d8b 100644 --- a/skyrl/backends/skyrl_train/utils/replay_utils.py +++ b/skyrl/backends/skyrl_train/utils/replay_utils.py @@ -2,14 +2,14 @@ Utility functions for MoE Router Replay. """ +from contextlib import contextmanager from typing import List import torch -from skyrl.backends.skyrl_train.distributed.megatron.packing_utils import ( - get_packed_seq_align_size, - get_unpacked_seq_align_size, - is_fp8_enabled, +from skyrl.utils.token_metadata import ( + TokenMetadataLayout, + align_token_metadata, ) @@ -45,41 +45,26 @@ def patched_set_layer_number(self, layer_number: int): TopKRouter._set_layer_number_patched = True -def _patch_alltoall_dispatcher_for_replay(): - """Monkey-patch MoEAlltoAllTokenDispatcher.preprocess to handle router replay. - - When router replay is enabled, duplicate indices in top_indices can cause - routing_map.sum() < num_tokens * topk, leading to a split size mismatch - in the alltoall collective. We fix this by deriving num_out_tokens from - the routing map instead of the static num_tokens * topk formula. - - Reference: https://github.com/verl-project/verl/pull/4986 - """ +def patch_topk_router_expert_bias_padding_mask(): + """Fix the token-mask broadcast in pinned Megatron's expert-bias accounting.""" try: - from megatron.core.transformer.moe.token_dispatcher import ( - MoEAlltoAllTokenDispatcher, - ) + from megatron.core.transformer.moe.router import TopKRouter except ImportError: return - if getattr(MoEAlltoAllTokenDispatcher, "_preprocess_patched", False): + if getattr(TopKRouter, "_expert_bias_padding_mask_patched", False): return - original_preprocess = MoEAlltoAllTokenDispatcher.preprocess + original_apply_expert_bias = TopKRouter._apply_expert_bias - def patched_preprocess(self, routing_map): - result = original_preprocess(self, routing_map) - if ( - getattr(self.config, "moe_enable_routing_replay", False) - and not self.drop_and_pad - and self.config.moe_expert_capacity_factor is None - and not self.config.moe_router_padding_for_quantization - ): - self.num_out_tokens = int(routing_map.sum().item()) - return result + def patched_apply_expert_bias(self, routing_map: torch.Tensor, padding_mask: torch.Tensor | None = None): + # Megatron combines [tokens, experts] with a token-only mask. + if padding_mask is not None and padding_mask.ndim == 1: + padding_mask = padding_mask.unsqueeze(-1) + return original_apply_expert_bias(self, routing_map, padding_mask) - MoEAlltoAllTokenDispatcher.preprocess = patched_preprocess - MoEAlltoAllTokenDispatcher._preprocess_patched = True + TopKRouter._apply_expert_bias = patched_apply_expert_bias + TopKRouter._expert_bias_padding_mask_patched = True def _split_replay_indices(rollout_expert_indices: torch.Tensor) -> List[torch.Tensor]: @@ -92,120 +77,30 @@ def _split_replay_indices(rollout_expert_indices: torch.Tensor) -> List[torch.Te return [per_layer[i].reshape(-1, per_layer.shape[-1]) for i in range(per_layer.shape[0])] -def _remove_left_padding_from_indices( - rollout_expert_indices: torch.Tensor, - attention_mask: torch.Tensor, - fp8_enabled: bool = False, -) -> torch.Tensor: - """Apply the same left-padding removal as remove_left_padding to routing indices. - - Args: - rollout_expert_indices: [batch, padded_seq_len, layers, topk] - attention_mask: [batch, padded_seq_len] (int or bool) - - Returns: - [batch, effective_seq_len, layers, topk] with real tokens packed left. - """ - import megatron.core.parallel_state as mpu - - seq_lens = attention_mask.sum(dim=1) - effective_seq_len = seq_lens.max().item() - tp_size = mpu.get_tensor_model_parallel_world_size() - align_size = get_unpacked_seq_align_size(tp_size, fp8_enabled=fp8_enabled) - if align_size > 1: - pad_size = (align_size - effective_seq_len % align_size) % align_size - effective_seq_len += pad_size - - batch_size = rollout_expert_indices.shape[0] - new_rii = torch.zeros( - batch_size, - effective_seq_len, - rollout_expert_indices.shape[2], - rollout_expert_indices.shape[3], - dtype=rollout_expert_indices.dtype, - device=rollout_expert_indices.device, - ) - for i in range(batch_size): - mask = attention_mask[i].bool() - new_rii[i, : seq_lens[i]] = rollout_expert_indices[i, mask] - return new_rii - - -def _pack_replay_indices( - rollout_expert_indices: torch.Tensor, - attention_mask: torch.Tensor, - fp8_enabled: bool = False, -) -> torch.Tensor: - """Pack routing indices to match the token layout produced by preprocess_packed_seqs. - - With sample packing, Megatron concatenates all sequences into one packed - sequence with per-sample alignment padding. The MoE router sees tokens in - this packed order, so replay indices must follow the same layout. - - Returns: - [1, total_packed_len, layers, topk] matching the packed model input. - """ - import megatron.core.parallel_state as mpu - - batch_size = rollout_expert_indices.shape[0] - num_layers = rollout_expert_indices.shape[2] - topk = rollout_expert_indices.shape[3] - - seq_lens = attention_mask.sum(dim=-1, dtype=torch.int32) - tp_size = mpu.get_tensor_model_parallel_world_size() - cp_size = mpu.get_context_parallel_world_size() - align_size = get_packed_seq_align_size(tp_size, cp_size, fp8_enabled=fp8_enabled) - - pad_sizes = (align_size - seq_lens % align_size) % align_size - seqlens_padded = seq_lens + pad_sizes - - total_packed_len = int(seqlens_padded.sum().item()) - - packed = torch.zeros( - total_packed_len, - num_layers, - topk, - dtype=rollout_expert_indices.dtype, - device=rollout_expert_indices.device, +def scatter_router_padding_mask_for_model( + router_padding_mask: torch.Tensor | None, + model, + model_config, +) -> torch.Tensor | None: + """Match the mask layout to sequence-parallel hidden states at model entry.""" + if router_padding_mask is None or not model_config.sequence_parallel: + return router_padding_mask + + from megatron.core.models.hybrid.hybrid_model import HybridModel + from megatron.core.tensor_parallel import scatter_to_sequence_parallel_region + from megatron.core.utils import unwrap_model + + unwrapped_model = unwrap_model(model) + # GPTModel scatters its mask beside the embedding on the first PP stage. HybridModel + # scatters only the embedding, so its mask must always be scattered here. + if not isinstance(unwrapped_model, HybridModel) and unwrapped_model.pre_process: + return router_padding_mask + return ( + scatter_to_sequence_parallel_region(router_padding_mask.transpose(0, 1).contiguous()) + .transpose(0, 1) + .contiguous() ) - seq_lens_cpu = seq_lens.tolist() - seqlens_padded_cpu = seqlens_padded.tolist() - offset = 0 - for i in range(batch_size): - n = seq_lens_cpu[i] - mask = attention_mask[i].bool() - d = rollout_expert_indices[i, mask] - packed[offset : offset + n] = d - offset += seqlens_padded_cpu[i] - - if cp_size > 1: - cp_rank = mpu.get_context_parallel_rank() - out = torch.zeros( - total_packed_len // cp_size, - num_layers, - topk, - dtype=packed.dtype, - device=packed.device, - ) - src_offset = 0 - dst_offset = 0 - for i in range(batch_size): - seqlen_padded_i = seqlens_padded_cpu[i] - seqlen_per_cp = seqlen_padded_i // cp_size - half = seqlen_per_cp // 2 - out[dst_offset : dst_offset + half] = packed[ - src_offset + half * cp_rank : src_offset + half * (cp_rank + 1) - ] - back_start = src_offset + seqlen_padded_i - half * (cp_rank + 1) - back_end = src_offset + seqlen_padded_i - half * cp_rank - out[dst_offset + half : dst_offset + seqlen_per_cp] = packed[back_start:back_end] - src_offset += seqlen_padded_i - dst_offset += seqlen_per_cp - packed = out - - return packed.unsqueeze(0) # [1, packed_len_per_cp, layers, topk] - def _get_current_pp_stage_layer_range(model_config) -> tuple[int, int]: """Return the current PP rank's transformer-layer range as (start_layer, @@ -226,12 +121,20 @@ def _get_current_pp_stage_layer_range(model_config) -> tuple[int, int]: def setup_per_microbatch_replay_forward( rollout_expert_indices: torch.Tensor, + router_padding_mask: torch.Tensor | None, attention_mask: torch.Tensor, + model, model_config, + metadata_layout: TokenMetadataLayout, remove_microbatch_padding: bool = False, -) -> None: - """Set up RouterReplay for a single micro-batch, aligning indices - with the left-padding-removed token layout that the MoE layer sees. +) -> dict[str, torch.Tensor]: + """Set up router replay and return its model-facing keyword arguments. + + Replay indices and the router padding mask start in the same batch layout and + undergo matching padding removal or packing and CP sharding. Their destinations + then differ: indices are TP-sliced and installed into per-layer ``RouterReplay`` + instances, while the mask follows Megatron's model-specific sequence-parallel + path and is passed to the model as ``padding_mask``. Handles context parallelism: when CP > 1, the sequence is split into 2*cp_size chunks with each CP rank receiving a front chunk and a back @@ -247,8 +150,8 @@ def setup_per_microbatch_replay_forward( layers. We use each instance's global layer_number (set by the patched TopKRouter.set_layer_number) to index into the correct slice of the data. - Handles pipeline parallelism: when PP > 1, the sequence is split across - PP ranks, so each rank only sees its local RouterReplay instances. In cases + Handles pipeline parallelism: when PP > 1, transformer layers are split + across PP ranks, so each rank only sees its local RouterReplay instances. In cases where the number of local RouterReplay instances does not match the local layer count, indicating that the model has dense layers before MoE layers, we use the global layer_number to index into the correct slice of the data. @@ -260,26 +163,41 @@ def setup_per_microbatch_replay_forward( RouterReplayAction, ) - _patch_alltoall_dispatcher_for_replay() - fp8_enabled = is_fp8_enabled(getattr(model_config, "fp8", None)) + if router_padding_mask is None: + raise ValueError("router_padding_mask is required with rollout_expert_indices") - if remove_microbatch_padding: - aligned = _pack_replay_indices(rollout_expert_indices, attention_mask, fp8_enabled=fp8_enabled) - else: - aligned = _remove_left_padding_from_indices( - rollout_expert_indices, - attention_mask, - fp8_enabled=fp8_enabled, + if router_padding_mask.shape != attention_mask.shape: + raise ValueError( + f"router_padding_mask shape {router_padding_mask.shape} does not match " + f"attention_mask shape {attention_mask.shape}" ) + if router_padding_mask.device != rollout_expert_indices.device: + raise ValueError("rollout_expert_indices and router_padding_mask must be on the same device") + + if (metadata_layout.padded_sequence_lengths is not None) != remove_microbatch_padding: + raise ValueError("Shared token metadata layout does not match the model packing mode") + aligned_router_padding_mask = align_token_metadata(router_padding_mask.to(torch.bool), metadata_layout, True) + route_padding = torch.arange( + rollout_expert_indices.shape[-1], + dtype=rollout_expert_indices.dtype, + device=rollout_expert_indices.device, + ) + aligned_rollout_expert_indices = align_token_metadata( + rollout_expert_indices, + metadata_layout, + route_padding, + ) # TP splitting: sequence parallelism across the tensor model parallel region tp_size = mpu.get_tensor_model_parallel_world_size() if tp_size > 1: tp_rank = mpu.get_tensor_model_parallel_rank() - seq_len = aligned.shape[1] + seq_len = aligned_rollout_expert_indices.shape[1] chunk_size = seq_len // tp_size - aligned = aligned[:, tp_rank * chunk_size : (tp_rank + 1) * chunk_size, :, :] - per_layer_data = _split_replay_indices(aligned) + aligned_rollout_expert_indices = aligned_rollout_expert_indices[ + :, tp_rank * chunk_size : (tp_rank + 1) * chunk_size, :, : + ] + per_layer_data = _split_replay_indices(aligned_rollout_expert_indices) global_num_layers_in_data = len(per_layer_data) instances = RouterReplay.global_router_replay_instances num_instances = len(instances) @@ -307,6 +225,13 @@ def setup_per_microbatch_replay_forward( router_instance.set_target_indices(per_layer_data[layer_idx]) RouterReplay.set_global_router_replay_action(RouterReplayAction.REPLAY_FORWARD) + model_router_padding_mask = scatter_router_padding_mask_for_model( + aligned_router_padding_mask, + model, + model_config, + ) + return {"padding_mask": model_router_padding_mask} + def setup_per_microbatch_replay_backward() -> None: """Switch RouterReplay to backward mode so that activation-checkpoint @@ -327,3 +252,22 @@ def clear_router_replay(): RouterReplay.clear_global_indices() RouterReplay.clear_global_router_replay_action() + + +@contextmanager +def router_replay_schedule(enabled: bool): + """Isolate global RouterReplay state to one Megatron pipeline schedule. + + The backward FIFO spans all microbatches in a training schedule, so it must + only be cleared at schedule boundaries. Forward-only schedules leave that + FIFO unconsumed, and failed schedules may leave it partially consumed. + """ + if not enabled: + yield + return + + clear_router_replay() + try: + yield + finally: + clear_router_replay() diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py index 61e5d71ff7..a49fa3644d 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py @@ -43,6 +43,7 @@ compute_approx_kl, ) from skyrl.backends.skyrl_train.utils.replay_utils import ( + router_replay_schedule, setup_per_microbatch_replay_backward, setup_per_microbatch_replay_forward, ) @@ -51,6 +52,7 @@ compute_minibatch_rollout_logprob_diff_metrics, ) from skyrl.train.config import TrainerConfig +from skyrl.utils.token_metadata import build_token_metadata_layout def _build_packed_targets( @@ -329,13 +331,7 @@ def forward_step(batch_iter, model): model_config = get_model_config(model) fp8_enabled = is_fp8_enabled(getattr(model_config, "fp8", None)) rollout_expert_indices = batch.pop("rollout_expert_indices", None) - if rollout_expert_indices is not None: - setup_per_microbatch_replay_forward( - rollout_expert_indices, - batch["attention_mask"], - model_config=model_config, - remove_microbatch_padding=self.remove_microbatch_padding, - ) + router_padding_mask = batch.pop("router_padding_mask", None) sequences = batch["sequences"] attention_mask = batch["attention_mask"].to(bool) @@ -378,6 +374,27 @@ def forward_step(batch_iter, model): if self.is_vlm: new_position_ids = None + metadata_layout = None + if rollout_expert_indices is not None: + metadata_layout = build_token_metadata_layout( + attention_mask, + attention_mask.device, + packed=packed_seq_params is not None, + fp8_enabled=fp8_enabled, + ) + + model_replay_kwargs = {} + if rollout_expert_indices is not None: + model_replay_kwargs = setup_per_microbatch_replay_forward( + rollout_expert_indices, + router_padding_mask, + attention_mask, + model=model, + model_config=model_config, + metadata_layout=metadata_layout, + remove_microbatch_padding=self.remove_microbatch_padding, + ) + if self._fused_lm_head: # Fused LM-head inference: the output_processor returns decoder # hidden states (not logits) and stashes the LM-head weight, so @@ -393,6 +410,7 @@ def forward_step(batch_iter, model): packed_seq_params=packed_seq_params, output_processor=_fused_lm_head_output_processor, output_processor_context=_op_ctx, + **model_replay_kwargs, **vlm_inputs, ) batch["lm_head_weight"] = _op_ctx.get("lm_head_weight") @@ -402,6 +420,7 @@ def forward_step(batch_iter, model): new_position_ids, to_te_attention_mask(new_attention_mask), packed_seq_params=packed_seq_params, + **model_replay_kwargs, **vlm_inputs, ) @@ -418,15 +437,17 @@ def forward_step(batch_iter, model): batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.actor_module)) - output = forward_backward_func( - forward_step_func=forward_step, - data_iterator=batch_generator, - model=self.actor_module, - num_microbatches=len(micro_batches), - seq_length=seq_len, - micro_batch_size=micro_batch_size, - forward_only=True, - ) + replay_enabled = any(batch["rollout_expert_indices"] is not None for batch in micro_batches) + with router_replay_schedule(replay_enabled): + output = forward_backward_func( + forward_step_func=forward_step, + data_iterator=batch_generator, + model=self.actor_module, + num_microbatches=len(micro_batches), + seq_length=seq_len, + micro_batch_size=micro_batch_size, + forward_only=True, + ) if mpu.is_pipeline_last_stage(ignore_virtual=True): log_probs = [o["log_probs"] for o in output] @@ -895,13 +916,7 @@ def forward_step(batch_iter, model): model_config = get_model_config(model) fp8_enabled = is_fp8_enabled(getattr(model_config, "fp8", None)) rollout_expert_indices = batch.pop("rollout_expert_indices", None) - if rollout_expert_indices is not None: - setup_per_microbatch_replay_forward( - rollout_expert_indices, - batch["attention_mask"], - model_config=model_config, - remove_microbatch_padding=self.remove_microbatch_padding, - ) + router_padding_mask = batch.pop("router_padding_mask", None) sequences = batch["sequences"] attention_mask = batch["attention_mask"].to(bool) @@ -965,6 +980,27 @@ def forward_step(batch_iter, model): is_last_stage = mpu.is_pipeline_last_stage(ignore_virtual=True) + metadata_layout = None + if rollout_expert_indices is not None: + metadata_layout = build_token_metadata_layout( + attention_mask, + attention_mask.device, + packed=packed_seq_params is not None, + fp8_enabled=fp8_enabled, + ) + + model_replay_kwargs = {} + if rollout_expert_indices is not None: + model_replay_kwargs = setup_per_microbatch_replay_forward( + rollout_expert_indices, + router_padding_mask, + attention_mask, + model=model, + model_config=model_config, + metadata_layout=metadata_layout, + remove_microbatch_padding=self.remove_microbatch_padding, + ) + # Recover [batch, seq_len, ...] from Megatron's internal (left-removed) layout. Only used # on the non-packed path: with sample packing (remove_microbatch_padding) the logits stay # packed ([1, T, vocab]) and loss_func consumes packed_targets instead. MTP draft training @@ -1017,6 +1053,7 @@ def depad(tensor): packed_seq_params=packed_seq_params, output_processor=_fused_lm_head_output_processor, output_processor_context=_op_ctx, + **model_replay_kwargs, **vlm_inputs, ) batch["lm_head_weight"] = _op_ctx.get("lm_head_weight") @@ -1026,6 +1063,7 @@ def depad(tensor): new_position_ids, to_te_attention_mask(new_attention_mask), packed_seq_params=packed_seq_params, + **model_replay_kwargs, **vlm_inputs, ) # Replay the MTP block on *detached* trunk hidden states (decoupled draft forward) @@ -1059,15 +1097,17 @@ def depad(tensor): # batch should be a list of micro-batches batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.actor_module)) - metrics_list = forward_backward_func( - forward_step_func=forward_step, - data_iterator=batch_generator, - model=self.actor_module, - num_microbatches=len(micro_batches), - seq_length=seq_len, - micro_batch_size=micro_batch_size, - forward_only=forward_only, - ) + replay_enabled = any(batch["rollout_expert_indices"] is not None for batch in micro_batches) + with router_replay_schedule(replay_enabled): + metrics_list = forward_backward_func( + forward_step_func=forward_step, + data_iterator=batch_generator, + model=self.actor_module, + num_microbatches=len(micro_batches), + seq_length=seq_len, + micro_batch_size=micro_batch_size, + forward_only=forward_only, + ) # The decoupled MTP/draft loss is computed and logged per-microbatch inside loss_func # (metric key "mtp_loss"); no MTPLossLoggingHelper plumbing is needed. diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py index 9f2facaf61..825c2f798b 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py @@ -73,6 +73,7 @@ from skyrl.env_vars import SKYRL_WORKER_NCCL_TIMEOUT_IN_S from skyrl.train.config.config import MegatronDDPConfig, get_config_as_dict from skyrl.train.utils.utils import str_to_torch_dtype, update_model_config +from skyrl.utils.routed_experts import make_replay_padding_indices from skyrl.utils.tok import get_tokenizer if TYPE_CHECKING: @@ -672,6 +673,7 @@ def _forward_logprobs(self, data: TrainingInputBatch) -> torch.Tensor: "position_ids": position_ids, "num_actions": micro.metadata["response_length"], "rollout_expert_indices": (rollout_expert_indices if self.enable_router_replay else None), + "router_padding_mask": micro.get("router_padding_mask") if self.enable_router_replay else None, "sub_seq_lengths": micro.get("sub_seq_lengths"), **vlm_inputs, } @@ -787,6 +789,14 @@ def _pad_microbatch_to_size(self, micro_dict: dict, target_batch_size: int) -> d # position_ids for padded samples seq_len = value.shape[1] pad_tensor = torch.arange(seq_len, device=device).unsqueeze(0).expand(pad_count, -1) + elif key == "router_padding_mask": + pad_tensor = torch.ones((pad_count, *value.shape[1:]), dtype=torch.bool, device=device) + elif key == "rollout_expert_indices": + pad_tensor = make_replay_padding_indices( + (pad_count, *value.shape[1:]), + dtype=value.dtype, + device=device, + ) elif key == "action_mask": # action_mask should be zeros for padded samples pad_tensor = torch.zeros((pad_count, *value.shape[1:]), dtype=value.dtype, device=device) @@ -907,9 +917,11 @@ def init_model(self, model_path, num_training_steps: int = 1e9): if self.enable_router_replay: from skyrl.backends.skyrl_train.utils.replay_utils import ( + patch_topk_router_expert_bias_padding_mask, patch_topk_router_layer_number, ) + patch_topk_router_expert_bias_padding_mask() patch_topk_router_layer_number() # Freeze MoE router params before optimizer build. @@ -1052,6 +1064,7 @@ def forward( "rollout_action_logprobs": experience.rollout_logprobs, "action_mask": experience.action_mask, "rollout_expert_indices": rollout_expert_indices if self.enable_router_replay else None, + "router_padding_mask": experience.router_padding_mask if self.enable_router_replay else None, "sub_seq_lengths": experience.sub_seq_lengths, **vlm_inputs, } @@ -1180,6 +1193,7 @@ def forward_backward( "rollout_action_logprobs": experience.rollout_logprobs, "action_mask": experience.action_mask, "rollout_expert_indices": rollout_expert_indices if self.enable_router_replay else None, + "router_padding_mask": experience.router_padding_mask if self.enable_router_replay else None, # used with global sequence packing (None when token-based batching is active) "sub_seq_lengths": experience.sub_seq_lengths, "is_padding_batch": ( diff --git a/skyrl/backends/skyrl_train/workers/worker_utils.py b/skyrl/backends/skyrl_train/workers/worker_utils.py index 5642aa8f12..1d2078f853 100644 --- a/skyrl/backends/skyrl_train/workers/worker_utils.py +++ b/skyrl/backends/skyrl_train/workers/worker_utils.py @@ -9,6 +9,7 @@ from skyrl.backends.skyrl_train.utils.torch_utils import masked_mean from skyrl.train.dataset.bin_packing import make_seq_packer from skyrl.train.dataset.replay_buffer import Experience +from skyrl.utils.routed_experts import make_replay_padding_indices # Metrics that end in `_loss` but are plain per-token MEANS, not pre-scaled minibatch sums. # The `sum_loss_metrics` convention sums every `_loss` key because the *policy* losses are @@ -162,6 +163,7 @@ def batch_to_experience(batch: TrainingInputBatch): num_actions=batch.metadata["response_length"], # int rollout_logprobs=batch.get("rollout_logprobs"), rollout_expert_indices=batch.get("rollout_expert_indices"), + router_padding_mask=batch.get("router_padding_mask"), # additional info # can be used to log metrics etc for micro-batches in the worker info={}, @@ -325,9 +327,13 @@ def _create_padding_microbatch(self) -> TrainingInputBatch: data["rollout_logprobs"] = torch.zeros((batch_size, num_actions), dtype=ref_tensor.dtype, device=device) if self.data.get("rollout_expert_indices") is not None: ref_tensor = self.data["rollout_expert_indices"] - data["rollout_expert_indices"] = torch.zeros( - (batch_size, *ref_tensor.shape[1:]), dtype=ref_tensor.dtype, device=device + data["rollout_expert_indices"] = make_replay_padding_indices( + (batch_size, *ref_tensor.shape[1:]), + dtype=ref_tensor.dtype, + device=device, ) + if self.data.get("router_padding_mask") is not None: + data["router_padding_mask"] = torch.ones((batch_size, seq_len), dtype=torch.bool, device=device) data.metadata = {} if self.data.metadata: data.metadata.update(self.data.metadata) diff --git a/skyrl/train/dataset/preprocess.py b/skyrl/train/dataset/preprocess.py index 2b34521d1c..f5af10741b 100644 --- a/skyrl/train/dataset/preprocess.py +++ b/skyrl/train/dataset/preprocess.py @@ -2,12 +2,55 @@ from typing import List, Optional, Tuple import torch -from jaxtyping import Float, Integer +from jaxtyping import Bool, Float, Integer from transformers import AutoTokenizer +from skyrl.utils.routed_experts import make_replay_padding_indices + logger = logging.getLogger(__name__) +def make_router_padding_mask( + attention_mask: torch.Tensor, + captured_route_lengths: List[int], +) -> Bool[torch.Tensor, "batch seq_len"]: + """Build Megatron's router-only padding mask for a ragged vLLM route prefix. + + vLLM records routes only for tokens it evaluates. The final training sequence can be + longer because the last sampled token has no subsequent decode forward, and SkyRL may + append a synthetic EOS. In multi-turn generation, observations join the captured prefix + only when a later turn evaluates them. Captured route rows therefore align with a prefix + of each real, left-padded sequence; the remaining suffix needs dummy routes. + + This cannot be derived from the loss mask. A loss-masked prompt or observation may still + condition later trained actions and must replay its captured route. ``True`` marks only + left padding and tokens without a captured route so Megatron excludes their dummy routes + from router accounting. + """ + if attention_mask.ndim != 2: + raise ValueError(f"Expected 2D attention_mask, got shape {attention_mask.shape}") + if len(captured_route_lengths) != attention_mask.shape[0]: + raise ValueError( + f"Expected one captured route length per trajectory, got {len(captured_route_lengths)} " + f"for batch size {attention_mask.shape[0]}" + ) + + captured = torch.as_tensor(captured_route_lengths, dtype=torch.long, device=attention_mask.device) + sequence_lengths = attention_mask.sum(dim=1, dtype=torch.long) + if torch.any(captured < 0) or torch.any(captured > sequence_lengths): + raise ValueError( + f"Captured route lengths must be within trajectory lengths, got " + f"captured={captured.tolist()} and lengths={sequence_lengths.tolist()}" + ) + + sequence_starts = attention_mask.shape[1] - sequence_lengths + positions = torch.arange(attention_mask.shape[1], device=attention_mask.device).unsqueeze(0) + captured_positions = (positions >= sequence_starts.unsqueeze(1)) & ( + positions < (sequence_starts + captured).unsqueeze(1) + ) + return ~captured_positions + + def _verify_inputs( prompts: List[List[int]], responses: List[List[int]], @@ -160,24 +203,41 @@ def convert_prompts_responses_to_batch_tensors( logprobs_tensor[i, max_response - len(sample_logprobs) :] = lp rollout_expert_indices_tensor = None - if rollout_expert_indices: - first_non_empty = next((x for x in rollout_expert_indices if x), None) - if first_non_empty: - num_layers = len(first_non_empty[0]) - topk = len(first_non_empty[0][0]) if num_layers > 0 else 0 - padded = torch.zeros(len(rollout_expert_indices), max_total, num_layers, topk, dtype=torch.int32) - for i, sample_indices in enumerate(rollout_expert_indices): - if sample_indices: - left_pad = max_total - (prompt_token_lens[i] + response_token_lens[i]) - n = min(len(sample_indices), max_total - left_pad) - padded[i, left_pad : left_pad + n] = torch.tensor(sample_indices[:n], dtype=torch.int32) - rollout_expert_indices_tensor = padded - - # downcast to uint8 if possible, otherwise int16 to save memory - if rollout_expert_indices_tensor.max().item() < 2**8: - rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.uint8) - elif rollout_expert_indices_tensor.max().item() < 2**15: - rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.int16) + if rollout_expert_indices is not None: + num_samples = len(prompts) + if len(rollout_expert_indices) != num_samples or any(not indices for indices in rollout_expert_indices): + raise ValueError("rollout_expert_indices must contain routes for every trajectory") + + num_layers = len(rollout_expert_indices[0][0]) + topk = len(rollout_expert_indices[0][0][0]) if num_layers > 0 else 0 + if topk < 1: + raise ValueError("rollout_expert_indices must contain at least one expert per layer") + + padded = make_replay_padding_indices( + (num_samples, max_total, num_layers, topk), + dtype=torch.int32, + ) + for sample_index, sample_indices in enumerate(rollout_expert_indices): + sample_indices_tensor = torch.as_tensor(sample_indices, dtype=torch.int32) + if sample_indices_tensor.ndim != 3 or sample_indices_tensor.shape[1:] != (num_layers, topk): + raise ValueError( + "rollout_expert_indices entries must share [layers, topk], " + f"got shape {tuple(sample_indices_tensor.shape)} at sample {sample_index}" + ) + left_pad = max_total - (prompt_token_lens[sample_index] + response_token_lens[sample_index]) + available = max_total - left_pad + if len(sample_indices) > available: + raise ValueError( + f"Trajectory {sample_index} has {len(sample_indices)} route rows for {available} tokens" + ) + padded[sample_index, left_pad : left_pad + len(sample_indices)] = sample_indices_tensor + rollout_expert_indices_tensor = padded + + max_expert_id = int(rollout_expert_indices_tensor.max().item()) + if max_expert_id < 2**8: + rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.uint8) + elif max_expert_id < 2**15: + rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.int16) return ( sequences, diff --git a/skyrl/train/dataset/replay_buffer.py b/skyrl/train/dataset/replay_buffer.py index 77dd1b1aca..37b8db1966 100644 --- a/skyrl/train/dataset/replay_buffer.py +++ b/skyrl/train/dataset/replay_buffer.py @@ -12,7 +12,7 @@ import torch import torch.nn.functional as F -from jaxtyping import Float, Integer +from jaxtyping import Bool, Float, Integer from skyrl.backends.skyrl_train.training_batch import TensorList @@ -70,6 +70,7 @@ class Experience: rollout_expert_indices: Optional[Integer[torch.Tensor, "batch seq_len layer_num topk"]] num_actions: int info: Optional[dict] + router_padding_mask: Optional[Bool[torch.Tensor, "batch seq_len"]] = None kl: Optional[Float[torch.Tensor, "batch response_len"]] = None metadata: Optional[Dict[str, Any]] = None pixel_values: Optional[TensorList] = None @@ -101,6 +102,8 @@ def to_device(self, device: torch.device) -> None: self.rollout_logprobs = to(self.rollout_logprobs, device) if self.rollout_expert_indices is not None: self.rollout_expert_indices = to(self.rollout_expert_indices, device) + if self.router_padding_mask is not None: + self.router_padding_mask = to(self.router_padding_mask, device) if self.pixel_values is not None: self.pixel_values = self.pixel_values.to(device) if self.image_grid_thw is not None: @@ -130,6 +133,8 @@ def pin_memory(self): self.rollout_logprobs = self.rollout_logprobs.pin_memory() if self.rollout_expert_indices is not None: self.rollout_expert_indices = self.rollout_expert_indices.pin_memory() + if self.router_padding_mask is not None: + self.router_padding_mask = self.router_padding_mask.pin_memory() return self diff --git a/skyrl/train/generators/skyrl_gym_generator.py b/skyrl/train/generators/skyrl_gym_generator.py index cc9bc95e80..71a3b81477 100644 --- a/skyrl/train/generators/skyrl_gym_generator.py +++ b/skyrl/train/generators/skyrl_gym_generator.py @@ -91,25 +91,8 @@ class TurnOutput: added_eos: bool = False def get_turn_rollout_expert_indices(self) -> Optional[List[List[List[int]]]]: - """ - Get rollout inference indices for this turn's tokens (output tokens + observation tokens). - - Returns indices for generated output tokens, with padding entries (all 0) - for any manually-added EOS token and observation tokens - Returns None if rollout_expert_indices is None. - """ - if self.rollout_expert_indices is None: - return None - if not self.rollout_expert_indices: - return self.rollout_expert_indices - layer_num = len(self.rollout_expert_indices[0]) - topk = len(self.rollout_expert_indices[0][0]) if layer_num > 0 else 0 - pad_entry = [[0] * topk for _ in range(layer_num)] - indices = list(self.rollout_expert_indices) - if self.added_eos: - indices.append(pad_entry) - indices.extend(pad_entry for _ in range(len(self.obs_ids))) - return indices + """Return only routes that the inference model actually executed.""" + return self.rollout_expert_indices def get_turn_loss_mask(self) -> List[int]: """ @@ -569,15 +552,9 @@ async def agent_loop( assert response_ids is not None and loss_mask is not None if stop_reason != "length" and response_ids and response_ids[-1] != self.tokenizer.eos_token_id: response_ids.append(self.tokenizer.eos_token_id) - # TODO(Charlie): this should be 0? Otherwise logprobs will be extremely off. But if it is loss - # masked with 0, why bother adding it? loss_mask.append(1) if rollout_logprobs is not None: rollout_logprobs.append(0.0) - if rollout_expert_indices_out is not None and rollout_expert_indices_out: - layer_num = len(rollout_expert_indices_out[0]) - topk = len(rollout_expert_indices_out[0][0]) if layer_num > 0 else 0 - rollout_expert_indices_out.append([[0] * topk for _ in range(layer_num)]) appended_eos_token = True if self.generator_cfg.step_wise_trajectories: diff --git a/skyrl/train/generators/utils.py b/skyrl/train/generators/utils.py index 4c5228f189..0edc93ca17 100644 --- a/skyrl/train/generators/utils.py +++ b/skyrl/train/generators/utils.py @@ -745,19 +745,22 @@ def _is_prefix(maybe_prefix: List[int], candidate: List[int]) -> bool: return maybe_prefix == candidate[: len(maybe_prefix)] -def _slice_generator_output(generator_output: GeneratorOutput, indices: List[int]) -> GeneratorOutput: +def slice_generator_output( + generator_output: GeneratorOutput, indices: List[int], *, preserve_metrics: bool = True +) -> GeneratorOutput: """Slice a GeneratorOutput to keep only the entries at the given indices. - All sliced entries must share the same TrajectoryID — this helper is used by - prefix-aware merging which operates on one trajectory at a time. + Generator-specific per-trajectory fields are sliced without naming them here. + Prefix-aware merging passes entries that all share one ``TrajectoryID``; + dynamic sampling may intentionally select entries from different trajectories. """ assert len(indices) > 0, "indices must be non-empty" # Every key except `rollout_metrics` is either a per-entry list to slice, or None. sliced: GeneratorOutput = {} for key, value in generator_output.items(): if key == "rollout_metrics": - # Skip since metrics are already recorded before calling `merge_stepwise_output()`. - continue + if preserve_metrics: + sliced[key] = value elif value is None: sliced[key] = None else: @@ -913,7 +916,9 @@ def merge_stepwise_output(generator_output: GeneratorOutput) -> GeneratorOutput: start = 0 for i in range(num_samples): if is_last_step[i]: - trajectory_slices.append(_slice_generator_output(generator_output, list(range(start, i + 1)))) + trajectory_slices.append( + slice_generator_output(generator_output, list(range(start, i + 1)), preserve_metrics=False) + ) start = i + 1 merged_slices = [_merge_single_trajectory(s) for s in trajectory_slices] diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index 78dab4e84f..2ad4c0a50d 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -55,6 +55,7 @@ compute_prompt_boundaries, compute_prompt_mini_batch_boundaries, convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, ) from skyrl.train.evaluate import evaluate, evaluate_step_wise from skyrl.train.generators.base import ( @@ -898,6 +899,12 @@ def convert_to_training_input(self, generator_output: GeneratorOutput, uids: Lis rollout_expert_indices, max_seq_len=self.cfg.trainer.algorithm.max_seq_len, ) + router_padding_mask = None + if rollout_expert_indices is not None: + router_padding_mask = make_router_padding_mask( + attention_masks_tensor, + [len(indices) for indices in rollout_expert_indices], + ) # sanity check for off_policy_correction off_policy_correction = self.cfg.trainer.algorithm.off_policy_correction @@ -919,6 +926,7 @@ def convert_to_training_input(self, generator_output: GeneratorOutput, uids: Lis "loss_mask": loss_masks_tensor, "rollout_logprobs": rollout_logprobs_tensor, "rollout_expert_indices": rollout_expert_indices_tensor, + "router_padding_mask": router_padding_mask, "pixel_values": pixel_values, "image_grid_thw": image_grid_thw, }, @@ -1300,6 +1308,8 @@ def fwd_logprobs_values_reward( fwd_keys = ["sequences", "attention_mask"] if training_input.get("rollout_expert_indices") is not None: fwd_keys.append("rollout_expert_indices") + if training_input.get("router_padding_mask") is not None: + fwd_keys.append("router_padding_mask") if training_input.get("pixel_values") is not None: fwd_keys.append("pixel_values") if training_input.get("image_grid_thw") is not None: diff --git a/skyrl/train/utils/trainer_utils.py b/skyrl/train/utils/trainer_utils.py index 84cb15b9be..0d91977ecb 100644 --- a/skyrl/train/utils/trainer_utils.py +++ b/skyrl/train/utils/trainer_utils.py @@ -27,6 +27,7 @@ from skyrl.train.generators.utils import ( concatenate_generator_outputs, get_metrics_from_generator_output, + slice_generator_output, ) BasicType = Union[int, float, str, bool, type(None)] @@ -489,20 +490,11 @@ def handle_replace_sampling( for uid in bad_uids: bad_indices.extend(uid2indices[uid]) - # Replace bad samples with good ones (modify in place because replacement_idx and bad_idx should not overlap) + source_indices = list(range(len(uids))) for bad_idx, replacement_idx in zip(bad_indices, replacement_indices): - generator_output["prompt_token_ids"][bad_idx] = generator_output["prompt_token_ids"][replacement_idx].copy() - generator_output["response_ids"][bad_idx] = generator_output["response_ids"][replacement_idx].copy() - replacement_reward = generator_output["rewards"][replacement_idx] - generator_output["rewards"][bad_idx] = ( - replacement_reward.copy() if isinstance(replacement_reward, list) else replacement_reward - ) - generator_output["loss_masks"][bad_idx] = generator_output["loss_masks"][replacement_idx].copy() - if generator_output["stop_reasons"]: - generator_output["stop_reasons"][bad_idx] = generator_output["stop_reasons"][replacement_idx] - - if generator_output["rollout_logprobs"]: - generator_output["rollout_logprobs"][bad_idx] = generator_output["rollout_logprobs"][replacement_idx] + source_indices[bad_idx] = replacement_idx + if bad_indices: + generator_output = slice_generator_output(generator_output, source_indices) # Update UIDs accordingly replaced_uids = uids.copy() @@ -631,22 +623,7 @@ def get_bad_sample_replacements(good_uids: List[str], bad_uids: List[str]) -> Li def filter_generator_output(output: GeneratorOutput, kept_indices: List[int]) -> GeneratorOutput: """Filter GeneratorOutput based on kept indices.""" - filtered = { - "prompt_token_ids": [output["prompt_token_ids"][i] for i in kept_indices], - "response_ids": [output["response_ids"][i] for i in kept_indices], - "rewards": [output["rewards"][i] for i in kept_indices], - "loss_masks": [output["loss_masks"][i] for i in kept_indices], - "stop_reasons": None, - "rollout_metrics": output.get("rollout_metrics"), - "rollout_logprobs": ( - [output["rollout_logprobs"][i] for i in kept_indices] if output["rollout_logprobs"] else None - ), - } - - if output.get("stop_reasons"): - filtered["stop_reasons"] = [output["stop_reasons"][i] for i in kept_indices] - - return filtered + return slice_generator_output(output, kept_indices) def zero_variance_filter( diff --git a/skyrl/utils/routed_experts.py b/skyrl/utils/routed_experts.py new file mode 100644 index 0000000000..9275c330e6 --- /dev/null +++ b/skyrl/utils/routed_experts.py @@ -0,0 +1,14 @@ +import torch + + +def make_replay_padding_indices( + shape: tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device | str | int | None = None, +) -> torch.Tensor: + """Return dummy routes with ``topk`` distinct experts in every row.""" + if not shape or shape[-1] < 1: + raise ValueError(f"Replay route padding requires a positive topk dimension, got {shape}") + padding_row = torch.arange(shape[-1], dtype=dtype, device=device) + return padding_row.expand(shape).clone() diff --git a/skyrl/utils/token_metadata.py b/skyrl/utils/token_metadata.py new file mode 100644 index 0000000000..91d70a9e0a --- /dev/null +++ b/skyrl/utils/token_metadata.py @@ -0,0 +1,204 @@ +"""Token-aligned metadata layout transforms shared by training features.""" + +from dataclasses import dataclass + +import torch + +from skyrl.backends.skyrl_train.distributed.megatron.packing_utils import ( + get_packed_seq_align_size, + get_unpacked_seq_align_size, +) + + +def _new_metadata_tensor( + source: torch.Tensor, + shape: tuple[int, ...], + padding_value: torch.Tensor | bool | int, +) -> torch.Tensor: + output = torch.empty(shape, dtype=source.dtype, device=source.device) + output[...] = padding_value + return output + + +@dataclass(frozen=True) +class TokenMetadataLayout: + """One shared description of Megatron's token padding and CP sharding.""" + + attention_mask: torch.Tensor + sequence_lengths: list[int] + aligned_sequence_length: int + padded_sequence_lengths: list[int] | None = None + # Retained to reconstruct CP-sharded packed outputs in canonical batch order. + cu_seqlens_padded: torch.Tensor | None = None + context_parallel_size: int = 1 + context_parallel_rank: int = 0 + + +def build_token_metadata_layout( + attention_mask: torch.Tensor, + device: torch.device, + *, + packed: bool, + fp8_enabled: bool, +) -> TokenMetadataLayout: + """Compute the shared layout once for all replayed token metadata.""" + import megatron.core.parallel_state as mpu + + aligned_attention_mask = attention_mask.to(device=device, dtype=torch.bool) + sequence_lengths_tensor = aligned_attention_mask.sum(dim=1, dtype=torch.int32) + sequence_lengths = sequence_lengths_tensor.tolist() + tp_size = mpu.get_tensor_model_parallel_world_size() + + if not packed: + align_size = get_unpacked_seq_align_size(tp_size, fp8_enabled=fp8_enabled) + max_sequence_length = max(sequence_lengths) + aligned_sequence_length = max_sequence_length + (-max_sequence_length % align_size) + return TokenMetadataLayout( + attention_mask=aligned_attention_mask, + sequence_lengths=sequence_lengths, + aligned_sequence_length=aligned_sequence_length, + ) + + cp_size = mpu.get_context_parallel_world_size() + align_size = get_packed_seq_align_size(tp_size, cp_size, fp8_enabled=fp8_enabled) + padded_sequence_lengths_tensor = sequence_lengths_tensor + (-sequence_lengths_tensor % align_size) + padded_sequence_lengths = padded_sequence_lengths_tensor.tolist() + cu_seqlens_padded = torch.cat( + ( + torch.zeros(1, dtype=torch.int32, device=device), + padded_sequence_lengths_tensor.cumsum(dim=0), + ) + ) + return TokenMetadataLayout( + attention_mask=aligned_attention_mask, + sequence_lengths=sequence_lengths, + aligned_sequence_length=sum(padded_sequence_lengths), + padded_sequence_lengths=padded_sequence_lengths, + cu_seqlens_padded=cu_seqlens_padded, + context_parallel_size=cp_size, + context_parallel_rank=mpu.get_context_parallel_rank() if cp_size > 1 else 0, + ) + + +def align_token_metadata( + metadata: torch.Tensor, + layout: TokenMetadataLayout, + padding_value: torch.Tensor | bool | int, + *, + next_token: bool = False, +) -> torch.Tensor: + """Apply padding, optional next-token shifting, and CP sharding.""" + if metadata.device != layout.attention_mask.device: + raise ValueError("Token-aligned metadata and attention_mask must be on the same device") + if metadata.shape[:2] != layout.attention_mask.shape: + raise ValueError( + f"Token-aligned metadata shape {metadata.shape[:2]} does not match " + f"attention_mask shape {layout.attention_mask.shape}" + ) + + if layout.padded_sequence_lengths is None: + if next_token: + raise ValueError("next-token metadata alignment is only used for packed sequences") + aligned = _new_metadata_tensor( + metadata, + (metadata.shape[0], layout.aligned_sequence_length, *metadata.shape[2:]), + padding_value, + ) + for row_index, sequence_length in enumerate(layout.sequence_lengths): + aligned[row_index, :sequence_length] = metadata[row_index, layout.attention_mask[row_index]] + return aligned + + packed = _new_metadata_tensor( + metadata, + (layout.aligned_sequence_length, *metadata.shape[2:]), + padding_value, + ) + offset = 0 + for row_index, (sequence_length, padded_length) in enumerate( + zip(layout.sequence_lengths, layout.padded_sequence_lengths, strict=True) + ): + packed[offset : offset + sequence_length] = metadata[row_index, layout.attention_mask[row_index]] + # Match Megatron's [seq0, pad0, seq1, pad1, ...] microbatch layout. + offset += padded_length + + if next_token: + # Each packed logit predicts the next token within its own padded sequence. + shifted = _new_metadata_tensor(metadata, packed.shape, padding_value) + offset = 0 + for padded_length in layout.padded_sequence_lengths: + shifted[offset : offset + padded_length - 1] = packed[offset + 1 : offset + padded_length] + offset += padded_length + packed = shifted + + if layout.context_parallel_size > 1: + out = _new_metadata_tensor( + metadata, + (packed.shape[0] // layout.context_parallel_size, *packed.shape[1:]), + padding_value, + ) + src_offset = 0 + dst_offset = 0 + for padded_length in layout.padded_sequence_lengths: + # CP uses matching front/back chunks of each padded sequence. + length_per_cp = padded_length // layout.context_parallel_size + half = length_per_cp // 2 + front_start = src_offset + half * layout.context_parallel_rank + back_start = src_offset + padded_length - half * (layout.context_parallel_rank + 1) + out[dst_offset : dst_offset + half] = packed[front_start : front_start + half] + out[dst_offset + half : dst_offset + length_per_cp] = packed[back_start : back_start + half] + src_offset += padded_length + dst_offset += length_per_cp + packed = out + + return packed.unsqueeze(0) + + +def scatter_packed_token_values_to_batch( + model_values: torch.Tensor, + layout: TokenMetadataLayout, + padding_value: bool | int, +) -> torch.Tensor: + """Scatter packed model outputs into canonical ``[batch, seq_len - 1]`` positions.""" + if layout.padded_sequence_lengths is None or layout.cu_seqlens_padded is None: + raise ValueError("Scattering packed token values requires a packed metadata layout") + if model_values.ndim != 2 or model_values.shape[0] != 1: + raise ValueError(f"Expected packed model values with shape [1, tokens], got {model_values.shape}") + + values = model_values.squeeze(0) + if layout.context_parallel_size > 1: + import megatron.core.parallel_state as mpu + + from skyrl.backends.skyrl_train.distributed.megatron.model_utils import ( + allgather_cp_sharded_packed_tensor, + ) + + values = allgather_cp_sharded_packed_tensor( + values, + layout.cu_seqlens_padded, + mpu.get_context_parallel_group(), + ) + + from skyrl.backends.skyrl_train.distributed.megatron.model_utils import ( + _packed_sequence_indices, + ) + + _, _, sequence_indices, sequence_offsets, _ = _packed_sequence_indices( + layout.cu_seqlens_padded, + values.shape[0], + values.device, + ) + valid_counts = torch.tensor(layout.sequence_lengths, dtype=torch.long, device=values.device) - 1 + packed_mask = sequence_offsets < valid_counts[sequence_indices] + + attention_mask = layout.attention_mask + token_ordinals = attention_mask.to(torch.long).cumsum(dim=1) + output_mask = attention_mask[:, :-1] & ( + token_ordinals[:, :-1] < torch.tensor(layout.sequence_lengths, device=values.device).unsqueeze(1) + ) + batch_values = _new_metadata_tensor( + model_values, + (attention_mask.shape[0], attention_mask.shape[1] - 1), + padding_value, + ) + batch_values[output_mask] = values[packed_mask] + return batch_values diff --git a/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py b/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py index 2e5db6b265..dbb53642a9 100644 --- a/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py +++ b/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py @@ -17,7 +17,10 @@ ) from skyrl.backends.skyrl_train.training_batch import TrainingInputBatch from skyrl.train.config import SamplingParams, SkyRLTrainConfig -from skyrl.train.dataset.preprocess import convert_prompts_responses_to_batch_tensors +from skyrl.train.dataset.preprocess import ( + convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, +) from skyrl.train.generators.base import GeneratorInput from skyrl.train.generators.skyrl_gym_generator import SkyRLGymGenerator from skyrl.train.utils.utils import validate_cfg @@ -201,6 +204,7 @@ async def test_logprobs(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs): ) assert rii_tensor is not None + router_padding_mask = make_router_padding_mask(attention_mask, [len(sample) for sample in indices]) num_actions = response_mask.shape[1] batch_size = sequences.shape[0] training_input = TrainingInputBatch( @@ -216,6 +220,7 @@ async def test_logprobs(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs): else torch.zeros((batch_size, num_actions), dtype=torch.float32) ), "rollout_expert_indices": rii_tensor, + "router_padding_mask": router_padding_mask, "action_log_probs": torch.zeros((batch_size, num_actions), dtype=torch.float32), "base_action_log_probs": torch.zeros((batch_size, num_actions), dtype=torch.float32), "advantages": torch.zeros((batch_size, num_actions), dtype=torch.float32), @@ -333,10 +338,15 @@ def test_forward_backward(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs MOONLIGHT_NUM_LAYERS = 27 MOONLIGHT_TOPK = 6 MOONLIGHT_NUM_EXPERTS = 64 - rollout_expert_indices = torch.randint( - 0, MOONLIGHT_NUM_EXPERTS, (batch_size, seq_len, MOONLIGHT_NUM_LAYERS, MOONLIGHT_TOPK), dtype=torch.int32 + route_start = torch.randint( + 0, + MOONLIGHT_NUM_EXPERTS, + (batch_size, seq_len, MOONLIGHT_NUM_LAYERS, 1), + dtype=torch.int32, ) - rollout_expert_indices[attention_mask == 0] = 0 + route_offsets = torch.arange(MOONLIGHT_TOPK, dtype=torch.int32) + rollout_expert_indices = (route_start + route_offsets) % MOONLIGHT_NUM_EXPERTS + rollout_expert_indices[attention_mask == 0] = route_offsets gen = torch.Generator().manual_seed(42) training_input = TrainingInputBatch( @@ -348,6 +358,7 @@ def test_forward_backward(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs "loss_mask": loss_mask_t, "rollout_logprobs": -torch.rand((batch_size, num_actions), generator=gen) * 2.0, "rollout_expert_indices": rollout_expert_indices, + "router_padding_mask": ~attention_mask.bool(), "action_log_probs": -torch.rand((batch_size, num_actions), generator=gen) * 2.0, "base_action_log_probs": -torch.rand((batch_size, num_actions), generator=gen) * 2.0, "advantages": torch.randn((batch_size, num_actions), generator=gen), diff --git a/tests/backends/skyrl_train/test_token_based_batching_utils.py b/tests/backends/skyrl_train/test_token_based_batching_utils.py index 26a2b3c561..4fac8a6679 100644 --- a/tests/backends/skyrl_train/test_token_based_batching_utils.py +++ b/tests/backends/skyrl_train/test_token_based_batching_utils.py @@ -199,6 +199,18 @@ def test_padding_microbatch_matches_seq_len(self): # Padding rows must not contribute to the loss. assert padding["loss_mask"].sum().item() == 0 + def test_padding_microbatch_uses_unique_dummy_routes(self): + batch = self._make_batch([4, 4], num_actions=2) + batch["rollout_expert_indices"] = torch.full((2, 4, 2, 3), 7, dtype=torch.int16) + batch["router_padding_mask"] = torch.zeros((2, 4), dtype=torch.bool) + iterator = TokenBasedBatchIterator(batch, max_tokens_per_microbatch=8) + + padding = iterator._create_padding_microbatch() + + expected = torch.tensor([0, 1, 2], dtype=torch.int16).expand_as(padding["rollout_expert_indices"]) + assert torch.equal(padding["rollout_expert_indices"], expected) + assert torch.all(padding["router_padding_mask"]) + def test_multimodal_tensorlist_microbatching(self): """Token-based microbatching must gather TensorList fields (multi-modal pixel_values / image_grid_thw) via the same index gather used for regular tensors.""" diff --git a/tests/backends/skyrl_train/test_train_batch.py b/tests/backends/skyrl_train/test_train_batch.py index b105750957..e562f5ed48 100644 --- a/tests/backends/skyrl_train/test_train_batch.py +++ b/tests/backends/skyrl_train/test_train_batch.py @@ -551,6 +551,7 @@ def test_tensor_batch_none_tensor_list(): "rewards", "rollout_logprobs", "rollout_expert_indices", + "router_padding_mask", "pixel_values", "image_grid_thw", } @@ -576,6 +577,7 @@ def _make_full_training_batch(batch_size: int = 4, seq_len: int = 5) -> Training "rewards": torch.randn(batch_size, seq_len), "rollout_logprobs": torch.randn(batch_size, seq_len), "rollout_expert_indices": torch.randint(0, 8, (batch_size, seq_len, 2, 3), dtype=torch.long), + "router_padding_mask": torch.zeros((batch_size, seq_len), dtype=torch.bool), "pixel_values": TensorList([torch.randn(i + 1, 3) for i in range(batch_size)]), # batch_size * (i + 1) * 3 "image_grid_thw": TensorList([torch.tensor([[1, 2, 3]]) for _ in range(batch_size)]), # batch_size * 1 * 3 } @@ -636,7 +638,19 @@ def test_pad_batch_all_fields(): # Regular tensor fields (not loss_mask, not TensorList): original rows untouched, # padding rows are copies of row 0. - regular_tensor_keys = EXPECTED_TRAINING_INPUT_FIELDS - {"loss_mask", "pixel_values", "image_grid_thw"} + assert torch.equal(padded["router_padding_mask"][:batch_size], batch["router_padding_mask"]) + assert torch.all(padded["router_padding_mask"][batch_size:]) + assert torch.equal(padded["rollout_expert_indices"][:batch_size], batch["rollout_expert_indices"]) + expected_routes = torch.tensor([0, 1, 2]).expand_as(padded["rollout_expert_indices"][batch_size:]) + assert torch.equal(padded["rollout_expert_indices"][batch_size:], expected_routes) + + regular_tensor_keys = EXPECTED_TRAINING_INPUT_FIELDS - { + "loss_mask", + "rollout_expert_indices", + "router_padding_mask", + "pixel_values", + "image_grid_thw", + } for key in regular_tensor_keys: assert torch.equal(padded[key][:batch_size], batch[key]), f"Original rows changed for {key!r}" for i in range(batch_size, batch_size + pad_size): diff --git a/tests/backends/skyrl_train/utils/test_replay_utils.py b/tests/backends/skyrl_train/utils/test_replay_utils.py new file mode 100644 index 0000000000..5f9be495a7 --- /dev/null +++ b/tests/backends/skyrl_train/utils/test_replay_utils.py @@ -0,0 +1,206 @@ +import inspect +import sys +import types +from types import SimpleNamespace + +import pytest +import torch + +from skyrl.backends.skyrl_train.utils import replay_utils +from skyrl.utils.routed_experts import make_replay_padding_indices +from skyrl.utils.token_metadata import build_token_metadata_layout + + +@pytest.fixture +def parallel_state(monkeypatch): + try: + import megatron.core.parallel_state as mpu + except ModuleNotFoundError: + megatron = types.ModuleType("megatron") + core = types.ModuleType("megatron.core") + mpu = types.ModuleType("megatron.core.parallel_state") + megatron.core = core + core.parallel_state = mpu + monkeypatch.setitem(sys.modules, "megatron", megatron) + monkeypatch.setitem(sys.modules, "megatron.core", core) + monkeypatch.setitem(sys.modules, "megatron.core.parallel_state", mpu) + + monkeypatch.setattr(mpu, "get_tensor_model_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_rank", lambda: 0, raising=False) + return mpu + + +def test_patch_topk_router_expert_bias_excludes_padding(monkeypatch): + router_module = types.ModuleType("megatron.core.transformer.moe.router") + + class TopKRouter: + def __init__(self): + self.local_tokens_per_expert = torch.zeros(3, dtype=torch.int64) + + def _apply_expert_bias(self, routing_map, padding_mask=None): + if padding_mask is not None: + routing_map = routing_map & (~padding_mask) + self.local_tokens_per_expert += routing_map.sum(dim=0) + + router_module.TopKRouter = TopKRouter + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router", router_module) + + replay_utils.patch_topk_router_expert_bias_padding_mask() + router = TopKRouter() + router._apply_expert_bias( + torch.tensor([[1, 0, 1], [0, 1, 1]], dtype=torch.bool), + torch.tensor([False, True]), + ) + + assert torch.equal(router.local_tokens_per_expert, torch.tensor([1, 0, 1])) + + +@pytest.mark.parametrize("dtype", [torch.uint8, torch.int16, torch.int32]) +def test_replay_padding_indices_are_unique(dtype): + padding = make_replay_padding_indices((2, 3, 4, 3), dtype=dtype) + + assert padding.shape == (2, 3, 4, 3) + assert torch.equal(padding, torch.tensor([0, 1, 2], dtype=dtype).expand_as(padding)) + + +def test_replay_has_no_dispatcher_specific_patch(): + assert "TokenDispatcher" not in inspect.getsource(replay_utils) + + +def test_setup_replay_installs_indices_and_returns_model_mask(monkeypatch, parallel_state): + router_replay_module = types.ModuleType("megatron.core.transformer.moe.router_replay") + + class RouterReplay: + global_router_replay_instances = [object()] + replay_data = None + action = None + + @classmethod + def set_replay_data(cls, replay_data): + cls.replay_data = replay_data + + @classmethod + def set_global_router_replay_action(cls, action): + cls.action = action + + class RouterReplayAction: + REPLAY_FORWARD = "replay_forward" + + router_replay_module.RouterReplay = RouterReplay + router_replay_module.RouterReplayAction = RouterReplayAction + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router_replay", router_replay_module) + monkeypatch.setattr(replay_utils, "_get_current_pp_stage_layer_range", lambda model_config: (0, 1)) + monkeypatch.setattr( + replay_utils, + "scatter_router_padding_mask_for_model", + lambda mask, model, model_config: mask, + ) + routes = torch.tensor([[[[0, 1]], [[1, 2]], [[3, 4]], [[5, 6]]]], dtype=torch.int16) + attention_mask = torch.tensor([[0, 1, 1, 1]]) + router_padding_mask = torch.tensor([[1, 0, 0, 1]], dtype=torch.bool) + metadata_layout = build_token_metadata_layout( + attention_mask, + routes.device, + packed=False, + fp8_enabled=False, + ) + + model_kwargs = replay_utils.setup_per_microbatch_replay_forward( + routes, + router_padding_mask, + attention_mask, + model=object(), + model_config=SimpleNamespace(fp8=None), + metadata_layout=metadata_layout, + ) + + assert RouterReplay.replay_data[0].tolist() == [[1, 2], [3, 4], [5, 6]] + assert RouterReplay.action == RouterReplayAction.REPLAY_FORWARD + assert model_kwargs["padding_mask"].tolist() == [[False, False, True]] + + +@pytest.mark.parametrize( + ("model_kind", "pre_process", "expected"), + [ + ("gpt", True, [[False, False, True, True]]), + ("gpt", False, [[True, True]]), + ("hybrid", True, [[True, True]]), + ], +) +def test_sequence_parallel_mask_layout(monkeypatch, model_kind, pre_process, expected): + hybrid_model = types.ModuleType("megatron.core.models.hybrid.hybrid_model") + tensor_parallel = types.ModuleType("megatron.core.tensor_parallel") + utils = types.ModuleType("megatron.core.utils") + + class HybridModel: + def __init__(self): + self.pre_process = pre_process + + class GPTModel: + def __init__(self): + self.pre_process = pre_process + + hybrid_model.HybridModel = HybridModel + tensor_parallel.scatter_to_sequence_parallel_region = lambda value: value.chunk(2, dim=0)[1] + utils.unwrap_model = lambda model: model + monkeypatch.setitem(sys.modules, "megatron.core.models.hybrid.hybrid_model", hybrid_model) + monkeypatch.setitem(sys.modules, "megatron.core.tensor_parallel", tensor_parallel) + monkeypatch.setitem(sys.modules, "megatron.core.utils", utils) + + mask = torch.tensor([[0, 0, 1, 1]], dtype=torch.bool) + model = HybridModel() if model_kind == "hybrid" else GPTModel() + scattered = replay_utils.scatter_router_padding_mask_for_model( + mask, + model, + SimpleNamespace(sequence_parallel=True), + ) + + assert scattered.tolist() == expected + + +@pytest.fixture +def router_replay_module(monkeypatch): + module = types.ModuleType("megatron.core.transformer.moe.router_replay") + router = SimpleNamespace(replay_backward_list=[], action=None) + + class RouterReplay: + global_router_replay_instances = [router] + + @classmethod + def clear_global_indices(cls): + for instance in cls.global_router_replay_instances: + instance.replay_backward_list = [] + + @classmethod + def clear_global_router_replay_action(cls): + for instance in cls.global_router_replay_instances: + instance.action = None + + module.RouterReplay = RouterReplay + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router_replay", module) + return router + + +def test_router_replay_schedule_clears_stale_forward_only_fifo(router_replay_module): + router_replay_module.replay_backward_list = ["stale-forward-only"] + + with replay_utils.router_replay_schedule(enabled=True): + assert router_replay_module.replay_backward_list == [] + router_replay_module.replay_backward_list.extend(["microbatch-0", "microbatch-1"]) + assert router_replay_module.replay_backward_list.pop(0) == "microbatch-0" + assert router_replay_module.replay_backward_list.pop(0) == "microbatch-1" + + assert router_replay_module.replay_backward_list == [] + assert router_replay_module.action is None + + +def test_router_replay_schedule_clears_after_exception(router_replay_module): + with pytest.raises(RuntimeError, match="schedule failed"): + with replay_utils.router_replay_schedule(enabled=True): + router_replay_module.replay_backward_list.append("partially-consumed-schedule") + router_replay_module.action = "replay-backward" + raise RuntimeError("schedule failed") + + assert router_replay_module.replay_backward_list == [] + assert router_replay_module.action is None diff --git a/tests/train/dataset/test_preprocess.py b/tests/train/dataset/test_preprocess.py index 1df1d002a1..61d5ccb757 100644 --- a/tests/train/dataset/test_preprocess.py +++ b/tests/train/dataset/test_preprocess.py @@ -9,6 +9,7 @@ from skyrl.train.dataset.preprocess import ( convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, ) @@ -56,6 +57,40 @@ def fake_tokenizer_decode_list(ids, **kwargs): return mock_tokenizer +def test_router_padding_mask_marks_left_padding_and_uncaptured_suffix(): + attention_mask = torch.tensor([[0, 1, 1, 1], [1, 1, 1, 1]]) + + mask = make_router_padding_mask(attention_mask, [2, 4]) + + assert mask.tolist() == [[True, False, False, True], [False, False, False, False]] + + +def test_routed_expert_tensor_uses_unique_dummy_routes(tokenizer): + routes = [ + [ + [[2, 3], [4, 5]], + [[6, 7], [0, 1]], + ], + [ + [[1, 2], [3, 4]], + [[5, 6], [7, 0]], + [[2, 4], [6, 7]], + ], + ] + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10], [20]], + responses=[[11, 12], [21, 22]], + rewards=[[0.0, 0.0], [0.0, 0.0]], + loss_masks=[[1, 1], [1, 1]], + rollout_expert_indices=routes, + ) + + assert routed.shape == (2, 3, 2, 2) + assert routed[0, 2].tolist() == [[0, 1], [0, 1]] + + def test_convert_prompts_responses_to_batch_tensors_exact(tokenizer): """ Test with inputs of exact lengths. @@ -306,20 +341,21 @@ def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): # Shape: [batch=2, max_total=6, layers=2, topk=2] assert rei_tensor.shape == (2, 6, num_layers, topk) - # Sample 0 has total=5, so 1 left-pad position → first position should be zeros - assert rei_tensor[0, 0].tolist() == [[0, 0]] * num_layers # padding + dummy_routes = [[0, 1]] * num_layers + # Sample 0 has total=5, so the first position uses unique dummy routes. + assert rei_tensor[0, 0].tolist() == dummy_routes assert rei_tensor[0, 1].tolist() == [[1, 2]] * num_layers # first real token # Sample 1 has total=6, no padding assert rei_tensor[1, 0].tolist() == [[3, 4]] * num_layers # first real token - # Non-zero positions in rei_tensor align exactly with attention_mask==1 + # Dummy positions in rei_tensor align exactly with attention_mask==0. for i in range(2): for pos in range(6): if attn[i, pos] == 0: - assert rei_tensor[i, pos].tolist() == [[0, 0]] * num_layers + assert rei_tensor[i, pos].tolist() == dummy_routes else: - assert rei_tensor[i, pos].tolist() != [[0, 0]] * num_layers + assert rei_tensor[i, pos].tolist() != dummy_routes def test_rollout_expert_indices_none_when_not_provided(tokenizer): diff --git a/tests/train/generators/test_skyrl_gym_generator.py b/tests/train/generators/test_skyrl_gym_generator.py index c29afbe920..f0c7ffdcfd 100644 --- a/tests/train/generators/test_skyrl_gym_generator.py +++ b/tests/train/generators/test_skyrl_gym_generator.py @@ -14,7 +14,7 @@ GeneratorInput, GeneratorOutput, ) -from skyrl.train.generators.skyrl_gym_generator import SkyRLGymGenerator +from skyrl.train.generators.skyrl_gym_generator import SkyRLGymGenerator, TurnOutput from skyrl_gym.envs.base_text_env import BaseTextEnv, BaseTextEnvStepOutput # Mock constants, where 4 is the eos token id @@ -22,6 +22,23 @@ MOCK_TOKENIZER_ENCODED_IDS = [1, 2, 3, 4] +def test_turn_output_keeps_uncaptured_suffix_out_of_routes(): + routes = [[[2, 3]], [[4, 5]]] + output = TurnOutput( + output="answer", + output_ids=[10, 11, 4], + output_logprobs=None, + new_obs=[], + obs_ids=[20, 21], + rollout_expert_indices=routes, + reward=1.0, + added_eos=True, + ) + + assert output.get_turn_rollout_expert_indices() is routes + assert output.get_turn_loss_mask() == [1, 1, 0, 0, 0] + + # TODO (erictang000): clean up the mocking for tests in this file @pytest.fixture def mock_tokenizer(): @@ -376,7 +393,7 @@ def mock_generate(_, model=None): # No EOS: just add it expected_response_ids = mock_llm_output_ids + [mock_tokenizer.eos_token_id] - expected_loss_mask = [1] * (len(expected_response_ids)) + expected_loss_mask = [1] * len(expected_response_ids) if logprobs_setting is not None: assert output.rollout_logprobs is not None diff --git a/tests/train/test_trainer_utils.py b/tests/train/test_trainer_utils.py index 6a70e8ca66..919a068141 100644 --- a/tests/train/test_trainer_utils.py +++ b/tests/train/test_trainer_utils.py @@ -393,6 +393,7 @@ def test_handle_replace_sampling_sufficient_good_samples(): "stop_reasons": ["length"] * 6, "rollout_metrics": None, "rollout_logprobs": [[0.1, 0.2], [0.3, 0.4], [0.5, 0.25], [0.15, 0.25], [0.1, 0.2], [0.3, 0.4]], + "rollout_expert_indices": [[[[i, i + 1]]] for i in range(6)], } uids = ["uid1", "uid1", "uid2", "uid2", "uid3", "uid3"] # 2 samples per prompt sampling_config = {"n_samples_per_prompt": 2, "min_replace_ratio": 0.3} @@ -408,6 +409,12 @@ def test_handle_replace_sampling_sufficient_good_samples(): assert len(result_output["rewards"]) == 6 assert len(result_output["rollout_logprobs"]) == 6 assert len(result_uids) == 6 + route_by_response = { + tuple(response): routes + for response, routes in zip(generator_output["response_ids"], generator_output["rollout_expert_indices"]) + } + for response, routes in zip(result_output["response_ids"], result_output["rollout_expert_indices"]): + assert routes == route_by_response[tuple(response)] # Check that bad uid2 samples were replaced with good samples uid2_indices = [i for i, uid in enumerate(result_uids) if uid == "uid2"] @@ -644,6 +651,7 @@ def test_filter_generator_output(): "stop_reasons": ["length", "length", "stop"], "rollout_metrics": {"metric": "value"}, "rollout_logprobs": [[0.16, 0.4], [0.1, 0.2], [0.3, 0.4]], + "rollout_expert_indices": ["routes-0", "routes-1", "routes-2"], } kept_indices = [0, 2] # Keep first and third samples @@ -656,6 +664,7 @@ def test_filter_generator_output(): assert filtered["stop_reasons"] == ["length", "stop"] assert filtered["rollout_metrics"] == {"metric": "value"} assert filtered["rollout_logprobs"] == [[0.16, 0.4], [0.3, 0.4]] + assert filtered["rollout_expert_indices"] == ["routes-0", "routes-2"] def test_zero_variance_filter_mixed_groups(): diff --git a/tests/utils/test_token_metadata.py b/tests/utils/test_token_metadata.py new file mode 100644 index 0000000000..901533a6dd --- /dev/null +++ b/tests/utils/test_token_metadata.py @@ -0,0 +1,88 @@ +import sys +import types + +import pytest +import torch + +from skyrl.utils import token_metadata + + +@pytest.fixture +def parallel_state(monkeypatch): + try: + import megatron.core.parallel_state as mpu + except ModuleNotFoundError: + megatron = types.ModuleType("megatron") + core = types.ModuleType("megatron.core") + mpu = types.ModuleType("megatron.core.parallel_state") + megatron.core = core + core.parallel_state = mpu + monkeypatch.setitem(sys.modules, "megatron", megatron) + monkeypatch.setitem(sys.modules, "megatron.core", core) + monkeypatch.setitem(sys.modules, "megatron.core.parallel_state", mpu) + + monkeypatch.setattr(mpu, "get_tensor_model_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_rank", lambda: 0, raising=False) + return mpu + + +def test_microbatch_rows_share_one_packed_layout(monkeypatch, parallel_state): + monkeypatch.setattr(token_metadata, "get_packed_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1], [0, 0, 1, 1]]) + routes = torch.tensor( + [ + [[[0, 1]], [[10, 11]], [[12, 13]], [[14, 15]]], + [[[0, 1]], [[0, 1]], [[20, 21]], [[22, 23]]], + ], + dtype=torch.int16, + ) + router_mask = torch.tensor([[1, 0, 0, 1], [1, 1, 0, 0]], dtype=torch.bool) + + layout = token_metadata.build_token_metadata_layout( + attention_mask, + routes.device, + packed=True, + fp8_enabled=False, + ) + packed_routes = token_metadata.align_token_metadata( + routes, + layout, + torch.tensor([0, 1], dtype=routes.dtype), + ) + packed_mask = token_metadata.align_token_metadata(router_mask, layout, True) + + assert packed_routes[0, :, 0].tolist() == [ + [10, 11], + [12, 13], + [14, 15], + [0, 1], + [20, 21], + [22, 23], + [0, 1], + [0, 1], + ] + assert packed_mask.tolist() == [[False, False, True, True, False, False, True, True]] + assert layout.cu_seqlens_padded.tolist() == [0, 4, 8] + + +def test_packed_layout_aligns_next_token_metadata_and_scatters_rows(monkeypatch, parallel_state): + monkeypatch.setattr(token_metadata, "get_packed_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1], [0, 0, 1, 1]]) + metadata = torch.tensor([[0, 10, 11, 12], [0, 0, 20, 21]], dtype=torch.int32) + layout = token_metadata.build_token_metadata_layout( + attention_mask, + metadata.device, + packed=True, + fp8_enabled=False, + ) + + aligned = token_metadata.align_token_metadata(metadata, layout, -1, next_token=True) + batch_values = token_metadata.scatter_packed_token_values_to_batch( + torch.arange(1, 9, dtype=torch.float32).unsqueeze(0), + layout, + 0, + ) + + assert aligned.tolist() == [[11, 12, -1, -1, 21, -1, -1, -1]] + assert batch_values.tolist() == [[0.0, 1.0, 2.0], [0.0, 0.0, 5.0]] From fb86814d45a57b193ffce64b8c457693944986f0 Mon Sep 17 00:00:00 2001 From: dyurk-lila Date: Thu, 16 Jul 2026 21:38:03 +0000 Subject: [PATCH 2/2] perf(r3): accelerate routed-expert transport with packed arrays Store routed-expert (R3) generation data as compact NumPy arrays instead of large nested Python lists, and send it over the network base64-encoded alongside its shape and dtype. Expert IDs are compacted to the smallest safe uint8/int16/int32 dtype, vLLM responses and client responses use orjson, and preprocessing accepts the decoded NumPy route arrays directly. Co-Authored-By: Claude Opus 4.8 (1M context) --- pyproject.toml | 2 + .../skyrl_train/inference_servers/base.py | 4 +- .../remote_inference_client.py | 22 ++-- .../inference_servers/routed_experts_wire.py | 45 ++++++++ .../inference_servers/vllm_server_actor.py | 21 ++-- skyrl/train/dataset/preprocess.py | 62 +++++++---- skyrl/train/generators/base.py | 3 +- skyrl/train/generators/skyrl_gym_generator.py | 16 ++- skyrl/train/trainer.py | 4 +- skyrl/utils/routed_experts.py | 34 ++++++ .../test_remote_inference_client.py | 48 +++++++++ .../test_routed_experts_wire.py | 92 ++++++++++++++++ tests/train/dataset/test_preprocess.py | 102 ++++++++++++++++-- .../generators/test_generator_output_utils.py | 2 +- .../generators/test_skyrl_gym_generator.py | 3 +- tests/train/test_trainer_utils.py | 11 +- uv.lock | 10 ++ 17 files changed, 414 insertions(+), 67 deletions(-) create mode 100644 skyrl/backends/skyrl_train/inference_servers/routed_experts_wire.py create mode 100644 tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py diff --git a/pyproject.toml b/pyproject.toml index 331aaf1b02..edf058fa01 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -92,6 +92,8 @@ skyrl-train = [ "polars", "s3fs", "fastapi", + "orjson>=3.11.9", + "pybase64>=1.4.2", "uvicorn", "vllm-router; sys_platform == 'linux'", "pybind11", diff --git a/skyrl/backends/skyrl_train/inference_servers/base.py b/skyrl/backends/skyrl_train/inference_servers/base.py index e0152df5d4..26f87f54a2 100644 --- a/skyrl/backends/skyrl_train/inference_servers/base.py +++ b/skyrl/backends/skyrl_train/inference_servers/base.py @@ -1,6 +1,8 @@ from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Dict, Hashable, List, Optional, Tuple, TypedDict +from skyrl.utils.routed_experts import RoutedExpertIndices + if TYPE_CHECKING: from skyrl.backends.skyrl_train.weight_sync import WeightUpdateRequest from skyrl.backends.skyrl_train.weight_sync.transfer_strategy import ( @@ -47,7 +49,7 @@ class InferenceEngineOutput(TypedDict): stop_reasons: List[str] response_logprobs: Optional[List[List[float]]] prompt_logprobs: Optional[List[List[float]]] # per-prompt-token logprobs under the current model - rollout_expert_indices: Optional[List[List[List[int]]]] # [seq_len, layer_num, topk] + rollout_expert_indices: Optional[List[RoutedExpertIndices]] class InferenceEngineInterface(ABC): diff --git a/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py b/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py index b2d1e4e812..4612396afd 100644 --- a/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py +++ b/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py @@ -64,6 +64,7 @@ ) import aiohttp +import orjson from skyrl.backends.skyrl_train.inference_servers.base import ( InferenceEngineInput, @@ -72,6 +73,9 @@ MMPlaceholderRangeInfo, MultiModalFeatures, ) +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + decode_packed_routed_experts, +) from skyrl.env_vars import ( SKYRL_GENERATE_CONCURRENCY_PER_ENGINE, SKYRL_HTTP_CONNECTION_LIMIT, @@ -312,8 +316,8 @@ async def _post(self, url: str, json: Dict[str, Any], headers: Optional[Dict[str try: async with session.post(url, json=json, headers=headers) as resp: try: - body = await resp.json(content_type=None) - except Exception as e: + body = orjson.loads(await resp.read()) + except orjson.JSONDecodeError as e: if 400 <= resp.status < 500: # Non-JSON client error (e.g. plain text 422 from vllm-router). # Raise immediately — client errors won't succeed on retry. @@ -446,15 +450,16 @@ async def _throttled_detokenize(token_ids: List[int]) -> str: raw_results = await asyncio.gather(*[_throttled_generate(idx) for idx in range(batch_size)]) responses = await asyncio.gather(*[_throttled_detokenize(r["response_ids"]) for r in raw_results]) - rollout_expert_indices = [r.get("routed_experts") for r in raw_results] - has_routed_experts = any(x is not None for x in rollout_expert_indices) + rollout_expert_indices = ( + [result["routed_experts"] for result in raw_results] if self.enable_return_routed_experts else None + ) return InferenceEngineOutput( responses=responses, stop_reasons=[r["stop_reason"] for r in raw_results], response_ids=[r["response_ids"] for r in raw_results], response_logprobs=[r["response_logprobs"] for r in raw_results] if get_logprobs else None, - rollout_expert_indices=rollout_expert_indices if has_routed_experts else None, + rollout_expert_indices=rollout_expert_indices, ) async def _generate_single( @@ -511,7 +516,12 @@ async def _generate_single( if logprobs_content: response_logprobs = [logprob_info["logprob"] for logprob_info in logprobs_content] - routed_experts = choice.get("routed_experts") + routed_experts = None + if self.enable_return_routed_experts: + packed_routed_experts = choice.get("routed_experts") + if not isinstance(packed_routed_experts, dict): + raise ValueError("/skyrl/v1/generate must return packed routed_experts") + routed_experts = decode_packed_routed_experts(packed_routed_experts) return { "stop_reason": stop_reason, diff --git a/skyrl/backends/skyrl_train/inference_servers/routed_experts_wire.py b/skyrl/backends/skyrl_train/inference_servers/routed_experts_wire.py new file mode 100644 index 0000000000..c4b9911c39 --- /dev/null +++ b/skyrl/backends/skyrl_train/inference_servers/routed_experts_wire.py @@ -0,0 +1,45 @@ +"""Compact routed-expert HTTP payloads.""" + +import math +from typing import Any + +import numpy as np +import pybase64 + +from skyrl.utils.routed_experts import ( + ROUTED_EXPERT_DTYPES, + RoutedExpertIndices, + compact_routed_expert_indices, +) + +_DTYPES = {dtype.name: dtype for dtype in ROUTED_EXPERT_DTYPES} + + +def pack_routed_experts(routed_experts: RoutedExpertIndices) -> dict[str, Any]: + compact = compact_routed_expert_indices(routed_experts) + return { + "data": pybase64.b64encode(memoryview(compact)).decode("ascii"), + "shape": list(compact.shape), + "dtype": compact.dtype.name, + } + + +def decode_packed_routed_experts(payload: dict[str, Any]) -> RoutedExpertIndices: + if not isinstance(payload, dict): + raise TypeError("packed routed expert indices must be an object") + try: + dtype = _DTYPES[payload["dtype"]] + shape = tuple(payload["shape"]) + data = pybase64.b64decode_as_bytearray(payload["data"], validate=True) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("invalid packed routed_experts payload") from exc + if len(shape) != 3 or any(type(dim) is not int or dim < 0 for dim in shape): + raise ValueError(f"invalid packed routed_experts shape: {shape}") + expected_size = math.prod(shape) * dtype.itemsize + if len(data) != expected_size: + raise ValueError(f"packed routed_experts has {len(data)} bytes, expected {expected_size}") + decoded = np.frombuffer(data, dtype=dtype).reshape(shape) + compact = compact_routed_expert_indices(decoded) + if compact.dtype != dtype: + raise ValueError(f"packed routed_experts uses non-canonical dtype {dtype.name}; expected {compact.dtype.name}") + return compact diff --git a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py index 35d7057012..aa50c0790a 100644 --- a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py +++ b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py @@ -4,15 +4,18 @@ import asyncio import logging +import math import os import time from argparse import Namespace from typing import List, Optional, Tuple import httpx +import numpy as np +import orjson import uvicorn import vllm.envs as envs -from fastapi import HTTPException, Request +from fastapi import HTTPException, Request, Response from ray.util.placement_group import PlacementGroup from vllm.engine.arg_utils import AsyncEngineArgs from vllm.engine.async_llm_engine import AsyncLLMEngine @@ -34,6 +37,9 @@ get_node_ip, ) from skyrl.backends.skyrl_train.inference_servers.protocols import ServerActorProtocol +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + pack_routed_experts, +) from skyrl.env_vars import ( SKYRL_HTTP_CONNECTION_LIMIT, SKYRL_VLLM_DP_PORT_OFFSET, @@ -424,7 +430,10 @@ async def _skyrl_generate(request: Request): content = [] for tid, lp_dict in zip(token_ids_out, resp.logprobs): if lp_dict and tid in lp_dict: - content.append({"logprob": lp_dict[tid].logprob}) + logprob = lp_dict[tid].logprob + if not math.isfinite(logprob): + raise ValueError("Out of range float values are not JSON compliant") + content.append({"logprob": logprob}) else: # -9999.0 is the default in vLLM's ChatCompletionLogProb content.append({"logprob": -9999.0}) @@ -432,12 +441,9 @@ async def _skyrl_generate(request: Request): routed_experts = None if resp.routed_experts is not None: - if hasattr(resp.routed_experts, "tolist"): - routed_experts = resp.routed_experts.tolist() - else: - routed_experts = resp.routed_experts + routed_experts = pack_routed_experts(np.asarray(resp.routed_experts)) - return { + payload = { "choices": [ { "token_ids": token_ids_out, @@ -447,6 +453,7 @@ async def _skyrl_generate(request: Request): } ] } + return Response(content=orjson.dumps(payload), media_type="application/json") async def shutdown(self) -> None: """Gracefully shutdown the server.""" diff --git a/skyrl/train/dataset/preprocess.py b/skyrl/train/dataset/preprocess.py index f5af10741b..90b987257d 100644 --- a/skyrl/train/dataset/preprocess.py +++ b/skyrl/train/dataset/preprocess.py @@ -1,11 +1,16 @@ import logging from typing import List, Optional, Tuple +import numpy as np import torch from jaxtyping import Bool, Float, Integer from transformers import AutoTokenizer -from skyrl.utils.routed_experts import make_replay_padding_indices +from skyrl.utils.routed_experts import ( + ROUTED_EXPERT_DTYPES, + RoutedExpertIndices, + compact_routed_expert_indices, +) logger = logging.getLogger(__name__) @@ -79,7 +84,7 @@ def convert_prompts_responses_to_batch_tensors( rewards: List[List[float]], loss_masks: List[List[int]], logprobs: Optional[List[List[float]]] = None, - rollout_expert_indices: Optional[List[List[List[List[int]]]]] = None, + rollout_expert_indices: Optional[List[RoutedExpertIndices]] = None, max_seq_len: Optional[int] = None, ) -> Tuple[ Float[torch.Tensor, "batch seq_len"], @@ -205,39 +210,50 @@ def convert_prompts_responses_to_batch_tensors( rollout_expert_indices_tensor = None if rollout_expert_indices is not None: num_samples = len(prompts) - if len(rollout_expert_indices) != num_samples or any(not indices for indices in rollout_expert_indices): + if not isinstance(rollout_expert_indices, list): + raise TypeError("rollout_expert_indices must be a list of NumPy arrays") + if len(rollout_expert_indices) != num_samples: raise ValueError("rollout_expert_indices must contain routes for every trajectory") - num_layers = len(rollout_expert_indices[0][0]) - topk = len(rollout_expert_indices[0][0][0]) if num_layers > 0 else 0 + canonical_indices = [] + for sample_index, sample_indices in enumerate(rollout_expert_indices): + if not isinstance(sample_indices, np.ndarray): + raise TypeError( + f"rollout_expert_indices entries must be NumPy arrays, got {type(sample_indices).__name__} " + f"at sample {sample_index}" + ) + if sample_indices.dtype not in ROUTED_EXPERT_DTYPES: + raise TypeError( + f"Unsupported routed expert dtype {sample_indices.dtype} at sample {sample_index}; " + "expected uint8, int16, or int32" + ) + canonical_indices.append(compact_routed_expert_indices(sample_indices)) + + first_shape = canonical_indices[0].shape + if len(first_shape) != 3 or first_shape[0] == 0: + raise ValueError("rollout_expert_indices must contain routes for every trajectory") + num_layers, topk = first_shape[1:] if topk < 1: raise ValueError("rollout_expert_indices must contain at least one expert per layer") - padded = make_replay_padding_indices( - (num_samples, max_total, num_layers, topk), - dtype=torch.int32, - ) - for sample_index, sample_indices in enumerate(rollout_expert_indices): - sample_indices_tensor = torch.as_tensor(sample_indices, dtype=torch.int32) - if sample_indices_tensor.ndim != 3 or sample_indices_tensor.shape[1:] != (num_layers, topk): + batch_dtype = max((indices.dtype for indices in canonical_indices), key=lambda dtype: dtype.itemsize) + padded = np.empty((num_samples, max_total, num_layers, topk), dtype=batch_dtype) + padded[...] = np.arange(topk, dtype=batch_dtype) + for sample_index, sample_indices in enumerate(canonical_indices): + if sample_indices.ndim != 3 or sample_indices.shape[1:] != (num_layers, topk): raise ValueError( "rollout_expert_indices entries must share [layers, topk], " - f"got shape {tuple(sample_indices_tensor.shape)} at sample {sample_index}" + f"got shape {sample_indices.shape} at sample {sample_index}" ) left_pad = max_total - (prompt_token_lens[sample_index] + response_token_lens[sample_index]) available = max_total - left_pad - if len(sample_indices) > available: + if sample_indices.shape[0] == 0 or sample_indices.shape[0] > available: raise ValueError( - f"Trajectory {sample_index} has {len(sample_indices)} route rows for {available} tokens" + f"Trajectory {sample_index} has {sample_indices.shape[0]} route rows for {available} tokens" ) - padded[sample_index, left_pad : left_pad + len(sample_indices)] = sample_indices_tensor - rollout_expert_indices_tensor = padded - - max_expert_id = int(rollout_expert_indices_tensor.max().item()) - if max_expert_id < 2**8: - rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.uint8) - elif max_expert_id < 2**15: - rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.int16) + route_end = left_pad + sample_indices.shape[0] + padded[sample_index, left_pad:route_end] = sample_indices + rollout_expert_indices_tensor = torch.from_numpy(padded) return ( sequences, diff --git a/skyrl/train/generators/base.py b/skyrl/train/generators/base.py index 26d95868b9..938d9f73a8 100644 --- a/skyrl/train/generators/base.py +++ b/skyrl/train/generators/base.py @@ -5,6 +5,7 @@ import torch from skyrl.backends.skyrl_train.inference_servers.base import ConversationType +from skyrl.utils.routed_experts import RoutedExpertIndices TrainingPhase = Literal["train", "eval"] @@ -46,7 +47,7 @@ class GeneratorOutput(TypedDict): # trajectory in the input batch (i.e. per ``agent_loop`` call). Used by the fully # async trainer to compute per-group / intra-group completion-time metrics. trajectory_generation_times: Optional[List[float]] - rollout_expert_indices: Optional[List[List[List[List[int]]]]] # [batch_size, seq_len, layer_num, topk] + rollout_expert_indices: Optional[List[RoutedExpertIndices]] # Applicable only for step-wise training is_last_step: Optional[List[bool]] # Per-row env metrics (one dict per row in the flattened batch). Used by diff --git a/skyrl/train/generators/skyrl_gym_generator.py b/skyrl/train/generators/skyrl_gym_generator.py index 71a3b81477..c0239225d7 100644 --- a/skyrl/train/generators/skyrl_gym_generator.py +++ b/skyrl/train/generators/skyrl_gym_generator.py @@ -36,6 +36,7 @@ get_generation_prompt_ids, get_rollout_metrics, ) +from skyrl.utils.routed_experts import RoutedExpertIndices from skyrl_gym.envs.base_text_env import BaseTextEnvStepOutput @@ -50,7 +51,7 @@ class TrajectoryOutput: prompt_ids: List[int] rollout_logprobs: Optional[List[float]] env_metrics: Dict[str, Any] - rollout_expert_indices: Optional[List[List[List[int]]]] = None + rollout_expert_indices: Optional[RoutedExpertIndices] = None pixel_values: Optional[torch.Tensor] = None image_grid_thw: Optional[torch.Tensor] = None # End-to-end wall-clock time (seconds) to generate this trajectory. Optional: agent loops may @@ -76,7 +77,7 @@ class AgentLoopState: rollout_logprobs: Optional[List[float]] response_end_idx: Optional[int] done: bool - rollout_expert_indices: Optional[List[List[List[int]]]] = None + rollout_expert_indices: Optional[RoutedExpertIndices] = None @dataclass @@ -86,11 +87,11 @@ class TurnOutput: output_logprobs: Optional[List[float]] new_obs: ConversationType obs_ids: List[int] - rollout_expert_indices: Optional[List[List[List[int]]]] # [seq_len, layer_num, topk] + rollout_expert_indices: Optional[RoutedExpertIndices] reward: Optional[float] added_eos: bool = False - def get_turn_rollout_expert_indices(self) -> Optional[List[List[List[int]]]]: + def get_turn_rollout_expert_indices(self) -> Optional[RoutedExpertIndices]: """Return only routes that the inference model actually executed.""" return self.rollout_expert_indices @@ -455,9 +456,6 @@ async def agent_loop( rollout_expert_indices=rollout_expert_indices, ) - if turn_output.rollout_expert_indices is not None and agent_loop_state.rollout_expert_indices is None: - agent_loop_state.rollout_expert_indices = [] - if is_step_wise: # current response + observation ids turn_response_ids = turn_output.output_ids + turn_output.obs_ids @@ -747,7 +745,7 @@ async def generate_batched( loss_masks = [] env_metrics = [] truncated_logprobs: Optional[List[List[float]]] = [] if logprobs is not None else None - truncated_indices: Optional[List] = [] if raw_rollout_expert_indices is not None else None + truncated_indices: Optional[List[RoutedExpertIndices]] = [] if raw_rollout_expert_indices is not None else None for i, (output, response, env, env_class) in enumerate(zip(outputs, responses, envs, env_classes)): # step on environment and compute reward @@ -1095,7 +1093,7 @@ def _update_agent_loop_state_with_multiturn_chat_template( agent_loop_state.loss_mask += loss_mask_for_turn if agent_loop_state.rollout_logprobs is not None and rollout_logprobs_for_turn is not None: agent_loop_state.rollout_logprobs += rollout_logprobs_for_turn - if agent_loop_state.rollout_expert_indices is not None and rollout_expert_indices_for_turn is not None: + if rollout_expert_indices_for_turn is not None: # overwrite the existing rollout inference indices, since the inference engine should # return the expert indices for the entire sequence including each turn's input # and the final response should not have an observation appended to it diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index 2ad4c0a50d..d137bcc02a 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -864,9 +864,7 @@ def convert_to_training_input(self, generator_output: GeneratorOutput, uids: Lis loss_masks: List[List[int]] = generator_output["loss_masks"] logprobs: Optional[List[List[float]]] = generator_output.get("rollout_logprobs", None) - rollout_expert_indices: Optional[List[List[List[List[int]]]]] = generator_output.get( - "rollout_expert_indices", None - ) + rollout_expert_indices = generator_output.get("rollout_expert_indices", None) pixel_values = generator_output.get("pixel_values", None) image_grid_thw = generator_output.get("image_grid_thw", None) diff --git a/skyrl/utils/routed_experts.py b/skyrl/utils/routed_experts.py index 9275c330e6..df75681964 100644 --- a/skyrl/utils/routed_experts.py +++ b/skyrl/utils/routed_experts.py @@ -1,5 +1,39 @@ +from typing import TypeAlias + +import numpy as np import torch +RoutedExpertIndices: TypeAlias = np.ndarray +ROUTED_EXPERT_DTYPES = frozenset({np.dtype(np.uint8), np.dtype(np.int16), np.dtype(np.int32)}) + + +def compact_routed_expert_indices(routed_experts: RoutedExpertIndices) -> RoutedExpertIndices: + """Validate and compact a routed-expert array to the canonical integer dtype.""" + if not isinstance(routed_experts, np.ndarray): + raise TypeError("routed expert indices must be a NumPy array") + if routed_experts.ndim != 3 or not np.issubdtype(routed_experts.dtype, np.integer): + raise ValueError( + "routed expert indices must be an integer [tokens, layers, topk] array, " + f"got shape {routed_experts.shape} and dtype {routed_experts.dtype}" + ) + if int(routed_experts.min(initial=0)) < 0: + raise ValueError("routed expert indices must be non-negative") + + max_expert_id = int(routed_experts.max(initial=0)) + if max_expert_id < 2**8: + dtype = np.dtype(np.uint8) + elif max_expert_id < 2**15: + dtype = np.dtype(np.int16) + elif max_expert_id < 2**31: + dtype = np.dtype(np.int32) + else: + raise ValueError(f"routed expert index exceeds signed int32: {max_expert_id}") + + compact = np.asarray(routed_experts, dtype=dtype, order="C") + if not compact.flags.writeable: + compact = compact.copy(order="C") + return compact + def make_replay_padding_indices( shape: tuple[int, ...], diff --git a/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py b/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py index 65158253ba..2fa028eda9 100644 --- a/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py +++ b/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py @@ -8,6 +8,7 @@ import aiohttp import httpx +import numpy as np import pytest import pytest_asyncio import uvicorn @@ -20,6 +21,9 @@ PauseMode, RemoteInferenceClient, ) +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + pack_routed_experts, +) from skyrl.backends.skyrl_train.inference_servers.setup import ( build_new_inference_client, ) @@ -113,6 +117,9 @@ async def generate(request: Request): for i in range(num_choices) ] } + if request.url.path == "/skyrl/v1/generate": + routes = np.arange(12).reshape(3, 2, 2) + response["choices"][0]["routed_experts"] = pack_routed_experts(routes) features = body.get("features") app.state.last_generate_features = features @@ -450,6 +457,47 @@ async def test_generate_with_session_id(self, client): result = await client.generate(input_batch) assert len(result["responses"]) == 1 + @pytest.mark.asyncio + async def test_generate_decodes_packed_routed_experts(self, mock_servers): + client = RemoteInferenceClient( + proxy_url=mock_servers["proxy_url"], + server_urls=mock_servers["server_urls"], + data_parallel_size=1, + enable_return_routed_experts=True, + ) + try: + result = await client.generate({"prompt_token_ids": [[1, 2, 3]]}) + finally: + await client.teardown() + + assert len(result["rollout_expert_indices"]) == 1 + assert result["rollout_expert_indices"][0].dtype == np.uint8 + assert np.array_equal(result["rollout_expert_indices"][0], np.arange(12).reshape(3, 2, 2)) + + @pytest.mark.asyncio + async def test_generate_rejects_list_routed_experts(self, monkeypatch): + client = RemoteInferenceClient( + proxy_url="http://unused", + server_urls=["http://unused"], + data_parallel_size=1, + enable_return_routed_experts=True, + ) + + async def return_list_routes(*args, **kwargs): + return { + "choices": [ + { + "token_ids": [1], + "finish_reason": "stop", + "routed_experts": [[[0, 1]]], + } + ] + } + + monkeypatch.setattr(client, "_post", return_list_routes) + with pytest.raises(ValueError, match="must return packed"): + await client._generate_single([1], {}, None, "model") + @pytest.mark.asyncio async def test_chat_completion(self, client): """Test chat completion method.""" diff --git a/tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py b/tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py new file mode 100644 index 0000000000..92655372e0 --- /dev/null +++ b/tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py @@ -0,0 +1,92 @@ +import base64 + +import numpy as np +import pytest + +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + decode_packed_routed_experts, + pack_routed_experts, +) +from skyrl.utils.routed_experts import compact_routed_expert_indices + + +@pytest.mark.parametrize( + "routes,expected_dtype", + [ + (np.arange(12).reshape(3, 2, 2), "uint8"), + (np.array([[[2**8 - 1]]]), "uint8"), + (np.array([[[0, 2**8]]]), "int16"), + (np.array([[[0, 2**15 - 1]]]), "int16"), + (np.array([[[0, 2**15]]]), "int32"), + (np.array([[[0, 2**31 - 1]]], dtype=np.int64), "int32"), + (np.empty((0, 2, 2), dtype=np.int64), "uint8"), + (np.arange(24).reshape(6, 2, 2)[::2], "uint8"), + ], +) +def test_packed_routed_experts_round_trip(routes, expected_dtype): + payload = pack_routed_experts(routes) + decoded = decode_packed_routed_experts(payload) + + assert payload["dtype"] == expected_dtype + assert decoded.dtype.name == expected_dtype + assert decoded.flags.c_contiguous + assert np.array_equal(decoded, routes) + + +def test_packed_routed_experts_uses_raw_base64(): + assert pack_routed_experts(np.array([[[1, 2, 3]]]))["data"] == "AQID" + + +@pytest.mark.parametrize( + "routes", + [np.array([1, 2]), np.array([[[-1]]]), np.array([[[2**31]]], dtype=np.uint64)], +) +def test_pack_rejects_invalid_routes(routes): + with pytest.raises(ValueError): + pack_routed_experts(routes) + + +def test_pack_rejects_nested_lists(): + with pytest.raises(TypeError, match="NumPy array"): + pack_routed_experts([[[1, 2]]]) + + +def test_compaction_makes_read_only_arrays_writable(): + routes = np.arange(12, dtype=np.uint8).reshape(3, 2, 2) + routes.flags.writeable = False + + compact = compact_routed_expert_indices(routes) + + assert compact.dtype == np.uint8 + assert compact.flags.c_contiguous + assert compact.flags.writeable + + +def test_decode_rejects_incorrect_byte_count(): + with pytest.raises(ValueError, match="bytes"): + decode_packed_routed_experts({"data": "AQ==", "shape": [2, 1, 1], "dtype": "uint8"}) + + +@pytest.mark.parametrize( + "payload", + [ + {"data": "AQ==", "shape": [1, 1, 1], "dtype": "uint16"}, + {"data": "!", "shape": [1, 1, 1], "dtype": "uint8"}, + {"data": "AQ==", "shape": [True, 1, 1], "dtype": "uint8"}, + ], +) +def test_decode_rejects_malformed_payloads(payload): + with pytest.raises(ValueError): + decode_packed_routed_experts(payload) + + +def test_decode_rejects_noncanonical_dtype(): + routes = np.array([[[300]]], dtype=np.int32) + payload = { + "data": base64.b64encode(routes.tobytes()).decode("ascii"), + "shape": [1, 1, 1], + "dtype": "int32", + } + + with pytest.raises(ValueError, match="non-canonical dtype"): + decode_packed_routed_experts(payload) diff --git a/tests/train/dataset/test_preprocess.py b/tests/train/dataset/test_preprocess.py index 61d5ccb757..a3e4a121d6 100644 --- a/tests/train/dataset/test_preprocess.py +++ b/tests/train/dataset/test_preprocess.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock +import numpy as np import pytest import torch @@ -67,15 +68,21 @@ def test_router_padding_mask_marks_left_padding_and_uncaptured_suffix(): def test_routed_expert_tensor_uses_unique_dummy_routes(tokenizer): routes = [ - [ - [[2, 3], [4, 5]], - [[6, 7], [0, 1]], - ], - [ - [[1, 2], [3, 4]], - [[5, 6], [7, 0]], - [[2, 4], [6, 7]], - ], + np.asarray( + [ + [[2, 3], [4, 5]], + [[6, 7], [0, 1]], + ], + dtype=np.uint8, + ), + np.asarray( + [ + [[1, 2], [3, 4]], + [[5, 6], [7, 0]], + [[2, 4], [6, 7]], + ], + dtype=np.uint8, + ), ] *_, routed = convert_prompts_responses_to_batch_tensors( @@ -88,9 +95,82 @@ def test_routed_expert_tensor_uses_unique_dummy_routes(tokenizer): ) assert routed.shape == (2, 3, 2, 2) + assert routed.dtype == torch.uint8 assert routed[0, 2].tolist() == [[0, 1], [0, 1]] +@pytest.mark.parametrize( + ("max_expert_id", "source_dtype", "expected_dtype"), + [(2**8, np.int16, torch.int16), (2**15, np.int32, torch.int32)], +) +def test_routed_expert_tensor_promotes_mixed_batch_dtype( + tokenizer, + max_expert_id, + source_dtype, + expected_dtype, +): + routes = [ + np.asarray([[[1, 2]]], dtype=np.uint8), + np.asarray([[[max_expert_id, max_expert_id + 1]]], dtype=source_dtype), + ] + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10], [20]], + responses=[[11], [21]], + rewards=[[0.0], [0.0]], + loss_masks=[[1], [1]], + rollout_expert_indices=routes, + ) + + assert routed.dtype == expected_dtype + assert routed[1, 0].tolist() == [[max_expert_id, max_expert_id + 1]] + + +def test_routed_expert_tensor_accepts_read_only_arrays(tokenizer): + routes = np.asarray([[[1, 2]], [[3, 4]]], dtype=np.uint8) + routes.flags.writeable = False + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[routes], + ) + + assert routed.dtype == torch.uint8 + assert routed.tolist() == [[[[1, 2]], [[3, 4]]]] + + +def test_routed_expert_tensor_rejects_nested_lists(tokenizer): + with pytest.raises(TypeError, match="NumPy arrays"): + convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[[[[1, 2]], [[3, 4]]]], + ) + + +@pytest.mark.parametrize("dtype", [np.uint16, np.int64]) +def test_routed_expert_tensor_rejects_unsupported_dtypes(tokenizer, dtype): + routes = np.asarray([[[1, 2]], [[3, 4]]], dtype=dtype) + + with pytest.raises(TypeError, match="Unsupported routed expert dtype"): + convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[routes], + ) + + def test_convert_prompts_responses_to_batch_tensors_exact(tokenizer): """ Test with inputs of exact lengths. @@ -325,8 +405,8 @@ def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): topk = 2 # rollout_expert_indices[i] has shape [prompt_len_i + response_len_i, num_layers, topk] # Sample 0: 5 tokens, sample 1: 6 tokens - rei_0 = [[[1, 2]] * num_layers for _ in range(5)] # 5 tokens - rei_1 = [[[3, 4]] * num_layers for _ in range(6)] # 6 tokens + rei_0 = np.asarray([[[1, 2]] * num_layers for _ in range(5)], dtype=np.uint8) # 5 tokens + rei_1 = np.asarray([[[3, 4]] * num_layers for _ in range(6)], dtype=np.uint8) # 6 tokens seq, attn, action, rew, lm, lp, rei_tensor = convert_prompts_responses_to_batch_tensors( tokenizer, diff --git a/tests/train/generators/test_generator_output_utils.py b/tests/train/generators/test_generator_output_utils.py index 7546858bb8..db1570d02b 100644 --- a/tests/train/generators/test_generator_output_utils.py +++ b/tests/train/generators/test_generator_output_utils.py @@ -653,7 +653,7 @@ def test_asserts_no_expert_indices(self): "rollout_metrics": None, "rollout_logprobs": None, "trajectory_ids": [tid], - "rollout_expert_indices": [[[[1, 2]]]], + "rollout_expert_indices": [np.asarray([[[1, 2]]], dtype=np.uint8)], "is_last_step": [True], } with pytest.raises(AssertionError, match="rollout_expert_indices not supported"): diff --git a/tests/train/generators/test_skyrl_gym_generator.py b/tests/train/generators/test_skyrl_gym_generator.py index f0c7ffdcfd..460e1b9755 100644 --- a/tests/train/generators/test_skyrl_gym_generator.py +++ b/tests/train/generators/test_skyrl_gym_generator.py @@ -5,6 +5,7 @@ from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock, patch +import numpy as np import pytest from skyrl.train.config import ChatTemplateConfig, GeneratorConfig @@ -23,7 +24,7 @@ def test_turn_output_keeps_uncaptured_suffix_out_of_routes(): - routes = [[[2, 3]], [[4, 5]]] + routes = np.asarray([[[2, 3]], [[4, 5]]], dtype=np.uint8) output = TurnOutput( output="answer", output_ids=[10, 11, 4], diff --git a/tests/train/test_trainer_utils.py b/tests/train/test_trainer_utils.py index 919a068141..4fa401f1cf 100644 --- a/tests/train/test_trainer_utils.py +++ b/tests/train/test_trainer_utils.py @@ -10,6 +10,7 @@ from typing import Union from unittest.mock import Mock, mock_open, patch +import numpy as np import pytest import ray @@ -393,7 +394,7 @@ def test_handle_replace_sampling_sufficient_good_samples(): "stop_reasons": ["length"] * 6, "rollout_metrics": None, "rollout_logprobs": [[0.1, 0.2], [0.3, 0.4], [0.5, 0.25], [0.15, 0.25], [0.1, 0.2], [0.3, 0.4]], - "rollout_expert_indices": [[[[i, i + 1]]] for i in range(6)], + "rollout_expert_indices": [np.asarray([[[i, i + 1]]], dtype=np.uint8) for i in range(6)], } uids = ["uid1", "uid1", "uid2", "uid2", "uid3", "uid3"] # 2 samples per prompt sampling_config = {"n_samples_per_prompt": 2, "min_replace_ratio": 0.3} @@ -414,7 +415,7 @@ def test_handle_replace_sampling_sufficient_good_samples(): for response, routes in zip(generator_output["response_ids"], generator_output["rollout_expert_indices"]) } for response, routes in zip(result_output["response_ids"], result_output["rollout_expert_indices"]): - assert routes == route_by_response[tuple(response)] + assert np.array_equal(routes, route_by_response[tuple(response)]) # Check that bad uid2 samples were replaced with good samples uid2_indices = [i for i, uid in enumerate(result_uids) if uid == "uid2"] @@ -643,6 +644,7 @@ def test_handle_filter_sampling_single_sample_per_prompt(): def test_filter_generator_output(): """Test the filter_generator_output utility function.""" + routes = [np.asarray([[[i, i + 1]]], dtype=np.uint8) for i in range(3)] generator_output = { "prompt_token_ids": [[1, 2], [3, 4], [5, 6]], "response_ids": [[7, 8], [9, 10], [11, 12]], @@ -651,7 +653,7 @@ def test_filter_generator_output(): "stop_reasons": ["length", "length", "stop"], "rollout_metrics": {"metric": "value"}, "rollout_logprobs": [[0.16, 0.4], [0.1, 0.2], [0.3, 0.4]], - "rollout_expert_indices": ["routes-0", "routes-1", "routes-2"], + "rollout_expert_indices": routes, } kept_indices = [0, 2] # Keep first and third samples @@ -664,7 +666,8 @@ def test_filter_generator_output(): assert filtered["stop_reasons"] == ["length", "stop"] assert filtered["rollout_metrics"] == {"metric": "value"} assert filtered["rollout_logprobs"] == [[0.16, 0.4], [0.3, 0.4]] - assert filtered["rollout_expert_indices"] == ["routes-0", "routes-2"] + assert filtered["rollout_expert_indices"][0] is routes[0] + assert filtered["rollout_expert_indices"][1] is routes[2] def test_zero_variance_filter_mixed_groups(): diff --git a/uv.lock b/uv.lock index 8fba78c710..01c249a8fb 100644 --- a/uv.lock +++ b/uv.lock @@ -8380,8 +8380,10 @@ fsdp = [ { name = "ninja" }, { name = "nixl", marker = "sys_platform == 'linux' or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-miniswe') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-tpu') or (extra == 'extra-5-skyrl-miniswe' and extra == 'extra-5-skyrl-tpu')" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8440,8 +8442,10 @@ megatron = [ { name = "nixl", marker = "sys_platform == 'linux'" }, { name = "nvidia-modelopt", marker = "sys_platform == 'linux'" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8481,8 +8485,10 @@ miniswe = [ { name = "ninja" }, { name = "nixl", marker = "sys_platform == 'linux' or (extra == 'extra-5-skyrl-fsdp' and extra == 'extra-5-skyrl-jax')" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8516,8 +8522,10 @@ skyrl-train = [ { name = "ninja" }, { name = "nixl", marker = "sys_platform == 'linux' or (extra == 'extra-5-skyrl-fsdp' and extra == 'extra-5-skyrl-jax') or (extra == 'extra-5-skyrl-fsdp' and extra == 'extra-5-skyrl-megatron') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-megatron') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-miniswe') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-tpu') or (extra == 'extra-5-skyrl-jax' and extra == 'extra-5-skyrl-megatron') or (extra == 'extra-5-skyrl-megatron' and extra == 'extra-5-skyrl-miniswe') or (extra == 'extra-5-skyrl-megatron' and extra == 'extra-5-skyrl-tpu') or (extra == 'extra-5-skyrl-miniswe' and extra == 'extra-5-skyrl-tpu')" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8609,12 +8617,14 @@ requires-dist = [ { name = "nvidia-modelopt", marker = "sys_platform == 'linux' and extra == 'megatron'" }, { name = "omegaconf", marker = "extra == 'skyrl-train'" }, { name = "optax", marker = "extra == 'jax'", specifier = ">=0.2.5" }, + { name = "orjson", marker = "extra == 'skyrl-train'", specifier = ">=3.11.9" }, { name = "peft", specifier = "==0.18.1" }, { name = "peft", marker = "extra == 'skyrl-train'", specifier = "==0.18.1" }, { name = "pillow", specifier = ">=11.3.0" }, { name = "polars", marker = "extra == 'skyrl-train'" }, { name = "pre-commit", marker = "extra == 'dev'" }, { name = "psycopg2-binary", marker = "extra == 'tinker'" }, + { name = "pybase64", marker = "extra == 'skyrl-train'", specifier = ">=1.4.2" }, { name = "pybind11", marker = "extra == 'skyrl-train'" }, { name = "pymdown-extensions", marker = "extra == 'dev'", specifier = ">=10.7" }, { name = "pytest", marker = "extra == 'dev'" },