diff --git a/skyrl/backends/skyrl_train/distributed/megatron/model_utils.py b/skyrl/backends/skyrl_train/distributed/megatron/model_utils.py index af72391a2c..f985c75f53 100644 --- a/skyrl/backends/skyrl_train/distributed/megatron/model_utils.py +++ b/skyrl/backends/skyrl_train/distributed/megatron/model_utils.py @@ -24,6 +24,8 @@ import torch import torch.distributed as dist +from skyrl.backends.skyrl_train.utils.packed_tensor import lengths_from_offsets + @torch.no_grad() def _compute_distributed_log_softmax( @@ -924,7 +926,7 @@ def _packed_sequence_indices( token_indices = torch.arange(total_tokens, device=device) seq_indices = torch.searchsorted(cu_seqlens_padded[1:], token_indices, right=True) seq_offsets = token_indices - cu_seqlens_padded[seq_indices] - seq_lens_padded = cu_seqlens_padded[1:] - cu_seqlens_padded[:-1] + seq_lens_padded = lengths_from_offsets(cu_seqlens_padded) return cu_seqlens_padded, token_indices, seq_indices, seq_offsets, seq_lens_padded diff --git a/skyrl/backends/skyrl_train/distributed/megatron/token_metadata.py b/skyrl/backends/skyrl_train/distributed/megatron/token_metadata.py index cb3e526c95..3257bd57e8 100644 --- a/skyrl/backends/skyrl_train/distributed/megatron/token_metadata.py +++ b/skyrl/backends/skyrl_train/distributed/megatron/token_metadata.py @@ -1,13 +1,16 @@ """Token-aligned metadata layout transforms shared by training features.""" +from collections.abc import Callable, Sequence from dataclasses import dataclass +import numpy as np import torch from skyrl.backends.skyrl_train.distributed.megatron.packing_utils import ( get_packed_seq_align_size, get_unpacked_seq_align_size, ) +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor # Megatron is imported lazily inside functions so that non-Megatron backends can # import this module for its layout dataclass and padding transforms. @@ -99,34 +102,113 @@ def align_token_metadata( f"attention_mask shape {layout.attention_mask.shape}" ) + return _align_token_rows( + lambda row_index: metadata[row_index, layout.attention_mask[row_index]], + metadata, + metadata.shape[2:], + layout, + padding_value, + next_token=next_token, + ) + + +def align_packed_token_metadata( + metadata: PackedTensor, + layout: TokenMetadataLayout, + padding_value: torch.Tensor | bool | int, + *, + next_token: bool = False, + segment_starts: Sequence[int] | None = None, +) -> torch.Tensor: + """Align metadata that already arrives packed as ``[sum(seqlen), *row_shape]``. + + This relies on left padding, which makes each trajectory's real tokens contiguous. + Without ``segment_starts``, segments must match ``layout.sequence_lengths``; + otherwise each segment is placed at its specified real-token offset. + """ + if metadata.device != layout.attention_mask.device: + raise ValueError("Token-aligned metadata and attention_mask must be on the same device") + if len(metadata) != len(layout.sequence_lengths): + raise ValueError( + f"Packed metadata holds {len(metadata)} segments for {len(layout.sequence_lengths)} trajectories" + ) + segment_lengths = metadata.sequence_lengths.tolist() + if segment_starts is None: + if segment_lengths != list(layout.sequence_lengths): + raise ValueError( + f"Packed metadata segments {segment_lengths} do not match " + f"trajectory lengths {list(layout.sequence_lengths)}" + ) + else: + if len(segment_starts) != len(metadata): + raise ValueError(f"Got {len(segment_starts)} segment starts for {len(metadata)} segments") + for row_index, (start, length) in enumerate(zip(segment_starts, segment_lengths, strict=True)): + if start < 0 or start + length > layout.sequence_lengths[row_index]: + raise ValueError( + f"Segment {row_index} spans real tokens [{start}, {start + length}) of a " + f"{layout.sequence_lengths[row_index]}-token trajectory" + ) + + return _align_token_rows( + metadata.segment, + metadata.values, + metadata.row_shape, + layout, + padding_value, + next_token=next_token, + segment_starts=segment_starts, + ) + + +def _align_token_rows( + rows_for: Callable[[int], torch.Tensor], + source: torch.Tensor, + row_shape: tuple[int, ...] | torch.Size, + layout: TokenMetadataLayout, + padding_value: torch.Tensor | bool | int, + *, + next_token: bool = False, + segment_starts: Sequence[int] | None = None, +) -> torch.Tensor: + """Place each trajectory's real-token rows into Megatron's layout and CP-shard them. + + ``rows_for(row_index)`` yields one trajectory's rows. They land at the front of its + padded region unless ``segment_starts`` names a per-trajectory destination offset. + """ 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:]), + source, + (len(layout.sequence_lengths), layout.aligned_sequence_length, *row_shape), 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]] + rows = rows_for(row_index) + start = 0 if segment_starts is None else segment_starts[row_index] + end = sequence_length if segment_starts is None else start + rows.shape[0] + aligned[row_index, start:end] = rows return aligned packed = _new_metadata_tensor( - metadata, - (layout.aligned_sequence_length, *metadata.shape[2:]), + source, + (layout.aligned_sequence_length, *row_shape), 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]] + rows = rows_for(row_index) + start = offset if segment_starts is None else offset + segment_starts[row_index] + end = offset + sequence_length if segment_starts is None else start + rows.shape[0] + packed[start:end] = rows # 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) + shifted = _new_metadata_tensor(source, 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] @@ -135,7 +217,7 @@ def align_token_metadata( if layout.context_parallel_size > 1: out = _new_metadata_tensor( - metadata, + source, (packed.shape[0] // layout.context_parallel_size, *packed.shape[1:]), padding_value, ) @@ -205,3 +287,51 @@ def scatter_packed_token_values_to_batch( ) batch_values[output_mask] = values[packed_mask] return batch_values + + +class TokenMetadataTrace: + """Accumulate arrays whose first dimension is aligned to tokens.""" + + def __init__(self) -> None: + self._chunks: list[np.ndarray] = [] + self._schema: tuple[tuple[int, ...], np.dtype] | None = None + self._num_rows = 0 + self._finalized = False + + @property + def num_rows(self) -> int: + return self._num_rows + + def append(self, rows: np.ndarray, *, expected_rows: int) -> None: + if self._finalized: + raise RuntimeError("token metadata trace is already finalized") + if isinstance(expected_rows, bool) or not isinstance(expected_rows, int) or expected_rows < 0: + raise ValueError(f"expected_rows must be a non-negative integer, got {expected_rows!r}") + if not isinstance(rows, np.ndarray): + raise TypeError("token metadata rows must be a NumPy array") + if rows.ndim < 1: + raise ValueError("token metadata must have a token-row dimension") + if rows.shape[0] != expected_rows: + raise ValueError(f"token metadata has {rows.shape[0]} rows, expected {expected_rows}") + if not rows.flags.c_contiguous: + raise ValueError("token metadata rows must be contiguous") + + schema = (rows.shape[1:], rows.dtype) + if self._schema is None: + self._schema = schema + elif schema != self._schema: + raise ValueError(f"token metadata schema changed from {self._schema} to {schema}") + + self._chunks.append(rows) + self._num_rows += expected_rows + + def finalize(self, *, expected_rows: int) -> np.ndarray: + if self._finalized: + raise RuntimeError("token metadata trace is already finalized") + if self._num_rows != expected_rows: + raise ValueError(f"token metadata trace has {self._num_rows} rows, expected {expected_rows}") + if not self._chunks: + raise ValueError("token metadata trace has no chunks") + + self._finalized = True + return self._chunks[0] if len(self._chunks) == 1 else np.concatenate(self._chunks, axis=0) diff --git a/skyrl/backends/skyrl_train/inference_servers/base.py b/skyrl/backends/skyrl_train/inference_servers/base.py index aac9ee6f6d..4a8de4cdf4 100644 --- a/skyrl/backends/skyrl_train/inference_servers/base.py +++ b/skyrl/backends/skyrl_train/inference_servers/base.py @@ -34,6 +34,7 @@ class InferenceEngineInput(TypedDict): # Optional prefix-cache salt forwarded to vLLM as the request ``cache_salt`` so cache blocks are # only shared between requests carrying the same salt. See ``GeneratorConfig.use_cache_salt``. cache_salt: Optional[str] + routed_experts_prompt_starts: Optional[List[int]] class InferenceEngineOutput(TypedDict): diff --git a/skyrl/backends/skyrl_train/inference_servers/generate_wire.py b/skyrl/backends/skyrl_train/inference_servers/generate_wire.py index a9b8f43a2c..a5f7aac23c 100644 --- a/skyrl/backends/skyrl_train/inference_servers/generate_wire.py +++ b/skyrl/backends/skyrl_train/inference_servers/generate_wire.py @@ -3,14 +3,21 @@ ``VLLMServerActor`` writes these payloads and ``RemoteInferenceClient`` reads them; nothing else depends on the encoding. Both sides serialize with orjson, which rejects non-finite floats and has no notion of NumPy arrays, so the -helpers here exist to get sampled logprobs and routed-expert IDs across that +helpers here exist to get sampled logprobs and NumPy side channels across that boundary intact. + +Side-channel arrays use ``{data: , shape: [...], dtype: }`` +envelopes. Keeping ``data`` first lets ``load_packed_body`` decode it from the +raw response without materializing a large Python ``str``. """ import math -from typing import Any, Iterable, Mapping, Optional, Tuple +from collections import deque +from enum import StrEnum +from typing import Any, Collection, Iterable, Mapping, Optional, Tuple import numpy as np +import orjson import pybase64 from skyrl.backends.skyrl_train.utils.routed_experts import ( @@ -22,7 +29,32 @@ # Matches the floor vLLM applies at its own serving boundaries. CLAMPED_LOGPROB = -9999.0 -_DTYPES = {dtype.name: dtype for dtype in ROUTED_EXPERT_DTYPES} + +class PackedArrayKey(StrEnum): + """Envelope keys, with ``DATA`` first for ``load_packed_body``.""" + + DATA = "data" + SHAPE = "shape" + DTYPE = "dtype" + + +class PackedField(StrEnum): + """Response-body fields whose value is a packed-array envelope.""" + + ROUTED_EXPERTS = "routed_experts" + ROLLOUT_SAMPLE_SUPPORT = "rollout_sample_support" + + +PACKED_SIDE_CHANNEL_FIELDS: tuple[str, ...] = tuple(PackedField) + +_ENVELOPE_KEYS = frozenset(PackedArrayKey) + +_ROUTED_EXPERTS_NDIM = 3 + +_QUOTE = b'"' + +# Base64 cannot contain this scan anchor. +_PACKED_DATA_ANCHOR = f':{{"{PackedArrayKey.DATA}":"'.encode() def build_logprobs_content( @@ -72,35 +104,151 @@ def _to_host_array(routed_experts: Any) -> Any: return routed_experts -def pack_routed_experts(routed_experts: RoutedExpertIndices) -> dict[str, Any]: - compact = compact_routed_expert_indices(_to_host_array(routed_experts)) - return { - "data": pybase64.b64encode(memoryview(compact)).decode("ascii"), - "shape": list(compact.shape), - "dtype": compact.dtype.name, +def pack_ndarray( + arr: np.ndarray, + *, + allowed_dtypes: Collection[np.dtype], + extra: Optional[Mapping[str, Any]] = None, +) -> dict[str, Any]: + """Encode ``arr`` as a base64 envelope carrying ``extra`` as sidecar fields.""" + if not isinstance(arr, np.ndarray): + raise TypeError("packed array must be a NumPy array") + if arr.dtype not in allowed_dtypes: + allowed = sorted(dtype.name for dtype in allowed_dtypes) + raise ValueError(f"packed array {PackedArrayKey.DTYPE} {arr.dtype.name!r} is not one of {allowed}") + if extra is not None: + collisions = sorted(set(extra) & _ENVELOPE_KEYS) + if collisions: + raise ValueError(f"sidecar fields collide with envelope keys: {collisions}") + + contiguous = np.ascontiguousarray(arr) + # `.value` keys: orjson rejects str subclasses as dict keys. + payload = { + PackedArrayKey.DATA.value: pybase64.b64encode(memoryview(contiguous)).decode("ascii"), + PackedArrayKey.SHAPE.value: list(contiguous.shape), + PackedArrayKey.DTYPE.value: contiguous.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") + if extra is not None: + payload.update(extra) + return payload + + +def unpack_ndarray( + payload: Mapping[str, Any], + *, + allowed_dtypes: Collection[np.dtype], + ndim: int, +) -> Tuple[np.ndarray, dict[str, Any]]: + """Decode an envelope whose base64 ``data`` may be a string or buffer.""" + if not isinstance(payload, Mapping): + raise TypeError("packed array payload must be an object") try: - dtype = _DTYPES[payload["dtype"]] - shape = tuple(payload["shape"]) - data = pybase64.b64decode_as_bytearray(payload["data"], validate=True) + dtype_name = payload[PackedArrayKey.DTYPE] + shape = tuple(payload[PackedArrayKey.SHAPE]) + data = pybase64.b64decode_as_bytearray(payload[PackedArrayKey.DATA], validate=True) except (KeyError, TypeError, ValueError) as exc: - raise ValueError("invalid packed routed_experts payload") from exc - # bool is a subclass of int, so it needs an explicit rejection; np.integer is - # accepted for in-process callers, since orjson only ever yields plain ints. - if len(shape) != 3 or any( + raise ValueError(f"invalid packed array envelope: {exc}") from exc + + dtypes = {dtype.name: dtype for dtype in allowed_dtypes} + if not isinstance(dtype_name, str) or dtype_name not in dtypes: + raise ValueError(f"packed array {PackedArrayKey.DTYPE} {dtype_name!r} is not one of {sorted(dtypes)}") + dtype = dtypes[dtype_name] + # Reject bool, an int subclass; accept np.integer for in-process callers. + if len(shape) != ndim or any( not isinstance(dim, (int, np.integer)) or isinstance(dim, bool) or dim < 0 for dim in shape ): - raise ValueError(f"invalid packed routed_experts shape: {shape}") + raise ValueError(f"packed array {PackedArrayKey.SHAPE} {shape} is not {ndim} non-negative dimensions") 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) + raise ValueError( + f"packed array {PackedArrayKey.DATA} has {len(data)} bytes, " + f"expected {expected_size} for {dtype_name}{list(shape)}" + ) + + array = np.frombuffer(data, dtype=dtype).reshape(shape) + sidecar = {key: value for key, value in payload.items() if key not in _ENVELOPE_KEYS} + return array, sidecar + + +def pack_routed_experts(routed_experts: RoutedExpertIndices) -> dict[str, Any]: + compact = compact_routed_expert_indices(_to_host_array(routed_experts)) + return pack_ndarray(compact, allowed_dtypes=ROUTED_EXPERT_DTYPES) + + +def decode_packed_routed_experts(payload: dict[str, Any]) -> RoutedExpertIndices: + decoded, _ = unpack_ndarray(payload, allowed_dtypes=ROUTED_EXPERT_DTYPES, ndim=_ROUTED_EXPERTS_NDIM) 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}") + if compact.dtype != decoded.dtype: + raise ValueError( + f"packed routed_experts uses non-canonical dtype {decoded.dtype.name}; expected {compact.dtype.name}" + ) return compact + + +def _data_prefix(field: str) -> bytes: + """The bytes an orjson-serialized packed ``field`` opens with.""" + return f'"{field}"'.encode() + _PACKED_DATA_ANCHOR + + +def load_packed_body(raw: bytes, *, fields: tuple[str, ...] = PACKED_SIDE_CHANNEL_FIELDS) -> dict[str, Any]: + """Parse a response after replacing registered base64 blobs with views. + + Null fields pass through. An envelope layout the scan cannot splice raises + instead of falling back to materializing the base64 as a Python string. + """ + prefixes = {field: _data_prefix(field) for field in fields} + blobs: dict[str, deque[memoryview]] = {field: deque() for field in fields} + view = memoryview(raw) + pieces: list[memoryview] = [] + copied = 0 + scan = 0 + while (anchor := raw.find(_PACKED_DATA_ANCHOR, scan)) >= 0: + field = _match_packed_field(raw, anchor, prefixes) + if field is None: + scan = anchor + len(_PACKED_DATA_ANCHOR) + continue + start = anchor + len(_PACKED_DATA_ANCHOR) + end = raw.find(_QUOTE, start) + if end < 0: + raise ValueError(f"unterminated base64 {PackedArrayKey.DATA} for {field} in the response body") + pieces.append(view[copied:start]) + blobs[field].append(view[start:end]) + copied = scan = end + + if pieces: + pieces.append(view[copied:]) + body = orjson.loads(b"".join(pieces)) + else: + body = orjson.loads(raw) + _restore_packed_data(body, blobs) + + unplaced = {field: len(queue) for field, queue in blobs.items() if queue} + if unplaced: + raise ValueError(f"spliced packed blobs found no envelope in the response body: {unplaced}") + return body + + +def _match_packed_field(raw: bytes, anchor: int, prefixes: Mapping[str, bytes]) -> Optional[str]: + """Name the registered field whose prefix ends at ``anchor``, if any.""" + for field, prefix in prefixes.items(): + begin = anchor + len(_PACKED_DATA_ANCHOR) - len(prefix) + if begin >= 0 and raw.startswith(prefix, begin): + return field + return None + + +def _restore_packed_data(node: Any, blobs: Mapping[str, deque[memoryview]]) -> None: + """Put each blob back on its envelope's ``data`` key, in document order.""" + if isinstance(node, dict): + for key, value in node.items(): + queue = blobs.get(key) + if queue is not None and isinstance(value, dict) and PackedArrayKey.DATA in value: + if not queue: + raise ValueError(f"packed {key} survived the scan unspliced; the response-body layout drifted") + value[PackedArrayKey.DATA.value] = queue.popleft() + elif isinstance(value, (dict, list)): + _restore_packed_data(value, blobs) + elif isinstance(node, list): + for item in node: + if isinstance(item, (dict, list)): + _restore_packed_data(item, blobs) 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 7bfbd42a64..f6137b8e49 100644 --- a/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py +++ b/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py @@ -74,8 +74,11 @@ MultiModalFeatures, ) from skyrl.backends.skyrl_train.inference_servers.generate_wire import ( + PackedField, decode_packed_routed_experts, + load_packed_body, ) +from skyrl.backends.skyrl_train.utils.routed_experts import RoutedExpertIndices from skyrl.backends.utils import convert_vllm_prompt_logprobs from skyrl.env_vars import ( SKYRL_GENERATE_CONCURRENCY_PER_ENGINE, @@ -167,6 +170,168 @@ class SampleResponse(TypedDict): topk_prompt_logprobs: Optional[List[Optional[List[Tuple[int, float]]]]] +@dataclass(frozen=True) +class RemoteGenerateResult: + """Raw token generation result returned by ``RemoteGenerateClient``.""" + + raw_response: Dict[str, Any] + response_ids: List[int] + response_logprobs: Optional[List[float]] + stop_reason: str + routed_experts: Optional[RoutedExpertIndices] + + +@dataclass +class RemoteGenerateClient: + """Reusable HTTP client for one raw-token generation request.""" + + proxy_url: str + _session: Optional[aiohttp.ClientSession] = field(default=None, init=False, repr=False) + + async def _get_session(self) -> aiohttp.ClientSession: + current_loop = asyncio.get_running_loop() + if self._session is not None and not self._session.closed and self._session.loop != current_loop: + self._session = None + if self._session is None or self._session.closed: + connector = aiohttp.TCPConnector( + limit=SKYRL_HTTP_CONNECTION_LIMIT, + keepalive_timeout=2, + ) + self._session = aiohttp.ClientSession( + connector=connector, + timeout=aiohttp.ClientTimeout(total=None), + ) + return self._session + + async def _post( + self, + url: str, + json: Dict[str, Any], + headers: Optional[Dict[str, str]] = None, + *, + packed_side_channels: bool = False, + ) -> Any: + """POST JSON with retries, optionally splicing packed arrays before parsing.""" + session = await self._get_session() + last_exc: Optional[Exception] = None + for attempt in range(_DATA_PLANE_RETRIES): + try: + async with session.post(url, json=json, headers=headers) as resp: + try: + raw = await resp.read() + body = load_packed_body(raw) if packed_side_channels else orjson.loads(raw) + except orjson.JSONDecodeError as exc: + if 400 <= resp.status < 500: + text = await resp.text() + raise aiohttp.ClientResponseError( + resp.request_info, + resp.history, + status=resp.status, + message=text or resp.reason, + headers=resp.headers, + ) from exc + last_exc = exc + logger.debug(f"retry {attempt + 1}/{_DATA_PLANE_RETRIES} for {url=}: {exc}") + await asyncio.sleep(1) + continue + raise_for_status(resp, body) + return body + except (aiohttp.ServerDisconnectedError, aiohttp.ClientOSError) as exc: + last_exc = exc + logger.debug(f"POST retry {attempt + 1}/{_DATA_PLANE_RETRIES} for {url=}: {exc}") + await asyncio.sleep(1) + if last_exc is None: + raise RuntimeError(f"POST failed without an exception for {url=}") + raise last_exc + + async def generate( + self, + *, + prompt_token_ids: List[int], + sampling_params: Dict[str, Any], + session_id: Optional[Any], + model: str, + return_routed_experts: bool = False, + routed_experts_prompt_start: Optional[int] = None, + mm_features: Optional[MultiModalFeatures] = None, + cache_salt: Optional[str] = None, + ) -> RemoteGenerateResult: + """Generate one raw-token completion, optionally returning R3 routes.""" + if routed_experts_prompt_start is not None: + if not return_routed_experts: + raise ValueError("routed_experts_prompt_start requires return_routed_experts=True") + if ( + isinstance(routed_experts_prompt_start, bool) + or not isinstance(routed_experts_prompt_start, int) + or not 0 <= routed_experts_prompt_start <= len(prompt_token_ids) + ): + raise ValueError("routed_experts_prompt_start must be an integer within the prompt") + + path = "/skyrl/v1/generate" if return_routed_experts else "/inference/v1/generate" + request_sampling_params = dict(sampling_params) + if routed_experts_prompt_start is not None: + request_sampling_params["routed_experts_prompt_start"] = routed_experts_prompt_start + payload: Dict[str, Any] = { + "sampling_params": request_sampling_params, + "model": model, + "token_ids": prompt_token_ids, + } + if mm_features: + payload["features"] = mm_features + # `cache_salt` is a top-level request field (forwarded to vLLM's TokensPrompt), not a sampling + # param. + if cache_salt is not None: + payload["cache_salt"] = cache_salt + + headers = {"Content-Type": "application/json"} + if session_id: + headers["X-Session-ID"] = str(session_id) + + response = await self._post( + f"{self.proxy_url}{path}", + json=payload, + headers=headers, + packed_side_channels=return_routed_experts, + ) + choice = response["choices"][0] + token_ids = choice["token_ids"] + logprobs = choice.get("logprobs") + response_logprobs = None + if logprobs is not None: + logprobs_content = logprobs.get("content", []) + if logprobs_content: + response_logprobs = [logprob_info["logprob"] for logprob_info in logprobs_content] + + routed_experts = None + if return_routed_experts: + packed_routed_experts = choice.get(PackedField.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 RemoteGenerateResult( + raw_response=response, + response_ids=token_ids, + response_logprobs=response_logprobs, + stop_reason=choice["finish_reason"], + routed_experts=routed_experts, + ) + + async def aclose(self) -> None: + if self._session is not None and not self._session.closed: + await self._session.close() + self._session = None + + def __getstate__(self) -> Dict[str, Any]: + state = self.__dict__.copy() + state["_session"] = None + return state + + def __setstate__(self, state: Dict[str, Any]) -> None: + self.__dict__.update(state) + self._session = None + + @dataclass class RemoteInferenceClient(InferenceEngineInterface): """ @@ -225,7 +390,7 @@ class RemoteInferenceClient(InferenceEngineInterface): """Optional HF tokenizer for local tokenize/detokenize (avoids HTTP round-trips).""" # Private fields excluded from repr for cleaner output - _session: Optional[aiohttp.ClientSession] = field(default=None, repr=False) + _generate_client: Optional[RemoteGenerateClient] = field(default=None, repr=False) _world_size: Optional[Tuple[int, int]] = field(default=None, repr=False) _gen_sem: Optional[asyncio.Semaphore] = field(default=None, repr=False) _detok_sem: Optional[asyncio.Semaphore] = field(default=None, repr=False) @@ -282,66 +447,16 @@ def _get_semaphores(self) -> Tuple[Optional[asyncio.Semaphore], Optional[asyncio self._sem_loop = current_loop return self._gen_sem, self._detok_sem + def _get_generate_client(self) -> RemoteGenerateClient: + if self._generate_client is None: + self._generate_client = RemoteGenerateClient(proxy_url=self.proxy_url) + return self._generate_client + async def _get_session(self) -> aiohttp.ClientSession: - """Get or create the aiohttp session.""" - # Re-use the existing session object if it is not closed. - # Note that we also create a new session object if the event loop has changed, since - # aiohttp.ClientSession is tied to the event loop. - current_loop = asyncio.get_running_loop() - if self._session is not None and not self._session.closed and self._session.loop != current_loop: - # Event loop changed - the old session is unusable (bound to a dead loop). - self._session = None - if self._session is None or self._session.closed: - # keepalive_timeout must be shorter than the server's timeout_keep_alive - # (uvicorn default: 5s). Otherwise aiohttp reuses connections the server - # has already closed, causing ECONNRESET under high concurrency. - connector = aiohttp.TCPConnector( - limit=SKYRL_HTTP_CONNECTION_LIMIT, - keepalive_timeout=2, - ) - self._session = aiohttp.ClientSession(connector=connector, timeout=aiohttp.ClientTimeout(total=None)) - return self._session + return await self._get_generate_client()._get_session() async def _post(self, url: str, json: Dict[str, Any], headers: Optional[Dict[str, str]] = None) -> Any: - """POST with retry + backoff on transient connection errors. - - Between generate bursts the pool's keep-alive connections go stale - (server closes them after ``timeout_keep_alive``). An immediate - retry would grab another stale connection from the same pool, so we - sleep briefly to let the connector detect and purge dead sockets - before the next attempt. - """ - session = await self._get_session() - last_exc: Optional[Exception] = None - for attempt in range(_DATA_PLANE_RETRIES): - try: - async with session.post(url, json=json, headers=headers) as resp: - try: - 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. - text = await resp.text() - raise aiohttp.ClientResponseError( - resp.request_info, - resp.history, - status=resp.status, - message=text or resp.reason, - headers=resp.headers, - ) - last_exc = e - logger.debug(f"retry {attempt + 1}/{_DATA_PLANE_RETRIES} for {url=}: {e}") - await asyncio.sleep(1) - continue - raise_for_status(resp, body) - return body - except (aiohttp.ServerDisconnectedError, aiohttp.ClientOSError) as e: - last_exc = e - logger.debug(f"POST retry {attempt + 1}/{_DATA_PLANE_RETRIES} for {url=}: {e}") - await asyncio.sleep(1) - continue - raise last_exc # type: ignore[misc] + return await self._get_generate_client()._post(url, json=json, headers=headers) # --------------------------- # Data Plane @@ -406,6 +521,12 @@ async def generate( session_ids = input_batch.get("session_ids") mm_features = input_batch.get("mm_features") cache_salt = input_batch.get("cache_salt") + routed_experts_prompt_starts = input_batch.get("routed_experts_prompt_starts") + if routed_experts_prompt_starts is not None: + if not self.enable_return_routed_experts: + raise ValueError("routed_experts_prompt_starts requires enable_return_routed_experts=True") + if len(routed_experts_prompt_starts) != len(prompt_token_ids): + raise ValueError("routed_experts_prompt_starts must have one entry per prompt") get_logprobs = sampling_params.get("logprobs") is not None # Two semaphores decouple the generate and detokenize stages: @@ -429,6 +550,9 @@ async def _throttled_generate(idx: int) -> Dict[str, Any]: sampling_params=sampling_params, session_id=session_ids[idx] if session_ids and idx < len(session_ids) else None, mm_features=mm_features[idx] if mm_features and idx < len(mm_features) else None, + routed_experts_prompt_start=( + routed_experts_prompt_starts[idx] if routed_experts_prompt_starts is not None else None + ), model=model, cache_salt=cache_salt, ) @@ -438,6 +562,9 @@ async def _throttled_generate(idx: int) -> Dict[str, Any]: sampling_params=sampling_params, session_id=session_ids[idx] if session_ids and idx < len(session_ids) else None, mm_features=mm_features[idx] if mm_features and idx < len(mm_features) else None, + routed_experts_prompt_start=( + routed_experts_prompt_starts[idx] if routed_experts_prompt_starts is not None else None + ), model=model, cache_salt=cache_salt, ) @@ -471,64 +598,23 @@ async def _generate_single( model: str, mm_features: Optional[MultiModalFeatures] = None, cache_salt: Optional[str] = None, + routed_experts_prompt_start: Optional[int] = None, ) -> Dict[str, Any]: - """ - Generate completion for a single prompt. - - With keep-mode pause, in-flight requests are frozen by the vLLM - scheduler and resume where they left off after /resume. No retry - logic is needed. - - Returns: - Dict with keys: stop_reason, response_ids, response_logprobs - """ - url = ( - f"{self.proxy_url}/skyrl/v1/generate" - if self.enable_return_routed_experts - else f"{self.proxy_url}/inference/v1/generate" + result = await self._get_generate_client().generate( + prompt_token_ids=prompt_token_ids, + sampling_params=sampling_params, + session_id=session_id, + model=model, + return_routed_experts=self.enable_return_routed_experts, + routed_experts_prompt_start=routed_experts_prompt_start, + mm_features=mm_features, + cache_salt=cache_salt, ) - - payload: dict[str, Any] = { - "sampling_params": sampling_params, - "model": model, - "token_ids": prompt_token_ids, - } - if mm_features: - payload["features"] = mm_features - # `cache_salt` is a top-level request field (forwarded to vLLM's TokensPrompt), not a sampling - # param. - if cache_salt is not None: - payload["cache_salt"] = cache_salt - - headers = {"Content-Type": "application/json"} - if session_id: - headers["X-Session-ID"] = str(session_id) - - response = await self._post(url, json=payload, headers=headers) - - choice = response["choices"][0] - token_ids = choice["token_ids"] - stop_reason = choice["finish_reason"] - - response_logprobs: Optional[List[float]] = None - logprobs = choice.get("logprobs") - if logprobs is not None: - logprobs_content = logprobs.get("content", []) - if logprobs_content: - response_logprobs = [logprob_info["logprob"] for logprob_info in logprobs_content] - - 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, - "response_ids": token_ids, - "response_logprobs": response_logprobs, - "routed_experts": routed_experts, + "stop_reason": result.stop_reason, + "response_ids": result.response_ids, + "response_logprobs": result.response_logprobs, + "routed_experts": result.routed_experts, } async def _render_for_sample( @@ -1405,9 +1491,8 @@ async def get_world_size(self) -> Tuple[int, int]: async def teardown(self) -> None: """Close HTTP session.""" - if self._session and not self._session.closed: - await self._session.close() - self._session = None + if self._generate_client is not None: + await self._generate_client.aclose() async def __aenter__(self) -> "RemoteInferenceClient": """Async context manager entry.""" @@ -1424,7 +1509,6 @@ async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: def __getstate__(self) -> dict: """Exclude non-serializable fields from pickle.""" state = self.__dict__.copy() - state["_session"] = None state["_gen_sem"] = None state["_detok_sem"] = None state["_sem_loop"] = None @@ -1433,19 +1517,12 @@ def __getstate__(self) -> dict: def __setstate__(self, state: dict) -> None: """Restore state after unpickling.""" self.__dict__.update(state) - self._session = None self._gen_sem = None self._detok_sem = None self._sem_loop = None - async def aclose(self): - if self._session is not None: - try: - await self._session.close() - except Exception as e: - logger.warning(f"Encountered exception {e} while closing client session") - pass - self._session = None + async def aclose(self) -> None: + await self.teardown() def raise_for_status(resp: aiohttp.ClientResponse, body: Optional[Any] = None) -> None: 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 dd350ea586..dab33322e5 100644 --- a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py +++ b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py @@ -37,6 +37,7 @@ ) from skyrl.backends.skyrl_train.inference_servers.generate_wire import ( CLAMPED_LOGPROB, + PackedField, build_logprobs_content, pack_routed_experts, ) @@ -469,7 +470,7 @@ async def _skyrl_generate(request: Request): "token_ids": token_ids_out, "finish_reason": finish_reason, "logprobs": logprobs, - "routed_experts": routed_experts, + PackedField.ROUTED_EXPERTS.value: routed_experts, } ] } diff --git a/skyrl/backends/skyrl_train/training_batch.py b/skyrl/backends/skyrl_train/training_batch.py index 1542adeff3..db6c85d88f 100644 --- a/skyrl/backends/skyrl_train/training_batch.py +++ b/skyrl/backends/skyrl_train/training_batch.py @@ -3,24 +3,35 @@ import copy import io import pickle -from typing import Any, Dict, Generic, List, Optional, TypedDict, TypeVar +from enum import StrEnum +from typing import Any, Dict, Generic, List, Optional, TypedDict, TypeVar, Union import numpy as np import torch from jaxtyping import Bool, Float, Integer -from skyrl.backends.skyrl_train.utils.replay_utils import make_replay_padding_indices +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor +from skyrl.backends.skyrl_train.utils.replay_utils import append_packed_replay_padding DictType = TypeVar("DictType") +class TensorFormat(StrEnum): + """How one serialized batch field is encoded in the pickle stream.""" + + NUMPY = "numpy" + TORCH = "torch" + TENSOR_LIST = "tensor_list" + PACKED_TENSOR = "packed_tensor" + + def _serialize_tensor(value: torch.Tensor) -> dict: """Serialize a single tensor for pickle protocol.""" try: # Fast path: direct memory copy via numpy (works for most dtypes) arr = value.numpy() return { - "format": "numpy", + "format": TensorFormat.NUMPY, "data": arr.tobytes(), "shape": arr.shape, "dtype": str(arr.dtype), @@ -30,14 +41,14 @@ def _serialize_tensor(value: torch.Tensor) -> dict: buffer = io.BytesIO() torch.save(value, buffer) return { - "format": "torch", + "format": TensorFormat.TORCH, "data": buffer.getvalue(), } def _deserialize_tensor(value: dict) -> torch.Tensor: """Deserialize a single tensor from pickle format.""" - if value.get("format") == "torch": + if value.get("format") == TensorFormat.TORCH: # Fallback path: torch.load for unsupported dtypes buffer = io.BytesIO(value["data"]) return torch.load(buffer, weights_only=True) @@ -110,6 +121,13 @@ def cat(lists: list["TensorList"]) -> "TensorList": return TensorList([t for tl in lists for t in tl.tensors]) +# Value types a batch field may hold: a dense tensor, a ragged list of tensors, or a ragged +# token-aligned field packed to one buffer plus offsets. All three index by batch position. +BATCH_FIELD_TYPES = (torch.Tensor, TensorList, PackedTensor) +BatchField = Union[torch.Tensor, TensorList, PackedTensor] +_BATCH_FIELD_ERROR = f"must be a tensor, {TensorList.__name__}, or {PackedTensor.__name__}" + + def _rebuild_tensor_batch(cls, state: Dict[str, Any]): """Module-level helper for unpickling TensorBatch (must be importable by name).""" obj = dict.__new__(cls) @@ -169,8 +187,8 @@ def _check_consistency(self): value = self[key] if value is None: continue - if not isinstance(value, (torch.Tensor, TensorList)): - raise ValueError(f"Field {key} must be a tensor or TensorList, got {type(value)}") + if not isinstance(value, BATCH_FIELD_TYPES): + raise ValueError(f"Field {key} {_BATCH_FIELD_ERROR}, got {type(value)}") self._device = value.device if self._device is None else self._device if len(value) != batch_size: raise ValueError(f"Batch size mismatch in {key}") @@ -185,13 +203,13 @@ def __getitem__(self, index) -> "TensorBatch[DictType]": else: return super().__getitem__(index) - def __setitem__(self, key: str, value: Optional[torch.Tensor | TensorList]) -> None: + def __setitem__(self, key: str, value: Optional[BatchField]) -> None: if value is None: super().__setitem__(key, value) return - if not isinstance(value, (torch.Tensor, TensorList)): - raise ValueError(f"Field {key} must be a tensor or TensorList, got {type(value)}") + if not isinstance(value, BATCH_FIELD_TYPES): + raise ValueError(f"Field {key} {_BATCH_FIELD_ERROR}, got {type(value)}") if hasattr(self, "_batch_size") and self._batch_size is not None and len(value) != self._batch_size: raise ValueError(f"Batch size mismatch in {key}. Expected size {self._batch_size}, got {len(value)}.") @@ -214,9 +232,7 @@ def to( for key, value in self.items(): if value is None: continue - assert isinstance( - value, (torch.Tensor, TensorList) - ), f"Field {key} must be a tensor or TensorList, got {type(value)}" + assert isinstance(value, BATCH_FIELD_TYPES), f"Field {key} {_BATCH_FIELD_ERROR}, got {type(value)}" self[key] = value.to(device=device, dtype=dtype, non_blocking=non_blocking) return self @@ -225,9 +241,7 @@ def contiguous(self) -> "TensorBatch": for key, value in self.items(): if value is None: continue - assert isinstance( - value, (torch.Tensor, TensorList) - ), f"Field {key} must be a tensor or TensorList, got {type(value)}" + assert isinstance(value, BATCH_FIELD_TYPES), f"Field {key} {_BATCH_FIELD_ERROR}, got {type(value)}" self[key] = value.contiguous() return self @@ -267,9 +281,15 @@ def __getstate__(self): batch_dict[key] = None elif isinstance(value, TensorList): batch_dict[key] = { - "format": "tensor_list", + "format": TensorFormat.TENSOR_LIST, "items": [_serialize_tensor(t) for t in value.tensors], } + elif isinstance(value, PackedTensor): + batch_dict[key] = { + "format": TensorFormat.PACKED_TENSOR, + "values": _serialize_tensor(value.values), + "cu_seqlens": _serialize_tensor(value.cu_seqlens), + } else: batch_dict[key] = _serialize_tensor(value) @@ -288,8 +308,13 @@ def __setstate__(self, state): for key, value in state["batch_dict"].items(): if value is None: self[key] = None - elif value.get("format") == "tensor_list": + elif value.get("format") == TensorFormat.TENSOR_LIST: self[key] = TensorList([_deserialize_tensor(item) for item in value["items"]]) + elif value.get("format") == TensorFormat.PACKED_TENSOR: + self[key] = PackedTensor( + _deserialize_tensor(value["values"]), + _deserialize_tensor(value["cu_seqlens"]), + ) else: self[key] = _deserialize_tensor(value) @@ -314,10 +339,8 @@ def repeat(self, repeats: int) -> "TensorBatch[DictType]": for key, value in self.items(): if value is None: new_batch[key] = value - elif isinstance(value, TensorList): - new_batch[key] = value.repeat(repeats) else: - assert isinstance(value, torch.Tensor), f"Field {key} must be a tensor, got {type(value)}" + assert isinstance(value, BATCH_FIELD_TYPES), f"Field {key} {_BATCH_FIELD_ERROR}, got {type(value)}" new_batch[key] = value.repeat(repeats) new_batch = self.__class__(new_batch) new_batch.metadata = self.metadata @@ -338,10 +361,8 @@ def repeat_interleave(self, repeats: int) -> "TensorBatch[DictType]": for key, value in self.items(): if value is None: new_batch[key] = value - elif isinstance(value, TensorList): - new_batch[key] = value.repeat_interleave(repeats) else: - assert isinstance(value, torch.Tensor), f"Field {key} must be a tensor, got {type(value)}" + assert isinstance(value, BATCH_FIELD_TYPES), f"Field {key} {_BATCH_FIELD_ERROR}, got {type(value)}" new_batch[key] = value.repeat_interleave(repeats) new_batch = self.__class__(new_batch) new_batch.metadata = self.metadata @@ -354,7 +375,7 @@ def chunk(self, chunk_size: int) -> List["TensorBatch[DictType]"]: chunk_data = {} for key, value in self.items(): if value is not None: - if isinstance(value, (torch.Tensor, TensorList)): + if isinstance(value, BATCH_FIELD_TYPES): chunk_data[key] = value[i : i + chunk_size] else: raise ValueError(f"Unsupported type {type(value)} for key {key}") @@ -381,7 +402,7 @@ def slice(self, start: int, end: int, step: int = 1) -> "TensorBatch[DictType]": sliced_data = {} for key, value in self.items(): if value is not None: - if isinstance(value, (torch.Tensor, TensorList)): + if isinstance(value, BATCH_FIELD_TYPES): sliced_data[key] = value[slice_obj] else: raise ValueError(f"Unsupported type {type(value)} for key {key}") @@ -418,6 +439,8 @@ def cat(cls, shards: List["TensorBatch[DictType]"]) -> "TensorBatch[DictType]": if value is not None: if isinstance(value, TensorList): cat_data[key] = TensorList.cat([shard[key] for shard in shards]) + elif isinstance(value, PackedTensor): + cat_data[key] = PackedTensor.cat([shard[key] for shard in shards]) elif isinstance(value, torch.Tensor): cat_data[key] = torch.cat([shard[key] for shard in shards]) else: @@ -483,7 +506,8 @@ class TrainingInput(TypedDict, total=False): kl: Float[torch.Tensor, "batch_size response_len"] # per-token KL, current vs reference policy rewards: Optional[Float[torch.Tensor, "batch_size response_len"]] # env reward, typically only on the last token rollout_logprobs: Optional[Float[torch.Tensor, "batch_size response_len"]] # sampling policy; off-policy corr. - rollout_expert_indices: Optional[Integer[torch.Tensor, "batch_size seq_len layer_num topk"]] # MoE router replay + # MoE router replay, packed to real tokens: values [sum(seq_len_i), layer_num, topk] + cu_seqlens + rollout_expert_indices: Optional[PackedTensor] router_padding_mask: Optional[Bool[torch.Tensor, "batch_size seq_len"]] # True = no captured route (skip in replay) 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 @@ -534,13 +558,9 @@ def pad_training_input_batch(unpadded_batch: TrainingInputBatch, pad_size: int) 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) + # Every other field copies row 0 into the padding rows, so each padded row holds + # as many real tokens as row 0 and needs a route segment of that length. + new_tensors[key] = append_packed_replay_padding(tensor, segment_lengths=[len(tensor.segment(0))] * pad_size) elif key == "router_padding_mask": additional_dims = tensor.shape[1:] padding_tensor = torch.ones(pad_size, *additional_dims, dtype=torch.bool, device=tensor.device) diff --git a/skyrl/backends/skyrl_train/utils/packed_tensor.py b/skyrl/backends/skyrl_train/utils/packed_tensor.py new file mode 100644 index 0000000000..c5a0c20983 --- /dev/null +++ b/skyrl/backends/skyrl_train/utils/packed_tensor.py @@ -0,0 +1,171 @@ +"""One packed buffer plus segment offsets for ragged token-aligned batch fields.""" + +from collections.abc import Sequence + +import torch + +# Megatron's own cu_seqlens dtype; a packed global batch stays far inside int32. +CU_SEQLENS_DTYPE = torch.int32 + + +def cu_seqlens_from_lengths( + sequence_lengths: Sequence[int] | torch.Tensor, + *, + device: torch.device | str | int | None = None, +) -> torch.Tensor: + """Return the ``[batch + 1]`` exclusive prefix sum of ``sequence_lengths``.""" + lengths = torch.as_tensor(sequence_lengths, dtype=CU_SEQLENS_DTYPE, device=device) + if lengths.ndim != 1: + raise ValueError(f"sequence lengths must be 1-D, got shape {lengths.shape}") + if lengths.numel() and int(lengths.min()) < 0: + raise ValueError(f"sequence lengths must be non-negative, got {lengths.tolist()}") + offsets = torch.zeros(lengths.numel() + 1, dtype=CU_SEQLENS_DTYPE, device=lengths.device) + # torch.cumsum promotes to int64; accumulate into the target dtype instead. + torch.cumsum(lengths, dim=0, out=offsets[1:]) + return offsets + + +def lengths_from_offsets(cu_seqlens: torch.Tensor) -> torch.Tensor: + """Return the ``[batch]`` segment lengths that ``cu_seqlens`` encodes.""" + return cu_seqlens[1:] - cu_seqlens[:-1] + + +def row_index_from_offsets( + starts: torch.Tensor, + lengths: torch.Tensor, +) -> torch.Tensor: + """Return row indices that lay the requested segments back to back.""" + starts = starts.to(torch.long) + lengths = lengths.to(torch.long) + total_rows = int(lengths.sum()) + # output_size lets repeat_interleave skip its own device-side sum of `lengths`. + destination_starts = torch.repeat_interleave( + cu_seqlens_from_lengths(lengths, device=lengths.device)[:-1].to(torch.long), + lengths, + output_size=total_rows, + ) + within_segment = torch.arange(total_rows, device=lengths.device) - destination_starts + return torch.repeat_interleave(starts, lengths, output_size=total_rows) + within_segment + + +class PackedTensor: + """A ragged batch of token-aligned rows held as one buffer plus ``cu_seqlens``. + + ``values`` is ``[sum(sequence_lengths), *row_shape]`` in canonical batch order and + ``cu_seqlens`` is the ``[batch + 1]`` exclusive prefix sum of the segment lengths. + Indexing and batch operations address segments rather than individual rows. + """ + + def __init__(self, values: torch.Tensor, cu_seqlens: torch.Tensor): + if values.ndim < 1: + raise ValueError("packed values must have a token-row dimension") + if cu_seqlens.ndim != 1 or cu_seqlens.numel() < 2: + raise ValueError(f"cu_seqlens must hold at least two offsets, got shape {cu_seqlens.shape}") + if cu_seqlens.dtype != CU_SEQLENS_DTYPE: + raise ValueError(f"cu_seqlens must be {CU_SEQLENS_DTYPE}, got {cu_seqlens.dtype}") + if cu_seqlens.device != values.device: + raise ValueError( + f"packed values and cu_seqlens must share a device, got {values.device} and {cu_seqlens.device}" + ) + if int(cu_seqlens[0]) != 0 or int(cu_seqlens[-1]) != values.shape[0]: + raise ValueError( + f"cu_seqlens must run from 0 to the {values.shape[0]} packed rows, " + f"got {int(cu_seqlens[0])} to {int(cu_seqlens[-1])}" + ) + self.values = values + self.cu_seqlens = cu_seqlens + + @classmethod + def from_segments(cls, segments: Sequence[torch.Tensor]) -> "PackedTensor": + """Concatenate per-batch-entry row blocks into one packed buffer.""" + if not segments: + raise ValueError("cannot pack an empty list of segments") + cu_seqlens = cu_seqlens_from_lengths([segment.shape[0] for segment in segments], device=segments[0].device) + return cls(torch.cat(segments, dim=0), cu_seqlens) + + @property + def sequence_lengths(self) -> torch.Tensor: + return lengths_from_offsets(self.cu_seqlens) + + @property + def row_shape(self) -> torch.Size: + return self.values.shape[1:] + + @property + def device(self) -> torch.device: + return self.values.device + + @property + def dtype(self) -> torch.dtype: + return self.values.dtype + + def __len__(self) -> int: + return self.cu_seqlens.numel() - 1 + + def __getitem__(self, index) -> "torch.Tensor | PackedTensor": + if isinstance(index, slice): + if index.step in (None, 1): + start, stop, _ = index.indices(len(self)) + stop = max(start, stop) + offsets = self.cu_seqlens[start : stop + 1] + return PackedTensor(self.values[int(offsets[0]) : int(offsets[-1])], offsets - offsets[0]) + return self._gather(range(*index.indices(len(self)))) + if isinstance(index, torch.Tensor): + if index.ndim == 0: + return self.segment(int(index)) + return self._gather(index.tolist()) + if isinstance(index, (list, tuple, range)): + return self._gather(index) + return self.segment(index) + + def segment(self, index: int) -> torch.Tensor: + """Return one batch entry's row block as a view.""" + position = index + len(self) if index < 0 else index + if not 0 <= position < len(self): + raise IndexError(f"segment {index} is out of range for a packed batch of {len(self)}") + return self.values[int(self.cu_seqlens[position]) : int(self.cu_seqlens[position + 1])] + + def _gather(self, indices: Sequence[int]) -> "PackedTensor": + """Select segments in the requested order into a freshly allocated buffer.""" + selected = torch.as_tensor(list(indices), dtype=torch.long, device=self.values.device) + selected_starts = self.cu_seqlens[:-1].to(torch.long)[selected] + selected_lengths = self.sequence_lengths.to(torch.long)[selected] + row_index = row_index_from_offsets(selected_starts, selected_lengths) + cu_seqlens = cu_seqlens_from_lengths(selected_lengths, device=self.values.device) + return PackedTensor(self.values.index_select(0, row_index), cu_seqlens) + + def to(self, device=None, dtype=None, non_blocking: bool = False) -> "PackedTensor": + return PackedTensor( + self.values.to(device=device, dtype=dtype, non_blocking=non_blocking), + self.cu_seqlens.to(device=device, non_blocking=non_blocking), + ) + + def contiguous(self) -> "PackedTensor": + return PackedTensor(self.values.contiguous(), self.cu_seqlens.contiguous()) + + def pin_memory(self) -> "PackedTensor": + return PackedTensor(self.values.pin_memory(), self.cu_seqlens.pin_memory()) + + def repeat(self, repeats: int) -> "PackedTensor": + return self._gather(list(range(len(self))) * repeats) + + def repeat_interleave(self, repeats: int) -> "PackedTensor": + return self._gather([index for index in range(len(self)) for _ in range(repeats)]) + + def __eq__(self, other: object) -> bool: + if not isinstance(other, PackedTensor): + return False + return torch.equal(self.values, other.values) and torch.equal(self.cu_seqlens, other.cu_seqlens) + + def __repr__(self) -> str: + return f"PackedTensor(batch={len(self)}, values={tuple(self.values.shape)}, dtype={self.values.dtype})" + + @staticmethod + def cat(batches: Sequence["PackedTensor"]) -> "PackedTensor": + if not batches: + raise ValueError("cannot cat an empty list of packed batches") + lengths = torch.cat([batch.sequence_lengths for batch in batches]) + return PackedTensor( + torch.cat([batch.values for batch in batches], dim=0), + cu_seqlens_from_lengths(lengths, device=batches[0].device), + ) diff --git a/skyrl/backends/skyrl_train/utils/replay_utils.py b/skyrl/backends/skyrl_train/utils/replay_utils.py index 400376c65a..3e22c7afd7 100644 --- a/skyrl/backends/skyrl_train/utils/replay_utils.py +++ b/skyrl/backends/skyrl_train/utils/replay_utils.py @@ -2,18 +2,23 @@ Utility functions for MoE Router Replay. """ +from collections.abc import Sequence from contextlib import contextmanager -import numpy as np import torch from skyrl.backends.skyrl_train.distributed.megatron.token_metadata import ( TokenMetadataLayout, + align_packed_token_metadata, align_token_metadata, ) +from skyrl.backends.skyrl_train.utils.packed_tensor import ( + PackedTensor, + cu_seqlens_from_lengths, +) -def _replay_padding_row( +def replay_padding_row( topk: int, *, dtype: torch.dtype, @@ -40,25 +45,35 @@ def make_replay_padding_indices( """Return dummy routes with ``topk`` distinct experts in every row.""" if not shape: raise ValueError(f"Replay route padding requires a positive topk dimension, got {shape}") - padding_row = _replay_padding_row(shape[-1], dtype=dtype, device=device) + padding_row = replay_padding_row(shape[-1], dtype=dtype, device=device) return padding_row.expand(shape).clone() -def make_replay_padding_indices_np(shape: tuple[int, ...], *, dtype: np.dtype) -> np.ndarray: - """NumPy sibling of :func:`make_replay_padding_indices`. +def make_packed_replay_padding( + reference: PackedTensor, + *, + segment_lengths: Sequence[int], +) -> PackedTensor: + """Return dummy-route segments matching ``reference``'s row shape and dtype. - Preprocessing builds the padded route array in NumPy before handing it to - ``torch.from_numpy``, so it needs the same distinct-expert padding rows - without a round trip through torch. + Batch padding rows exist only to give Megatron a uniform micro-batch size; their + tokens are loss-masked, so one dummy route per token is all they need. """ - if not shape: - raise ValueError(f"Replay route padding requires a positive topk dimension, got {shape}") - topk = shape[-1] - if topk < 1: - raise ValueError(f"Replay route padding requires a positive topk dimension, got {topk}") - padded = np.empty(shape, dtype=dtype) - padded[...] = np.arange(topk, dtype=dtype) - return padded + padding = make_replay_padding_indices( + (sum(segment_lengths), *reference.row_shape), + dtype=reference.dtype, + device=reference.device, + ) + return PackedTensor(padding, cu_seqlens_from_lengths(segment_lengths, device=reference.device)) + + +def append_packed_replay_padding( + routes: PackedTensor, + *, + segment_lengths: Sequence[int], +) -> PackedTensor: + """Extend ``routes`` with one dummy-route segment per batch padding row.""" + return PackedTensor.cat([routes, make_packed_replay_padding(routes, segment_lengths=segment_lengths)]) def patch_topk_router_layer_number(): @@ -187,7 +202,7 @@ def _get_local_router_layer_indices(model_config, global_num_layers: int, instan def setup_per_microbatch_replay_forward( - rollout_expert_indices: torch.Tensor, + rollout_expert_indices: PackedTensor, router_padding_mask: torch.Tensor | None, attention_mask: torch.Tensor, model, @@ -197,8 +212,9 @@ def setup_per_microbatch_replay_forward( ) -> 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 + Replay indices arrive packed to their real tokens (``[sum(seqlen), layers, topk]`` plus + ``cu_seqlens``) while the router padding mask arrives batch-major; both undergo matching + padding removal or packing and CP sharding against the shared layout. 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``. @@ -232,8 +248,8 @@ def setup_per_microbatch_replay_forward( if router_padding_mask is None: raise ValueError("router_padding_mask is required with rollout_expert_indices") - if rollout_expert_indices.dim() != 4: - raise ValueError(f"Expected 4D replay indices, got shape {rollout_expert_indices.shape}") + if len(rollout_expert_indices.row_shape) != 2: + raise ValueError(f"Expected [tokens, layers, topk] replay indices, got {rollout_expert_indices!r}") if router_padding_mask.shape != attention_mask.shape: raise ValueError( @@ -243,24 +259,28 @@ def setup_per_microbatch_replay_forward( if router_padding_mask.device != rollout_expert_indices.device: raise ValueError("rollout_expert_indices and router_padding_mask must be on the same device") + num_captured_layers, topk = rollout_expert_indices.row_shape instances = RouterReplay.global_router_replay_instances local_layer_indices = _get_local_router_layer_indices( model_config, - rollout_expert_indices.shape[2], + num_captured_layers, instances, ) layer_index = torch.tensor(local_layer_indices, dtype=torch.long, device=rollout_expert_indices.device) - local_rollout_expert_indices = rollout_expert_indices.index_select(2, layer_index) + local_rollout_expert_indices = PackedTensor( + rollout_expert_indices.values.index_select(1, layer_index), + rollout_expert_indices.cu_seqlens, + ) 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 = _replay_padding_row( - rollout_expert_indices.shape[-1], + route_padding = replay_padding_row( + topk, dtype=rollout_expert_indices.dtype, - device=local_rollout_expert_indices.device, + device=rollout_expert_indices.device, ) - aligned_rollout_expert_indices = align_token_metadata( + aligned_rollout_expert_indices = align_packed_token_metadata( local_rollout_expert_indices, metadata_layout, route_padding, diff --git a/skyrl/backends/skyrl_train/utils/routed_experts.py b/skyrl/backends/skyrl_train/utils/routed_experts.py index b3caec7c8b..a35bcc543a 100644 --- a/skyrl/backends/skyrl_train/utils/routed_experts.py +++ b/skyrl/backends/skyrl_train/utils/routed_experts.py @@ -1,11 +1,68 @@ +from collections.abc import Sequence from typing import TypeAlias import numpy as np +from skyrl.backends.skyrl_train.distributed.megatron.token_metadata import ( + TokenMetadataTrace, +) + RoutedExpertIndices: TypeAlias = np.ndarray ROUTED_EXPERT_DTYPES = frozenset({np.dtype(np.uint8), np.dtype(np.int16), np.dtype(np.int32)}) +class RoutedExpertTrace: + """Accumulate routed experts across incremental generation calls.""" + + def __init__(self) -> None: + self._metadata = TokenMetadataTrace() + self._schema: tuple[int, int, np.dtype] | None = None + + @property + def prompt_start(self) -> int: + return self._metadata.num_rows + + def record_generation( + self, + *, + prompt_token_count: int, + generated_token_count: int, + routed_experts: RoutedExpertIndices, + ) -> None: + if prompt_token_count < self.prompt_start: + raise ValueError("routed-expert prompt start exceeds prompt length") + if generated_token_count < 1: + raise ValueError("routed-expert generation must produce at least one token") + + expected_rows = prompt_token_count - self.prompt_start + generated_token_count - 1 + compact = compact_routed_expert_indices(routed_experts) + if self._schema is None: + self._schema = (*compact.shape[1:], compact.dtype) + self._metadata.append(compact, expected_rows=expected_rows) + + def finalize(self, *, token_count: int, loss_mask: Sequence[int]) -> RoutedExpertIndices: + if len(loss_mask) != token_count: + raise ValueError(f"loss mask has {len(loss_mask)} entries, expected {token_count}") + if self.prompt_start > token_count: + raise ValueError(f"routed-expert trace has {self.prompt_start} rows for {token_count} tokens") + + if any(loss_mask[self.prompt_start + 1 : token_count]): + for source_index in range(self.prompt_start, token_count - 1): + if loss_mask[source_index + 1] != 0: + raise ValueError(f"missing routed-expert row for loss-active target at token {source_index + 1}") + + padding_count = token_count - self.prompt_start + if padding_count: + if self._schema is None: + raise ValueError("cannot pad routed-expert trace before any routes are captured") + num_layers, topk, dtype = self._schema + padding_row = np.arange(topk, dtype=dtype) + padding = np.broadcast_to(padding_row, (padding_count, num_layers, topk)).copy() + self._metadata.append(padding, expected_rows=padding_count) + + return self._metadata.finalize(expected_rows=token_count) + + 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): diff --git a/skyrl/backends/skyrl_train/weight_sync/delta_checkpoint.py b/skyrl/backends/skyrl_train/weight_sync/delta_checkpoint.py index 41c3a55a64..3a36dbf876 100644 --- a/skyrl/backends/skyrl_train/weight_sync/delta_checkpoint.py +++ b/skyrl/backends/skyrl_train/weight_sync/delta_checkpoint.py @@ -32,6 +32,7 @@ decompress_bytes, uint8_tensor_to_bytes, ) +from skyrl.utils.cpu_topology import pool_workers logger = logging.getLogger(__name__) @@ -760,7 +761,8 @@ def apply_one(item: tuple[DeltaTensorRecord, str, bytes]) -> None: mismatches.append(record.name) del region, patch - workers = min(len(payloads), max(1, min(32, os.cpu_count() or 8))) + # Weight sync is a barrier, so it need not reserve cores for colocated work. + workers = min(len(payloads), pool_workers(cap=32, reserved=0)) with ThreadPoolExecutor(max_workers=workers, thread_name_prefix="skyrl-delta-mmap-apply") as executor: list(executor.map(apply_one, payloads)) finally: @@ -1077,7 +1079,8 @@ def _empty_stats() -> dict[str, float]: } def _num_publish_workers(self) -> int: - default = min(8, os.cpu_count() or 1) + # Publishing is a barrier, like the apply path above. + default = pool_workers(cap=8, reserved=0) return self.publish_num_workers or default def _publish_executor_for(self, num_workers: int) -> ThreadPoolExecutor: 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 e66c42841e..0a70b4a45e 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py @@ -42,6 +42,7 @@ unpadded_vocab_shard_width, ) from skyrl.backends.skyrl_train.training_batch import TensorList +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor from skyrl.backends.skyrl_train.utils.ppo_utils import ( PolicyLossRegistry, compute_approx_kl, @@ -132,7 +133,7 @@ def _build_packed_valid_mask( def _copy_tensor_tree_to_device(value: Any, device: int) -> Any: """Move all tensors in a nested microbatch to a CUDA device.""" - if torch.is_tensor(value) or isinstance(value, TensorList): + if torch.is_tensor(value) or isinstance(value, (TensorList, PackedTensor)): return value.to(device=device, non_blocking=True) if isinstance(value, dict): return {key: _copy_tensor_tree_to_device(item, device) for key, item in value.items()} diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py index d116213abd..3b99902d56 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py @@ -46,8 +46,9 @@ TrainingInputBatch, TrainingOutputBatch, ) +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor from skyrl.backends.skyrl_train.utils.profiler import build_profiler_from_policy_cfg -from skyrl.backends.skyrl_train.utils.replay_utils import make_replay_padding_indices +from skyrl.backends.skyrl_train.utils.replay_utils import append_packed_replay_padding from skyrl.backends.skyrl_train.weight_sync import ( LoraLoadRequest, WeightChunk, @@ -493,6 +494,15 @@ def init_configs( for k, v in transformer_config_kwargs.items(): setattr(provider, k, v) + # Check the resolved provider because it may supply its own VPP default. Interleaved + # chunks desynchronise each RouterReplay instance's backward FIFO. + vpp_size = provider.virtual_pipeline_model_parallel_size + if provider.moe_enable_routing_replay and vpp_size is not None and vpp_size > 1: + raise ValueError( + f"moe_enable_routing_replay is incompatible with virtual_pipeline_model_parallel_size={vpp_size}: " + "interleaved chunks desync the replay FIFO. Unset virtual_pipeline_model_parallel_size." + ) + # MTP head count: megatron-bridge infers provider.mtp_num_layers from the model's HF config. if not enable_mtp: provider.mtp_num_layers = None @@ -774,6 +784,13 @@ def _pad_microbatch_to_size(self, micro_dict: dict, target_batch_size: int) -> d if value is None: padded[key] = None continue + if key == "rollout_expert_indices": + # The dummy attention_mask row below marks one valid token, so each padded + # row's route segment holds one row. + padded[key] = append_packed_replay_padding(value, segment_lengths=[1] * pad_count) + continue + if isinstance(value, PackedTensor): + raise ValueError(f"Micro-batch field {key!r} is packed and has no padding rule") if isinstance(value, torch.Tensor): if key == "loss_mask": # Pad with zeros so padded samples don't contribute to loss @@ -791,12 +808,6 @@ def _pad_microbatch_to_size(self, micro_dict: dict, target_batch_size: int) -> d 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 == "response_mask": # response_mask should be zeros for padded samples pad_tensor = torch.zeros((pad_count, *value.shape[1:]), dtype=value.dtype, device=device) diff --git a/skyrl/backends/skyrl_train/workers/worker_utils.py b/skyrl/backends/skyrl_train/workers/worker_utils.py index 2efa37ad6c..d3f04b5e79 100644 --- a/skyrl/backends/skyrl_train/workers/worker_utils.py +++ b/skyrl/backends/skyrl_train/workers/worker_utils.py @@ -6,7 +6,7 @@ from skyrl.backends.skyrl_train.distributed.strategy import DistributedStrategy from skyrl.backends.skyrl_train.training_batch import TensorBatch, TrainingInputBatch -from skyrl.backends.skyrl_train.utils.replay_utils import make_replay_padding_indices +from skyrl.backends.skyrl_train.utils.replay_utils import make_packed_replay_padding 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 @@ -326,11 +326,10 @@ def _create_padding_microbatch(self) -> TrainingInputBatch: ref_tensor = self.data["rollout_logprobs"] 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"] = make_replay_padding_indices( - (batch_size, *ref_tensor.shape[1:]), - dtype=ref_tensor.dtype, - device=device, + # The dummy attention_mask row marks one valid token, so its route segment holds one row. + data["rollout_expert_indices"] = make_packed_replay_padding( + self.data["rollout_expert_indices"], + segment_lengths=[1] * batch_size, ) 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) diff --git a/skyrl/benchmarks/bench_packed_route_collation.py b/skyrl/benchmarks/bench_packed_route_collation.py new file mode 100644 index 0000000000..668a9d157d --- /dev/null +++ b/skyrl/benchmarks/bench_packed_route_collation.py @@ -0,0 +1,210 @@ +"""Compare padded and packed route collation with serial and pooled fills. + +Production shapes need ~110 GiB of host RAM, so drive this from a cluster harness:: + + uv run --isolated --extra skyrl-train python -m \ + skyrl.benchmarks.bench_packed_route_collation --num-moe-layers 40 + +Scaled-down smoke:: + + uv run --isolated --extra skyrl-train python -m \ + skyrl.benchmarks.bench_packed_route_collation \ + --num-sequences 64 --max-seqlen 4096 --num-moe-layers 4 --iterations 1 +""" + +import argparse +import functools +import gc +import os +import resource +import statistics +import time + +import numpy as np +import torch + +from skyrl.backends.skyrl_train.utils.packed_tensor import cu_seqlens_from_lengths +from skyrl.backends.skyrl_train.utils.replay_utils import replay_padding_row +from skyrl.train.dataset.parallel_fill import default_fill_workers, fill_batch_rows + +DEFAULT_NUM_MOE_LAYERS = 40 +DEFAULT_TOPK = 22 +DEFAULT_NUM_SEQUENCES = 1024 +DEFAULT_MAX_SEQLEN = 32768 +ROUTE_DTYPE = torch.int16 + +# Fraction of ``max_seqlen`` each distribution draws its shortest sequence from. "uniform" +# has no padding at all; "typical_rl" is the measured production spread. +LENGTH_DISTRIBUTIONS = { + "uniform": 1.0, + "mild_ragged": 0.5, + "typical_rl": 1 / 16, + "heavy_tail": 1 / 64, +} + + +def _sequence_lengths(distribution: str, num_sequences: int, max_seqlen: int, seed: int) -> np.ndarray: + """Draw per-trajectory total lengths, always including one full-length sequence.""" + minimum = max(1, round(max_seqlen * LENGTH_DISTRIBUTIONS[distribution])) + if minimum >= max_seqlen: + return np.full(num_sequences, max_seqlen, dtype=np.int64) + rng = np.random.default_rng(seed) + lengths = rng.integers(minimum, max_seqlen + 1, size=num_sequences).astype(np.int64) + # max_total is set by the longest trajectory, so pin one to the cap for a stable rectangle. + lengths[0] = max_seqlen + return lengths + + +def _make_trajectories(lengths: np.ndarray, num_layers: int, topk: int, seed: int) -> list[np.ndarray]: + """One route array per trajectory, sized to its full sequence length. + + Arrays are views over one template so source allocation does not dominate the benchmark. + """ + rng = np.random.default_rng(seed + 1) + template = rng.integers(0, 128, size=(int(lengths.max()), num_layers, topk), dtype=np.int16) + return [template[: int(length)] for length in lengths] + + +def _write_padded_row( + padded: torch.Tensor, + trajectories: list[np.ndarray], + lengths: np.ndarray, + sample_index: int, +) -> None: + """One trajectory's slot in the rectangle: left dummy rows, routes, trailing dummy rows.""" + padding_row = replay_padding_row(padded.shape[-1], dtype=padded.dtype) + sample_indices = trajectories[sample_index] + left_pad = padded.shape[1] - int(lengths[sample_index]) + route_end = left_pad + sample_indices.shape[0] + padded[sample_index, :left_pad] = padding_row + padded[sample_index, left_pad:route_end] = torch.from_numpy(sample_indices) + padded[sample_index, route_end:] = padding_row + + +def _write_packed_segment( + packed: torch.Tensor, + cu_seqlens: torch.Tensor, + trajectories: list[np.ndarray], + sample_index: int, +) -> None: + """One trajectory's segment of the packed buffer: routes, then any trailing dummy rows.""" + sample_indices = trajectories[sample_index] + segment = packed[int(cu_seqlens[sample_index]) : int(cu_seqlens[sample_index + 1])] + captured = sample_indices.shape[0] + segment[:captured] = torch.from_numpy(sample_indices) + segment[captured:] = replay_padding_row(segment.shape[-1], dtype=packed.dtype) + + +def _make_fill(packed: bool, workers: int): + """Build a fill callable over the trainer's own pool helper.""" + + def fill(buffer: torch.Tensor, trajectories: list[np.ndarray], lengths: np.ndarray) -> None: + if packed: + cu_seqlens = cu_seqlens_from_lengths(lengths) + write = functools.partial(_write_packed_segment, buffer, cu_seqlens, trajectories) + else: + write = functools.partial(_write_padded_row, buffer, trajectories, lengths) + fill_batch_rows(write, len(trajectories), workers=workers) + + return fill + + +def _padded_shape(lengths: np.ndarray, num_layers: int, topk: int) -> tuple[int, ...]: + return (len(lengths), int(lengths.max()), num_layers, topk) + + +def _packed_shape(lengths: np.ndarray, num_layers: int, topk: int) -> tuple[int, ...]: + return (int(lengths.sum()), num_layers, topk) + + +def _peak_rss_bytes() -> int: + """Process high-water RSS. Monotone, so only ever read as a whole-run ceiling.""" + return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * 1024 + + +def _time_cold(fill, shape, trajectories, lengths, iterations: int) -> tuple[float, int]: + """Median wall clock over a freshly allocated buffer each iteration. + + Fresh buffers retain the first-touch allocation cost measured in production. + """ + durations = [] + for _ in range(iterations): + gc.collect() + start = time.perf_counter() + buffer = torch.empty(shape, dtype=ROUTE_DTYPE) + fill(buffer, trajectories, lengths) + durations.append(time.perf_counter() - start) + buffer_bytes = buffer.numel() * buffer.element_size() + del buffer + return statistics.median(durations), buffer_bytes + + +def _format_gib(num_bytes: int) -> str: + return f"{num_bytes / 1024**3:8.2f}" + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--num-sequences", type=int, default=DEFAULT_NUM_SEQUENCES) + parser.add_argument("--max-seqlen", type=int, default=DEFAULT_MAX_SEQLEN) + parser.add_argument("--num-moe-layers", type=int, default=DEFAULT_NUM_MOE_LAYERS) + parser.add_argument("--topk", type=int, default=DEFAULT_TOPK) + parser.add_argument("--iterations", type=int, default=3) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--workers", type=int, default=default_fill_workers()) + parser.add_argument( + "--distributions", + nargs="+", + default=list(LENGTH_DISTRIBUTIONS), + choices=list(LENGTH_DISTRIBUTIONS), + ) + args = parser.parse_args() + + print( + f"num_sequences={args.num_sequences} max_seqlen={args.max_seqlen} " + f"moe_layers={args.num_moe_layers} topk={args.topk} dtype={ROUTE_DTYPE} " + f"iterations={args.iterations}" + ) + print(f"OMP_NUM_THREADS={os.environ.get('OMP_NUM_THREADS', '')} torch_threads={torch.get_num_threads()}") + print(f"pooled arms use {args.workers} workers (the trainer's autoscaled pool size)") + header = ( + f"{'distribution':<20} {'pad_1t':>8} {'pad_pool':>9} {'pack_1t':>8} {'pack_pool':>10} " + f"{'best_pad':>9} {'vs_best':>8} {'pad_GiB':>9} {'packed_GiB':>11} {'saved':>7}" + ) + print(header) + print("-" * len(header)) + + for distribution in args.distributions: + lengths = _sequence_lengths(distribution, args.num_sequences, args.max_seqlen, args.seed) + trajectories = _make_trajectories(lengths, args.num_moe_layers, args.topk, args.seed) + + padded_shape = _padded_shape(lengths, args.num_moe_layers, args.topk) + packed_shape = _packed_shape(lengths, args.num_moe_layers, args.topk) + arms = { + "pad_1t": (_make_fill(False, 1), padded_shape), + "pad_pool": (_make_fill(False, args.workers), padded_shape), + "pack_1t": (_make_fill(True, 1), packed_shape), + "pack_pool": (_make_fill(True, args.workers), packed_shape), + } + timings = {} + buffer_bytes = {} + for name, (fill, shape) in arms.items(): + timings[name], buffer_bytes[name] = _time_cold(fill, shape, trajectories, lengths, args.iterations) + del trajectories + gc.collect() + + best_padded = min(timings["pad_1t"], timings["pad_pool"]) + best_packed = min(timings["pack_1t"], timings["pack_pool"]) + print( + f"{distribution:<20} {timings['pad_1t'] * 1000:8.1f} {timings['pad_pool'] * 1000:9.1f} " + f"{timings['pack_1t'] * 1000:8.1f} {timings['pack_pool'] * 1000:10.1f} " + f"{best_padded * 1000:9.1f} {best_padded / best_packed:7.2f}x " + f"{_format_gib(buffer_bytes['pad_1t'])} {_format_gib(buffer_bytes['pack_1t']):>11} " + f"{1 - buffer_bytes['pack_1t'] / buffer_bytes['pad_1t']:6.1%}" + ) + + print(f"process peak RSS: {_format_gib(_peak_rss_bytes()).strip()} GiB") + + +if __name__ == "__main__": + main() diff --git a/skyrl/train/dataset/parallel_fill.py b/skyrl/train/dataset/parallel_fill.py new file mode 100644 index 0000000000..813884604d --- /dev/null +++ b/skyrl/train/dataset/parallel_fill.py @@ -0,0 +1,51 @@ +"""Fill a controller-side batch buffer with a locally sized thread pool. + +The pool parallelises first-touch page faults without changing process-wide torch settings. +Callbacks own disjoint row ranges, and their copies release the GIL. +""" + +import functools +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor + +from skyrl.utils.cpu_topology import pool_workers + +# Extra threads beyond this cap take cores from colocated actors without improving throughput. +MAX_FILL_WORKERS = 32 +# Leave room for Ray services in the same cgroup. +RESERVED_FILL_CORES = 8 + + +@functools.cache +def default_fill_workers() -> int: + return pool_workers(cap=MAX_FILL_WORKERS, reserved=RESERVED_FILL_CORES) + + +def fill_batch_rows( + fill_row: Callable[[int], None], + num_rows: int, + *, + workers: int | None = None, +) -> None: + """Call ``fill_row`` for every row, possibly in parallel. + + Each callback must write to a disjoint row range. + """ + if num_rows < 0: + raise ValueError(f"row count must be non-negative, got {num_rows}") + if num_rows == 0: + return + if workers is None: + workers = default_fill_workers() + if workers < 1: + raise ValueError(f"worker count must be positive, got {workers}") + + workers = min(workers, num_rows) + if workers == 1: + for index in range(num_rows): + fill_row(index) + return + + with ThreadPoolExecutor(max_workers=workers, thread_name_prefix="skyrl-batch-fill") as pool: + # Eagerly consume the map so worker exceptions surface here. + list(pool.map(fill_row, range(num_rows))) diff --git a/skyrl/train/dataset/preprocess.py b/skyrl/train/dataset/preprocess.py index fb7ac6e2f7..23aa0a123c 100644 --- a/skyrl/train/dataset/preprocess.py +++ b/skyrl/train/dataset/preprocess.py @@ -1,18 +1,31 @@ +import functools import logging -from typing import List, Optional, Tuple, Union +from typing import Dict, List, Optional, Tuple, Union import numpy as np import torch -from jaxtyping import Bool, Float, Integer +from jaxtyping import Bool, Float -from skyrl.backends.skyrl_train.utils.replay_utils import make_replay_padding_indices_np +from skyrl.backends.skyrl_train.utils.packed_tensor import ( + PackedTensor, + cu_seqlens_from_lengths, +) +from skyrl.backends.skyrl_train.utils.replay_utils import replay_padding_row from skyrl.backends.skyrl_train.utils.routed_experts import ( + ROUTED_EXPERT_DTYPES, RoutedExpertIndices, - compact_routed_expert_indices, ) +from skyrl.train.dataset.parallel_fill import fill_batch_rows logger = logging.getLogger(__name__) +# Torch counterparts of the canonical routed-expert dtypes. +ROUTED_EXPERT_TORCH_DTYPES: Dict[np.dtype, torch.dtype] = { + np.dtype(np.uint8): torch.uint8, + np.dtype(np.int16): torch.int16, + np.dtype(np.int32): torch.int32, +} + def make_router_padding_mask( attention_mask: torch.Tensor, @@ -87,6 +100,86 @@ def _reward_to_numpy(custom_reward: Union[List[float], torch.Tensor]) -> np.ndar return reward_arr +def _fill_routed_expert_segment( + packed: torch.Tensor, + cu_seqlens: torch.Tensor, + rollout_expert_indices: List[RoutedExpertIndices], + sample_index: int, +) -> None: + """Write one route segment, using distinct dummy routes for uncaptured trailing tokens.""" + sample_indices = rollout_expert_indices[sample_index] + flags = sample_indices.flags + # torch.from_numpy refuses a non-writeable buffer, and decoded wire routes may be read-only. + if not flags.c_contiguous or not flags.writeable: + sample_indices = sample_indices.copy(order="C") + segment = packed[int(cu_seqlens[sample_index]) : int(cu_seqlens[sample_index + 1])] + captured = sample_indices.shape[0] + segment[:captured] = torch.from_numpy(sample_indices) + segment[captured:] = replay_padding_row(segment.shape[-1], dtype=packed.dtype) + + +def _collate_rollout_expert_indices( + rollout_expert_indices: List[RoutedExpertIndices], + total_real: np.ndarray, +) -> PackedTensor: + """Pack per-trajectory routes into one ``[sum(seq_len_i), layers, topk]`` buffer. + + Entries already have a canonical dtype and are filled from the trainer's local thread pool. + """ + num_samples = len(rollout_expert_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: + supported = ", ".join( + dtype.name for dtype in sorted(ROUTED_EXPERT_DTYPES, key=lambda dtype: dtype.itemsize) + ) + raise ValueError( + f"rollout_expert_indices entries must use a canonical routed-expert dtype ({supported}), " + f"got {sample_indices.dtype} at sample {sample_index}" + ) + + first_shape = rollout_expert_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") + + # Validate serially so an invalid trajectory raises deterministically rather than from a worker. + for sample_index, sample_indices in enumerate(rollout_expert_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 {sample_indices.shape} at sample {sample_index}" + ) + available = int(total_real[sample_index]) + if sample_indices.shape[0] == 0 or sample_indices.shape[0] > available: + raise ValueError( + f"Trajectory {sample_index} has {sample_indices.shape[0]} route rows for {available} tokens" + ) + + batch_dtype = max((indices.dtype for indices in rollout_expert_indices), key=lambda dtype: dtype.itemsize) + if batch_dtype == np.dtype(np.int32): + logger.warning( + "Collating rollout_expert_indices as int32, which doubles this buffer. No supported expert count " + "needs more than int16, so the inference server is not compacting its routes." + ) + cu_seqlens = cu_seqlens_from_lengths(total_real) + packed = torch.empty( + (int(total_real.sum()), num_layers, topk), + dtype=ROUTED_EXPERT_TORCH_DTYPES[batch_dtype], + ) + fill_batch_rows( + functools.partial(_fill_routed_expert_segment, packed, cu_seqlens, rollout_expert_indices), + num_samples, + ) + return PackedTensor(packed, cu_seqlens) + + def convert_prompts_responses_to_batch_tensors( pad_token_id: int, prompts: List[List[int]], @@ -103,7 +196,7 @@ def convert_prompts_responses_to_batch_tensors( Float[torch.Tensor, "batch response_len"], Float[torch.Tensor, "batch response_len"], Optional[Float[torch.Tensor, "batch response_len"]], - Optional[Integer[torch.Tensor, "batch seq_len layer_num topk"]], + Optional[PackedTensor], ]: """ Convert prompts and responses to batch tensors for training. @@ -160,6 +253,9 @@ def convert_prompts_responses_to_batch_tensors( rewards: ``(batch, max_response)`` — right-aligned. loss_masks: ``(batch, max_response)`` — right-aligned. logprobs: ``(batch, max_response)`` — right-aligned, or ``None``. + rollout_expert_indices: ``PackedTensor`` whose values are + ``(sum(prompt_i + response_i), layers, topk)`` in canonical batch order, with + ``cu_seqlens`` naming each trajectory's segment, or ``None``. """ _verify_inputs(prompts, responses, rewards, loss_masks) @@ -235,42 +331,7 @@ def convert_prompts_responses_to_batch_tensors( if len(rollout_expert_indices) != num_samples: raise ValueError("rollout_expert_indices must contain routes for every trajectory") - 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}" - ) - 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") - - batch_dtype = max((indices.dtype for indices in canonical_indices), key=lambda dtype: dtype.itemsize) - padded = make_replay_padding_indices_np( - (num_samples, max_total, num_layers, 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 {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 sample_indices.shape[0] == 0 or sample_indices.shape[0] > available: - raise ValueError( - f"Trajectory {sample_index} has {sample_indices.shape[0]} route rows for {available} tokens" - ) - 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) + rollout_expert_indices_tensor = _collate_rollout_expert_indices(rollout_expert_indices, total_real) return ( sequences, diff --git a/skyrl/train/dataset/replay_buffer.py b/skyrl/train/dataset/replay_buffer.py index 6f627e31ed..35668a310a 100644 --- a/skyrl/train/dataset/replay_buffer.py +++ b/skyrl/train/dataset/replay_buffer.py @@ -15,23 +15,24 @@ from jaxtyping import Bool, Float, Integer from skyrl.backends.skyrl_train.training_batch import TensorList +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor BasicType = Union[int, float, str, bool] -def to(tensor: Union[torch.Tensor, List[torch.Tensor], BasicType], device): +def to(tensor: Union[torch.Tensor, PackedTensor, List[torch.Tensor], BasicType], device): if isinstance(tensor, list): return [to(t, device) for t in tensor] - elif isinstance(tensor, torch.Tensor): + elif isinstance(tensor, (torch.Tensor, PackedTensor)): return tensor.to(device) else: return tensor -def pin_memory(tensor: Union[torch.Tensor, List[torch.Tensor], BasicType]): +def pin_memory(tensor: Union[torch.Tensor, PackedTensor, List[torch.Tensor], BasicType]): if isinstance(tensor, list): return [pin_memory(t) for t in tensor] - elif isinstance(tensor, torch.Tensor): + elif isinstance(tensor, (torch.Tensor, PackedTensor)): return tensor.pin_memory() else: return tensor @@ -67,7 +68,8 @@ class Experience: loss_mask: Optional[Integer[torch.LongTensor, "batch response_len"]] response_mask: Optional[Integer[torch.Tensor, "batch response_len"]] rollout_logprobs: Optional[Float[torch.Tensor, "batch response_len"]] - rollout_expert_indices: Optional[Integer[torch.Tensor, "batch seq_len layer_num topk"]] + # Routes packed to real tokens: values [sum(seq_len_i), layer_num, topk] + cu_seqlens. + rollout_expert_indices: Optional[PackedTensor] num_actions: int info: Optional[dict] router_padding_mask: Optional[Bool[torch.Tensor, "batch seq_len"]] = None diff --git a/skyrl/train/generators/skyrl_gym_generator.py b/skyrl/train/generators/skyrl_gym_generator.py index 26b1102e6c..75c5d07671 100644 --- a/skyrl/train/generators/skyrl_gym_generator.py +++ b/skyrl/train/generators/skyrl_gym_generator.py @@ -23,7 +23,10 @@ InferenceEngineInput, InferenceEngineInterface, ) -from skyrl.backends.skyrl_train.utils.routed_experts import RoutedExpertIndices +from skyrl.backends.skyrl_train.utils.routed_experts import ( + RoutedExpertIndices, + RoutedExpertTrace, +) from skyrl.train.config import GeneratorConfig, SkyRLGymConfig from skyrl.train.generators.base import ( GeneratorInput, @@ -83,7 +86,7 @@ class AgentLoopState: rollout_logprobs: Optional[List[float]] response_end_idx: Optional[int] done: bool - rollout_expert_indices: Optional[RoutedExpertIndices] = None + routed_expert_trace: Optional[RoutedExpertTrace] = None @dataclass @@ -93,14 +96,9 @@ class TurnOutput: output_logprobs: Optional[List[float]] new_obs: ConversationType obs_ids: List[int] - rollout_expert_indices: Optional[RoutedExpertIndices] reward: Optional[float] added_eos: bool = False - def get_turn_rollout_expert_indices(self) -> Optional[RoutedExpertIndices]: - """Return only routes that the inference model actually executed.""" - return self.rollout_expert_indices - def get_turn_loss_mask(self) -> List[int]: """ Get loss mask for this turn's tokens. @@ -379,6 +377,9 @@ async def agent_loop( rollout_logprobs=[] if get_logprobs else None, response_end_idx=None, done=False, + routed_expert_trace=( + RoutedExpertTrace() if self.generator_cfg.inference_engine.enable_return_routed_experts else None + ), ) while not agent_loop_state.done: @@ -401,11 +402,13 @@ async def agent_loop( agent_loop_state.loss_mask = [] agent_loop_state.rollout_logprobs = None + routed_expert_trace = agent_loop_state.routed_expert_trace engine_input = InferenceEngineInput( prompt_token_ids=[agent_loop_state.input_ids], session_ids=[session_id], sampling_params=sampling_params, cache_salt=cache_salt, + routed_experts_prompt_starts=[routed_expert_trace.prompt_start] if routed_expert_trace else None, ) llm_call_start_time = time.monotonic() engine_output = await self.inference_engine_client.generate(engine_input, model=self.policy_model_name) @@ -426,6 +429,14 @@ async def agent_loop( raise ValueError( "Rollout expert indices bookkeeping is not supported with custom chat template" ) + if routed_expert_trace is not None: + if rollout_expert_indices is None: + raise ValueError("R3 generation did not return routed expert indices") + routed_expert_trace.record_generation( + prompt_token_count=len(agent_loop_state.input_ids), + generated_token_count=len(output_ids), + routed_experts=rollout_expert_indices, + ) # Append eos when sampling_params.stop is not None. Does not affect 3.a as chat templates add eos_token. # sampling_params is not None for eval, but None for training (which uses engine.sampling_params which are from cfg) stop_strs = current_sampling_params.get("stop", None) @@ -459,6 +470,8 @@ async def agent_loop( ) output = env_step_output["postprocessed_action"] output_ids = self.tokenizer.encode(output, add_special_tokens=False) + if routed_expert_trace is not None: + raise ValueError("R3 bookkeeping is incompatible with postprocessed_action") obs_ids = self.get_obs_ids_from_obs(new_obs, agent_loop_state.done) @@ -471,7 +484,6 @@ async def agent_loop( reward=step_reward, obs_ids=obs_ids, added_eos=added_eos, - rollout_expert_indices=rollout_expert_indices, ) if is_step_wise: @@ -491,7 +503,6 @@ async def agent_loop( rollout_logprobs=turn_response_logprobs, stop_reason=stop_reason, env_metrics=env.get_metrics() if agent_loop_state.done else {}, - rollout_expert_indices=turn_output.get_turn_rollout_expert_indices(), ) agent_loop_output.step_outputs.append(per_step_output) @@ -553,10 +564,6 @@ async def agent_loop( rollout_logprobs = agent_loop_state.rollout_logprobs[ : agent_loop_state.response_end_idx - initial_prompt_length + 1 ] - if agent_loop_state.rollout_expert_indices is not None: - rollout_expert_indices_out = agent_loop_state.rollout_expert_indices[ - : agent_loop_state.response_end_idx + 1 - ] # fix index for per_step_rewards per_step_rewards = [(reward, idx - initial_prompt_length) for reward, idx in per_step_rewards] assert len(loss_mask) == len( @@ -573,6 +580,12 @@ async def agent_loop( rollout_logprobs.append(0.0) appended_eos_token = True + if agent_loop_state.routed_expert_trace is not None and agent_loop_state.routed_expert_trace.prompt_start: + rollout_expert_indices_out = agent_loop_state.routed_expert_trace.finalize( + token_count=len(prompt_ids) + len(response_ids), + loss_mask=[0] * len(prompt_ids) + loss_mask, + ) + if self.generator_cfg.step_wise_trajectories: for per_step_output, (reward, resp_end_idx) in zip(agent_loop_output.step_outputs, per_step_rewards): per_token_reward = [0.0] * len(per_step_output.response_ids) @@ -1047,8 +1060,6 @@ def _update_agent_state_by_retokenizing_chat_history( agent_loop_state.response_end_idx = None # `logprobs` are not computed because retokenizing breaks token-in-token-out agent_loop_state.rollout_logprobs = None - # indices are not meaningful when retokenizing - agent_loop_state.rollout_expert_indices = None return agent_loop_state def _update_agent_loop_state_with_multiturn_chat_template( @@ -1100,17 +1111,12 @@ def _update_agent_loop_state_with_multiturn_chat_template( loss_mask_for_turn = turn_output.get_turn_loss_mask() rollout_logprobs_for_turn = turn_output.get_turn_rollout_logprobs() - # use the raw rollout expert indices without any appending of observation tokens - # this will be overwritten each turn, so we don't need to append observation tokens to it - rollout_expert_indices_for_turn = turn_output.rollout_expert_indices - if self.generator_cfg.step_wise_trajectories: # cumulative input_ids is not tracked for step wise training agent_loop_state.response_end_idx = len(turn_output.output_ids) - 1 - # no running loss_mask, `rollout_logprobs`, or `rollout_expert_indices` are tracked for step-wise training + # no running loss_mask or rollout logprobs are tracked for step-wise training agent_loop_state.loss_mask = None agent_loop_state.rollout_logprobs = None - agent_loop_state.rollout_expert_indices = None else: # Directly append turn output turn_ids = turn_output.output_ids + turn_output.obs_ids @@ -1119,11 +1125,6 @@ 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 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 - agent_loop_state.rollout_expert_indices = rollout_expert_indices_for_turn return agent_loop_state @@ -1194,13 +1195,4 @@ def _update_agent_loop_state_with_singleturn_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 ( - self.generator_cfg.inference_engine.enable_return_routed_experts - and turn_output.rollout_expert_indices 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 observation tokens - # and the final response should not have an observation appended to it - agent_loop_state.rollout_expert_indices = turn_output.rollout_expert_indices - return agent_loop_state diff --git a/skyrl/train/utils/utils.py b/skyrl/train/utils/utils.py index e90c59630e..a1961e5a0c 100644 --- a/skyrl/train/utils/utils.py +++ b/skyrl/train/utils/utils.py @@ -217,6 +217,12 @@ def validate_megatron_cfg(cfg: SkyRLTrainConfig): f"{worker_type}.megatron_config: moe_enable_routing_replay is incompatible with " "moe_router_fusion=True -- the fused router bypasses replay. Set moe_router_fusion=False." ) + # Interleaved chunks desynchronise each RouterReplay instance's backward FIFO. + assert not config.megatron_config.transformer_config_kwargs.get("virtual_pipeline_model_parallel_size"), ( + f"{worker_type}.megatron_config: moe_enable_routing_replay is incompatible with " + "virtual_pipeline_model_parallel_size -- interleaved chunks desync the replay FIFO. " + "Unset virtual_pipeline_model_parallel_size." + ) # context, expert, and expert tensor parallel are not yet supported for megatron if config.megatron_config.context_parallel_size > 1: assert ( diff --git a/skyrl/utils/cpu_topology.py b/skyrl/utils/cpu_topology.py new file mode 100644 index 0000000000..0afb8a6d27 --- /dev/null +++ b/skyrl/utils/cpu_topology.py @@ -0,0 +1,76 @@ +"""Determine usable CPUs from process affinity and cgroup quota.""" + +import os +from typing import Optional, Tuple + +# Container cgroup namespaces expose the current cgroup at these paths. Module-level constants +# let tests replace them with fixtures. +CGROUP_V2_CPU_MAX_PATH = "/sys/fs/cgroup/cpu.max" +CGROUP_V1_CPU_QUOTA_PATH = "/sys/fs/cgroup/cpu/cpu.cfs_quota_us" +CGROUP_V1_CPU_PERIOD_PATH = "/sys/fs/cgroup/cpu/cpu.cfs_period_us" + +# cgroup v2 uses this literal for an unlimited quota. +CGROUP_V2_CPU_MAX_UNLIMITED = "max" + + +def _read_cgroup_file(path: str) -> str: + with open(path, encoding="utf-8") as handle: + return handle.read() + + +def _read_cgroup_v2_cpu_max() -> Optional[Tuple[float, float]]: + """``(quota, period)`` from cgroup v2 ``cpu.max``, or ``None`` if absent or unlimited.""" + try: + quota_text, period_text = _read_cgroup_file(CGROUP_V2_CPU_MAX_PATH).split() + if quota_text == CGROUP_V2_CPU_MAX_UNLIMITED: + return None + return float(quota_text), float(period_text) + except (OSError, ValueError): + return None + + +def _read_cgroup_v1_cpu_max() -> Optional[Tuple[float, float]]: + """``(quota, period)`` from cgroup v1 ``cpu.cfs_*_us``, or ``None`` if absent or unlimited.""" + try: + quota = float(_read_cgroup_file(CGROUP_V1_CPU_QUOTA_PATH).strip()) + period = float(_read_cgroup_file(CGROUP_V1_CPU_PERIOD_PATH).strip()) + except (OSError, ValueError): + return None + if quota < 0: + return None + return quota, period + + +def cgroup_cpu_quota() -> Optional[int]: + """Return whole CPUs permitted by CFS, or ``None`` when no quota applies.""" + limits = _read_cgroup_v2_cpu_max() or _read_cgroup_v1_cpu_max() + if limits is None: + return None + quota, period = limits + if quota <= 0 or period <= 0: + return None + # Floor a fractional allowance, but a sub-CPU quota still gets one worker. + return max(1, int(quota // period)) + + +def permitted_cpu_cores() -> int: + """Return the lesser of the process affinity and cgroup quota.""" + try: + affinity = len(os.sched_getaffinity(0)) + except AttributeError: + affinity = os.cpu_count() or 1 + quota = cgroup_cpu_quota() + if quota is None: + return affinity + return min(affinity, quota) + + +def pool_workers(*, cap: int, reserved: int, cores: Optional[int] = None) -> int: + """Size a pool from permitted cores, a cap, and a reserve for colocated processes.""" + if cap < 1: + raise ValueError(f"pool cap must be positive, got {cap}") + if reserved < 0: + raise ValueError(f"reserved cores must be non-negative, got {reserved}") + if cores is None: + cores = permitted_cpu_cores() + return max(1, min(cap, cores - reserved)) diff --git a/tests/backends/skyrl_train/distributed/test_token_metadata.py b/tests/backends/skyrl_train/distributed/test_token_metadata.py index ff8f3dc7de..e3e7660567 100644 --- a/tests/backends/skyrl_train/distributed/test_token_metadata.py +++ b/tests/backends/skyrl_train/distributed/test_token_metadata.py @@ -1,10 +1,16 @@ import sys import types +import numpy as np import pytest import torch from skyrl.backends.skyrl_train.distributed.megatron import token_metadata +from skyrl.backends.skyrl_train.distributed.megatron.token_metadata import ( + TokenMetadataTrace, +) +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor +from skyrl.backends.skyrl_train.utils.routed_experts import RoutedExpertTrace @pytest.fixture @@ -86,3 +92,133 @@ def test_packed_layout_aligns_next_token_metadata_and_scatters_rows(monkeypatch, 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]] + + +def test_token_metadata_trace_chunks_and_independent_schema() -> None: + trace, other = TokenMetadataTrace(), TokenMetadataTrace() + trace.append(np.ones((2, 3), dtype=np.int32), expected_rows=2) + trace.append(np.zeros((1, 3), dtype=np.int32), expected_rows=1) + other.append(np.empty((0, 4), dtype=np.float32), expected_rows=0) + + with pytest.raises(ValueError, match="expected 4"): + trace.finalize(expected_rows=4) + result = trace.finalize(expected_rows=3) + assert result.shape == (3, 3) + assert other.finalize(expected_rows=0).shape == (0, 4) + with pytest.raises(RuntimeError, match="already finalized"): + trace.finalize(expected_rows=3) + + +@pytest.mark.parametrize( + ("rows", "expected", "match"), + [ + (np.ones((2, 2), dtype=np.int32), 1, "has 2 rows"), + (np.ones((2, 2), dtype=np.int32)[:, ::2], 2, "contiguous"), + (np.ones((1, 3), dtype=np.int32), 1, "schema changed"), + (np.ones((1, 2), dtype=np.int16), 1, "schema changed"), + ], +) +def test_token_metadata_trace_rejects_invalid_chunks(rows, expected, match) -> None: + trace = TokenMetadataTrace() + if rows.shape[0] == 1: + trace.append(np.ones((1, 2), dtype=np.int32), expected_rows=1) + with pytest.raises(ValueError, match=match): + trace.append(rows, expected_rows=expected) + + +def routes(rows: int) -> np.ndarray: + return np.arange(rows * 4, dtype=np.int32).reshape(rows, 2, 2) % 8 + + +def test_routed_expert_trace_tracks_multiturn_suffix_and_terminal_gap() -> None: + trace = RoutedExpertTrace() + trace.record_generation(prompt_token_count=3, generated_token_count=2, routed_experts=routes(4)) + assert trace.prompt_start == 4 + trace.record_generation(prompt_token_count=7, generated_token_count=2, routed_experts=routes(4)) + + result = trace.finalize(token_count=9, loss_mask=[0, 0, 0, 1, 1, 0, 0, 1, 1]) + assert result.shape == (9, 2, 2) and result.dtype == np.uint8 + assert np.array_equal(result[-1, 0], [0, 1]) + + +@pytest.mark.parametrize("active", [False, True]) +def test_routed_expert_trace_only_pads_masked_suffix(active: bool) -> None: + trace = RoutedExpertTrace() + trace.record_generation(prompt_token_count=3, generated_token_count=1, routed_experts=routes(3)) + mask = [0, 0, 0, 0, int(active)] + if active: + with pytest.raises(ValueError, match="loss-active target"): + trace.finalize(token_count=5, loss_mask=mask) + else: + result = trace.finalize(token_count=5, loss_mask=mask) + assert np.array_equal(result[-2:, 0], [[0, 1], [0, 1]]) + + +@pytest.mark.parametrize("packed", [False, True]) +def test_align_token_rows_places_each_trajectory_from_its_own_row_source(monkeypatch, parallel_state, packed): + """``_align_token_rows`` is the one placement loop both alignment entry points share.""" + monkeypatch.setattr(token_metadata, "get_packed_seq_align_size", lambda *args, **kwargs: 4) + monkeypatch.setattr(token_metadata, "get_unpacked_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1], [0, 0, 1, 1]]) + rows = [torch.tensor([10, 11, 12], dtype=torch.int32), torch.tensor([20, 21], dtype=torch.int32)] + layout = token_metadata.build_token_metadata_layout( + attention_mask, + rows[0].device, + packed=packed, + fp8_enabled=False, + ) + + aligned = token_metadata._align_token_rows( + rows.__getitem__, + rows[0], + (), + layout, + -1, + ) + + if packed: + assert aligned.tolist() == [[10, 11, 12, -1, 20, 21, -1, -1]] + else: + assert aligned.tolist() == [[10, 11, 12, -1], [20, 21, -1, -1]] + + +@pytest.mark.parametrize("packed", [False, True]) +def test_align_packed_token_metadata_honours_per_segment_starts(monkeypatch, parallel_state, packed): + """A response-suffix channel covers part of a trajectory and needs its own start.""" + monkeypatch.setattr(token_metadata, "get_packed_seq_align_size", lambda *args, **kwargs: 4) + monkeypatch.setattr(token_metadata, "get_unpacked_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1], [0, 0, 1, 1]]) + # Trajectory 0 keeps its last 2 of 3 real tokens; trajectory 1 keeps its last 1 of 2. + suffix = PackedTensor.from_segments( + [torch.tensor([11, 12], dtype=torch.int32), torch.tensor([21], dtype=torch.int32)] + ) + layout = token_metadata.build_token_metadata_layout( + attention_mask, + suffix.device, + packed=packed, + fp8_enabled=False, + ) + + aligned = token_metadata.align_packed_token_metadata(suffix, layout, -1, segment_starts=[1, 1]) + + if packed: + assert aligned.tolist() == [[-1, 11, 12, -1, -1, 21, -1, -1]] + else: + assert aligned.tolist() == [[-1, 11, 12, -1], [-1, 21, -1, -1]] + + +def test_align_packed_token_metadata_rejects_segments_that_leave_the_trajectory(monkeypatch, parallel_state): + monkeypatch.setattr(token_metadata, "get_unpacked_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1]]) + suffix = PackedTensor.from_segments([torch.tensor([11, 12], dtype=torch.int32)]) + layout = token_metadata.build_token_metadata_layout( + attention_mask, + suffix.device, + packed=False, + fp8_enabled=False, + ) + + with pytest.raises(ValueError, match="spans real tokens"): + token_metadata.align_packed_token_metadata(suffix, layout, -1, segment_starts=[2]) + with pytest.raises(ValueError, match="do not match"): + token_metadata.align_packed_token_metadata(suffix, layout, -1) 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 3f97e0ed50..89bfd85515 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 @@ -16,6 +16,7 @@ get_sampling_params_for_backend, ) from skyrl.backends.skyrl_train.training_batch import TrainingInputBatch +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor from skyrl.train.config import SamplingParams, SkyRLTrainConfig from skyrl.train.dataset.preprocess import ( convert_prompts_responses_to_batch_tensors, @@ -36,6 +37,25 @@ NUM_PROMPTS = 10 N_SAMPLES_PER_PROMPT = 4 MAX_GENERATE_LENGTH = 128 +# Moonlight 16B: 27 MoE layers, top_k=6, 64 routed experts. +MOONLIGHT_NUM_LAYERS = 27 +MOONLIGHT_TOPK = 6 +MOONLIGHT_NUM_EXPERTS = 64 + + +def _packed_moonlight_routes(attention_mask: torch.Tensor) -> PackedTensor: + """Moonlight-shaped routes packed to each trajectory's real tokens.""" + route_offsets = torch.arange(MOONLIGHT_TOPK, dtype=torch.int32) + segments = [] + for real_tokens in attention_mask.sum(dim=1).tolist(): + route_start = torch.randint( + 0, + MOONLIGHT_NUM_EXPERTS, + (real_tokens, MOONLIGHT_NUM_LAYERS, 1), + dtype=torch.int32, + ) + segments.append((route_start + route_offsets) % MOONLIGHT_NUM_EXPERTS) + return PackedTensor.from_segments(segments) def _extra_env_vars_for_model(model_name: str) -> dict[str, str] | None: @@ -353,22 +373,9 @@ def test_forward_backward(tp, pp, cp, ep, etp, extra_tf_kwargs): ) batch_size = sequences.shape[0] - seq_len = sequences.shape[1] num_actions = response_mask.shape[1] - # Moonlight 16B: 27 MoE layers, top_k=6, 64 routed experts - MOONLIGHT_NUM_LAYERS = 27 - MOONLIGHT_TOPK = 6 - MOONLIGHT_NUM_EXPERTS = 64 - route_start = torch.randint( - 0, - MOONLIGHT_NUM_EXPERTS, - (batch_size, seq_len, MOONLIGHT_NUM_LAYERS, 1), - dtype=torch.int32, - ) - 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 + rollout_expert_indices = _packed_moonlight_routes(attention_mask) gen = torch.Generator().manual_seed(42) training_input = TrainingInputBatch( @@ -419,3 +426,145 @@ def test_forward_backward(tp, pp, cp, ep, etp, extra_tf_kwargs): for actor in actor_group._actor_handlers: ray.kill(actor) + + +@pytest.mark.h100 +@pytest.mark.parametrize( + "tp,pp,cp,ep,etp,extra_tf_kwargs", + [ + pytest.param(2, 2, 1, 2, 1, {"num_layers_in_first_pipeline_stage": 13}, id="tp2_pp2_ep2"), + ], +) +def test_forward_backward_variable_length_full_recompute(tp, pp, cp, ep, etp, extra_tf_kwargs): + """Replayed routes must stay paired with their own microbatch. + + Each forward microbatch appends its expert routes to a FIFO that + activation-checkpoint recomputation drains once during backward. If a + microbatch queues its routes more than once (or not at all), backward + replays a *different* microbatch's routes. Megatron's MoE all-to-all then + computes split sizes from a token count that doesn't match the tensors in + flight and the collective fails. + + Two conditions are needed to expose that, and both are set here: + + * ``recompute_granularity="full"`` so backward actually recomputes the + MoE layers and consumes the queue (the default only recomputes + ``core_attn``, which never replays routes). + * Sequence lengths that differ *between* microbatches, so a mispaired + replay changes the token count rather than silently reusing a + same-shaped tensor. Uniform lengths can mask the bug entirely. + + Prompts are built with deliberately spread lengths and ``micro_*=2`` over + 8 samples, giving 4 microbatches whose padded widths differ. + """ + with ray_init(extra_env_vars=_extra_env_vars_for_model(MOE_MODEL_NAME)): + cfg = get_test_actor_config(model_name=MOE_MODEL_NAME) + cfg.trainer.strategy = "megatron" + + tokenizer = AutoTokenizer.from_pretrained(MOE_MODEL_NAME, trust_remote_code=True) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + + # Lengths chosen so consecutive microbatches (pairs, micro_bs=2) have + # different maxima: a stale replay is then shape-visible, not benign. + filler_words = [1, 6, 2, 14, 3, 25, 4, 40] + prompts, responses, rewards, loss_masks = [], [], [], [] + for i, filler in enumerate(filler_words): + prompt_ids = tokenizer.encode( + "Question: " + ("token " * filler) + f"what is {i} plus {i}?", + add_special_tokens=False, + ) + response_ids = tokenizer.encode( + ("because " * filler) + f"the answer is {i + i}.", + add_special_tokens=False, + ) + if tokenizer.eos_token_id is not None and (not response_ids or response_ids[-1] != tokenizer.eos_token_id): + response_ids.append(tokenizer.eos_token_id) + prompts.append(prompt_ids) + responses.append(response_ids) + rewards.append([1.0] * len(response_ids)) + loss_masks.append([1] * len(response_ids)) + + sequences, attention_mask, response_mask, rewards_t, loss_mask_t, _, _ = ( + convert_prompts_responses_to_batch_tensors( + tokenizer=tokenizer, + prompts=prompts, + responses=responses, + rewards=rewards, + loss_masks=loss_masks, + ) + ) + + # Guard the premise: if padding collapsed the spread, the test would + # pass for the wrong reason. + real_token_counts = attention_mask.sum(dim=-1) + assert ( + real_token_counts.unique().numel() > 1 + ), f"variable-length premise broken: every sample has {real_token_counts[0].item()} real tokens" + + batch_size = sequences.shape[0] + num_actions = response_mask.shape[1] + + rollout_expert_indices = _packed_moonlight_routes(attention_mask) + + gen = torch.Generator().manual_seed(42) + training_input = TrainingInputBatch( + { + "sequences": sequences, + "attention_mask": attention_mask, + "response_mask": response_mask, + "rewards": rewards_t, + "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), + "action_mask": response_mask.to(dtype=torch.int64), + } + ) + training_input.metadata = {"response_length": num_actions} + + cfg.trainer.placement.policy_num_gpus_per_node = 4 + if extra_tf_kwargs is not None: + cfg.trainer.policy.megatron_config.transformer_config_kwargs.update(extra_tf_kwargs) + # Recompute whole layers so backward re-runs the MoE routers and drains + # the replay queue; ``core_attn`` alone never replays routes. + cfg.trainer.policy.megatron_config.transformer_config_kwargs.update( + { + "recompute_granularity": "full", + "recompute_method": "uniform", + "recompute_num_layers": 1, + } + ) + cfg.trainer.policy.megatron_config.transformer_config_kwargs.pop("recompute_modules", None) + cfg.trainer.policy.megatron_config.tensor_model_parallel_size = tp + cfg.trainer.policy.megatron_config.pipeline_model_parallel_size = pp + cfg.trainer.policy.megatron_config.context_parallel_size = cp + cfg.trainer.policy.megatron_config.expert_model_parallel_size = ep + cfg.trainer.policy.megatron_config.expert_tensor_parallel_size = etp + # 8 samples / 2 per microbatch = 4 microbatches of differing widths. + cfg.trainer.micro_forward_batch_size_per_gpu = 2 + cfg.trainer.micro_train_batch_size_per_gpu = 2 + cfg.trainer.policy.megatron_config.moe_enable_routing_replay = True + + actor_group = init_worker_with_type( + "policy", + num_gpus_per_node=4, + cfg=cfg, + ) + + # Two steps: the first leaves any surplus queue entry behind, so a + # mispairing shows up on the second even if the first survives. + ray.get(actor_group.async_run_ray_method("mesh", "forward_backward", data=training_input)) + ray.get(actor_group.async_run_ray_method("pass_through", "optim_step")) + results = ray.get(actor_group.async_run_ray_method("mesh", "forward_backward", data=training_input)) + + loss = results[0].metrics["policy_loss"] + print(f"Variable-length replay forward_backward - loss: {loss:.6f}") + assert loss is not None and not torch.isnan(torch.tensor(loss)), "Loss should be valid (not NaN)" + assert loss != 0.0, "Loss should be non-zero with non-zero advantages" + + for actor in actor_group._actor_handlers: + ray.kill(actor) diff --git a/tests/backends/skyrl_train/inference_servers/test_generate_wire.py b/tests/backends/skyrl_train/inference_servers/test_generate_wire.py index ac2a6219ae..c93a4d4cfd 100644 --- a/tests/backends/skyrl_train/inference_servers/test_generate_wire.py +++ b/tests/backends/skyrl_train/inference_servers/test_generate_wire.py @@ -1,6 +1,7 @@ """Tests for the /skyrl/v1/generate payload contract.""" import base64 +import json import math from dataclasses import dataclass @@ -11,11 +12,27 @@ from skyrl.backends.skyrl_train.inference_servers.generate_wire import ( CLAMPED_LOGPROB, + PackedArrayKey, + PackedField, build_logprobs_content, decode_packed_routed_experts, + load_packed_body, + pack_ndarray, pack_routed_experts, + unpack_ndarray, ) +_FLOAT32 = frozenset({np.dtype(np.float32)}) +_INT16 = frozenset({np.dtype(np.int16)}) + + +def _support_envelope(support: np.ndarray) -> dict: + return pack_ndarray(support, allowed_dtypes=_FLOAT32) + + +def _body(**choice_fields) -> dict: + return {"choices": [{"token_ids": [1, 2, 3], "finish_reason": "stop", **choice_fields}]} + @dataclass class _Logprob: @@ -183,3 +200,189 @@ def test_decode_rejects_noncanonical_dtype(): with pytest.raises(ValueError, match="non-canonical dtype"): decode_packed_routed_experts(payload) + + +@pytest.mark.parametrize( + "arr,allowed_dtypes,extra", + [ + (np.arange(6, dtype=np.float32).reshape(2, 3), _FLOAT32, None), + (np.arange(6, dtype=np.float32).reshape(2, 3), _FLOAT32, {"prompt_start": 4, "labels": ["a", "b"]}), + (np.arange(12, dtype=np.int16).reshape(3, 2, 2), _INT16, {"prompt_start": 0}), + (np.empty((0, 3), dtype=np.float32), _FLOAT32, None), + ], +) +def test_ndarray_round_trip_with_sidecar_fields(arr, allowed_dtypes, extra): + payload = pack_ndarray(arr, allowed_dtypes=allowed_dtypes, extra=extra) + decoded, sidecar = unpack_ndarray(payload, allowed_dtypes=allowed_dtypes, ndim=arr.ndim) + + assert np.array_equal(decoded, arr) + assert decoded.dtype == arr.dtype + assert decoded.flags.c_contiguous + assert sidecar == (extra or {}) + + +def test_packed_envelope_leads_with_data(): + payload = pack_ndarray(np.zeros((2, 2), np.float32), allowed_dtypes=_FLOAT32, extra={"prompt_start": 1}) + + assert list(payload) == [PackedArrayKey.DATA, PackedArrayKey.SHAPE, PackedArrayKey.DTYPE, "prompt_start"] + assert orjson.dumps(payload).startswith(b'{"data":"') + + +def test_pack_routed_experts_is_byte_identical_to_the_hand_built_envelope(): + routes = np.arange(12).reshape(3, 2, 2) + + assert orjson.dumps(pack_routed_experts(routes)) == b'{"data":"AAECAwQFBgcICQoL","shape":[3,2,2],"dtype":"uint8"}' + + +@pytest.mark.parametrize( + "wrap", [str, lambda data: memoryview(data.encode("ascii")), lambda data: data.encode("ascii")] +) +def test_unpack_accepts_str_and_buffers(wrap): + support = np.arange(6, dtype=np.float32).reshape(2, 3) + payload = dict(_support_envelope(support)) + payload[PackedArrayKey.DATA.value] = wrap(payload[PackedArrayKey.DATA.value]) + + decoded, _ = unpack_ndarray(payload, allowed_dtypes=_FLOAT32, ndim=2) + assert np.array_equal(decoded, support) + + +def test_pack_rejects_disallowed_dtype(): + with pytest.raises(ValueError, match="dtype"): + pack_ndarray(np.zeros((2, 2), np.float64), allowed_dtypes=_FLOAT32) + + +def test_pack_rejects_sidecar_collision_with_envelope_keys(): + with pytest.raises(ValueError, match="collide"): + pack_ndarray(np.zeros((2, 2), np.float32), allowed_dtypes=_FLOAT32, extra={"dtype": "float64"}) + + +def test_unpack_rejects_disallowed_dtype(): + payload = pack_ndarray(np.zeros((2, 2), np.float32), allowed_dtypes=_FLOAT32) + + with pytest.raises(ValueError, match="dtype"): + unpack_ndarray(payload, allowed_dtypes=_INT16, ndim=2) + + +@pytest.mark.parametrize("ndim", [1, 3]) +def test_unpack_rejects_wrong_ndim(ndim): + payload = pack_ndarray(np.zeros((2, 2), np.float32), allowed_dtypes=_FLOAT32) + + with pytest.raises(ValueError, match="dimensions"): + unpack_ndarray(payload, allowed_dtypes=_FLOAT32, ndim=ndim) + + +def test_unpack_rejects_byte_count_mismatched_with_declared_shape(): + payload = pack_ndarray(np.zeros((2, 3), np.float32), allowed_dtypes=_FLOAT32) + payload[PackedArrayKey.SHAPE.value] = [2, 4] + + with pytest.raises(ValueError, match="24 bytes, expected 32"): + unpack_ndarray(payload, allowed_dtypes=_FLOAT32, ndim=2) + + +def test_load_packed_body_splices_both_blobs_in_one_body(): + routes = np.arange(12).reshape(3, 2, 2) + support = np.arange(6, dtype=np.float32).reshape(2, 3) + raw = orjson.dumps( + _body( + logprobs={"content": [{"logprob": -0.5}]}, + routed_experts=pack_routed_experts(routes), + rollout_sample_support=_support_envelope(support), + ) + ) + + choice = load_packed_body(raw)["choices"][0] + + assert all( + isinstance(choice[field][PackedArrayKey.DATA], memoryview) + for field in (PackedField.ROUTED_EXPERTS, PackedField.ROLLOUT_SAMPLE_SUPPORT) + ) + assert np.array_equal(decode_packed_routed_experts(choice[PackedField.ROUTED_EXPERTS]), routes) + decoded_support, _ = unpack_ndarray(choice[PackedField.ROLLOUT_SAMPLE_SUPPORT], allowed_dtypes=_FLOAT32, ndim=2) + assert np.array_equal(decoded_support, support) + assert choice["logprobs"] == {"content": [{"logprob": -0.5}]} + + +def test_load_packed_body_keeps_sidecar_fields(): + support = np.arange(4, dtype=np.float32).reshape(2, 2) + envelope = pack_ndarray(support, allowed_dtypes=_FLOAT32, extra={"prompt_start": 7}) + raw = orjson.dumps(_body(rollout_sample_support=envelope)) + + decoded, sidecar = unpack_ndarray( + load_packed_body(raw)["choices"][0][PackedField.ROLLOUT_SAMPLE_SUPPORT], + allowed_dtypes=_FLOAT32, + ndim=2, + ) + + assert np.array_equal(decoded, support) + assert sidecar == {"prompt_start": 7} + + +@pytest.mark.parametrize("value", [None, "absent"]) +def test_load_packed_body_passes_through_absent_and_null_fields(value): + fields = {} if value == "absent" else {PackedField.ROUTED_EXPERTS.value: None} + body = _body(logprobs=None, **fields) + + assert load_packed_body(orjson.dumps(body)) == body + + +def test_load_packed_body_rejects_reordered_envelope_keys(): + routes = np.arange(12).reshape(3, 2, 2) + envelope = pack_routed_experts(routes) + reordered = {key: envelope[key] for key in reversed(list(envelope))} + + with pytest.raises(ValueError, match="layout drifted"): + load_packed_body(orjson.dumps(_body(routed_experts=reordered))) + + +def test_load_packed_body_rejects_a_reserialized_body(): + # stdlib json spaces its separators; the blob would silently land in a + # ~121 MiB Python str instead, which is exactly the cost this avoids. + raw = json.dumps(_body(routed_experts=pack_routed_experts(np.arange(12).reshape(3, 2, 2)))).encode() + + with pytest.raises(ValueError, match="layout drifted"): + load_packed_body(raw) + + +def test_load_packed_body_is_not_spoofable_from_a_string_value(): + routes = np.arange(12).reshape(3, 2, 2) + spoof = '"routed_experts":{"data":"AAAA","shape":[1,1,1],"dtype":"uint8"}' + raw = orjson.dumps(_body(note=spoof, routed_experts=pack_routed_experts(routes))) + + choice = load_packed_body(raw)["choices"][0] + + assert choice["note"] == spoof + assert np.array_equal(decode_packed_routed_experts(choice[PackedField.ROUTED_EXPERTS]), routes) + + +def test_load_packed_body_ignores_unregistered_packed_fields(): + envelope = pack_ndarray(np.zeros((2, 2), np.float32), allowed_dtypes=_FLOAT32) + body = _body(some_other_array=envelope) + + assert load_packed_body(orjson.dumps(body)) == body + + +def test_load_packed_body_honours_a_narrowed_field_registry(): + routes = np.arange(12).reshape(3, 2, 2) + raw = orjson.dumps(_body(routed_experts=pack_routed_experts(routes))) + + body = load_packed_body(raw, fields=(PackedField.ROLLOUT_SAMPLE_SUPPORT,)) + + assert isinstance(body["choices"][0][PackedField.ROUTED_EXPERTS][PackedArrayKey.DATA], str) + + +def test_load_packed_body_splices_one_blob_per_choice(): + first, second = np.arange(12).reshape(3, 2, 2), np.arange(12, 24).reshape(3, 2, 2) + raw = orjson.dumps({"choices": [{"routed_experts": pack_routed_experts(routes)} for routes in (first, second)]}) + + choices = load_packed_body(raw)["choices"] + + assert np.array_equal(decode_packed_routed_experts(choices[0][PackedField.ROUTED_EXPERTS]), first) + assert np.array_equal(decode_packed_routed_experts(choices[1][PackedField.ROUTED_EXPERTS]), second) + + +def test_load_packed_body_rejects_an_unterminated_blob(): + raw = orjson.dumps(_body(routed_experts=pack_routed_experts(np.arange(12).reshape(3, 2, 2)))) + truncated = raw[: raw.index(b'"shape"') - 2] + + with pytest.raises(ValueError, match="unterminated"): + load_packed_body(truncated) 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 75267428b6..85b4f09151 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 @@ -1,6 +1,7 @@ """Tests for RemoteInferenceClient.""" import asyncio +import json import pickle import threading import time @@ -9,15 +10,24 @@ import aiohttp import httpx import numpy as np +import orjson import pytest import pytest_asyncio import uvicorn from fastapi import FastAPI, Query, Request -from fastapi.responses import JSONResponse, PlainTextResponse +from fastapi.responses import JSONResponse, PlainTextResponse, Response +from skyrl.backends.skyrl_train.inference_servers import ( + remote_inference_client as remote_client_module, +) from skyrl.backends.skyrl_train.inference_servers.common import get_open_port from skyrl.backends.skyrl_train.inference_servers.generate_wire import ( + PackedArrayKey, + PackedField, + decode_packed_routed_experts, + pack_ndarray, pack_routed_experts, + unpack_ndarray, ) from skyrl.backends.skyrl_train.inference_servers.remote_inference_client import ( SKYRL_LORA_ADAPTER_NAME, @@ -29,12 +39,28 @@ ) from skyrl.train.config import SkyRLTrainConfig +_SUPPORT_DTYPES = frozenset({np.dtype(np.float32)}) +_ROUTES = np.arange(12).reshape(3, 2, 2) +_SUPPORT = np.arange(6, dtype=np.float32).reshape(2, 3) + + +def _packed_generate_body(*, two_blobs: bool) -> dict: + choice: dict = { + "token_ids": [1, 2, 3], + "finish_reason": "stop", + PackedField.ROUTED_EXPERTS.value: pack_routed_experts(_ROUTES), + } + if two_blobs: + choice[PackedField.ROLLOUT_SAMPLE_SUPPORT.value] = pack_ndarray(_SUPPORT, allowed_dtypes=_SUPPORT_DTYPES) + return {"choices": [choice]} + def create_mock_vllm_server(server_id: int) -> FastAPI: """Create a mock vLLM server with standard endpoints.""" app = FastAPI() app.state.last_generate_features = None app.state.last_generate_model = None + app.state.last_generate_sampling_params = None app.state.last_chat_model = None app.state.last_completion_model = None app.state.last_render_model = None @@ -45,11 +71,45 @@ def create_mock_vllm_server(server_id: int) -> FastAPI: app.state.finished_sessions = [] # Number of /get_world_size hits, used to assert client-side caching. app.state.world_size_calls = 0 + app.state.drifted_body_calls = 0 + app.state.flaky_body_calls = 0 @app.get("/health") async def health(): return {"status": "ok"} + @app.post("/test/packed_body") + async def packed_body(two_blobs: bool = False): + return Response(content=orjson.dumps(_packed_generate_body(two_blobs=two_blobs)), media_type="application/json") + + @app.post("/test/drifted_packed_body") + async def drifted_packed_body(): + app.state.drifted_body_calls += 1 + # stdlib json spaces its separators, so the splice prefix no longer matches. + content = json.dumps(_packed_generate_body(two_blobs=False)).encode() + return Response(content=content, media_type="application/json") + + @app.post("/test/flaky_packed_body") + async def flaky_packed_body(): + app.state.flaky_body_calls += 1 + if app.state.flaky_body_calls == 1: + return Response(content=b"gateway hiccup", media_type="application/json", status_code=502) + return Response(content=orjson.dumps(_packed_generate_body(two_blobs=True)), media_type="application/json") + + @app.post("/test/reset_packed_body_calls") + async def reset_packed_body_calls(): + app.state.drifted_body_calls = 0 + app.state.flaky_body_calls = 0 + return {"status": "ok"} + + @app.get("/test/packed_body_calls") + async def packed_body_calls(): + return {"drifted": app.state.drifted_body_calls, "flaky": app.state.flaky_body_calls} + + @app.post("/test/bad_request_text") + async def bad_request_text(): + return PlainTextResponse("prompt too long", status_code=400) + @app.post("/finish_session") async def finish_session(session_id: str = Query(...)): app.state.finished_sessions.append(session_id) @@ -63,6 +123,10 @@ async def get_finished(): async def get_last_generate_features(): return {"features": app.state.last_generate_features} + @app.get("/test/last_generate_sampling_params") + async def get_last_generate_sampling_params(): + return app.state.last_generate_sampling_params + @app.get("/test/last_models") async def get_last_models(): return { @@ -104,6 +168,7 @@ async def completions(request: Request): async def generate(request: Request): body = await request.json() # Consume body sp = body.get("sampling_params", {}) + app.state.last_generate_sampling_params = sp input_token_ids = body.get("token_ids", []) app.state.last_generate_model = body.get("model") n = sp.get("n", 1) @@ -438,8 +503,7 @@ def test_serialization(self, mock_servers): assert restored.proxy_url == client.proxy_url assert restored.server_urls == client.server_urls assert restored.model_name == client.model_name - # Session should be None after unpickling - assert restored._session is None + assert restored._generate_client is None class TestDataPlane: @@ -480,13 +544,16 @@ async def test_generate_decodes_packed_routed_experts(self, mock_servers): enable_return_routed_experts=True, ) try: - result = await client.generate({"prompt_token_ids": [[1, 2, 3]]}) + result = await client.generate({"prompt_token_ids": [[1, 2, 3]], "routed_experts_prompt_starts": [1]}) + async with httpx.AsyncClient() as http: + captured = (await http.get(f"{mock_servers['proxy_url']}/test/last_generate_sampling_params")).json() 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)) + assert captured["routed_experts_prompt_start"] == 1 @pytest.mark.asyncio async def test_generate_rejects_list_routed_experts(self, monkeypatch): @@ -508,7 +575,7 @@ async def return_list_routes(*args, **kwargs): ] } - monkeypatch.setattr(client, "_post", return_list_routes) + monkeypatch.setattr(client._get_generate_client(), "_post", return_list_routes) with pytest.raises(ValueError, match="must return packed"): await client._generate_single([1], {}, None, "model") @@ -550,6 +617,90 @@ async def test_detokenize(self, client): assert result[0] == "hello world" # Mock response +class TestPackedSideChannelBodies: + """Tests for response parsing with packed side channels.""" + + async def _post_packed(self, client, mock_servers, path: str, **kwargs): + return await client._get_generate_client()._post( + f"{mock_servers['proxy_url']}{path}", json={}, packed_side_channels=True, **kwargs + ) + + @pytest.mark.asyncio + async def test_splices_both_registered_fields(self, client, mock_servers): + body = await self._post_packed(client, mock_servers, "/test/packed_body?two_blobs=true") + choice = body["choices"][0] + + assert isinstance(choice[PackedField.ROUTED_EXPERTS][PackedArrayKey.DATA], memoryview) + assert isinstance(choice[PackedField.ROLLOUT_SAMPLE_SUPPORT][PackedArrayKey.DATA], memoryview) + assert np.array_equal(decode_packed_routed_experts(choice[PackedField.ROUTED_EXPERTS]), _ROUTES) + support, _ = unpack_ndarray(choice[PackedField.ROLLOUT_SAMPLE_SUPPORT], allowed_dtypes=_SUPPORT_DTYPES, ndim=2) + assert np.array_equal(support, _SUPPORT) + + @pytest.mark.asyncio + async def test_drifted_layout_raises_without_retrying(self, client, mock_servers): + await client._post(f"{mock_servers['proxy_url']}/test/reset_packed_body_calls", json={}) + + with pytest.raises(ValueError, match="layout drifted"): + await self._post_packed(client, mock_servers, "/test/drifted_packed_body") + + async with httpx.AsyncClient() as http: + counts = (await http.get(f"{mock_servers['proxy_url']}/test/packed_body_calls")).json() + assert counts["drifted"] == 1 + + @pytest.mark.asyncio + async def test_undecodable_body_is_retried_then_spliced(self, client, mock_servers): + await client._post(f"{mock_servers['proxy_url']}/test/reset_packed_body_calls", json={}) + + body = await self._post_packed(client, mock_servers, "/test/flaky_packed_body") + + async with httpx.AsyncClient() as http: + counts = (await http.get(f"{mock_servers['proxy_url']}/test/packed_body_calls")).json() + assert counts["flaky"] == 2 + assert np.array_equal( + decode_packed_routed_experts(body["choices"][0][PackedField.ROUTED_EXPERTS]), + _ROUTES, + ) + + @pytest.mark.asyncio + async def test_client_error_with_non_json_body_surfaces_the_text(self, client, mock_servers): + with pytest.raises(aiohttp.ClientResponseError, match="prompt too long"): + await client._post(f"{mock_servers['proxy_url']}/test/bad_request_text", json={}) + + @pytest.mark.asyncio + async def test_non_routed_expert_generate_never_scans_the_body(self, client, monkeypatch): + def fail(*args, **kwargs): + raise AssertionError("the non-R3 path must not scan the response body") + + monkeypatch.setattr(remote_client_module, "load_packed_body", fail) + + result = await client.generate({"prompt_token_ids": [[1, 2, 3]]}) + assert len(result["responses"]) == 1 + + @pytest.mark.asyncio + async def test_routed_expert_generate_goes_through_the_splice(self, mock_servers, monkeypatch): + calls: List[int] = [] + original = remote_client_module.load_packed_body + + def counted(raw, **kwargs): + calls.append(len(raw)) + return original(raw, **kwargs) + + monkeypatch.setattr(remote_client_module, "load_packed_body", counted) + 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(calls) == 1 + assert np.array_equal(result["rollout_expert_indices"][0], _ROUTES) + + class TestControlPlane: """Test control plane methods (fan-out to all servers).""" @@ -1025,8 +1176,7 @@ async def test_async_context_manager(self, mock_servers): result = await client.resume() assert len(result) == 2 - # Session should be closed after exiting context - assert client._session is None or client._session.closed + assert client._generate_client is None or client._generate_client._session is None async def _get_lora_registries(server_urls: List[str]) -> List[Dict[str, str]]: 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 4fac8a6679..c78d0fee80 100644 --- a/tests/backends/skyrl_train/test_token_based_batching_utils.py +++ b/tests/backends/skyrl_train/test_token_based_batching_utils.py @@ -12,6 +12,10 @@ import torch from skyrl.backends.skyrl_train.training_batch import TensorList, TrainingInputBatch +from skyrl.backends.skyrl_train.utils.packed_tensor import ( + PackedTensor, + cu_seqlens_from_lengths, +) from skyrl.backends.skyrl_train.workers.worker_utils import ( TokenBasedBatchIterator, get_microbatch_iterator, @@ -201,16 +205,33 @@ def test_padding_microbatch_matches_seq_len(self): 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["rollout_expert_indices"] = PackedTensor( + torch.full((8, 2, 3), 7, dtype=torch.int16), + cu_seqlens_from_lengths([4, 4]), + ) 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) + padded_routes = padding["rollout_expert_indices"] + assert padded_routes.sequence_lengths.tolist() == [1] + expected = torch.tensor([0, 1, 2], dtype=torch.int16).expand_as(padded_routes.values) + assert torch.equal(padded_routes.values, expected) assert torch.all(padding["router_padding_mask"]) + def test_microbatch_selection_gathers_packed_route_segments(self): + batch = self._make_batch([4, 2], num_actions=2) + batch["rollout_expert_indices"] = PackedTensor.from_segments( + [torch.full((4, 2, 3), 1, dtype=torch.int16), torch.full((2, 2, 3), 2, dtype=torch.int16)] + ) + + microbatch = TokenBasedBatchIterator(batch, max_tokens_per_microbatch=8)._create_microbatch_from_indices([1]) + + routes = microbatch["rollout_expert_indices"] + assert routes.sequence_lengths.tolist() == [2] + assert torch.equal(routes.segment(0), torch.full((2, 2, 3), 2, dtype=torch.int16)) + 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 e562f5ed48..9aae1c5bb7 100644 --- a/tests/backends/skyrl_train/test_train_batch.py +++ b/tests/backends/skyrl_train/test_train_batch.py @@ -7,11 +7,16 @@ from skyrl.backends.skyrl_train.training_batch import ( TensorBatch, + TensorFormat, TensorList, TrainingInput, TrainingInputBatch, pad_training_input_batch, ) +from skyrl.backends.skyrl_train.utils.packed_tensor import ( + PackedTensor, + cu_seqlens_from_lengths, +) def test_train_batch_initialization(): @@ -576,7 +581,11 @@ def _make_full_training_batch(batch_size: int = 4, seq_len: int = 5) -> Training "kl": torch.randn(batch_size, seq_len), "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), + # The fixture is fully attended, so each route segment has seq_len rows. + "rollout_expert_indices": PackedTensor( + torch.randint(0, 8, (batch_size * seq_len, 2, 3), dtype=torch.long), + cu_seqlens_from_lengths([seq_len] * batch_size), + ), "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 @@ -640,9 +649,11 @@ def test_pad_batch_all_fields(): # padding rows are copies of row 0. 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) + assert padded["rollout_expert_indices"][:batch_size] == batch["rollout_expert_indices"] + padded_routes = padded["rollout_expert_indices"][batch_size:] + assert padded_routes.sequence_lengths.tolist() == [seq_len] * pad_size + expected_routes = torch.tensor([0, 1, 2]).expand_as(padded_routes.values) + assert torch.equal(padded_routes.values, expected_routes) regular_tensor_keys = EXPECTED_TRAINING_INPUT_FIELDS - { "loss_mask", @@ -711,3 +722,48 @@ def test_pad_batch_preserves_none_fields(): padded = pad_training_input_batch(batch, pad_size=1) assert padded["values"] is None assert padded.batch_size == 4 + + +def test_packed_tensor_field_survives_the_ray_pickle_round_trip(): + """Packed route buffers cross to the workers via TensorBatch's custom pickle path.""" + segment_lengths = [3, 1, 4] + routes = PackedTensor( + torch.randint(0, 128, (sum(segment_lengths), 2, 3), dtype=torch.int16), + cu_seqlens_from_lengths(segment_lengths), + ) + data = TensorBatch( + { + "sequences": torch.randn(len(segment_lengths), 4), + "rollout_expert_indices": routes, + } + ) + data.metadata = {"response_length": 4} + + unpickled = pickle.loads(pickle.dumps(data)) + + restored = unpickled["rollout_expert_indices"] + assert isinstance(restored, PackedTensor) + assert restored.values.dtype == torch.int16 + assert restored.cu_seqlens.dtype == routes.cu_seqlens.dtype + assert restored == routes + assert unpickled == data + + +def test_serialized_field_formats_are_stable(): + data = TensorBatch( + { + "sequences": torch.randn(2, 4), + "pixel_values": TensorList([torch.randn(1, 3), torch.randn(2, 3)]), + "rollout_expert_indices": PackedTensor( + torch.zeros((3, 2, 3), dtype=torch.int16), cu_seqlens_from_lengths([2, 1]) + ), + "bf16_logprobs": torch.randn(2, 4, dtype=torch.bfloat16), + } + ) + + state = data.__getstate__()["batch_dict"] + + assert state["sequences"]["format"] == TensorFormat.NUMPY + assert state["bf16_logprobs"]["format"] == TensorFormat.TORCH + assert state["pixel_values"]["format"] == TensorFormat.TENSOR_LIST + assert state["rollout_expert_indices"]["format"] == TensorFormat.PACKED_TENSOR diff --git a/tests/backends/skyrl_train/utils/test_packed_tensor.py b/tests/backends/skyrl_train/utils/test_packed_tensor.py new file mode 100644 index 0000000000..f0101cb23a --- /dev/null +++ b/tests/backends/skyrl_train/utils/test_packed_tensor.py @@ -0,0 +1,240 @@ +"""Tests for ``PackedTensor`` batch operations.""" + +import pytest +import torch + +from skyrl.backends.skyrl_train.utils.packed_tensor import ( + CU_SEQLENS_DTYPE, + PackedTensor, + cu_seqlens_from_lengths, + lengths_from_offsets, + row_index_from_offsets, +) + +SEGMENT_LENGTHS = [3, 1, 4, 2] + + +def _segments(lengths=SEGMENT_LENGTHS, *, row_shape=(2, 3)) -> list[torch.Tensor]: + """Distinct rows per segment so any misplacement is visible.""" + segments = [] + next_value = 0 + for length in lengths: + size = length * torch.Size(row_shape).numel() + segments.append(torch.arange(next_value, next_value + size, dtype=torch.int16).reshape(length, *row_shape)) + next_value += size + return segments + + +@pytest.mark.parametrize( + ("lengths", "expected_offsets"), + [ + (SEGMENT_LENGTHS, [0, 3, 4, 8, 10]), + ([0, 2, 0], [0, 0, 2, 2]), + ], +) +def test_offsets_round_trip_lengths_in_int32(lengths, expected_offsets): + offsets = cu_seqlens_from_lengths(lengths) + + assert offsets.tolist() == expected_offsets + assert offsets.dtype == CU_SEQLENS_DTYPE + assert lengths_from_offsets(offsets).tolist() == lengths + assert lengths_from_offsets(offsets).dtype == CU_SEQLENS_DTYPE + + +def test_cu_seqlens_reject_negative_lengths_and_extra_dimensions(): + with pytest.raises(ValueError, match="non-negative"): + cu_seqlens_from_lengths([3, -1]) + with pytest.raises(ValueError, match="must be 1-D"): + cu_seqlens_from_lengths(torch.zeros((2, 2), dtype=torch.int32)) + + +def test_row_index_from_offsets_lays_selected_segments_back_to_back(): + starts = torch.tensor([8, 0]) + lengths = torch.tensor([2, 3]) + + assert row_index_from_offsets(starts, lengths).tolist() == [8, 9, 0, 1, 2] + + +def test_from_segments_round_trips_every_segment(): + segments = _segments() + packed = PackedTensor.from_segments(segments) + + assert len(packed) == len(segments) + assert packed.values.shape == (sum(SEGMENT_LENGTHS), 2, 3) + assert packed.sequence_lengths.tolist() == SEGMENT_LENGTHS + assert packed.row_shape == torch.Size((2, 3)) + assert packed.dtype == torch.int16 + assert packed.device == segments[0].device + for index, segment in enumerate(segments): + assert torch.equal(packed.segment(index), segment) + + +def test_from_segments_rejects_an_empty_batch(): + with pytest.raises(ValueError, match="empty list of segments"): + PackedTensor.from_segments([]) + + +def test_negative_and_out_of_range_segment_indices(): + packed = PackedTensor.from_segments(_segments()) + + assert torch.equal(packed.segment(-1), packed.segment(len(packed) - 1)) + with pytest.raises(IndexError, match="out of range"): + packed.segment(len(packed)) + with pytest.raises(IndexError, match="out of range"): + packed.segment(-len(packed) - 1) + + +def test_integer_index_returns_that_segment(): + segments = _segments() + packed = PackedTensor.from_segments(segments) + + assert torch.equal(packed[2], segments[2]) + assert torch.equal(packed[torch.tensor(2)], segments[2]) + + +@pytest.mark.parametrize("bounds", [(0, 4), (1, 3), (2, 4)]) +def test_contiguous_slice_selects_the_same_segments(bounds): + segments = _segments() + packed = PackedTensor.from_segments(segments) + start, stop = bounds + + sliced = packed[start:stop] + + assert len(sliced) == stop - start + assert sliced.cu_seqlens[0] == 0 + for offset, segment in enumerate(segments[start:stop]): + assert torch.equal(sliced.segment(offset), segment) + + +@pytest.mark.parametrize("indices", [[2, 0], [3, 3, 1], [0, 1, 2, 3], [1]]) +def test_gather_selects_segments_in_the_requested_order(indices): + segments = _segments() + packed = PackedTensor.from_segments(segments) + + for gathered in (packed[torch.tensor(indices)], packed[indices], packed[tuple(indices)]): + assert gathered.sequence_lengths.tolist() == [SEGMENT_LENGTHS[index] for index in indices] + for position, index in enumerate(indices): + assert torch.equal(gathered.segment(position), segments[index]) + + +def test_empty_slice_is_rejected_like_an_empty_tensor_list(): + """``TensorBatch`` fields cannot hold zero batch entries, so neither can a slice.""" + packed = PackedTensor.from_segments(_segments()) + + with pytest.raises(ValueError, match="at least two offsets"): + packed[2:2] + + +def test_strided_slice_falls_back_to_a_gather(): + segments = _segments() + packed = PackedTensor.from_segments(segments) + + strided = packed[::2] + + assert len(strided) == 2 + assert torch.equal(strided.segment(0), segments[0]) + assert torch.equal(strided.segment(1), segments[2]) + + +def test_cat_joins_batches_end_to_end(): + left = PackedTensor.from_segments(_segments([3, 1])) + right = PackedTensor.from_segments(_segments([4, 2])) + + joined = PackedTensor.cat([left, right]) + + assert joined.sequence_lengths.tolist() == [3, 1, 4, 2] + assert torch.equal(joined.segment(0), left.segment(0)) + assert torch.equal(joined.segment(2), right.segment(0)) + + +def test_cat_rejects_an_empty_list(): + with pytest.raises(ValueError, match="empty list of packed batches"): + PackedTensor.cat([]) + + +def test_repeat_tiles_and_repeat_interleave_duplicates(): + packed = PackedTensor.from_segments(_segments([3, 1])) + + tiled = packed.repeat(2) + interleaved = packed.repeat_interleave(2) + + assert tiled.sequence_lengths.tolist() == [3, 1, 3, 1] + assert interleaved.sequence_lengths.tolist() == [3, 3, 1, 1] + assert torch.equal(tiled.segment(2), packed.segment(0)) + assert torch.equal(interleaved.segment(1), packed.segment(0)) + + +def test_to_and_contiguous_preserve_the_batch(): + packed = PackedTensor.from_segments(_segments()) + + widened = packed.to(dtype=torch.int32) + + assert widened.dtype == torch.int32 + assert widened.cu_seqlens.dtype == CU_SEQLENS_DTYPE + assert torch.equal(widened.values, packed.values.to(torch.int32)) + assert packed.contiguous() == packed + + +def test_equality_compares_values_and_offsets(): + packed = PackedTensor.from_segments(_segments()) + + assert packed == PackedTensor.from_segments(_segments()) + assert packed != PackedTensor.from_segments(_segments([3, 1, 4, 2], row_shape=(1, 3))) + assert packed != PackedTensor.from_segments(_segments([4, 4, 2])) + assert packed != packed.values + + +def test_rejects_mismatched_device_or_offset_dtype(): + values = torch.zeros((4, 2), dtype=torch.int16) + + with pytest.raises(ValueError, match="must be torch.int32"): + PackedTensor(values, torch.tensor([0, 4], dtype=torch.int64)) + with pytest.raises(ValueError, match="at least two offsets"): + PackedTensor(values, torch.tensor([0], dtype=CU_SEQLENS_DTYPE)) + with pytest.raises(ValueError, match="at least two offsets"): + PackedTensor(values, torch.zeros((2, 2), dtype=CU_SEQLENS_DTYPE)) + + +def test_rejects_offsets_that_do_not_span_the_buffer(): + values = torch.zeros((4, 2), dtype=torch.int16) + + with pytest.raises(ValueError, match="must run from 0"): + PackedTensor(values, torch.tensor([0, 3], dtype=CU_SEQLENS_DTYPE)) + with pytest.raises(ValueError, match="must run from 0"): + PackedTensor(values, torch.tensor([1, 4], dtype=CU_SEQLENS_DTYPE)) + + +def test_rejects_values_without_a_token_row_dimension(): + with pytest.raises(ValueError, match="token-row dimension"): + PackedTensor(torch.tensor(1), torch.tensor([0, 1], dtype=CU_SEQLENS_DTYPE)) + + +def test_segments_and_contiguous_slices_are_views(): + """Reading a batch entry must not copy: alignment walks every segment per micro-batch.""" + packed = PackedTensor.from_segments(_segments()) + + assert packed.segment(1).data_ptr() == packed.values[3].data_ptr() + assert packed[1:3].values.data_ptr() == packed.values[3].data_ptr() + + +@pytest.mark.parametrize( + "operation", + [ + lambda packed: packed[torch.tensor([1, 0])], + lambda packed: packed[::2], + lambda packed: packed.repeat(2), + lambda packed: packed.repeat_interleave(2), + lambda packed: PackedTensor.cat([packed, packed]), + lambda packed: packed.to(dtype=torch.int32), + ], + ids=["gather", "strided_slice", "repeat", "repeat_interleave", "cat", "to_dtype"], +) +def test_reordering_operations_allocate_rather_than_alias(operation): + """A duplicated or reordered segment must own its rows; the source stays untouched.""" + packed = PackedTensor.from_segments(_segments()) + original = packed.values.clone() + + produced = operation(packed) + produced.values[:] = -1 + + assert torch.equal(packed.values, original) diff --git a/tests/backends/skyrl_train/utils/test_replay_utils.py b/tests/backends/skyrl_train/utils/test_replay_utils.py index 304e11c625..f713908a92 100644 --- a/tests/backends/skyrl_train/utils/test_replay_utils.py +++ b/tests/backends/skyrl_train/utils/test_replay_utils.py @@ -3,7 +3,6 @@ import types from types import SimpleNamespace -import numpy as np import pytest import torch @@ -11,12 +10,19 @@ build_token_metadata_layout, ) from skyrl.backends.skyrl_train.utils import replay_utils +from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor from skyrl.backends.skyrl_train.utils.replay_utils import ( + append_packed_replay_padding, + make_packed_replay_padding, make_replay_padding_indices, - make_replay_padding_indices_np, ) +def _pack_routes(routes: torch.Tensor, attention_mask: torch.Tensor) -> PackedTensor: + """Pack a ``[batch, seq_len, layers, topk]`` fixture to its real tokens.""" + return PackedTensor.from_segments([routes[row][attention_mask[row].bool()] for row in range(routes.shape[0])]) + + @pytest.fixture def parallel_state(monkeypatch): try: @@ -70,22 +76,36 @@ def test_replay_padding_indices_are_unique(dtype): assert torch.equal(padding, torch.tensor([0, 1, 2], dtype=dtype).expand_as(padding)) -@pytest.mark.parametrize("dtype", [np.uint8, np.int16, np.int32]) -def test_numpy_replay_padding_matches_torch(dtype): - torch_dtype = getattr(torch, np.dtype(dtype).name) - padding = make_replay_padding_indices_np((2, 3, 4, 3), dtype=np.dtype(dtype)) +@pytest.mark.parametrize("shape", [(), (2, 3, 4, 0)]) +def test_replay_padding_rejects_missing_topk(shape): + with pytest.raises(ValueError, match="positive topk"): + make_replay_padding_indices(shape, dtype=torch.uint8) + + +@pytest.mark.parametrize("segment_lengths", [[1, 1], [3], [2, 5, 1]]) +def test_packed_replay_padding_matches_the_reference_row_shape(segment_lengths): + reference = PackedTensor.from_segments([torch.full((4, 2, 3), 9, dtype=torch.int16)]) - assert padding.dtype == dtype - assert torch.equal( - torch.from_numpy(padding), - make_replay_padding_indices((2, 3, 4, 3), dtype=torch_dtype), + padding = make_packed_replay_padding(reference, segment_lengths=segment_lengths) + + assert padding.sequence_lengths.tolist() == segment_lengths + assert padding.row_shape == reference.row_shape + assert padding.dtype == reference.dtype + assert torch.equal(padding.values, torch.tensor([0, 1, 2], dtype=torch.int16).expand_as(padding.values)) + + +@pytest.mark.parametrize("pad_count", [1, 3]) +def test_appending_replay_padding_keeps_the_real_segments_and_the_arange_invariant(pad_count): + routes = PackedTensor.from_segments( + [torch.full((4, 2, 3), 9, dtype=torch.int16), torch.full((2, 2, 3), 8, dtype=torch.int16)] ) + padded = append_packed_replay_padding(routes, segment_lengths=[1] * pad_count) -@pytest.mark.parametrize("shape", [(), (2, 3, 4, 0)]) -def test_numpy_replay_padding_rejects_missing_topk(shape): - with pytest.raises(ValueError, match="positive topk"): - make_replay_padding_indices_np(shape, dtype=np.dtype(np.uint8)) + assert padded.sequence_lengths.tolist() == [4, 2] + [1] * pad_count + assert padded[: len(routes)] == routes + appended = padded[len(routes) :] + assert torch.equal(appended.values, torch.tensor([0, 1, 2], dtype=torch.int16).expand_as(appended.values)) def test_replay_has_no_dispatcher_specific_patch(): @@ -121,15 +141,14 @@ class RouterReplayAction: "scatter_router_padding_mask_for_model", lambda mask, model, model_config: mask, ) - apply_layout = replay_utils.align_token_metadata + apply_layout = replay_utils.align_packed_token_metadata routed_layer_counts = [] def record_routed_layer_count(metadata, layout, padding_value): - if metadata.ndim == 4: - routed_layer_counts.append(metadata.shape[2]) + routed_layer_counts.append(metadata.row_shape[0]) return apply_layout(metadata, layout, padding_value) - monkeypatch.setattr(replay_utils, "align_token_metadata", record_routed_layer_count) + monkeypatch.setattr(replay_utils, "align_packed_token_metadata", record_routed_layer_count) routes = torch.tensor( [ @@ -152,7 +171,7 @@ def record_routed_layer_count(metadata, layout, padding_value): ) model_kwargs = replay_utils.setup_per_microbatch_replay_forward( - routes, + _pack_routes(routes, attention_mask), router_padding_mask, attention_mask, model=object(), @@ -215,7 +234,7 @@ def run(routes): fp8_enabled=False, ) replay_utils.setup_per_microbatch_replay_forward( - routes, + _pack_routes(routes, attention_mask), router_padding_mask, attention_mask, model=object(), diff --git a/tests/train/dataset/test_parallel_fill.py b/tests/train/dataset/test_parallel_fill.py new file mode 100644 index 0000000000..1b0e6daedd --- /dev/null +++ b/tests/train/dataset/test_parallel_fill.py @@ -0,0 +1,74 @@ +""" +uv run --isolated --extra dev pytest tests/train/dataset/test_parallel_fill.py +""" + +import threading + +import pytest + +from skyrl.train.dataset.parallel_fill import fill_batch_rows + + +@pytest.mark.parametrize("workers", [None, 1, 4, 64]) +def test_every_index_is_filled_exactly_once(workers): + num_rows = 32 + calls = [0] * num_rows + + def fill_row(index: int) -> None: + calls[index] += 1 + + fill_batch_rows(fill_row, num_rows, workers=workers) + + assert calls == [1] * num_rows + + +def test_single_worker_runs_serially_on_the_calling_thread(): + order = [] + threads = set() + + def fill_row(index: int) -> None: + order.append(index) + threads.add(threading.current_thread()) + + fill_batch_rows(fill_row, 4, workers=1) + + assert order == [0, 1, 2, 3] + assert threads == {threading.current_thread()} + + +def test_multiple_workers_leave_the_calling_thread(): + threads = set() + + def fill_row(index: int) -> None: + threads.add(threading.current_thread()) + + fill_batch_rows(fill_row, 8, workers=4) + + assert threading.current_thread() not in threads + + +def test_zero_rows_is_a_no_op(): + def fill_row(index: int) -> None: + raise AssertionError("fill_row must not be called for an empty batch") + + fill_batch_rows(fill_row, 0) + + +def test_negative_row_count_raises(): + with pytest.raises(ValueError, match="row count must be non-negative"): + fill_batch_rows(lambda index: None, -1) + + +def test_non_positive_worker_count_raises(): + with pytest.raises(ValueError, match="worker count must be positive"): + fill_batch_rows(lambda index: None, 4, workers=0) + + +@pytest.mark.parametrize("workers", [1, 4]) +def test_callback_exception_propagates(workers): + def fill_row(index: int) -> None: + if index == 2: + raise RuntimeError("row 2 failed") + + with pytest.raises(RuntimeError, match="row 2 failed"): + fill_batch_rows(fill_row, 8, workers=workers) diff --git a/tests/train/dataset/test_preprocess.py b/tests/train/dataset/test_preprocess.py index c5e885ebdd..91e7db7b1e 100644 --- a/tests/train/dataset/test_preprocess.py +++ b/tests/train/dataset/test_preprocess.py @@ -2,13 +2,17 @@ uv run --isolated --extra dev pytest tests/train/dataset/test_preprocess.py """ +import logging +from typing import List from unittest.mock import MagicMock import numpy as np import pytest import torch +from skyrl.backends.skyrl_train.utils.routed_experts import ROUTED_EXPERT_DTYPES from skyrl.train.dataset.preprocess import ( + ROUTED_EXPERT_TORCH_DTYPES, convert_prompts_responses_to_batch_tensors, make_router_padding_mask, ) @@ -94,9 +98,10 @@ def test_routed_expert_tensor_uses_unique_dummy_routes(tokenizer): rollout_expert_indices=routes, ) - assert routed.shape == (2, 3, 2, 2) + assert routed.values.shape == (6, 2, 2) + assert routed.cu_seqlens.tolist() == [0, 3, 6] assert routed.dtype == torch.uint8 - assert routed[0, 2].tolist() == [[0, 1], [0, 1]] + assert routed.segment(0)[2].tolist() == [[0, 1], [0, 1]] @pytest.mark.parametrize( @@ -124,7 +129,7 @@ def test_routed_expert_tensor_promotes_mixed_batch_dtype( ) assert routed.dtype == expected_dtype - assert routed[1, 0].tolist() == [[max_expert_id, max_expert_id + 1]] + assert routed.segment(1)[0].tolist() == [[max_expert_id, max_expert_id + 1]] def test_routed_expert_tensor_accepts_read_only_arrays(tokenizer): @@ -141,7 +146,7 @@ def test_routed_expert_tensor_accepts_read_only_arrays(tokenizer): ) assert routed.dtype == torch.uint8 - assert routed.tolist() == [[[[1, 2]], [[3, 4]]]] + assert routed.segment(0).tolist() == [[[1, 2]], [[3, 4]]] def test_routed_expert_tensor_rejects_nested_lists(tokenizer): @@ -156,15 +161,11 @@ def test_routed_expert_tensor_rejects_nested_lists(tokenizer): ) -@pytest.mark.parametrize("dtype", [np.uint16, np.int64]) -def test_routed_expert_tensor_narrows_wide_dtypes(tokenizer, dtype): - """Wide dtypes are compacted, not rejected. - - The wire decoder already restricts routes to uint8/int16/int32, so this path - only sees a wide dtype from a hand-built generator -- narrowing it is more - useful than refusing it. - """ - routes = np.asarray([[[1, 2]], [[3, 4]]], dtype=dtype) +def test_routed_expert_tensor_accepts_non_contiguous_arrays(tokenizer): + # Every other expert column, which leaves a non-contiguous view. + base = np.asarray([[[1, 9, 2, 9]], [[3, 9, 4, 9]]], dtype=np.uint8) + routes = base[:, :, ::2] + assert not routes.flags.c_contiguous *_, routed = convert_prompts_responses_to_batch_tensors( tokenizer.pad_token_id, @@ -176,12 +177,27 @@ def test_routed_expert_tensor_narrows_wide_dtypes(tokenizer, dtype): ) assert routed.dtype == torch.uint8 - assert routed.tolist() == [[[[1, 2]], [[3, 4]]]] + assert routed.segment(0).tolist() == [[[1, 2]], [[3, 4]]] + + +@pytest.mark.parametrize("dtype", [np.uint16, np.int64]) +def test_routed_expert_tensor_rejects_non_canonical_dtypes(tokenizer, dtype): + """The sender compacts to the canonical dtype, so collation validates instead of rescanning.""" + routes = np.asarray([[[1, 2]], [[3, 4]]], dtype=dtype) + + with pytest.raises(ValueError, match="canonical routed-expert dtype"): + convert_prompts_responses_to_batch_tensors( + tokenizer.pad_token_id, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[routes], + ) -def test_routed_expert_tensor_retightens_after_truncation(tokenizer): - """A truncated array can fit a narrower dtype than the wire declared.""" - # int16 on the wire because of the trailing 300, which truncation then drops. +def test_routed_expert_tensor_keeps_the_sender_dtype_after_truncation(tokenizer): + # The dropped trailing value required int16 on the sender. routes = np.asarray([[[1, 2]], [[3, 4]], [[300, 5]]], dtype=np.int16) *_, routed = convert_prompts_responses_to_batch_tensors( @@ -193,8 +209,86 @@ def test_routed_expert_tensor_retightens_after_truncation(tokenizer): rollout_expert_indices=[routes[:2]], ) - assert routed.dtype == torch.uint8 - assert routed.tolist() == [[[[1, 2]], [[3, 4]]]] + assert routed.dtype == torch.int16 + assert routed.segment(0).tolist() == [[[1, 2]], [[3, 4]]] + + +@pytest.mark.parametrize( + ("dtype", "expert_id", "expect_warning"), + [(np.int16, 300, False), (np.int32, 2**16, True)], +) +def test_routed_expert_tensor_warns_only_on_an_int32_batch(tokenizer, caplog, dtype, expert_id, expect_warning): + routes = np.asarray([[[expert_id, expert_id + 1]]], dtype=dtype) + + with caplog.at_level(logging.WARNING, logger="skyrl.train.dataset.preprocess"): + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer.pad_token_id, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[routes], + ) + + assert routed.dtype == ROUTED_EXPERT_TORCH_DTYPES[np.dtype(dtype)] + assert ("not compacting its routes" in caplog.text) is expect_warning + + +def test_routed_expert_torch_dtype_map_covers_the_canonical_dtypes(): + assert set(ROUTED_EXPERT_TORCH_DTYPES) == set(ROUTED_EXPERT_DTYPES) + + +def _numpy_padded_routes( + routes: List[np.ndarray], + prompts: List[List[int]], + responses: List[List[int]], +) -> np.ndarray: + """Reference NumPy implementation of route collation.""" + max_total = max(len(prompt) + len(response) for prompt, response in zip(prompts, responses)) + num_layers, topk = routes[0].shape[1:] + batch_dtype = max((sample.dtype for sample in routes), key=lambda dtype: dtype.itemsize) + padded = np.empty((len(routes), max_total, num_layers, topk), dtype=batch_dtype) + padded[...] = np.arange(topk, dtype=batch_dtype) + for index, sample in enumerate(routes): + left_pad = max_total - (len(prompts[index]) + len(responses[index])) + padded[index, left_pad : left_pad + sample.shape[0]] = sample + return padded + + +def test_routed_expert_tensor_is_bit_identical_to_numpy_collation(tokenizer): + """Packed collation matches the real rows of the padded NumPy reference.""" + prompts = [[1, 2], [3, 4, 5, 6]] + responses = [[10, 11, 12], [20, 21]] + num_layers, topk = 2, 3 + # Sample 0 has 5 tokens but only 4 captured route rows, so its segment pads at the end. + routes = [ + np.arange(4 * num_layers * topk, dtype=np.uint8).reshape(4, num_layers, topk), + (np.arange(6 * num_layers * topk, dtype=np.int16) + 300).reshape(6, num_layers, topk), + ] + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer.pad_token_id, + prompts, + responses, + rewards=[[0.0] * 3, [0.0] * 2], + loss_masks=[[1] * 3, [1] * 2], + rollout_expert_indices=routes, + ) + + padded = _numpy_padded_routes(routes, prompts, responses) + real_rows = np.concatenate( + [ + padded[index, padded.shape[1] - (len(prompt) + len(response)) :] + for index, (prompt, response) in enumerate(zip(prompts, responses)) + ] + ) + assert routed.dtype == torch.int16 + assert routed.cu_seqlens.tolist() == [0, 5, 11] + assert torch.equal(routed.values, torch.from_numpy(real_rows)) + # Padding routes use distinct experts for Megatron's dropless dispatcher. + padding_row = [[0, 1, 2]] * num_layers + assert routed.segment(0)[4].tolist() == padding_row + assert not torch.equal(routed.segment(0)[4], torch.zeros_like(routed.segment(0)[4])) def test_convert_prompts_responses_to_batch_tensors_exact(tokenizer): @@ -411,17 +505,8 @@ def test_max_seq_len_warns_but_does_not_truncate(tokenizer): assert action.shape == (2, 50) -# --------------------------------------------------------------------------- -# R3 (Router Replay) — rollout_expert_indices padding tests -# --------------------------------------------------------------------------- - - def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): - """rollout_expert_indices tensor should have shape [batch, max_total, layers, topk] - with left-padding aligned to the attention_mask.""" - # Sample 0: prompt=2, response=3 → total=5 - # Sample 1: prompt=4, response=2 → total=6 - # max_total=6 + """Routes pack to [sum(seq_len), layers, topk] with one cu_seqlens segment per trajectory.""" prompts = [[1, 2], [3, 4, 5, 6]] responses = [[10, 11, 12], [20, 21]] rewards = [[0.0] * 3, [0.0] * 2] @@ -429,10 +514,8 @@ def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): num_layers = 2 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 = 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 + rei_0 = np.asarray([[[1, 2]] * num_layers for _ in range(5)], dtype=np.uint8) + rei_1 = np.asarray([[[3, 4]] * num_layers for _ in range(6)], dtype=np.uint8) seq, attn, action, rew, lm, lp, rei_tensor = convert_prompts_responses_to_batch_tensors( tokenizer.pad_token_id, @@ -444,24 +527,11 @@ def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): ) assert rei_tensor is not None - # Shape: [batch=2, max_total=6, layers=2, topk=2] - assert rei_tensor.shape == (2, 6, num_layers, topk) - - 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 - - # 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() == dummy_routes - else: - assert rei_tensor[i, pos].tolist() != dummy_routes + assert rei_tensor.values.shape == (11, num_layers, topk) + assert rei_tensor.cu_seqlens.tolist() == [0, 5, 11] + assert rei_tensor.sequence_lengths.tolist() == attn.sum(dim=1).tolist() + assert rei_tensor.segment(0).tolist() == [[[1, 2]] * num_layers] * 5 + assert rei_tensor.segment(1).tolist() == [[[3, 4]] * num_layers] * 6 def test_rollout_expert_indices_none_when_not_provided(tokenizer): diff --git a/tests/train/generators/test_datatypes.py b/tests/train/generators/test_datatypes.py index 12c9efdba2..580cbd97f5 100644 --- a/tests/train/generators/test_datatypes.py +++ b/tests/train/generators/test_datatypes.py @@ -31,7 +31,6 @@ def test_turn_output(output_ids, observation_ids, output_logprobs, added_eos, ex output_logprobs=output_logprobs, new_obs=[], obs_ids=observation_ids, - rollout_expert_indices=None, added_eos=added_eos, reward=1.0, ) diff --git a/tests/train/generators/test_skyrl_gym_generator.py b/tests/train/generators/test_skyrl_gym_generator.py index 6d03d937b6..b253f02433 100644 --- a/tests/train/generators/test_skyrl_gym_generator.py +++ b/tests/train/generators/test_skyrl_gym_generator.py @@ -23,20 +23,17 @@ MOCK_TOKENIZER_ENCODED_IDS = [1, 2, 3, 4] -def test_turn_output_keeps_uncaptured_suffix_out_of_routes(): - routes = np.asarray([[[2, 3]], [[4, 5]]], dtype=np.uint8) +def test_turn_output_masks_uncaptured_suffix(): 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] @@ -413,6 +410,63 @@ def mock_generate(_, model=None): assert output.stop_reason == "stop" +@pytest.mark.asyncio +@patch("skyrl_gym.make") +async def test_agent_loop_uses_incremental_routed_expert_trace( + mock_make, + mock_tokenizer, + mock_llm, + mock_env, + generator_cfg, + mock_env_cfg, +): + generator_cfg.batched = False + generator_cfg.max_turns = 2 + generator_cfg.use_conversation_multi_turn = True + generator_cfg.inference_engine.enable_return_routed_experts = True + mock_make.return_value = mock_env + mock_env.init.return_value = ([{"role": "user", "content": "Initial input"}], {}) + + mock_env.step.side_effect = [ + BaseTextEnvStepOutput(observations=[{"role": "user", "content": "next"}], reward=1.0, done=done, metadata={}) + for done in (False, True) + ] + prompt_starts = [] + + def generate(input_batch, model=None): + prompt_tokens = input_batch["prompt_token_ids"][0] + prompt_start = input_batch["routed_experts_prompt_starts"][0] + prompt_starts.append(prompt_start) + output_ids = [10, 11] + num_route_rows = len(prompt_tokens) - prompt_start + len(output_ids) - 1 + routes = np.arange(num_route_rows * 4, dtype=np.int32).reshape(num_route_rows, 2, 2) % 8 + return { + "responses": ["mocked output"], + "response_ids": [output_ids], + "stop_reasons": ["stop"], + "rollout_expert_indices": [routes], + } + + mock_llm.generate = AsyncMock(side_effect=generate) + generator = SkyRLGymGenerator( + generator_cfg=generator_cfg, + skyrl_gym_cfg=mock_env_cfg, + inference_engine_client=mock_llm, + tokenizer=mock_tokenizer, + ) + generator.base_conversation_token_ids = [] + + await generator.agent_loop( + [{"role": "user", "content": "Start"}], + mock_env_cfg.env_class, + {}, + max_tokens=32, + max_input_length=64, + ) + + assert prompt_starts == [0, 5] + + @pytest.mark.asyncio @patch("skyrl_gym.make") async def test_generate_batched(mock_make, mock_tokenizer, mock_llm, mock_env, generator_cfg, mock_env_cfg): diff --git a/tests/train/test_config.py b/tests/train/test_config.py index 5afc726c43..7e4dc1a493 100644 --- a/tests/train/test_config.py +++ b/tests/train/test_config.py @@ -20,7 +20,11 @@ build_nested_dataclass, overrides_dict_to_dotlist, ) -from skyrl.train.utils.utils import validate_cfg, validate_inference_engine_cfg +from skyrl.train.utils.utils import ( + validate_cfg, + validate_inference_engine_cfg, + validate_megatron_cfg, +) from tests.train.util import example_dummy_config @@ -850,3 +854,35 @@ def test_delta_weight_sync_defaults(self): # `publish_staging_dir` and `local_checkpoint_dir` should be constructed based on `sync_dir` assert "my_sync_dir" in cfg.publish_staging_dir assert "my_sync_dir" in cfg.local_checkpoint_dir + + +class TestMegatronRouterReplayValidation: + @staticmethod + def _cfg(): + cfg = _make_validated_test_config() + cfg.trainer.strategy = "megatron" + cfg.generator.inference_engine.enable_return_routed_experts = True + cfg.trainer.policy.megatron_config.moe_enable_routing_replay = True + return cfg + + @pytest.mark.parametrize("vpp_size", [1, 2]) + def test_routing_replay_refuses_virtual_pipeline_parallelism(self, vpp_size): + cfg = self._cfg() + cfg.trainer.policy.megatron_config.transformer_config_kwargs["virtual_pipeline_model_parallel_size"] = vpp_size + + with pytest.raises(AssertionError, match="virtual_pipeline_model_parallel_size"): + validate_megatron_cfg(cfg) + + @pytest.mark.parametrize("vpp_size", [None, 0]) + def test_routing_replay_allows_unset_virtual_pipeline_parallelism(self, vpp_size): + cfg = self._cfg() + cfg.trainer.policy.megatron_config.transformer_config_kwargs["virtual_pipeline_model_parallel_size"] = vpp_size + + validate_megatron_cfg(cfg) + + def test_virtual_pipeline_parallelism_allowed_without_routing_replay(self): + cfg = self._cfg() + cfg.trainer.policy.megatron_config.moe_enable_routing_replay = False + cfg.trainer.policy.megatron_config.transformer_config_kwargs["virtual_pipeline_model_parallel_size"] = 2 + + validate_megatron_cfg(cfg) diff --git a/tests/train/test_packed_route_collation_equivalence.py b/tests/train/test_packed_route_collation_equivalence.py new file mode 100644 index 0000000000..5d8c9ac679 --- /dev/null +++ b/tests/train/test_packed_route_collation_equivalence.py @@ -0,0 +1,334 @@ +"""Compare packed route collation with the padded reference path end to end.""" + +import sys +import types +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from skyrl.backends.skyrl_train.distributed.megatron.token_metadata import ( + align_token_metadata, + build_token_metadata_layout, +) +from skyrl.backends.skyrl_train.training_batch import ( + TrainingInputBatch, + pad_training_input_batch, +) +from skyrl.backends.skyrl_train.utils import replay_utils +from skyrl.backends.skyrl_train.utils.replay_utils import ( + _split_replay_indices, + make_replay_padding_indices, + replay_padding_row, +) +from skyrl.train.dataset.preprocess import ( + convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, +) + +NUM_LAYERS = 3 +TOPK = 4 +PAD_TOKEN_ID = 0 +# Above 2**8 so the compact dtype is int16, matching the production route width. +MIN_EXPERT_ID = 300 + + +def _reference_padded_routes( + routes: list[np.ndarray], + prompt_lens: list[int], + response_lens: list[int], +) -> torch.Tensor: + """Collate routes into a left-padded batch-major reference tensor.""" + max_total = max(p + r for p, r in zip(prompt_lens, response_lens)) + dtype = max((entry.dtype for entry in routes), key=lambda d: d.itemsize) + torch_dtype = torch.from_numpy(np.empty(0, dtype=dtype)).dtype + padded = torch.empty((len(routes), max_total, NUM_LAYERS, TOPK), dtype=torch_dtype) + padding_row = torch.arange(TOPK, dtype=torch_dtype) + for sample_index, entry in enumerate(routes): + left_pad = max_total - (prompt_lens[sample_index] + response_lens[sample_index]) + route_end = left_pad + entry.shape[0] + padded[sample_index, :left_pad] = padding_row + padded[sample_index, left_pad:route_end] = torch.from_numpy(entry) + padded[sample_index, route_end:] = padding_row + return padded + + +def _reference_replay_data( + padded_routes: torch.Tensor, + attention_mask: torch.Tensor, + local_layers: list[int], + *, + packed: bool, + tp_size: int, + tp_rank: int, +) -> list[torch.Tensor]: + """Gather real-token routes from the padded reference tensor.""" + layout = build_token_metadata_layout( + attention_mask, + padded_routes.device, + packed=packed, + fp8_enabled=False, + ) + local = padded_routes.index_select(2, torch.tensor(local_layers, dtype=torch.long)) + aligned = align_token_metadata( + local, + layout, + replay_padding_row(TOPK, dtype=padded_routes.dtype), + ) + if tp_size > 1: + chunk = aligned.shape[1] // tp_size + aligned = aligned[:, tp_rank * chunk : (tp_rank + 1) * chunk, :, :] + return _split_replay_indices(aligned) + + +@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_tensor_model_parallel_rank", lambda: 0, 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 + + +@pytest.fixture +def router_replay(monkeypatch): + """Capture what ``setup_per_microbatch_replay_forward`` hands Megatron.""" + module = types.ModuleType("megatron.core.transformer.moe.router_replay") + + class RouterReplay: + global_router_replay_instances = [object() for _ in range(NUM_LAYERS)] + replay_data: list[torch.Tensor] | None = None + + @classmethod + def set_replay_data(cls, replay_data): + cls.replay_data = replay_data + + @classmethod + def set_global_router_replay_action(cls, action): + pass + + module.RouterReplay = RouterReplay + module.RouterReplayAction = SimpleNamespace(REPLAY_FORWARD="replay_forward") + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router_replay", module) + monkeypatch.setattr( + replay_utils, + "scatter_router_padding_mask_for_model", + lambda mask, model, model_config: mask, + ) + return RouterReplay + + +def _make_batch(lengths: list[tuple[int, int]], *, captured_shortfall: int = 0, seed: int = 0): + """Build one batch of trajectories with the given ``(prompt_len, response_len)`` pairs. + + ``captured_shortfall`` leaves that many trailing tokens of the last trajectory without a + captured route, exercising the dummy-row tail that both paths must fill identically. + """ + rng = np.random.default_rng(seed) + prompts, responses, rewards, loss_masks, routes = [], [], [], [], [] + for index, (prompt_len, response_len) in enumerate(lengths): + prompts.append(list(rng.integers(1, 1000, size=prompt_len))) + responses.append(list(rng.integers(1, 1000, size=response_len))) + rewards.append([0.0] * response_len) + loss_masks.append([1] * response_len) + captured = prompt_len + response_len + if index == len(lengths) - 1: + captured -= captured_shortfall + routes.append( + rng.integers( + MIN_EXPERT_ID, + MIN_EXPERT_ID + 2000, + size=(captured, NUM_LAYERS, TOPK), + dtype=np.int16, + ) + ) + return prompts, responses, rewards, loss_masks, routes + + +def _run_both_paths( + lengths: list[tuple[int, int]], + *, + packed: bool, + tp_size: int, + local_layers: list[int], + captured_shortfall: int = 0, + batch_pad_size: int = 0, + stage_range: tuple[int, int] = (0, NUM_LAYERS), + monkeypatch, + parallel_state, + router_replay, +) -> tuple[list[torch.Tensor], list[torch.Tensor]]: + prompts, responses, rewards, loss_masks, routes = _make_batch(lengths, captured_shortfall=captured_shortfall) + ( + sequences, + attention_mask, + response_mask, + _rewards, + loss_mask, + _logprobs, + packed_routes, + ) = convert_prompts_responses_to_batch_tensors( + PAD_TOKEN_ID, + prompts, + responses, + rewards, + loss_masks, + rollout_expert_indices=routes, + ) + router_padding_mask = make_router_padding_mask(attention_mask, [entry.shape[0] for entry in routes]) + padded_routes = _reference_padded_routes(routes, [len(p) for p in prompts], [len(r) for r in responses]) + + if batch_pad_size: + batch = TrainingInputBatch( + { + "sequences": sequences, + "attention_mask": attention_mask, + "response_mask": response_mask, + "loss_mask": loss_mask, + "rollout_expert_indices": packed_routes, + "router_padding_mask": router_padding_mask, + } + ) + batch.metadata = {"uids": [f"u{index}" for index in range(len(prompts))]} + batch = pad_training_input_batch(batch, batch_pad_size) + attention_mask = batch["attention_mask"] + router_padding_mask = batch["router_padding_mask"] + packed_routes = batch["rollout_expert_indices"] + # Match the dummy rows added by batch padding in the packed path. + padded_routes = torch.cat( + [ + padded_routes, + make_replay_padding_indices((batch_pad_size, *padded_routes.shape[1:]), dtype=padded_routes.dtype), + ], + dim=0, + ) + + tp_rank = tp_size - 1 + monkeypatch.setattr(parallel_state, "get_tensor_model_parallel_world_size", lambda: tp_size, raising=False) + monkeypatch.setattr(parallel_state, "get_tensor_model_parallel_rank", lambda: tp_rank, raising=False) + router_replay.global_router_replay_instances = [object() for _ in local_layers] + monkeypatch.setattr(replay_utils, "_get_current_pp_stage_layer_range", lambda model_config: stage_range) + + layout = build_token_metadata_layout(attention_mask, packed_routes.device, packed=packed, fp8_enabled=False) + replay_utils.setup_per_microbatch_replay_forward( + packed_routes, + router_padding_mask, + attention_mask, + model=object(), + model_config=SimpleNamespace(fp8=None, sequence_parallel=False), + metadata_layout=layout, + remove_microbatch_padding=packed, + ) + new_replay_data = [tensor.clone() for tensor in router_replay.replay_data] + + reference = _reference_replay_data( + padded_routes, + attention_mask, + local_layers, + packed=packed, + tp_size=tp_size, + tp_rank=tp_rank, + ) + return new_replay_data, reference + + +def _assert_bit_identical(new_data: list[torch.Tensor], reference: list[torch.Tensor]) -> None: + assert len(new_data) == len(reference) + for slot, (produced, expected) in enumerate(zip(new_data, reference, strict=True)): + assert produced.dtype == expected.dtype == torch.int32, slot + assert produced.shape == expected.shape, (slot, produced.shape, expected.shape) + assert torch.equal(produced, expected), slot + + +LENGTH_DISTRIBUTIONS = { + # No padding at all: the packed and padded layouts coincide. + "uniform": [(8, 8), (8, 8), (8, 8), (8, 8)], + "mild_ragged": [(8, 8), (7, 8), (8, 6), (6, 7)], + # Typical RL: an order of magnitude between the shortest and longest trajectory. + "typical_rl": [(2, 2), (8, 24), (4, 6), (1, 31)], + "heavy_tail": [(1, 1), (1, 2), (2, 1), (16, 32)], +} + + +@pytest.mark.parametrize("distribution", sorted(LENGTH_DISTRIBUTIONS)) +@pytest.mark.parametrize("packed", [False, True]) +@pytest.mark.parametrize("tp_size", [1, 2]) +def test_packed_routes_match_padded_rectangle( + monkeypatch, parallel_state, router_replay, distribution, packed, tp_size +): + new_data, reference = _run_both_paths( + LENGTH_DISTRIBUTIONS[distribution], + packed=packed, + tp_size=tp_size, + local_layers=list(range(NUM_LAYERS)), + monkeypatch=monkeypatch, + parallel_state=parallel_state, + router_replay=router_replay, + ) + _assert_bit_identical(new_data, reference) + + +@pytest.mark.parametrize( + ("distribution", "captured_shortfall", "batch_pad_size", "local_layers", "stage_range"), + [ + pytest.param("typical_rl", 3, 0, list(range(NUM_LAYERS)), (0, NUM_LAYERS), id="uncaptured_suffix"), + pytest.param("mild_ragged", 0, 2, list(range(NUM_LAYERS)), (0, NUM_LAYERS), id="batch_padding"), + pytest.param("typical_rl", 0, 0, [1, 2], (1, 2), id="pipeline_stage_subset"), + ], +) +@pytest.mark.parametrize("packed", [False, True]) +def test_packed_routes_match_edge_cases( + monkeypatch, + parallel_state, + router_replay, + distribution, + captured_shortfall, + batch_pad_size, + local_layers, + stage_range, + packed, +): + new_data, reference = _run_both_paths( + LENGTH_DISTRIBUTIONS[distribution], + packed=packed, + tp_size=1, + local_layers=local_layers, + captured_shortfall=captured_shortfall, + batch_pad_size=batch_pad_size, + stage_range=stage_range, + monkeypatch=monkeypatch, + parallel_state=parallel_state, + router_replay=router_replay, + ) + _assert_bit_identical(new_data, reference) + + +@pytest.mark.parametrize("cp_size", [2, 4]) +def test_packed_routes_match_under_context_parallelism(monkeypatch, parallel_state, router_replay, cp_size): + """CP shards each padded sequence into front/back halves per rank.""" + for cp_rank in range(cp_size): + monkeypatch.setattr(parallel_state, "get_context_parallel_world_size", lambda: cp_size, raising=False) + monkeypatch.setattr(parallel_state, "get_context_parallel_rank", lambda rank=cp_rank: rank, raising=False) + new_data, reference = _run_both_paths( + LENGTH_DISTRIBUTIONS["mild_ragged"], + packed=True, + tp_size=1, + local_layers=list(range(NUM_LAYERS)), + monkeypatch=monkeypatch, + parallel_state=parallel_state, + router_replay=router_replay, + ) + _assert_bit_identical(new_data, reference) diff --git a/tests/utils/test_cpu_topology.py b/tests/utils/test_cpu_topology.py new file mode 100644 index 0000000000..0530273fc2 --- /dev/null +++ b/tests/utils/test_cpu_topology.py @@ -0,0 +1,101 @@ +""" +uv run --isolated --extra dev pytest tests/utils/test_cpu_topology.py +""" + +import os +from pathlib import Path + +import pytest + +from skyrl.utils import cpu_topology +from skyrl.utils.cpu_topology import cgroup_cpu_quota, permitted_cpu_cores, pool_workers + + +@pytest.fixture +def cgroup_paths(monkeypatch, tmp_path: Path): + paths = { + "CGROUP_V2_CPU_MAX_PATH": tmp_path / "cpu.max", + "CGROUP_V1_CPU_QUOTA_PATH": tmp_path / "cpu.cfs_quota_us", + "CGROUP_V1_CPU_PERIOD_PATH": tmp_path / "cpu.cfs_period_us", + } + for name, path in paths.items(): + monkeypatch.setattr(cpu_topology, name, str(path)) + return paths + + +@pytest.mark.parametrize( + ("version", "quota", "expected"), + [(2, "400000", 4), (2, "max", None), (1, "200000", 2), (1, "-1", None)], +) +def test_cgroup_quota(cgroup_paths, version, quota, expected): + if version == 2: + cgroup_paths["CGROUP_V2_CPU_MAX_PATH"].write_text(f"{quota} 100000\n") + else: + cgroup_paths["CGROUP_V1_CPU_QUOTA_PATH"].write_text(f"{quota}\n") + cgroup_paths["CGROUP_V1_CPU_PERIOD_PATH"].write_text("100000\n") + + assert cgroup_cpu_quota() == expected + + +def test_missing_cgroup_files(cgroup_paths): + assert cgroup_cpu_quota() is None + + +def test_unreadable_cgroup_values(cgroup_paths): + cgroup_paths["CGROUP_V2_CPU_MAX_PATH"].write_text("not-a-quota 100000\n") + cgroup_paths["CGROUP_V1_CPU_QUOTA_PATH"].write_text("") + cgroup_paths["CGROUP_V1_CPU_PERIOD_PATH"].write_text("") + + assert cgroup_cpu_quota() is None + + +@pytest.mark.parametrize(("quota", "expected"), [(150000, 1), (50000, 1), (100000, 1), (250000, 2)]) +def test_fractional_quota_floors_to_at_least_one(cgroup_paths, quota, expected): + cgroup_paths["CGROUP_V2_CPU_MAX_PATH"].write_text(f"{quota} 100000\n") + + assert cgroup_cpu_quota() == expected + + +@pytest.mark.parametrize(("affinity", "quota_cpus", "expected"), [(16, 4, 4), (4, 16, 4), (8, 8, 8)]) +def test_permitted_cores_is_the_lesser_of_affinity_and_quota(cgroup_paths, monkeypatch, affinity, quota_cpus, expected): + monkeypatch.setattr(os, "sched_getaffinity", lambda pid: set(range(affinity))) + cgroup_paths["CGROUP_V2_CPU_MAX_PATH"].write_text(f"{quota_cpus * 100000} 100000\n") + + assert permitted_cpu_cores() == expected + + +def test_permitted_cores_without_quota_is_the_affinity_mask(cgroup_paths, monkeypatch): + monkeypatch.setattr(os, "sched_getaffinity", lambda pid: set(range(12))) + + assert permitted_cpu_cores() == 12 + + +@pytest.mark.parametrize( + ("cap", "reserved", "cores", "expected"), + [ + (32, 8, 64, 32), # cap binds + (32, 8, 24, 16), # reserve binds + (32, 8, 8, 1), # reserve would leave nothing, so keep one worker + (1, 0, 64, 1), # cap of one + ], +) +def test_pool_workers(cap, reserved, cores, expected): + assert pool_workers(cap=cap, reserved=reserved, cores=cores) == expected + + +def test_pool_workers_defaults_to_permitted_cores(cgroup_paths, monkeypatch): + monkeypatch.setattr(os, "sched_getaffinity", lambda pid: set(range(64))) + cgroup_paths["CGROUP_V2_CPU_MAX_PATH"].write_text("1200000 100000\n") + + assert pool_workers(cap=32, reserved=4) == 8 + + +@pytest.mark.parametrize("cap", [0, -1]) +def test_pool_workers_rejects_non_positive_cap(cap): + with pytest.raises(ValueError, match="cap must be positive"): + pool_workers(cap=cap, reserved=0, cores=8) + + +def test_pool_workers_rejects_negative_reserve(): + with pytest.raises(ValueError, match="reserved cores must be non-negative"): + pool_workers(cap=8, reserved=-1, cores=8)