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]]