diff --git a/pyproject.toml b/pyproject.toml index 331aaf1b02..edf058fa01 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -92,6 +92,8 @@ skyrl-train = [ "polars", "s3fs", "fastapi", + "orjson>=3.11.9", + "pybase64>=1.4.2", "uvicorn", "vllm-router; sys_platform == 'linux'", "pybind11", diff --git a/skyrl/backends/skyrl_train/inference_servers/base.py b/skyrl/backends/skyrl_train/inference_servers/base.py index e0152df5d4..1df6ab71c2 100644 --- a/skyrl/backends/skyrl_train/inference_servers/base.py +++ b/skyrl/backends/skyrl_train/inference_servers/base.py @@ -1,6 +1,8 @@ from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Dict, Hashable, List, Optional, Tuple, TypedDict +from skyrl.utils.routed_experts import RoutedExpertIndices + if TYPE_CHECKING: from skyrl.backends.skyrl_train.weight_sync import WeightUpdateRequest from skyrl.backends.skyrl_train.weight_sync.transfer_strategy import ( @@ -32,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): @@ -47,7 +50,8 @@ class InferenceEngineOutput(TypedDict): stop_reasons: List[str] response_logprobs: Optional[List[List[float]]] prompt_logprobs: Optional[List[List[float]]] # per-prompt-token logprobs under the current model - rollout_expert_indices: Optional[List[List[List[int]]]] # [seq_len, layer_num, topk] + rollout_expert_indices: Optional[List[RoutedExpertIndices]] + rollout_sample_support: Optional[List[List[List[int]]]] class InferenceEngineInterface(ABC): diff --git a/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py b/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py index b2d1e4e812..06b43a2750 100644 --- a/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py +++ b/skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py @@ -64,6 +64,7 @@ ) import aiohttp +import orjson from skyrl.backends.skyrl_train.inference_servers.base import ( InferenceEngineInput, @@ -72,10 +73,14 @@ MMPlaceholderRangeInfo, MultiModalFeatures, ) +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + decode_packed_routed_experts, +) from skyrl.env_vars import ( SKYRL_GENERATE_CONCURRENCY_PER_ENGINE, SKYRL_HTTP_CONNECTION_LIMIT, ) +from skyrl.utils.routed_experts import RoutedExpertIndices _DATA_PLANE_RETRIES = 30 @@ -162,6 +167,163 @@ class SampleResponse(TypedDict): topk_prompt_logprobs: Optional[List[Optional[List[Tuple[int, float]]]]] +@dataclass(frozen=True) +class RemoteGenerateResult: + """Raw token generation result returned by ``RemoteInferenceGenerator``.""" + + raw_response: Dict[str, Any] + response_ids: List[int] + response_logprobs: Optional[List[float]] + stop_reason: str + routed_experts: Optional[RoutedExpertIndices] + sample_support: Optional[List[List[int]]] + + +@dataclass +class RemoteInferenceGenerator: + """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) -> Any: + """POST JSON with retry on transient connection and response-decoding failures.""" + 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 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, + return_sample_support: bool = False, + mm_features: Optional[MultiModalFeatures] = None, + cache_salt: Optional[str] = None, + ) -> RemoteGenerateResult: + """Generate one raw-token completion with optional replay metadata.""" + 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") + + use_skyrl_endpoint = return_routed_experts or return_sample_support + path = "/skyrl/v1/generate" if use_skyrl_endpoint 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 return_sample_support: + payload["return_sample_support"] = True + 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) + 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("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) + + sample_support = choice["rollout_sample_support"] if return_sample_support else None + + return RemoteGenerateResult( + raw_response=response, + response_ids=token_ids, + response_logprobs=response_logprobs, + stop_reason=choice["finish_reason"], + routed_experts=routed_experts, + sample_support=sample_support, + ) + + 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): """ @@ -210,6 +372,9 @@ class RemoteInferenceClient(InferenceEngineInterface): enable_return_routed_experts: bool = False """Whether to return routed expert indices (R3 / rollout router replay).""" + enable_return_sample_support_set: bool = False + """Whether to return sampled-token support sets for replay.""" + uses_lora_weight_sync: bool = False """True when the trainer syncs LoRA adapters (rather than full/merged weights). When True, `sleep()` is forced to level=1: level=2 discards the base model from VRAM with no CPU backup, @@ -220,7 +385,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) + _generator: Optional[RemoteInferenceGenerator] = 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) @@ -277,66 +442,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_generator(self) -> RemoteInferenceGenerator: + if self._generator is None: + self._generator = RemoteInferenceGenerator(proxy_url=self.proxy_url) + return self._generator + 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_generator()._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 = await resp.json(content_type=None) - except Exception 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_generator()._post(url, json=json, headers=headers) # --------------------------- # Data Plane @@ -401,6 +516,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: @@ -424,6 +545,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, ) @@ -433,6 +557,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, ) @@ -446,15 +573,22 @@ async def _throttled_detokenize(token_ids: List[int]) -> str: raw_results = await asyncio.gather(*[_throttled_generate(idx) for idx in range(batch_size)]) responses = await asyncio.gather(*[_throttled_detokenize(r["response_ids"]) for r in raw_results]) - rollout_expert_indices = [r.get("routed_experts") for r in raw_results] - has_routed_experts = any(x is not None for x in rollout_expert_indices) + rollout_expert_indices = ( + [result["routed_experts"] for result in raw_results] if self.enable_return_routed_experts else None + ) + rollout_sample_support = ( + [result["rollout_sample_support"] for result in raw_results] + if self.enable_return_sample_support_set + else None + ) return InferenceEngineOutput( responses=responses, stop_reasons=[r["stop_reason"] for r in raw_results], response_ids=[r["response_ids"] for r in raw_results], response_logprobs=[r["response_logprobs"] for r in raw_results] if get_logprobs else None, - rollout_expert_indices=rollout_expert_indices if has_routed_experts else None, + rollout_expert_indices=rollout_expert_indices, + rollout_sample_support=rollout_sample_support, ) async def _generate_single( @@ -465,59 +599,25 @@ 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_generator().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, + return_sample_support=self.enable_return_sample_support_set, + 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 = choice.get("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, + "rollout_sample_support": result.sample_support, } async def _render_for_sample( @@ -1367,9 +1467,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._generator is not None: + await self._generator.aclose() async def __aenter__(self) -> "RemoteInferenceClient": """Async context manager entry.""" @@ -1386,7 +1485,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 @@ -1395,19 +1493,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/routed_experts_wire.py b/skyrl/backends/skyrl_train/inference_servers/routed_experts_wire.py new file mode 100644 index 0000000000..c4b9911c39 --- /dev/null +++ b/skyrl/backends/skyrl_train/inference_servers/routed_experts_wire.py @@ -0,0 +1,45 @@ +"""Compact routed-expert HTTP payloads.""" + +import math +from typing import Any + +import numpy as np +import pybase64 + +from skyrl.utils.routed_experts import ( + ROUTED_EXPERT_DTYPES, + RoutedExpertIndices, + compact_routed_expert_indices, +) + +_DTYPES = {dtype.name: dtype for dtype in ROUTED_EXPERT_DTYPES} + + +def pack_routed_experts(routed_experts: RoutedExpertIndices) -> dict[str, Any]: + compact = compact_routed_expert_indices(routed_experts) + return { + "data": pybase64.b64encode(memoryview(compact)).decode("ascii"), + "shape": list(compact.shape), + "dtype": compact.dtype.name, + } + + +def decode_packed_routed_experts(payload: dict[str, Any]) -> RoutedExpertIndices: + if not isinstance(payload, dict): + raise TypeError("packed routed expert indices must be an object") + try: + dtype = _DTYPES[payload["dtype"]] + shape = tuple(payload["shape"]) + data = pybase64.b64decode_as_bytearray(payload["data"], validate=True) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("invalid packed routed_experts payload") from exc + if len(shape) != 3 or any(type(dim) is not int or dim < 0 for dim in shape): + raise ValueError(f"invalid packed routed_experts shape: {shape}") + expected_size = math.prod(shape) * dtype.itemsize + if len(data) != expected_size: + raise ValueError(f"packed routed_experts has {len(data)} bytes, expected {expected_size}") + decoded = np.frombuffer(data, dtype=dtype).reshape(shape) + compact = compact_routed_expert_indices(decoded) + if compact.dtype != dtype: + raise ValueError(f"packed routed_experts uses non-canonical dtype {dtype.name}; expected {compact.dtype.name}") + return compact diff --git a/skyrl/backends/skyrl_train/inference_servers/setup.py b/skyrl/backends/skyrl_train/inference_servers/setup.py index 437d4804f3..84d481d596 100644 --- a/skyrl/backends/skyrl_train/inference_servers/setup.py +++ b/skyrl/backends/skyrl_train/inference_servers/setup.py @@ -293,6 +293,7 @@ def build_new_inference_client( server_urls=server_setup.server_urls, model_name=ie_cfg.served_model_name or cfg.trainer.policy.model.path, enable_return_routed_experts=ie_cfg.enable_return_routed_experts, + enable_return_sample_support_set=ie_cfg.enable_return_sample_support_set, uses_lora_weight_sync=_uses_lora_weight_sync(cfg), data_parallel_size=ie_cfg.data_parallel_size, tokenizer=tokenizer, diff --git a/skyrl/backends/skyrl_train/inference_servers/utils.py b/skyrl/backends/skyrl_train/inference_servers/utils.py index c9de06c4d1..b1173f7033 100644 --- a/skyrl/backends/skyrl_train/inference_servers/utils.py +++ b/skyrl/backends/skyrl_train/inference_servers/utils.py @@ -84,6 +84,7 @@ def build_vllm_cli_args(cfg: SkyRLTrainConfig) -> Namespace: args: Namespace = parser.parse_args(args=[]) ie_cfg = cfg.generator.inference_engine + sample_support_top_k = cfg.generator.sampling_params.top_k overrides = dict( model=cfg.trainer.policy.model.path, tensor_parallel_size=ie_cfg.tensor_parallel_size, @@ -118,6 +119,9 @@ def build_vllm_cli_args(cfg: SkyRLTrainConfig) -> Namespace: # Overridable via generator.inference_engine.engine_init_kwargs.trust_remote_code below. trust_remote_code=True, ) + if ie_cfg.enable_return_sample_support_set: + overrides["max_logprobs"] = sample_support_top_k + overrides["logprobs_mode"] = "processed_logprobs" for key, value in overrides.items(): setattr(args, key, value) diff --git a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py index 35d7057012..30ff4e56c9 100644 --- a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py +++ b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py @@ -4,15 +4,18 @@ import asyncio import logging +import math import os import time from argparse import Namespace from typing import List, Optional, Tuple import httpx +import numpy as np +import orjson import uvicorn import vllm.envs as envs -from fastapi import HTTPException, Request +from fastapi import HTTPException, Request, Response from ray.util.placement_group import PlacementGroup from vllm.engine.arg_utils import AsyncEngineArgs from vllm.engine.async_llm_engine import AsyncLLMEngine @@ -22,6 +25,7 @@ init_app_state, ) from vllm.inputs import TokensPrompt +from vllm.logprobs import FlatLogprobs from vllm.lora.request import LoRARequest from vllm.sampling_params import SamplingParams as VLLMSamplingParams from vllm.usage.usage_lib import UsageContext @@ -34,6 +38,9 @@ get_node_ip, ) from skyrl.backends.skyrl_train.inference_servers.protocols import ServerActorProtocol +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + pack_routed_experts, +) from skyrl.env_vars import ( SKYRL_HTTP_CONNECTION_LIMIT, SKYRL_VLLM_DP_PORT_OFFSET, @@ -43,6 +50,53 @@ logger = logging.getLogger(__name__) +def _sample_support_from_flat_logprobs( + logprobs: FlatLogprobs, + top_k: int, +) -> tuple[list[dict[str, float]], list[list[int]]]: + """Extract sampled scores and post-filter support from vLLM's flat rows.""" + # vLLM emits [sampled token, top-1, ..., top-k] for every generated token. + row_width = top_k + 1 + token_ids = np.asarray(logprobs.token_ids, dtype=np.int64).reshape(-1, row_width) + processed_logprobs = np.asarray(logprobs.logprobs).reshape(-1, row_width) + support_ids = np.where(np.isneginf(processed_logprobs[:, 1:]), -1, token_ids[:, 1:]) + sampled_logprobs = [{"logprob": value} for value in processed_logprobs[:, 0].tolist()] + + # Repair rows whose sampled token is absent from the captured support. + # + # The support row is columns 1: (the top_k neighbors); column 0 is the token + # vLLM actually sampled. vLLM's approximate Triton top-k/top-p pivot (in + # processed_logprobs mode) can leave slightly more than top_k survivors, so the + # sampled token can rank just beyond top_k and be absent from cols 1:. When that + # happens the downstream sample-support invariant + # (assert_sampled_tokens_in_sample_support_sets / build_sample_support_replay) + # hard-crashes the whole run. To preserve it we overwrite the row's WEAKEST valid + # member (vLLM returns top-k descending, so the trailing valid slot is weakest) + # with the sampled id. Overwriting keeps width == top_k, inserts the sampled id + # exactly once (so the trainer's renorm denominator is not double-counted), and + # leaves any trailing -1 padding intact (the weakest valid col precedes the pad). + sampled = token_ids[:, 0] + valid = support_ids >= 0 + present = np.any(support_ids == sampled[:, None], axis=1) + # Only repair rows that already have >=1 valid member (matches the downstream + # has_support condition); a fully -1 row is left untouched. + missing = (~present) & valid.any(axis=1) + if np.any(missing): + rows = np.flatnonzero(missing) + weakest_col = valid.sum(axis=1) - 1 # last valid (weakest) column per row + support_ids[rows, weakest_col[rows]] = sampled[rows] + logger.warning( + "sample-support repair: %d token(s) had the sampled id absent from top-%d " + "support; overwrote the weakest member to preserve the invariant (vLLM " + "approx top-k/top-p pivot artifact); example: sampled token %d at row %d", + rows.size, + top_k, + int(sampled[rows[0]]), + int(rows[0]), + ) + return sampled_logprobs, support_ids.tolist() + + class VLLMServerActor(ServerActorProtocol): """ Ray actor that runs a vLLM OpenAI-compatible API server. @@ -399,7 +453,10 @@ async def _skyrl_generate(request: Request): token_ids = body["token_ids"] sampling_params_dict = body.get("sampling_params", {}) cache_salt = body.get("cache_salt") - + capture_sample_support = body.get("return_sample_support", False) + if capture_sample_support: + sampling_params_dict["flat_logprobs"] = True + sampling_params_dict["logprobs"] = sampling_params_dict["top_k"] sampling_params = VLLMSamplingParams(**sampling_params_dict) # `cache_salt` salts vLLM's prefix cache; vLLM rejects an empty salt, so attach only when set. if cache_salt is not None: @@ -420,11 +477,21 @@ async def _skyrl_generate(request: Request): finish_reason = resp.finish_reason logprobs = None - if resp.logprobs is not None: + sample_support = None + if capture_sample_support: + content, sample_support = _sample_support_from_flat_logprobs( + resp.logprobs, + sampling_params_dict["top_k"], + ) + logprobs = {"content": content} + elif resp.logprobs is not None: content = [] for tid, lp_dict in zip(token_ids_out, resp.logprobs): if lp_dict and tid in lp_dict: - content.append({"logprob": lp_dict[tid].logprob}) + logprob = lp_dict[tid].logprob + if not math.isfinite(logprob): + raise ValueError("Out of range float values are not JSON compliant") + content.append({"logprob": logprob}) else: # -9999.0 is the default in vLLM's ChatCompletionLogProb content.append({"logprob": -9999.0}) @@ -432,21 +499,20 @@ async def _skyrl_generate(request: Request): routed_experts = None if resp.routed_experts is not None: - if hasattr(resp.routed_experts, "tolist"): - routed_experts = resp.routed_experts.tolist() - else: - routed_experts = resp.routed_experts + routed_experts = pack_routed_experts(np.asarray(resp.routed_experts)) - return { + payload = { "choices": [ { "token_ids": token_ids_out, "finish_reason": finish_reason, "logprobs": logprobs, "routed_experts": routed_experts, + "rollout_sample_support": sample_support, } ] } + return Response(content=orjson.dumps(payload), media_type="application/json") async def shutdown(self) -> None: """Gracefully shutdown the server.""" diff --git a/skyrl/backends/skyrl_train/training_batch.py b/skyrl/backends/skyrl_train/training_batch.py index f2bac3ab5a..ac6e631406 100644 --- a/skyrl/backends/skyrl_train/training_batch.py +++ b/skyrl/backends/skyrl_train/training_batch.py @@ -7,7 +7,9 @@ import numpy as np import torch -from jaxtyping import Float, Integer +from jaxtyping import Bool, Float, Integer + +from skyrl.utils.routed_experts import make_replay_padding_indices DictType = TypeVar("DictType") @@ -476,6 +478,8 @@ class TrainingInput(TypedDict, total=False): rewards: Optional[Float[torch.Tensor, "batch_size seq_len"]] rollout_logprobs: Optional[Float[torch.Tensor, "batch_size seq_len"]] rollout_expert_indices: Optional[Integer[torch.Tensor, "batch_size seq_len layer_num topk"]] + router_padding_mask: Optional[Bool[torch.Tensor, "batch_size seq_len"]] + sample_support_ids: Optional[Integer[torch.Tensor, "batch_size seq_len topk"]] pixel_values: Optional[TensorList] # list of `batch_size` [num_patches_i, dim] tensors image_grid_thw: Optional[TensorList] # list of `batch_size` [num_images_i, 3] tensors @@ -524,6 +528,27 @@ def pad_training_input_batch(unpadded_batch: TrainingInputBatch, pad_size: int) additional_dims = tensor.shape[1:] padding_tensor = torch.zeros(pad_size, *additional_dims, dtype=tensor.dtype, device=tensor.device) new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) + elif key == "rollout_expert_indices": + additional_dims = tensor.shape[1:] + padding_tensor = make_replay_padding_indices( + (pad_size, *additional_dims), + dtype=tensor.dtype, + device=tensor.device, + ) + new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) + elif key == "router_padding_mask": + additional_dims = tensor.shape[1:] + padding_tensor = torch.ones(pad_size, *additional_dims, dtype=torch.bool, device=tensor.device) + new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) + elif key == "sample_support_ids": + additional_dims = tensor.shape[1:] + padding_tensor = torch.full( + (pad_size, *additional_dims), + -1, + dtype=tensor.dtype, + device=tensor.device, + ) + new_tensors[key] = torch.cat([tensor, padding_tensor], dim=0) else: # Copy row 0 `pad_size` times. Loss masked so values don't affect the loss. Just need valid shape/dtype. assert tensor.shape[0] > 0, f"Cannot pad empty tensor field {key!r}" diff --git a/skyrl/backends/skyrl_train/utils/replay_utils.py b/skyrl/backends/skyrl_train/utils/replay_utils.py index 3d8ed7b56a..88c1c6e148 100644 --- a/skyrl/backends/skyrl_train/utils/replay_utils.py +++ b/skyrl/backends/skyrl_train/utils/replay_utils.py @@ -2,14 +2,13 @@ Utility functions for MoE Router Replay. """ -from typing import List +from contextlib import contextmanager import torch -from skyrl.backends.skyrl_train.distributed.megatron.packing_utils import ( - get_packed_seq_align_size, - get_unpacked_seq_align_size, - is_fp8_enabled, +from skyrl.utils.token_metadata import ( + TokenMetadataLayout, + align_token_metadata, ) @@ -45,167 +44,57 @@ def patched_set_layer_number(self, layer_number: int): TopKRouter._set_layer_number_patched = True -def _patch_alltoall_dispatcher_for_replay(): - """Monkey-patch MoEAlltoAllTokenDispatcher.preprocess to handle router replay. - - When router replay is enabled, duplicate indices in top_indices can cause - routing_map.sum() < num_tokens * topk, leading to a split size mismatch - in the alltoall collective. We fix this by deriving num_out_tokens from - the routing map instead of the static num_tokens * topk formula. - - Reference: https://github.com/verl-project/verl/pull/4986 - """ +def patch_topk_router_expert_bias_padding_mask(): + """Fix the token-mask broadcast in pinned Megatron's expert-bias accounting.""" try: - from megatron.core.transformer.moe.token_dispatcher import ( - MoEAlltoAllTokenDispatcher, - ) + from megatron.core.transformer.moe.router import TopKRouter except ImportError: return - if getattr(MoEAlltoAllTokenDispatcher, "_preprocess_patched", False): + if getattr(TopKRouter, "_expert_bias_padding_mask_patched", False): return - original_preprocess = MoEAlltoAllTokenDispatcher.preprocess - - def patched_preprocess(self, routing_map): - result = original_preprocess(self, routing_map) - if ( - getattr(self.config, "moe_enable_routing_replay", False) - and not self.drop_and_pad - and self.config.moe_expert_capacity_factor is None - and not self.config.moe_router_padding_for_quantization - ): - self.num_out_tokens = int(routing_map.sum().item()) - return result - - MoEAlltoAllTokenDispatcher.preprocess = patched_preprocess - MoEAlltoAllTokenDispatcher._preprocess_patched = True - - -def _split_replay_indices(rollout_expert_indices: torch.Tensor) -> List[torch.Tensor]: - if rollout_expert_indices is None: - return None - if rollout_expert_indices.dim() != 4: - raise ValueError(f"Expected 4D replay indices, got shape {rollout_expert_indices.shape}") - per_layer = rollout_expert_indices.permute(2, 0, 1, 3).contiguous() - # flatten [batch, seq, topk] to [batch * seq, topk] for each layer - return [per_layer[i].reshape(-1, per_layer.shape[-1]) for i in range(per_layer.shape[0])] - - -def _remove_left_padding_from_indices( - rollout_expert_indices: torch.Tensor, - attention_mask: torch.Tensor, - fp8_enabled: bool = False, -) -> torch.Tensor: - """Apply the same left-padding removal as remove_left_padding to routing indices. - - Args: - rollout_expert_indices: [batch, padded_seq_len, layers, topk] - attention_mask: [batch, padded_seq_len] (int or bool) - - Returns: - [batch, effective_seq_len, layers, topk] with real tokens packed left. - """ - import megatron.core.parallel_state as mpu - - seq_lens = attention_mask.sum(dim=1) - effective_seq_len = seq_lens.max().item() - tp_size = mpu.get_tensor_model_parallel_world_size() - align_size = get_unpacked_seq_align_size(tp_size, fp8_enabled=fp8_enabled) - if align_size > 1: - pad_size = (align_size - effective_seq_len % align_size) % align_size - effective_seq_len += pad_size - - batch_size = rollout_expert_indices.shape[0] - new_rii = torch.zeros( - batch_size, - effective_seq_len, - rollout_expert_indices.shape[2], - rollout_expert_indices.shape[3], - dtype=rollout_expert_indices.dtype, - device=rollout_expert_indices.device, - ) - for i in range(batch_size): - mask = attention_mask[i].bool() - new_rii[i, : seq_lens[i]] = rollout_expert_indices[i, mask] - return new_rii - - -def _pack_replay_indices( - rollout_expert_indices: torch.Tensor, - attention_mask: torch.Tensor, - fp8_enabled: bool = False, -) -> torch.Tensor: - """Pack routing indices to match the token layout produced by preprocess_packed_seqs. - - With sample packing, Megatron concatenates all sequences into one packed - sequence with per-sample alignment padding. The MoE router sees tokens in - this packed order, so replay indices must follow the same layout. + original_apply_expert_bias = TopKRouter._apply_expert_bias - Returns: - [1, total_packed_len, layers, topk] matching the packed model input. - """ - import megatron.core.parallel_state as mpu + def patched_apply_expert_bias(self, routing_map: torch.Tensor, padding_mask: torch.Tensor | None = None): + # Megatron combines [tokens, experts] with a token-only mask. + if padding_mask is not None and padding_mask.ndim == 1: + padding_mask = padding_mask.unsqueeze(-1) + return original_apply_expert_bias(self, routing_map, padding_mask) - batch_size = rollout_expert_indices.shape[0] - num_layers = rollout_expert_indices.shape[2] - topk = rollout_expert_indices.shape[3] + TopKRouter._apply_expert_bias = patched_apply_expert_bias + TopKRouter._expert_bias_padding_mask_patched = True - seq_lens = attention_mask.sum(dim=-1, dtype=torch.int32) - tp_size = mpu.get_tensor_model_parallel_world_size() - cp_size = mpu.get_context_parallel_world_size() - align_size = get_packed_seq_align_size(tp_size, cp_size, fp8_enabled=fp8_enabled) - pad_sizes = (align_size - seq_lens % align_size) % align_size - seqlens_padded = seq_lens + pad_sizes +def _split_replay_indices(rollout_expert_indices: torch.Tensor) -> list[torch.Tensor]: + per_layer = rollout_expert_indices.permute(2, 0, 1, 3).contiguous().to(torch.int32) + return list(per_layer.flatten(1, 2).unbind(0)) - total_packed_len = int(seqlens_padded.sum().item()) - packed = torch.zeros( - total_packed_len, - num_layers, - topk, - dtype=rollout_expert_indices.dtype, - device=rollout_expert_indices.device, +def scatter_router_padding_mask_for_model( + router_padding_mask: torch.Tensor | None, + model, + model_config, +) -> torch.Tensor | None: + """Match the mask layout to sequence-parallel hidden states at model entry.""" + if router_padding_mask is None or not model_config.sequence_parallel: + return router_padding_mask + + from megatron.core.models.hybrid.hybrid_model import HybridModel + from megatron.core.tensor_parallel import scatter_to_sequence_parallel_region + from megatron.core.utils import unwrap_model + + unwrapped_model = unwrap_model(model) + # GPTModel scatters its mask beside the embedding on the first PP stage. HybridModel + # scatters only the embedding, so its mask must always be scattered here. + if not isinstance(unwrapped_model, HybridModel) and unwrapped_model.pre_process: + return router_padding_mask + return ( + scatter_to_sequence_parallel_region(router_padding_mask.transpose(0, 1).contiguous()) + .transpose(0, 1) + .contiguous() ) - seq_lens_cpu = seq_lens.tolist() - seqlens_padded_cpu = seqlens_padded.tolist() - offset = 0 - for i in range(batch_size): - n = seq_lens_cpu[i] - mask = attention_mask[i].bool() - d = rollout_expert_indices[i, mask] - packed[offset : offset + n] = d - offset += seqlens_padded_cpu[i] - - if cp_size > 1: - cp_rank = mpu.get_context_parallel_rank() - out = torch.zeros( - total_packed_len // cp_size, - num_layers, - topk, - dtype=packed.dtype, - device=packed.device, - ) - src_offset = 0 - dst_offset = 0 - for i in range(batch_size): - seqlen_padded_i = seqlens_padded_cpu[i] - seqlen_per_cp = seqlen_padded_i // cp_size - half = seqlen_per_cp // 2 - out[dst_offset : dst_offset + half] = packed[ - src_offset + half * cp_rank : src_offset + half * (cp_rank + 1) - ] - back_start = src_offset + seqlen_padded_i - half * (cp_rank + 1) - back_end = src_offset + seqlen_padded_i - half * cp_rank - out[dst_offset + half : dst_offset + seqlen_per_cp] = packed[back_start:back_end] - src_offset += seqlen_padded_i - dst_offset += seqlen_per_cp - packed = out - - return packed.unsqueeze(0) # [1, packed_len_per_cp, layers, topk] - def _get_current_pp_stage_layer_range(model_config) -> tuple[int, int]: """Return the current PP rank's transformer-layer range as (start_layer, @@ -224,14 +113,43 @@ def _get_current_pp_stage_layer_range(model_config) -> tuple[int, int]: return offset, num_layers +def _get_local_router_layer_indices(model_config, global_num_layers: int, instances: list) -> list[int]: + local_layer_offset, local_num_layers = _get_current_pp_stage_layer_range(model_config) + if local_num_layers == len(instances): + layer_indices = list(range(local_layer_offset, local_layer_offset + local_num_layers)) + else: + layer_indices = [] + for local_router_index, router_instance in enumerate(instances): + layer_number = getattr(router_instance, "layer_number", None) + if layer_number is not None: + layer_index = layer_number - 1 + else: + layer_index = local_layer_offset + local_router_index + (local_num_layers - len(instances)) + layer_indices.append(layer_index) + + if any(layer_index < 0 or layer_index >= global_num_layers for layer_index in layer_indices): + raise ValueError( + f"Router replay layer indices {layer_indices} out of range for data with {global_num_layers} layers" + ) + return layer_indices + + def setup_per_microbatch_replay_forward( rollout_expert_indices: torch.Tensor, + router_padding_mask: torch.Tensor | None, attention_mask: torch.Tensor, + model, model_config, + metadata_layout: TokenMetadataLayout, remove_microbatch_padding: bool = False, -) -> None: - """Set up RouterReplay for a single micro-batch, aligning indices - with the left-padding-removed token layout that the MoE layer sees. +) -> dict[str, torch.Tensor]: + """Set up router replay and return its model-facing keyword arguments. + + Replay indices and the router padding mask start in the same batch layout and + undergo matching padding removal or packing and CP sharding. Their destinations + then differ: indices are TP-sliced and installed into per-layer ``RouterReplay`` + instances, while the mask follows Megatron's model-specific sequence-parallel + path and is passed to the model as ``padding_mask``. Handles context parallelism: when CP > 1, the sequence is split into 2*cp_size chunks with each CP rank receiving a front chunk and a back @@ -247,8 +165,8 @@ def setup_per_microbatch_replay_forward( layers. We use each instance's global layer_number (set by the patched TopKRouter.set_layer_number) to index into the correct slice of the data. - Handles pipeline parallelism: when PP > 1, the sequence is split across - PP ranks, so each rank only sees its local RouterReplay instances. In cases + Handles pipeline parallelism: when PP > 1, transformer layers are split + across PP ranks, so each rank only sees its local RouterReplay instances. In cases where the number of local RouterReplay instances does not match the local layer count, indicating that the model has dense layers before MoE layers, we use the global layer_number to index into the correct slice of the data. @@ -260,53 +178,61 @@ def setup_per_microbatch_replay_forward( RouterReplayAction, ) - _patch_alltoall_dispatcher_for_replay() - fp8_enabled = is_fp8_enabled(getattr(model_config, "fp8", None)) + if router_padding_mask is None: + raise ValueError("router_padding_mask is required with rollout_expert_indices") + if rollout_expert_indices.dim() != 4: + raise ValueError(f"Expected 4D replay indices, got shape {rollout_expert_indices.shape}") - if remove_microbatch_padding: - aligned = _pack_replay_indices(rollout_expert_indices, attention_mask, fp8_enabled=fp8_enabled) - else: - aligned = _remove_left_padding_from_indices( - rollout_expert_indices, - attention_mask, - fp8_enabled=fp8_enabled, + if router_padding_mask.shape != attention_mask.shape: + raise ValueError( + f"router_padding_mask shape {router_padding_mask.shape} does not match " + f"attention_mask shape {attention_mask.shape}" ) + if router_padding_mask.device != rollout_expert_indices.device: + raise ValueError("rollout_expert_indices and router_padding_mask must be on the same device") + + instances = RouterReplay.global_router_replay_instances + local_layer_indices = _get_local_router_layer_indices( + model_config, + rollout_expert_indices.shape[2], + 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) + + if (metadata_layout.padded_sequence_lengths is not None) != remove_microbatch_padding: + raise ValueError("Shared token metadata layout does not match the model packing mode") + aligned_router_padding_mask = align_token_metadata(router_padding_mask.to(torch.bool), metadata_layout, True) + route_padding = torch.arange( + rollout_expert_indices.shape[-1], + dtype=rollout_expert_indices.dtype, + device=local_rollout_expert_indices.device, + ) + aligned_rollout_expert_indices = align_token_metadata( + local_rollout_expert_indices, + metadata_layout, + route_padding, + ) # TP splitting: sequence parallelism across the tensor model parallel region tp_size = mpu.get_tensor_model_parallel_world_size() if tp_size > 1: tp_rank = mpu.get_tensor_model_parallel_rank() - seq_len = aligned.shape[1] + seq_len = aligned_rollout_expert_indices.shape[1] chunk_size = seq_len // tp_size - aligned = aligned[:, tp_rank * chunk_size : (tp_rank + 1) * chunk_size, :, :] - per_layer_data = _split_replay_indices(aligned) - global_num_layers_in_data = len(per_layer_data) - instances = RouterReplay.global_router_replay_instances - num_instances = len(instances) - local_layer_offset, local_num_layers = _get_current_pp_stage_layer_range(model_config) - - if local_num_layers == num_instances: - local_per_layer_data = per_layer_data[local_layer_offset : local_layer_offset + local_num_layers] - RouterReplay.set_replay_data(local_per_layer_data) - else: - # Dense-layer mismatch: map each MoE router to its global layer index. - # Prefer the patched layer_number; fall back to offset-based mapping - # (assumes dense layers precede MoE layers). - for local_router_idx, router_instance in enumerate(instances): - layer_number = getattr(router_instance, "layer_number", None) - if layer_number is not None: - layer_idx = layer_number - 1 # layer_number is 1-based - else: - layer_idx = local_layer_offset + local_router_idx + (local_num_layers - num_instances) - if layer_idx < 0 or layer_idx >= global_num_layers_in_data: - raise ValueError( - f"Router replay layer index {layer_idx} out of range " - f"for data with {global_num_layers_in_data} layers " - f"({num_instances} router instances)" - ) - router_instance.set_target_indices(per_layer_data[layer_idx]) + aligned_rollout_expert_indices = aligned_rollout_expert_indices[ + :, tp_rank * chunk_size : (tp_rank + 1) * chunk_size, :, : + ] + RouterReplay.set_replay_data(_split_replay_indices(aligned_rollout_expert_indices)) RouterReplay.set_global_router_replay_action(RouterReplayAction.REPLAY_FORWARD) + model_router_padding_mask = scatter_router_padding_mask_for_model( + aligned_router_padding_mask, + model, + model_config, + ) + return {"padding_mask": model_router_padding_mask} + def setup_per_microbatch_replay_backward() -> None: """Switch RouterReplay to backward mode so that activation-checkpoint @@ -327,3 +253,22 @@ def clear_router_replay(): RouterReplay.clear_global_indices() RouterReplay.clear_global_router_replay_action() + + +@contextmanager +def router_replay_schedule(enabled: bool): + """Isolate global RouterReplay state to one Megatron pipeline schedule. + + The backward FIFO spans all microbatches in a training schedule, so it must + only be cleared at schedule boundaries. Forward-only schedules leave that + FIFO unconsumed, and failed schedules may leave it partially consumed. + """ + if not enabled: + yield + return + + clear_router_replay() + try: + yield + finally: + clear_router_replay() diff --git a/skyrl/backends/skyrl_train/utils/sample_support_replay.py b/skyrl/backends/skyrl_train/utils/sample_support_replay.py new file mode 100644 index 0000000000..dd20179893 --- /dev/null +++ b/skyrl/backends/skyrl_train/utils/sample_support_replay.py @@ -0,0 +1,302 @@ +"""Support-conditioned logprobs for bounded sampler replay.""" + +import torch + +from skyrl.utils.token_metadata import ( + TokenMetadataLayout, + align_token_metadata, + scatter_packed_token_values_to_batch, +) + + +def _selected_hidden_projection( + hidden: torch.Tensor, + token_ids: torch.Tensor, + local_mask: torch.Tensor, + lm_head_weight: torch.Tensor, + temperature: float, + chunk_size: int | None, + invalid_value: float, +) -> torch.Tensor: + """Project selected candidate pairs without materializing vocabulary logits.""" + num_rows, width = token_ids.shape + row_ids = torch.arange(num_rows, device=hidden.device).unsqueeze(1).expand(-1, width).reshape(-1) + flat_token_ids = token_ids.reshape(-1) + flat_mask = local_mask.reshape(-1) + output = torch.empty(flat_token_ids.shape, dtype=torch.float32, device=hidden.device) + # Bound the temporary [candidate pairs, hidden] projection for wide supports. + pair_chunk_size = flat_token_ids.numel() if chunk_size is None else chunk_size + for start in range(0, flat_token_ids.numel(), pair_chunk_size): + end = min(start + pair_chunk_size, flat_token_ids.numel()) + selected_hidden = hidden.index_select(0, row_ids[start:end]).to(lm_head_weight.dtype) + selected_weight = lm_head_weight.index_select(0, flat_token_ids[start:end]) + projected = (selected_hidden * selected_weight).sum(dim=-1) / temperature + output[start:end] = torch.where( + flat_mask[start:end], + projected.to(torch.float32), + invalid_value, + ) + return output.reshape(num_rows, width) + + +def sample_support_logprobs( + logits_or_hidden: torch.Tensor, + sampled_ids: torch.Tensor, + support_ids: torch.Tensor, + *, + vocab_start_index: int, + vocab_end_index: int, + tp_group: torch.distributed.ProcessGroup | None, + lm_head_weight: torch.Tensor | None = None, + temperature: float = 1.0, + chunk_size: int | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Renormalize sampled-token scores over each recorded support row.""" + if logits_or_hidden.shape[:-1] != sampled_ids.shape or support_ids.shape[:-1] != sampled_ids.shape: + raise ValueError( + "logits, sampled_ids, and support_ids must have matching prefix shapes, got " + f"{logits_or_hidden.shape[:-1]}, {sampled_ids.shape}, and {support_ids.shape[:-1]}" + ) + if support_ids.dtype != torch.int32: + raise ValueError(f"sample support must use int32 vocab ids, got {support_ids.dtype}") + if temperature <= 0: + raise ValueError("temperature must be positive") + + flat_source = logits_or_hidden.reshape(-1, logits_or_hidden.shape[-1]) + flat_sampled = sampled_ids.reshape(-1).long() + flat_support = support_ids.reshape(-1, support_ids.shape[-1]).long() + valid_members = flat_support >= 0 + valid_rows = valid_members.any(dim=-1) + local_members = valid_members & (flat_support >= vocab_start_index) & (flat_support < vocab_end_index) + local_support_ids = (flat_support - vocab_start_index).clamp(0, vocab_end_index - vocab_start_index - 1) + local_sample_mask = (flat_sampled >= vocab_start_index) & (flat_sampled < vocab_end_index) + local_sample_ids = (flat_sampled - vocab_start_index).clamp(0, vocab_end_index - vocab_start_index - 1) + + compute_dtype = ( + torch.float32 if logits_or_hidden.dtype in (torch.float16, torch.bfloat16) else logits_or_hidden.dtype + ) + if lm_head_weight is None: + local_values = flat_source.gather(1, local_support_ids).to(compute_dtype) + local_values = torch.where(local_members, local_values, float("-inf")) + local_sampled = flat_source.gather(1, local_sample_ids.unsqueeze(1)).squeeze(1).to(compute_dtype) + local_sampled = torch.where(local_sample_mask, local_sampled, 0.0) + else: + if lm_head_weight.shape[0] != vocab_end_index - vocab_start_index: + raise ValueError("lm_head_weight rows do not match the configured vocabulary shard") + local_values = _selected_hidden_projection( + flat_source, + local_support_ids, + local_members, + lm_head_weight, + temperature, + chunk_size, + float("-inf"), + ) + local_sampled = _selected_hidden_projection( + flat_source, + local_sample_ids.unsqueeze(1), + local_sample_mask.unsqueeze(1), + lm_head_weight, + temperature, + chunk_size, + 0.0, + ).squeeze(1) + + local_max = local_values.detach().amax(dim=-1) + global_max = local_max.clone() + if tp_group is not None and torch.distributed.get_world_size(tp_group) > 1: + torch.distributed.all_reduce(global_max, op=torch.distributed.ReduceOp.MAX, group=tp_group) + safe_max = torch.where(valid_rows, global_max, 0.0) + + local_sum = torch.where(local_members, (local_values - safe_max.unsqueeze(1)).exp(), 0.0).sum(dim=-1) + # Numerator and denominator share one TP SUM collective. + local_stats = torch.stack((local_sum, local_sampled)) + global_stats = local_stats.detach().clone() + if tp_group is not None and torch.distributed.get_world_size(tp_group) > 1: + torch.distributed.all_reduce(global_stats, op=torch.distributed.ReduceOp.SUM, group=tp_group) + global_stats = global_stats + local_stats - local_stats.detach() + denominator, sampled_score = global_stats + logprobs = sampled_score - safe_max - torch.where(valid_rows, denominator, 1.0).log() + logprobs = torch.where(valid_rows, logprobs, 0.0) + return logprobs.reshape(sampled_ids.shape), valid_rows.reshape(sampled_ids.shape) + + +def synthetic_eos_logprobs( + logits_or_hidden: torch.Tensor, + sampled_ids: torch.Tensor, + synthetic_eos_mask: torch.Tensor, + *, + vocab_start_index: int, + vocab_end_index: int, + tp_group: torch.distributed.ProcessGroup, + inference_only: bool, + lm_head_weight: torch.Tensor | None = None, + temperature: float = 1.0, + chunk_size: int | None = None, + fused_backend: str = "torch", + metadata_layout: TokenMetadataLayout | None = None, +) -> torch.Tensor: + """Compute ordinary logprobs for EOS tokens appended after vLLM generation.""" + if synthetic_eos_mask.shape != sampled_ids.shape: + raise ValueError("synthetic_eos_mask and sampled_ids must have matching shapes") + + if metadata_layout is not None and metadata_layout.padded_sequence_lengths is not None: + if synthetic_eos_mask.shape[0] != 1: + raise ValueError("Packed synthetic EOS metadata must have a singleton batch dimension") + if metadata_layout.cu_seqlens_padded is None: + raise ValueError("Packed synthetic EOS fallback requires padded sequence boundaries") + if any(length <= 0 for length in metadata_layout.padded_sequence_lengths): + raise ValueError("Synthetic EOS fallback requires non-empty trajectory segments") + expected_tokens = metadata_layout.aligned_sequence_length // metadata_layout.context_parallel_size + if expected_tokens != synthetic_eos_mask.numel(): + raise ValueError("Synthetic EOS layout does not match the model token layout") + lengths = ( + metadata_layout.cu_seqlens_padded.to( + device=synthetic_eos_mask.device, + dtype=torch.long, + ).diff() + // metadata_layout.context_parallel_size + ) + else: + if synthetic_eos_mask.shape[0] == 0 or synthetic_eos_mask.shape[1] == 0: + raise ValueError("Synthetic EOS fallback requires non-empty trajectory segments") + lengths = torch.full( + (synthetic_eos_mask.shape[0],), + synthetic_eos_mask.shape[1], + dtype=torch.long, + device=synthetic_eos_mask.device, + ) + + # Preprocessing permits at most one unsupported loss-bearing EOS per + # trajectory. Select one fixed slot for every trajectory so TP collectives + # never depend on the number of EOS fallbacks in this microbatch. + offsets = lengths.cumsum(dim=0) - lengths + trajectory_ids = torch.repeat_interleave( + torch.arange(lengths.shape[0], device=lengths.device), + lengths, + output_size=synthetic_eos_mask.numel(), + ) + token_indices = torch.arange(synthetic_eos_mask.numel(), device=lengths.device) + sentinel = synthetic_eos_mask.numel() + candidate_indices = torch.where(synthetic_eos_mask.reshape(-1), token_indices, sentinel) + selected_indices = torch.full_like(offsets, sentinel).scatter_reduce( + 0, + trajectory_ids, + candidate_indices, + reduce="amin", + include_self=True, + ) + has_selection = selected_indices != sentinel + selected_indices = torch.where(has_selection, selected_indices, offsets) + + flat_source = logits_or_hidden.reshape(-1, logits_or_hidden.shape[-1]) + flat_targets = sampled_ids.reshape(-1) + selected_source = flat_source.index_select(0, selected_indices) + selected_targets = flat_targets.index_select(0, selected_indices) + if lm_head_weight is None: + from skyrl.backends.skyrl_train.distributed.megatron.model_utils import ( + DistributedLogprob, + ) + + selected = DistributedLogprob.apply( + selected_source.unsqueeze(0), + selected_targets.unsqueeze(0), + vocab_start_index, + vocab_end_index, + tp_group, + inference_only, + ).squeeze(0) + else: + from skyrl.backends.skyrl_train.distributed.megatron.model_utils import ( + _fused_lm_head_logprob_apply, + ) + + if temperature != 1.0: + lm_head_weight = lm_head_weight / temperature + selected_chunk_size = ( + selected_source.shape[0] if chunk_size is None else min(chunk_size, selected_source.shape[0]) + ) + selected = _fused_lm_head_logprob_apply( + fused_backend, + selected_source.unsqueeze(0), + lm_head_weight, + selected_targets.unsqueeze(0), + vocab_start_index, + vocab_end_index, + selected_chunk_size, + tp_group, + inference_only, + ).squeeze(0) + selected = torch.where(has_selection, selected, 0.0).to(torch.float32) + output = torch.zeros(sampled_ids.numel(), dtype=torch.float32, device=logits_or_hidden.device) + output = output.scatter_add(0, selected_indices, selected) + return output.reshape(sampled_ids.shape) + + +def compute_sample_support_logprobs( + logits_or_hidden: torch.Tensor, + sequences: torch.Tensor, + loss_mask: torch.Tensor | None, + sample_support_ids: torch.Tensor | None, + num_actions: int, + *, + packed: bool, + metadata_layout: TokenMetadataLayout | None, + vocab_start_index: int, + vocab_end_index: int, + tp_group: torch.distributed.ProcessGroup, + inference_only: bool, + lm_head_weight: torch.Tensor | None, + temperature: float, + chunk_size: int | None, + fused_backend: str, +) -> torch.Tensor: + """Compute support-conditioned logprobs in canonical trainer layout.""" + if sample_support_ids is None: + raise ValueError("sample-support replay is enabled but the microbatch has no recorded support") + if loss_mask is None: + raise ValueError("sample-support replay requires the response loss mask") + + target_loss_mask = torch.zeros_like(sequences, dtype=torch.bool) + target_loss_mask[:, -num_actions:] = loss_mask.to(torch.bool) + if packed: + if metadata_layout is None: + raise ValueError("Packed sample-support replay requires the shared token metadata layout") + aligned_sampled_ids = align_token_metadata(sequences, metadata_layout, 0, next_token=True) + aligned_support_ids = align_token_metadata(sample_support_ids, metadata_layout, -1, next_token=True) + aligned_loss_mask = align_token_metadata(target_loss_mask, metadata_layout, False, next_token=True) + else: + aligned_sampled_ids = sequences[:, 1:] + aligned_support_ids = sample_support_ids[:, 1:] + aligned_loss_mask = target_loss_mask[:, 1:] + + support_logprobs, valid_support = sample_support_logprobs( + logits_or_hidden if packed else logits_or_hidden[:, :-1], + aligned_sampled_ids, + aligned_support_ids, + vocab_start_index=vocab_start_index, + vocab_end_index=vocab_end_index, + tp_group=tp_group, + lm_head_weight=lm_head_weight, + temperature=temperature if lm_head_weight is not None else 1.0, + chunk_size=chunk_size, + ) + # Preprocessing permits an empty loss-bearing row only for an EOS that SkyRL + # appended after generation. vLLM never supplied a support set for that token. + synthetic_eos_mask = aligned_loss_mask & ~valid_support + eos_logprobs = synthetic_eos_logprobs( + logits_or_hidden if packed else logits_or_hidden[:, :-1], + aligned_sampled_ids, + synthetic_eos_mask, + vocab_start_index=vocab_start_index, + vocab_end_index=vocab_end_index, + tp_group=tp_group, + inference_only=inference_only, + lm_head_weight=lm_head_weight, + temperature=temperature if lm_head_weight is not None else 1.0, + chunk_size=chunk_size, + fused_backend=fused_backend, + metadata_layout=metadata_layout if packed else None, + ) + token_logprobs = torch.where(synthetic_eos_mask, eos_logprobs, support_logprobs) + return scatter_packed_token_values_to_batch(token_logprobs, metadata_layout, 0) if packed else token_logprobs diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py index 61e5d71ff7..c06c489619 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_model_wrapper.py @@ -43,14 +43,19 @@ compute_approx_kl, ) from skyrl.backends.skyrl_train.utils.replay_utils import ( + router_replay_schedule, setup_per_microbatch_replay_backward, setup_per_microbatch_replay_forward, ) +from skyrl.backends.skyrl_train.utils.sample_support_replay import ( + compute_sample_support_logprobs, +) from skyrl.backends.skyrl_train.utils.torch_utils import masked_mean from skyrl.backends.skyrl_train.workers.worker_utils import ( compute_minibatch_rollout_logprob_diff_metrics, ) from skyrl.train.config import TrainerConfig +from skyrl.utils.token_metadata import TokenMetadataLayout, build_token_metadata_layout def _build_packed_targets( @@ -242,7 +247,7 @@ def forward( self._assert_vlm_supported() forward_backward_func = get_forward_backward_func() - def collection_func(logits, data): + def collection_func(logits, *, data, metadata_layout: TokenMetadataLayout | None): sequences = data["sequences"] packed_seq_params = data.get("packed_seq_params") packed_targets = data.get("packed_targets") @@ -263,7 +268,26 @@ def collection_func(logits, data): if temperature != 1.0 and not fused_lm_head: logits.div_(temperature) - if fused_lm_head and packed_seq_params is not None and packed_targets is not None: + shard_vocab_size = lm_head_weight.shape[0] if fused_lm_head else logits.shape[-1] + if self.cfg.algorithm.enable_sample_support_replay: + token_logprobs = compute_sample_support_logprobs( + logits, + sequences, + data.get("loss_mask"), + data.get("sample_support_ids"), + data["num_actions"], + packed=packed_seq_params is not None, + metadata_layout=metadata_layout, + vocab_start_index=tp_rank * shard_vocab_size, + vocab_end_index=(tp_rank + 1) * shard_vocab_size, + tp_group=tp_grp, + lm_head_weight=lm_head_weight if fused_lm_head else None, + temperature=temperature, + inference_only=True, + chunk_size=self.cfg.logprobs_chunk_size, + fused_backend=self._fused_lm_head_backend, + ) + elif fused_lm_head and packed_seq_params is not None and packed_targets is not None: token_logprobs = from_parallel_hidden_to_logprobs_packed_sequences( logits, # decoder hidden states [1, T, H] lm_head_weight, @@ -329,13 +353,8 @@ def forward_step(batch_iter, model): model_config = get_model_config(model) fp8_enabled = is_fp8_enabled(getattr(model_config, "fp8", None)) rollout_expert_indices = batch.pop("rollout_expert_indices", None) - if rollout_expert_indices is not None: - setup_per_microbatch_replay_forward( - rollout_expert_indices, - batch["attention_mask"], - model_config=model_config, - remove_microbatch_padding=self.remove_microbatch_padding, - ) + router_padding_mask = batch.pop("router_padding_mask", None) + sample_support_ids = batch.get("sample_support_ids") sequences = batch["sequences"] attention_mask = batch["attention_mask"].to(bool) @@ -343,6 +362,12 @@ def forward_step(batch_iter, model): sub_seq_lengths_field = batch.get("sub_seq_lengths") sub_seq_lengths = [t.tolist() for t in sub_seq_lengths_field] if sub_seq_lengths_field is not None else None batch["sub_seq_lengths_list"] = sub_seq_lengths + if ( + sample_support_ids is not None + and sub_seq_lengths is not None + and any(len(row_lengths) > 1 for row_lengths in sub_seq_lengths) + ): + raise ValueError("sample-support replay does not support controller-packed multi-subsequence rows") vlm_inputs = {} if batch.get("pixel_values") is not None and mpu.get_pipeline_model_parallel_rank() == 0: @@ -378,6 +403,27 @@ def forward_step(batch_iter, model): if self.is_vlm: new_position_ids = None + metadata_layout = None + if rollout_expert_indices is not None or (sample_support_ids is not None and packed_seq_params is not None): + metadata_layout = build_token_metadata_layout( + attention_mask, + attention_mask.device, + packed=packed_seq_params is not None, + fp8_enabled=fp8_enabled, + ) + + model_replay_kwargs = {} + if rollout_expert_indices is not None: + model_replay_kwargs = setup_per_microbatch_replay_forward( + rollout_expert_indices, + router_padding_mask, + attention_mask, + model=model, + model_config=model_config, + metadata_layout=metadata_layout, + remove_microbatch_padding=self.remove_microbatch_padding, + ) + if self._fused_lm_head: # Fused LM-head inference: the output_processor returns decoder # hidden states (not logits) and stashes the LM-head weight, so @@ -393,6 +439,7 @@ def forward_step(batch_iter, model): packed_seq_params=packed_seq_params, output_processor=_fused_lm_head_output_processor, output_processor_context=_op_ctx, + **model_replay_kwargs, **vlm_inputs, ) batch["lm_head_weight"] = _op_ctx.get("lm_head_weight") @@ -402,6 +449,7 @@ def forward_step(batch_iter, model): new_position_ids, to_te_attention_mask(new_attention_mask), packed_seq_params=packed_seq_params, + **model_replay_kwargs, **vlm_inputs, ) @@ -414,19 +462,21 @@ def forward_step(batch_iter, model): post_process=mpu.is_pipeline_last_stage(ignore_virtual=True), ) - return outputs, partial(collection_func, data=batch) + return outputs, partial(collection_func, data=batch, metadata_layout=metadata_layout) batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.actor_module)) - output = forward_backward_func( - forward_step_func=forward_step, - data_iterator=batch_generator, - model=self.actor_module, - num_microbatches=len(micro_batches), - seq_length=seq_len, - micro_batch_size=micro_batch_size, - forward_only=True, - ) + replay_enabled = any(batch["rollout_expert_indices"] is not None for batch in micro_batches) + with router_replay_schedule(replay_enabled): + output = forward_backward_func( + forward_step_func=forward_step, + data_iterator=batch_generator, + model=self.actor_module, + num_microbatches=len(micro_batches), + seq_length=seq_len, + micro_batch_size=micro_batch_size, + forward_only=True, + ) if mpu.is_pipeline_last_stage(ignore_virtual=True): log_probs = [o["log_probs"] for o in output] @@ -516,7 +566,7 @@ def forward_backward_mini_batch( # NOTE: users can provide a custom loss config class, so we need to use the same class after applying overrides loss_config = type(loss_config).from_dict_config(new_loss_config) - def loss_func(logits, data): + def loss_func(logits, *, data, metadata_layout: TokenMetadataLayout | None): sequences = data["sequences"] packed_seq_params = data.get("packed_seq_params") packed_targets = data.get("packed_targets") @@ -557,7 +607,26 @@ def loss_func(logits, data): if temperature != 1.0 and not fused_lm_head: logits.div_(temperature) - if fused_lm_head and packed_seq_params is not None and packed_targets is not None: + shard_vocab_size = lm_head_weight.shape[0] if fused_lm_head else logits.shape[-1] + if self.cfg.algorithm.enable_sample_support_replay: + token_logprobs = compute_sample_support_logprobs( + logits, + sequences, + loss_mask, + data.get("sample_support_ids"), + num_actions, + packed=packed_seq_params is not None, + metadata_layout=metadata_layout, + vocab_start_index=tp_rank * shard_vocab_size, + vocab_end_index=(tp_rank + 1) * shard_vocab_size, + tp_group=tp_grp, + lm_head_weight=lm_head_weight if fused_lm_head else None, + temperature=temperature, + inference_only=forward_only, + chunk_size=self.cfg.logprobs_chunk_size, + fused_backend=self._fused_lm_head_backend, + ) + elif fused_lm_head and packed_seq_params is not None and packed_targets is not None: token_logprobs = from_parallel_hidden_to_logprobs_packed_sequences( logits, # decoder hidden states [1, T, H] lm_head_weight, @@ -895,13 +964,8 @@ def forward_step(batch_iter, model): model_config = get_model_config(model) fp8_enabled = is_fp8_enabled(getattr(model_config, "fp8", None)) rollout_expert_indices = batch.pop("rollout_expert_indices", None) - if rollout_expert_indices is not None: - setup_per_microbatch_replay_forward( - rollout_expert_indices, - batch["attention_mask"], - model_config=model_config, - remove_microbatch_padding=self.remove_microbatch_padding, - ) + router_padding_mask = batch.pop("router_padding_mask", None) + sample_support_ids = batch.get("sample_support_ids") sequences = batch["sequences"] attention_mask = batch["attention_mask"].to(bool) @@ -917,6 +981,12 @@ def forward_step(batch_iter, model): sub_seq_lengths_field = batch.get("sub_seq_lengths") sub_seq_lengths = [t.tolist() for t in sub_seq_lengths_field] if sub_seq_lengths_field is not None else None batch["sub_seq_lengths_list"] = sub_seq_lengths + if ( + sample_support_ids is not None + and sub_seq_lengths is not None + and any(len(row_lengths) > 1 for row_lengths in sub_seq_lengths) + ): + raise ValueError("sample-support replay does not support controller-packed multi-subsequence rows") vlm_inputs = {} if batch.get("pixel_values") is not None and mpu.get_pipeline_model_parallel_rank() == 0: @@ -965,6 +1035,27 @@ def forward_step(batch_iter, model): is_last_stage = mpu.is_pipeline_last_stage(ignore_virtual=True) + metadata_layout = None + if rollout_expert_indices is not None or (sample_support_ids is not None and packed_seq_params is not None): + metadata_layout = build_token_metadata_layout( + attention_mask, + attention_mask.device, + packed=packed_seq_params is not None, + fp8_enabled=fp8_enabled, + ) + + model_replay_kwargs = {} + if rollout_expert_indices is not None: + model_replay_kwargs = setup_per_microbatch_replay_forward( + rollout_expert_indices, + router_padding_mask, + attention_mask, + model=model, + model_config=model_config, + metadata_layout=metadata_layout, + remove_microbatch_padding=self.remove_microbatch_padding, + ) + # Recover [batch, seq_len, ...] from Megatron's internal (left-removed) layout. Only used # on the non-packed path: with sample packing (remove_microbatch_padding) the logits stay # packed ([1, T, vocab]) and loss_func consumes packed_targets instead. MTP draft training @@ -1017,6 +1108,7 @@ def depad(tensor): packed_seq_params=packed_seq_params, output_processor=_fused_lm_head_output_processor, output_processor_context=_op_ctx, + **model_replay_kwargs, **vlm_inputs, ) batch["lm_head_weight"] = _op_ctx.get("lm_head_weight") @@ -1026,6 +1118,7 @@ def depad(tensor): new_position_ids, to_te_attention_mask(new_attention_mask), packed_seq_params=packed_seq_params, + **model_replay_kwargs, **vlm_inputs, ) # Replay the MTP block on *detached* trunk hidden states (decoupled draft forward) @@ -1054,20 +1147,22 @@ def depad(tensor): if rollout_expert_indices is not None: setup_per_microbatch_replay_backward() - return outputs, partial(loss_func, data=batch) + return outputs, partial(loss_func, data=batch, metadata_layout=metadata_layout) # batch should be a list of micro-batches batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.actor_module)) - metrics_list = forward_backward_func( - forward_step_func=forward_step, - data_iterator=batch_generator, - model=self.actor_module, - num_microbatches=len(micro_batches), - seq_length=seq_len, - micro_batch_size=micro_batch_size, - forward_only=forward_only, - ) + replay_enabled = any(batch["rollout_expert_indices"] is not None for batch in micro_batches) + with router_replay_schedule(replay_enabled): + metrics_list = forward_backward_func( + forward_step_func=forward_step, + data_iterator=batch_generator, + model=self.actor_module, + num_microbatches=len(micro_batches), + seq_length=seq_len, + micro_batch_size=micro_batch_size, + forward_only=forward_only, + ) # The decoupled MTP/draft loss is computed and logged per-microbatch inside loss_func # (metric key "mtp_loss"); no MTPLossLoggingHelper plumbing is needed. diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py index 9f2facaf61..1f26f4b503 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py @@ -73,6 +73,7 @@ from skyrl.env_vars import SKYRL_WORKER_NCCL_TIMEOUT_IN_S from skyrl.train.config.config import MegatronDDPConfig, get_config_as_dict from skyrl.train.utils.utils import str_to_torch_dtype, update_model_config +from skyrl.utils.routed_experts import make_replay_padding_indices from skyrl.utils.tok import get_tokenizer if TYPE_CHECKING: @@ -656,8 +657,6 @@ def _forward_logprobs(self, data: TrainingInputBatch) -> torch.Tensor: position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 0) rollout_expert_indices = micro.get("rollout_expert_indices") - if rollout_expert_indices is not None: - rollout_expert_indices = rollout_expert_indices.to(torch.int32) vlm_inputs = {} if micro.get("pixel_values") is not None: @@ -672,6 +671,11 @@ def _forward_logprobs(self, data: TrainingInputBatch) -> torch.Tensor: "position_ids": position_ids, "num_actions": micro.metadata["response_length"], "rollout_expert_indices": (rollout_expert_indices if self.enable_router_replay else None), + "router_padding_mask": micro.get("router_padding_mask") if self.enable_router_replay else None, + "sample_support_ids": ( + micro.get("sample_support_ids") if self.cfg.algorithm.enable_sample_support_replay else None + ), + "loss_mask": micro.get("loss_mask"), "sub_seq_lengths": micro.get("sub_seq_lengths"), **vlm_inputs, } @@ -787,6 +791,21 @@ def _pad_microbatch_to_size(self, micro_dict: dict, target_batch_size: int) -> d # position_ids for padded samples seq_len = value.shape[1] pad_tensor = torch.arange(seq_len, device=device).unsqueeze(0).expand(pad_count, -1) + elif key == "router_padding_mask": + pad_tensor = torch.ones((pad_count, *value.shape[1:]), dtype=torch.bool, device=device) + elif key == "rollout_expert_indices": + pad_tensor = make_replay_padding_indices( + (pad_count, *value.shape[1:]), + dtype=value.dtype, + device=device, + ) + elif key == "sample_support_ids": + pad_tensor = torch.full( + (pad_count, *value.shape[1:]), + -1, + dtype=value.dtype, + device=device, + ) elif key == "action_mask": # action_mask should be zeros for padded samples pad_tensor = torch.zeros((pad_count, *value.shape[1:]), dtype=value.dtype, device=device) @@ -907,9 +926,11 @@ def init_model(self, model_path, num_training_steps: int = 1e9): if self.enable_router_replay: from skyrl.backends.skyrl_train.utils.replay_utils import ( + patch_topk_router_expert_bias_padding_mask, patch_topk_router_layer_number, ) + patch_topk_router_expert_bias_padding_mask() patch_topk_router_layer_number() # Freeze MoE router params before optimizer build. @@ -1030,8 +1051,6 @@ def forward( position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 0) rollout_expert_indices = experience.rollout_expert_indices - if rollout_expert_indices is not None: - rollout_expert_indices = rollout_expert_indices.to(torch.int32) vlm_inputs = {} if experience.pixel_values is not None: @@ -1052,6 +1071,10 @@ def forward( "rollout_action_logprobs": experience.rollout_logprobs, "action_mask": experience.action_mask, "rollout_expert_indices": rollout_expert_indices if self.enable_router_replay else None, + "router_padding_mask": experience.router_padding_mask if self.enable_router_replay else None, + "sample_support_ids": ( + experience.sample_support_ids if self.cfg.algorithm.enable_sample_support_replay else None + ), "sub_seq_lengths": experience.sub_seq_lengths, **vlm_inputs, } @@ -1158,8 +1181,6 @@ def forward_backward( position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 0) rollout_expert_indices = experience.rollout_expert_indices - if rollout_expert_indices is not None: - rollout_expert_indices = rollout_expert_indices.to(torch.int32) vlm_inputs = {} if experience.pixel_values is not None: @@ -1180,6 +1201,10 @@ def forward_backward( "rollout_action_logprobs": experience.rollout_logprobs, "action_mask": experience.action_mask, "rollout_expert_indices": rollout_expert_indices if self.enable_router_replay else None, + "router_padding_mask": experience.router_padding_mask if self.enable_router_replay else None, + "sample_support_ids": ( + experience.sample_support_ids if self.cfg.algorithm.enable_sample_support_replay else None + ), # used with global sequence packing (None when token-based batching is active) "sub_seq_lengths": experience.sub_seq_lengths, "is_padding_batch": ( diff --git a/skyrl/backends/skyrl_train/workers/worker_utils.py b/skyrl/backends/skyrl_train/workers/worker_utils.py index 5642aa8f12..51e9f2f697 100644 --- a/skyrl/backends/skyrl_train/workers/worker_utils.py +++ b/skyrl/backends/skyrl_train/workers/worker_utils.py @@ -9,6 +9,7 @@ from skyrl.backends.skyrl_train.utils.torch_utils import masked_mean from skyrl.train.dataset.bin_packing import make_seq_packer from skyrl.train.dataset.replay_buffer import Experience +from skyrl.utils.routed_experts import make_replay_padding_indices # Metrics that end in `_loss` but are plain per-token MEANS, not pre-scaled minibatch sums. # The `sum_loss_metrics` convention sums every `_loss` key because the *policy* losses are @@ -162,6 +163,8 @@ def batch_to_experience(batch: TrainingInputBatch): num_actions=batch.metadata["response_length"], # int rollout_logprobs=batch.get("rollout_logprobs"), rollout_expert_indices=batch.get("rollout_expert_indices"), + sample_support_ids=batch.get("sample_support_ids"), + router_padding_mask=batch.get("router_padding_mask"), # additional info # can be used to log metrics etc for micro-batches in the worker info={}, @@ -325,8 +328,20 @@ def _create_padding_microbatch(self) -> TrainingInputBatch: data["rollout_logprobs"] = torch.zeros((batch_size, num_actions), dtype=ref_tensor.dtype, device=device) if self.data.get("rollout_expert_indices") is not None: ref_tensor = self.data["rollout_expert_indices"] - data["rollout_expert_indices"] = torch.zeros( - (batch_size, *ref_tensor.shape[1:]), dtype=ref_tensor.dtype, device=device + data["rollout_expert_indices"] = make_replay_padding_indices( + (batch_size, *ref_tensor.shape[1:]), + dtype=ref_tensor.dtype, + device=device, + ) + if self.data.get("router_padding_mask") is not None: + data["router_padding_mask"] = torch.ones((batch_size, seq_len), dtype=torch.bool, device=device) + if self.data.get("sample_support_ids") is not None: + ref_tensor = self.data["sample_support_ids"] + data["sample_support_ids"] = torch.full( + (batch_size, *ref_tensor.shape[1:]), + -1, + dtype=ref_tensor.dtype, + device=device, ) data.metadata = {} if self.data.metadata: diff --git a/skyrl/train/config/config.py b/skyrl/train/config/config.py index bb73cb8372..f1ff975398 100644 --- a/skyrl/train/config/config.py +++ b/skyrl/train/config/config.py @@ -645,6 +645,8 @@ class AlgorithmConfig(BaseConfig): use_tis: bool = False """Deprecated: use ``off_policy_correction`` instead.""" off_policy_correction: OffPolicyCorrectionConfig = field(default_factory=OffPolicyCorrectionConfig) + enable_sample_support_replay: bool = False + """Renormalize policy logprobs over the sampler's recorded support.""" sapo: SAPOConfig = field(default_factory=SAPOConfig) value_clip: float = 0.2 dynamic_sampling: DynamicSamplingConfig = field(default_factory=DynamicSamplingConfig) @@ -760,6 +762,8 @@ class InferenceEngineConfig(BaseConfig): enable_prefix_caching: bool = True enable_chunked_prefill: bool = True enable_return_routed_experts: bool = False + enable_return_sample_support_set: bool = False + """Return the bounded post-filter sampler support for each generated token.""" max_num_batched_tokens: int = 8192 enforce_eager: bool = False """Disable CUDA graphs for stability. Set to ``False`` for higher performance, @@ -1194,6 +1198,28 @@ def __post_init__(self): if self.trainer.algorithm.temperature is None: self.trainer.algorithm.temperature = self.generator.sampling_params.temperature + capture_sample_support = self.generator.inference_engine.enable_return_sample_support_set + replay_sample_support = self.trainer.algorithm.enable_sample_support_replay + sampling_params = self.generator.sampling_params + if replay_sample_support and not capture_sample_support: + raise ValueError( + "trainer.algorithm.enable_sample_support_replay requires " + "generator.inference_engine.enable_return_sample_support_set" + ) + if replay_sample_support and self.trainer.strategy != "megatron": + raise ValueError("sample-support replay requires trainer.strategy=megatron") + if capture_sample_support: + if sampling_params.temperature <= 0: + raise ValueError("sample-support capture requires generator.sampling_params.temperature > 0") + if sampling_params.top_k <= 1: + raise ValueError("sample-support capture requires generator.sampling_params.top_k > 1") + if sampling_params.repetition_penalty != 1.0: + raise ValueError("sample-support capture requires repetition_penalty=1.0") + if sampling_params.additional_kwargs: + raise ValueError("sample-support capture does not support sampling_params.additional_kwargs") + if self.generator.vision_language_generator: + raise ValueError("sample-support capture does not support vision_language_generator") + if self.data.dataloader.num_workers is None: self.data.dataloader.num_workers = 8 if self.data.dataloader.persistent_workers and self.data.dataloader.num_workers == 0: diff --git a/skyrl/train/dataset/preprocess.py b/skyrl/train/dataset/preprocess.py index 2b34521d1c..42e4da10c1 100644 --- a/skyrl/train/dataset/preprocess.py +++ b/skyrl/train/dataset/preprocess.py @@ -1,13 +1,61 @@ import logging from typing import List, Optional, Tuple +import numpy as np import torch -from jaxtyping import Float, Integer +from jaxtyping import Bool, Float, Integer from transformers import AutoTokenizer +from skyrl.utils.routed_experts import ( + ROUTED_EXPERT_DTYPES, + RoutedExpertIndices, + compact_routed_expert_indices, +) + logger = logging.getLogger(__name__) +def make_router_padding_mask( + attention_mask: torch.Tensor, + captured_route_lengths: List[int], +) -> Bool[torch.Tensor, "batch seq_len"]: + """Build Megatron's router-only padding mask for a ragged vLLM route prefix. + + vLLM records routes only for tokens it evaluates. The final training sequence can be + longer because the last sampled token has no subsequent decode forward, and SkyRL may + append a synthetic EOS. In multi-turn generation, observations join the captured prefix + only when a later turn evaluates them. Captured route rows therefore align with a prefix + of each real, left-padded sequence; the remaining suffix needs dummy routes. + + This cannot be derived from the loss mask. A loss-masked prompt or observation may still + condition later trained actions and must replay its captured route. ``True`` marks only + left padding and tokens without a captured route so Megatron excludes their dummy routes + from router accounting. + """ + if attention_mask.ndim != 2: + raise ValueError(f"Expected 2D attention_mask, got shape {attention_mask.shape}") + if len(captured_route_lengths) != attention_mask.shape[0]: + raise ValueError( + f"Expected one captured route length per trajectory, got {len(captured_route_lengths)} " + f"for batch size {attention_mask.shape[0]}" + ) + + captured = torch.as_tensor(captured_route_lengths, dtype=torch.long, device=attention_mask.device) + sequence_lengths = attention_mask.sum(dim=1, dtype=torch.long) + if torch.any(captured < 0) or torch.any(captured > sequence_lengths): + raise ValueError( + f"Captured route lengths must be within trajectory lengths, got " + f"captured={captured.tolist()} and lengths={sequence_lengths.tolist()}" + ) + + sequence_starts = attention_mask.shape[1] - sequence_lengths + positions = torch.arange(attention_mask.shape[1], device=attention_mask.device).unsqueeze(0) + captured_positions = (positions >= sequence_starts.unsqueeze(1)) & ( + positions < (sequence_starts + captured).unsqueeze(1) + ) + return ~captured_positions + + def _verify_inputs( prompts: List[List[int]], responses: List[List[int]], @@ -36,7 +84,7 @@ def convert_prompts_responses_to_batch_tensors( rewards: List[List[float]], loss_masks: List[List[int]], logprobs: Optional[List[List[float]]] = None, - rollout_expert_indices: Optional[List[List[List[List[int]]]]] = None, + rollout_expert_indices: Optional[List[RoutedExpertIndices]] = None, max_seq_len: Optional[int] = None, ) -> Tuple[ Float[torch.Tensor, "batch seq_len"], @@ -160,24 +208,52 @@ def convert_prompts_responses_to_batch_tensors( logprobs_tensor[i, max_response - len(sample_logprobs) :] = lp rollout_expert_indices_tensor = None - if rollout_expert_indices: - first_non_empty = next((x for x in rollout_expert_indices if x), None) - if first_non_empty: - num_layers = len(first_non_empty[0]) - topk = len(first_non_empty[0][0]) if num_layers > 0 else 0 - padded = torch.zeros(len(rollout_expert_indices), max_total, num_layers, topk, dtype=torch.int32) - for i, sample_indices in enumerate(rollout_expert_indices): - if sample_indices: - left_pad = max_total - (prompt_token_lens[i] + response_token_lens[i]) - n = min(len(sample_indices), max_total - left_pad) - padded[i, left_pad : left_pad + n] = torch.tensor(sample_indices[:n], dtype=torch.int32) - rollout_expert_indices_tensor = padded - - # downcast to uint8 if possible, otherwise int16 to save memory - if rollout_expert_indices_tensor.max().item() < 2**8: - rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.uint8) - elif rollout_expert_indices_tensor.max().item() < 2**15: - rollout_expert_indices_tensor = rollout_expert_indices_tensor.to(torch.int16) + if rollout_expert_indices is not None: + num_samples = len(prompts) + if not isinstance(rollout_expert_indices, list): + raise TypeError("rollout_expert_indices must be a list of NumPy arrays") + if len(rollout_expert_indices) != num_samples: + raise ValueError("rollout_expert_indices must contain routes for every trajectory") + + canonical_indices = [] + for sample_index, sample_indices in enumerate(rollout_expert_indices): + if not isinstance(sample_indices, np.ndarray): + raise TypeError( + f"rollout_expert_indices entries must be NumPy arrays, got {type(sample_indices).__name__} " + f"at sample {sample_index}" + ) + if sample_indices.dtype not in ROUTED_EXPERT_DTYPES: + raise TypeError( + f"Unsupported routed expert dtype {sample_indices.dtype} at sample {sample_index}; " + "expected uint8, int16, or int32" + ) + canonical_indices.append(compact_routed_expert_indices(sample_indices)) + + first_shape = canonical_indices[0].shape + if len(first_shape) != 3 or first_shape[0] == 0: + raise ValueError("rollout_expert_indices must contain routes for every trajectory") + num_layers, topk = first_shape[1:] + if topk < 1: + raise ValueError("rollout_expert_indices must contain at least one expert per layer") + + batch_dtype = max((indices.dtype for indices in canonical_indices), key=lambda dtype: dtype.itemsize) + padded = np.empty((num_samples, max_total, num_layers, topk), dtype=batch_dtype) + padded[...] = np.arange(topk, dtype=batch_dtype) + for sample_index, sample_indices in enumerate(canonical_indices): + if sample_indices.ndim != 3 or sample_indices.shape[1:] != (num_layers, topk): + raise ValueError( + "rollout_expert_indices entries must share [layers, topk], " + f"got shape {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) return ( sequences, @@ -190,6 +266,75 @@ def convert_prompts_responses_to_batch_tensors( ) +def build_dense_sample_support( + rollout_sample_support: Optional[List[List[List[int]]]], + response_ids: List[List[int]], + loss_masks: List[List[int]], + sequence_length: int, + top_k: int, + eos_token_id: int, +) -> Optional[Integer[torch.Tensor, "batch seq_len topk"]]: + """Validate and left-pad per-token sampler support for replay.""" + if rollout_sample_support is None: + return None + if len(rollout_sample_support) != len(response_ids): + raise ValueError("rollout_sample_support must have one entry per trajectory") + if len(loss_masks) != len(response_ids): + raise ValueError("loss_masks must have one entry per trajectory") + + support = torch.full((len(response_ids), sequence_length, top_k), -1, dtype=torch.int32) + int32_max = int(np.iinfo(np.int32).max) + for sample_index, (sample_rows, sampled_tokens, sample_loss_mask) in enumerate( + zip(rollout_sample_support, response_ids, loss_masks, strict=True) + ): + if len(sample_rows) != len(sampled_tokens): + raise ValueError( + f"rollout_sample_support[{sample_index}] has {len(sample_rows)} rows for " + f"{len(sampled_tokens)} response tokens" + ) + if len(sample_loss_mask) != len(sampled_tokens): + raise ValueError( + f"loss_masks[{sample_index}] has {len(sample_loss_mask)} entries for " + f"{len(sampled_tokens)} response tokens" + ) + + sample_support = torch.full((len(sample_rows), top_k), -1, dtype=torch.int64) + for token_index, row in enumerate(sample_rows): + if row: + if len(row) != top_k: + raise ValueError("rollout_sample_support rows must match generator.sampling_params.top_k") + sample_support[token_index] = torch.as_tensor(row, dtype=torch.int64) + + valid = sample_support >= 0 + if torch.any((sample_support < -1) | (sample_support > int32_max)): + raise ValueError("rollout_sample_support vocab ids must fit non-negative int32") + if torch.any(valid & ((~valid).cumsum(dim=1) > 0)): + raise ValueError("rollout_sample_support padding must use trailing -1 values") + + sampled = torch.as_tensor(sampled_tokens, dtype=torch.int64).unsqueeze(1) + loss_bearing = torch.as_tensor(sample_loss_mask, dtype=torch.bool) + has_support = valid.any(dim=1) + unsupported_loss = loss_bearing & ~has_support + if torch.count_nonzero(unsupported_loss) > 1: + raise ValueError(f"rollout_sample_support[{sample_index}] has more than one loss-bearing unsupported token") + unsupported_non_eos = unsupported_loss & (sampled.squeeze(1) != eos_token_id) + if torch.any(unsupported_non_eos): + token_index = int(torch.where(unsupported_non_eos)[0][0]) + raise ValueError( + f"rollout_sample_support[{sample_index}][{token_index}] is empty for a loss-bearing non-EOS token" + ) + missing = loss_bearing & has_support & ~torch.any(sample_support == sampled, dim=1) + if torch.any(missing): + missing_token = sampled_tokens[int(torch.where(missing)[0][0])] + raise ValueError(f"sampled token {missing_token} is missing from rollout_sample_support") + + start = sequence_length - len(sampled_tokens) + if start < 0: + raise ValueError("response tokens exceed the sample-support sequence width") + support[sample_index, start:] = sample_support.to(torch.int32) + return support + + def compute_prompt_boundaries(uids: List[str]) -> List[Tuple[int, int]]: """Compute per-prompt ``(start, end)`` slices from a flat ``uids`` list. diff --git a/skyrl/train/dataset/replay_buffer.py b/skyrl/train/dataset/replay_buffer.py index 77dd1b1aca..be13be203f 100644 --- a/skyrl/train/dataset/replay_buffer.py +++ b/skyrl/train/dataset/replay_buffer.py @@ -12,7 +12,7 @@ import torch import torch.nn.functional as F -from jaxtyping import Float, Integer +from jaxtyping import Bool, Float, Integer from skyrl.backends.skyrl_train.training_batch import TensorList @@ -70,6 +70,7 @@ class Experience: rollout_expert_indices: Optional[Integer[torch.Tensor, "batch seq_len layer_num topk"]] num_actions: int info: Optional[dict] + router_padding_mask: Optional[Bool[torch.Tensor, "batch seq_len"]] = None kl: Optional[Float[torch.Tensor, "batch response_len"]] = None metadata: Optional[Dict[str, Any]] = None pixel_values: Optional[TensorList] = None @@ -77,6 +78,7 @@ class Experience: # Per-row sub-sequence lengths for sequence packing (one 1-D int tensor per # packed row); ``None`` when packing is off. sub_seq_lengths: Optional[TensorList] = None + sample_support_ids: Optional[Integer[torch.Tensor, "batch seq_len topk"]] = None @torch.no_grad() def to_device(self, device: torch.device) -> None: @@ -101,6 +103,10 @@ def to_device(self, device: torch.device) -> None: self.rollout_logprobs = to(self.rollout_logprobs, device) if self.rollout_expert_indices is not None: self.rollout_expert_indices = to(self.rollout_expert_indices, device) + if self.sample_support_ids is not None: + self.sample_support_ids = to(self.sample_support_ids, device) + if self.router_padding_mask is not None: + self.router_padding_mask = to(self.router_padding_mask, device) if self.pixel_values is not None: self.pixel_values = self.pixel_values.to(device) if self.image_grid_thw is not None: @@ -130,6 +136,10 @@ def pin_memory(self): self.rollout_logprobs = self.rollout_logprobs.pin_memory() if self.rollout_expert_indices is not None: self.rollout_expert_indices = self.rollout_expert_indices.pin_memory() + if self.sample_support_ids is not None: + self.sample_support_ids = self.sample_support_ids.pin_memory() + if self.router_padding_mask is not None: + self.router_padding_mask = self.router_padding_mask.pin_memory() return self diff --git a/skyrl/train/generators/base.py b/skyrl/train/generators/base.py index 26d95868b9..2659634bbe 100644 --- a/skyrl/train/generators/base.py +++ b/skyrl/train/generators/base.py @@ -5,6 +5,7 @@ import torch from skyrl.backends.skyrl_train.inference_servers.base import ConversationType +from skyrl.utils.routed_experts import RoutedExpertIndices TrainingPhase = Literal["train", "eval"] @@ -46,7 +47,8 @@ class GeneratorOutput(TypedDict): # trajectory in the input batch (i.e. per ``agent_loop`` call). Used by the fully # async trainer to compute per-group / intra-group completion-time metrics. trajectory_generation_times: Optional[List[float]] - rollout_expert_indices: Optional[List[List[List[List[int]]]]] # [batch_size, seq_len, layer_num, topk] + rollout_expert_indices: Optional[List[RoutedExpertIndices]] + rollout_sample_support: Optional[List[List[List[int]]]] # Applicable only for step-wise training is_last_step: Optional[List[bool]] # Per-row env metrics (one dict per row in the flattened batch). Used by diff --git a/skyrl/train/generators/skyrl_gym_generator.py b/skyrl/train/generators/skyrl_gym_generator.py index cc9bc95e80..a2697abb38 100644 --- a/skyrl/train/generators/skyrl_gym_generator.py +++ b/skyrl/train/generators/skyrl_gym_generator.py @@ -13,6 +13,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union from uuid import uuid4 +import numpy as np import torch from loguru import logger from tqdm.asyncio import tqdm @@ -36,6 +37,8 @@ get_generation_prompt_ids, get_rollout_metrics, ) +from skyrl.utils.routed_experts import RoutedExpertIndices, RoutedExpertTrace +from skyrl.utils.token_metadata import TokenMetadataTrace from skyrl_gym.envs.base_text_env import BaseTextEnvStepOutput @@ -50,7 +53,8 @@ class TrajectoryOutput: prompt_ids: List[int] rollout_logprobs: Optional[List[float]] env_metrics: Dict[str, Any] - rollout_expert_indices: Optional[List[List[List[int]]]] = None + rollout_expert_indices: Optional[RoutedExpertIndices] = None + rollout_sample_support: Optional[List[List[int]]] = None pixel_values: Optional[torch.Tensor] = None image_grid_thw: Optional[torch.Tensor] = None # End-to-end wall-clock time (seconds) to generate this trajectory. Optional: agent loops may @@ -76,7 +80,8 @@ class AgentLoopState: rollout_logprobs: Optional[List[float]] response_end_idx: Optional[int] done: bool - rollout_expert_indices: Optional[List[List[List[int]]]] = None + routed_expert_trace: Optional[RoutedExpertTrace] = None + sample_support_trace: Optional[TokenMetadataTrace] = None @dataclass @@ -86,30 +91,22 @@ class TurnOutput: output_logprobs: Optional[List[float]] new_obs: ConversationType obs_ids: List[int] - rollout_expert_indices: Optional[List[List[List[int]]]] # [seq_len, layer_num, topk] reward: Optional[float] + rollout_sample_support: Optional[np.ndarray] = None added_eos: bool = False - def get_turn_rollout_expert_indices(self) -> Optional[List[List[List[int]]]]: - """ - Get rollout inference indices for this turn's tokens (output tokens + observation tokens). - - Returns indices for generated output tokens, with padding entries (all 0) - for any manually-added EOS token and observation tokens - Returns None if rollout_expert_indices is None. - """ - if self.rollout_expert_indices is None: + def get_turn_rollout_sample_support(self) -> Optional[np.ndarray]: + if self.rollout_sample_support is None: return None - if not self.rollout_expert_indices: - return self.rollout_expert_indices - layer_num = len(self.rollout_expert_indices[0]) - topk = len(self.rollout_expert_indices[0][0]) if layer_num > 0 else 0 - pad_entry = [[0] * topk for _ in range(layer_num)] - indices = list(self.rollout_expert_indices) - if self.added_eos: - indices.append(pad_entry) - indices.extend(pad_entry for _ in range(len(self.obs_ids))) - return indices + padding_count = int(self.added_eos) + len(self.obs_ids) + if not padding_count: + return self.rollout_sample_support + padding = np.full( + (padding_count, self.rollout_sample_support.shape[1]), + -1, + dtype=self.rollout_sample_support.dtype, + ) + return np.concatenate((self.rollout_sample_support, padding), axis=0) def get_turn_loss_mask(self) -> List[int]: """ @@ -365,6 +362,8 @@ async def agent_loop( current_sampling_params: dict = ( sampling_params if sampling_params is not None else asdict(self.generator_cfg.sampling_params) ) + capture_sample_support = self.generator_cfg.inference_engine.enable_return_sample_support_set + sample_support_width = current_sampling_params["top_k"] if capture_sample_support else 0 # Accumulate per-step rewards. Format: (reward, response_end_token_idx) per_step_rewards: List[Tuple[float, Optional[int]]] = [] @@ -381,6 +380,10 @@ 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 + ), + sample_support_trace=TokenMetadataTrace() if capture_sample_support and not is_step_wise else None, ) while not agent_loop_state.done: @@ -403,11 +406,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, ) engine_output = await self.inference_engine_client.generate(engine_input, model=self.policy_model_name) output = engine_output["responses"][0] @@ -426,6 +431,27 @@ 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, + ) + sample_support_rows = None + if capture_sample_support: + sample_support_rows = np.asarray( + engine_output["rollout_sample_support"][0], + dtype=np.int32, + order="C", + ).reshape(-1, sample_support_width) + if self.custom_chat_template is not None: + raise ValueError("Sample-support bookkeeping is not supported with custom chat template") + if sample_support_rows.shape[0] != len(output_ids): + raise ValueError( + f"Sample support has {sample_support_rows.shape[0]} rows for {len(output_ids)} tokens" + ) # 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) @@ -457,6 +483,10 @@ 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") + if sample_support_rows is not None: + raise ValueError("Sample-support bookkeeping is incompatible with postprocessed_action") obs_ids = self.get_obs_ids_from_obs(new_obs, agent_loop_state.done) @@ -468,13 +498,10 @@ async def agent_loop( new_obs=new_obs, reward=step_reward, obs_ids=obs_ids, + rollout_sample_support=sample_support_rows, added_eos=added_eos, - rollout_expert_indices=rollout_expert_indices, ) - if turn_output.rollout_expert_indices is not None and agent_loop_state.rollout_expert_indices is None: - agent_loop_state.rollout_expert_indices = [] - if is_step_wise: # current response + observation ids turn_response_ids = turn_output.output_ids + turn_output.obs_ids @@ -483,6 +510,7 @@ async def agent_loop( # agent loop only tracks loss mask and rollout logprobs for this turn with step_wise training turn_loss_mask = turn_output.get_turn_loss_mask() turn_response_logprobs: Optional[List[float]] = turn_output.get_turn_rollout_logprobs() + turn_sample_support = turn_output.get_turn_rollout_sample_support() per_step_output = TrajectoryOutput( response_ids=turn_response_ids, @@ -492,7 +520,9 @@ 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(), + rollout_sample_support=( + turn_sample_support.tolist() if turn_sample_support is not None else None + ), ) agent_loop_output.step_outputs.append(per_step_output) @@ -524,6 +554,7 @@ async def agent_loop( prompt_ids = agent_loop_state.input_ids[:initial_prompt_length] rollout_logprobs = None rollout_expert_indices_out = None + rollout_sample_support_out = None response_ids = None # Prepare the final loss_mask, response_ids and rollout_logprobs . @@ -554,10 +585,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( @@ -569,17 +596,29 @@ async def agent_loop( assert response_ids is not None and loss_mask is not None if stop_reason != "length" and response_ids and response_ids[-1] != self.tokenizer.eos_token_id: response_ids.append(self.tokenizer.eos_token_id) - # TODO(Charlie): this should be 0? Otherwise logprobs will be extremely off. But if it is loss - # masked with 0, why bother adding it? loss_mask.append(1) if rollout_logprobs is not None: rollout_logprobs.append(0.0) - if rollout_expert_indices_out is not None and rollout_expert_indices_out: - layer_num = len(rollout_expert_indices_out[0]) - topk = len(rollout_expert_indices_out[0][0]) if layer_num > 0 else 0 - rollout_expert_indices_out.append([[0] * topk for _ in range(layer_num)]) + if agent_loop_state.sample_support_trace is not None: + padding = np.full((1, sample_support_width), -1, dtype=np.int32) + agent_loop_state.sample_support_trace.append(padding, expected_rows=1) 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 agent_loop_state.sample_support_trace is not None and agent_loop_state.sample_support_trace.num_rows: + sample_support_rows = agent_loop_state.sample_support_trace.finalize( + expected_rows=agent_loop_state.sample_support_trace.num_rows + ) + if sample_support_rows.shape[0] < len(response_ids): + raise ValueError( + f"Sample-support trace has {sample_support_rows.shape[0]} rows for {len(response_ids)} tokens" + ) + rollout_sample_support_out = sample_support_rows[: len(response_ids)].tolist() + 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) @@ -598,6 +637,7 @@ async def agent_loop( rollout_logprobs=rollout_logprobs, env_metrics=env_metrics, rollout_expert_indices=rollout_expert_indices_out, + rollout_sample_support=rollout_sample_support_out, ) agent_loop_output = self._post_process_agent_loop_output( @@ -764,13 +804,17 @@ async def generate_batched( stop_reasons = engine_output["stop_reasons"] logprobs = engine_output.get("response_logprobs", None) raw_rollout_expert_indices = engine_output.get("rollout_expert_indices", None) + raw_rollout_sample_support = engine_output.get("rollout_sample_support", None) truncated_responses = [] rewards = [] loss_masks = [] env_metrics = [] truncated_logprobs: Optional[List[List[float]]] = [] if logprobs is not None else None - truncated_indices: Optional[List] = [] if raw_rollout_expert_indices is not None else None + truncated_indices: Optional[List[RoutedExpertIndices]] = [] if raw_rollout_expert_indices is not None else None + truncated_sample_support: Optional[List[List[List[int]]]] = ( + [] if raw_rollout_sample_support is not None else None + ) for i, (output, response, env, env_class) in enumerate(zip(outputs, responses, envs, env_classes)): # step on environment and compute reward @@ -789,6 +833,8 @@ async def generate_batched( sample_indices = raw_rollout_expert_indices[i] prompt_len = len(prompt_token_ids[i]) truncated_indices.append(sample_indices[: prompt_len + len(response)]) + if raw_rollout_sample_support is not None: + truncated_sample_support.append(raw_rollout_sample_support[i][: len(response)]) # Get environment-specific metrics env_metrics.append(env.get_metrics()) @@ -810,6 +856,7 @@ async def generate_batched( "rollout_metrics": rollout_metrics, "rollout_logprobs": truncated_logprobs, "rollout_expert_indices": truncated_indices, + "rollout_sample_support": truncated_sample_support, } return generator_output @@ -944,6 +991,16 @@ async def generate(self, input_batch: GeneratorInput, disable_tqdm: bool = False else: rollout_expert_indices = None + if self.generator_cfg.step_wise_trajectories: + sample_support_values = [ + step_output.rollout_sample_support for output in all_outputs for step_output in output.step_outputs + ] + else: + sample_support_values = [output.rollout_sample_support for output in all_outputs] + rollout_sample_support = ( + sample_support_values if any(value is not None for value in sample_support_values) else None + ) + rollout_metrics = get_rollout_metrics( responses, rewards, @@ -975,6 +1032,7 @@ async def generate(self, input_batch: GeneratorInput, disable_tqdm: bool = False # NOTE: for completion metrics, we output the completion time "trajectory_generation_times": out_trajectory_generation_times, "rollout_expert_indices": rollout_expert_indices, + "rollout_sample_support": rollout_sample_support, "is_last_step": is_last_step, "env_metrics": env_metrics, } @@ -1046,8 +1104,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( @@ -1099,17 +1155,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 @@ -1118,11 +1169,9 @@ def _update_agent_loop_state_with_multiturn_chat_template( agent_loop_state.loss_mask += loss_mask_for_turn if agent_loop_state.rollout_logprobs is not None and rollout_logprobs_for_turn is not None: agent_loop_state.rollout_logprobs += rollout_logprobs_for_turn - if agent_loop_state.rollout_expert_indices is not None and rollout_expert_indices_for_turn is not None: - # 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 + turn_sample_support = turn_output.get_turn_rollout_sample_support() + if agent_loop_state.sample_support_trace is not None and turn_sample_support is not None: + agent_loop_state.sample_support_trace.append(turn_sample_support, expected_rows=len(turn_ids)) return agent_loop_state @@ -1186,6 +1235,15 @@ def _update_agent_loop_state_with_singleturn_chat_template( rollout_logprobs_for_turn = turn_output.output_logprobs[: len(new_resp_tokens)] + [0.0] * len( obs_ids_to_add ) + turn_sample_support = None + if turn_output.rollout_sample_support is not None: + generated_support = turn_output.rollout_sample_support[: len(new_resp_tokens)] + observation_support = np.full( + (len(obs_ids_to_add), generated_support.shape[1]), + -1, + dtype=generated_support.dtype, + ) + turn_sample_support = np.concatenate((generated_support, observation_support), axis=0) # Directly append turn output agent_loop_state.response_end_idx = len(agent_loop_state.input_ids) + len(new_resp_tokens) - 1 @@ -1193,13 +1251,6 @@ 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 - + if agent_loop_state.sample_support_trace is not None and turn_sample_support is not None: + agent_loop_state.sample_support_trace.append(turn_sample_support, expected_rows=len(turn_ids)) return agent_loop_state diff --git a/skyrl/train/generators/utils.py b/skyrl/train/generators/utils.py index 4c5228f189..ac101501a5 100644 --- a/skyrl/train/generators/utils.py +++ b/skyrl/train/generators/utils.py @@ -745,19 +745,22 @@ def _is_prefix(maybe_prefix: List[int], candidate: List[int]) -> bool: return maybe_prefix == candidate[: len(maybe_prefix)] -def _slice_generator_output(generator_output: GeneratorOutput, indices: List[int]) -> GeneratorOutput: +def slice_generator_output( + generator_output: GeneratorOutput, indices: List[int], *, preserve_metrics: bool = True +) -> GeneratorOutput: """Slice a GeneratorOutput to keep only the entries at the given indices. - All sliced entries must share the same TrajectoryID — this helper is used by - prefix-aware merging which operates on one trajectory at a time. + Generator-specific per-trajectory fields are sliced without naming them here. + Prefix-aware merging passes entries that all share one ``TrajectoryID``; + dynamic sampling may intentionally select entries from different trajectories. """ assert len(indices) > 0, "indices must be non-empty" # Every key except `rollout_metrics` is either a per-entry list to slice, or None. sliced: GeneratorOutput = {} for key, value in generator_output.items(): if key == "rollout_metrics": - # Skip since metrics are already recorded before calling `merge_stepwise_output()`. - continue + if preserve_metrics: + sliced[key] = value elif value is None: sliced[key] = None else: @@ -784,6 +787,7 @@ def _merge_single_trajectory(gen_out: GeneratorOutput) -> GeneratorOutput: is_token_level_rewards = isinstance(gen_out["rewards"][0], list) has_logprobs = gen_out.get("rollout_logprobs") is not None has_stop_reasons = gen_out.get("stop_reasons") is not None + has_sample_support = gen_out.get("rollout_sample_support") is not None # Per-field output accumulators. # Fields that we take from all the entries in the merge group @@ -791,6 +795,7 @@ def _merge_single_trajectory(gen_out: GeneratorOutput) -> GeneratorOutput: out_response_ids: List[List[int]] = [] out_loss_masks: List[List[int]] = [] out_logprobs: Optional[List[List[float]]] = [] if has_logprobs else None + out_sample_support: Optional[List[List[List[int]]]] = [] if has_sample_support else None # If per-token rewards, we keep appending. If per-turn rewards, we only take from the last turn. out_rewards: list = [] @@ -804,16 +809,21 @@ def _merge_single_trajectory(gen_out: GeneratorOutput) -> GeneratorOutput: acc_response: List[int] = list(gen_out["response_ids"][0]) acc_loss_mask: List[int] = list(gen_out["loss_masks"][0]) acc_logprobs: Optional[List[float]] = list(gen_out["rollout_logprobs"][0]) if has_logprobs else None + acc_sample_support: Optional[List[List[int]]] = ( + [list(row) for row in gen_out["rollout_sample_support"][0]] if has_sample_support else None + ) acc_rewards_tokens: Optional[List[float]] = list(gen_out["rewards"][0]) if is_token_level_rewards else None last = 0 def flush(): - nonlocal acc_prompt, acc_response, acc_loss_mask, acc_logprobs, acc_rewards_tokens, last + nonlocal acc_prompt, acc_response, acc_loss_mask, acc_logprobs, acc_sample_support, acc_rewards_tokens, last out_prompt_ids.append(acc_prompt) out_response_ids.append(acc_response) out_loss_masks.append(acc_loss_mask) if has_logprobs: out_logprobs.append(acc_logprobs) + if has_sample_support: + out_sample_support.append(acc_sample_support) out_rewards.append(acc_rewards_tokens if is_token_level_rewards else gen_out["rewards"][last]) if has_stop_reasons: out_stop_reasons.append(gen_out["stop_reasons"][last]) @@ -831,6 +841,9 @@ def flush(): acc_response = list(gen_out["response_ids"][i]) acc_loss_mask = list(gen_out["loss_masks"][i]) acc_logprobs = list(gen_out["rollout_logprobs"][i]) if has_logprobs else None + acc_sample_support = ( + [list(row) for row in gen_out["rollout_sample_support"][i]] if has_sample_support else None + ) acc_rewards_tokens = list(gen_out["rewards"][i]) if is_token_level_rewards else None last = i continue @@ -845,6 +858,8 @@ def flush(): acc_loss_mask.extend([0] * len(obs_delta)) if acc_logprobs is not None: acc_logprobs.extend([0.0] * len(obs_delta)) + if acc_sample_support is not None: + acc_sample_support.extend([] for _ in obs_delta) if acc_rewards_tokens is not None: acc_rewards_tokens.extend([0.0] * len(obs_delta)) @@ -853,6 +868,8 @@ def flush(): acc_loss_mask.extend(gen_out["loss_masks"][i]) if acc_logprobs is not None: acc_logprobs.extend(gen_out["rollout_logprobs"][i]) + if acc_sample_support is not None: + acc_sample_support.extend(gen_out["rollout_sample_support"][i]) if acc_rewards_tokens is not None: acc_rewards_tokens.extend(gen_out["rewards"][i]) @@ -867,6 +884,7 @@ def flush(): "loss_masks": out_loss_masks, "stop_reasons": out_stop_reasons, "rollout_logprobs": out_logprobs, + "rollout_sample_support": out_sample_support, "trajectory_ids": out_trajectory_ids, "rollout_expert_indices": None, "is_last_step": out_is_last_step, @@ -913,7 +931,9 @@ def merge_stepwise_output(generator_output: GeneratorOutput) -> GeneratorOutput: start = 0 for i in range(num_samples): if is_last_step[i]: - trajectory_slices.append(_slice_generator_output(generator_output, list(range(start, i + 1)))) + trajectory_slices.append( + slice_generator_output(generator_output, list(range(start, i + 1)), preserve_metrics=False) + ) start = i + 1 merged_slices = [_merge_single_trajectory(s) for s in trajectory_slices] diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index 78dab4e84f..a88f27c987 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -52,9 +52,11 @@ from skyrl.train.config import SkyRLTrainConfig from skyrl.train.dataset import PromptDataset from skyrl.train.dataset.preprocess import ( + build_dense_sample_support, compute_prompt_boundaries, compute_prompt_mini_batch_boundaries, convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, ) from skyrl.train.evaluate import evaluate, evaluate_step_wise from skyrl.train.generators.base import ( @@ -863,9 +865,8 @@ def convert_to_training_input(self, generator_output: GeneratorOutput, uids: Lis loss_masks: List[List[int]] = generator_output["loss_masks"] logprobs: Optional[List[List[float]]] = generator_output.get("rollout_logprobs", None) - rollout_expert_indices: Optional[List[List[List[List[int]]]]] = generator_output.get( - "rollout_expert_indices", None - ) + rollout_expert_indices = generator_output.get("rollout_expert_indices", None) + rollout_sample_support = generator_output.get("rollout_sample_support", None) pixel_values = generator_output.get("pixel_values", None) image_grid_thw = generator_output.get("image_grid_thw", None) @@ -898,6 +899,20 @@ def convert_to_training_input(self, generator_output: GeneratorOutput, uids: Lis rollout_expert_indices, max_seq_len=self.cfg.trainer.algorithm.max_seq_len, ) + router_padding_mask = None + if rollout_expert_indices is not None: + router_padding_mask = make_router_padding_mask( + attention_masks_tensor, + [len(indices) for indices in rollout_expert_indices], + ) + sample_support_ids = build_dense_sample_support( + rollout_sample_support, + response_ids, + loss_masks, + sequences_tensor.shape[1], + self.cfg.generator.sampling_params.top_k, + self.tokenizer.eos_token_id, + ) # sanity check for off_policy_correction off_policy_correction = self.cfg.trainer.algorithm.off_policy_correction @@ -919,6 +934,8 @@ def convert_to_training_input(self, generator_output: GeneratorOutput, uids: Lis "loss_mask": loss_masks_tensor, "rollout_logprobs": rollout_logprobs_tensor, "rollout_expert_indices": rollout_expert_indices_tensor, + "router_padding_mask": router_padding_mask, + "sample_support_ids": sample_support_ids, "pixel_values": pixel_values, "image_grid_thw": image_grid_thw, }, @@ -1300,6 +1317,10 @@ def fwd_logprobs_values_reward( fwd_keys = ["sequences", "attention_mask"] if training_input.get("rollout_expert_indices") is not None: fwd_keys.append("rollout_expert_indices") + if training_input.get("router_padding_mask") is not None: + fwd_keys.append("router_padding_mask") + if training_input.get("sample_support_ids") is not None: + fwd_keys.extend(["sample_support_ids", "loss_mask"]) if training_input.get("pixel_values") is not None: fwd_keys.append("pixel_values") if training_input.get("image_grid_thw") is not None: diff --git a/skyrl/train/utils/trainer_utils.py b/skyrl/train/utils/trainer_utils.py index 84cb15b9be..f28c9ba8c9 100644 --- a/skyrl/train/utils/trainer_utils.py +++ b/skyrl/train/utils/trainer_utils.py @@ -27,6 +27,7 @@ from skyrl.train.generators.utils import ( concatenate_generator_outputs, get_metrics_from_generator_output, + slice_generator_output, ) BasicType = Union[int, float, str, bool, type(None)] @@ -489,20 +490,11 @@ def handle_replace_sampling( for uid in bad_uids: bad_indices.extend(uid2indices[uid]) - # Replace bad samples with good ones (modify in place because replacement_idx and bad_idx should not overlap) + source_indices = list(range(len(uids))) for bad_idx, replacement_idx in zip(bad_indices, replacement_indices): - generator_output["prompt_token_ids"][bad_idx] = generator_output["prompt_token_ids"][replacement_idx].copy() - generator_output["response_ids"][bad_idx] = generator_output["response_ids"][replacement_idx].copy() - replacement_reward = generator_output["rewards"][replacement_idx] - generator_output["rewards"][bad_idx] = ( - replacement_reward.copy() if isinstance(replacement_reward, list) else replacement_reward - ) - generator_output["loss_masks"][bad_idx] = generator_output["loss_masks"][replacement_idx].copy() - if generator_output["stop_reasons"]: - generator_output["stop_reasons"][bad_idx] = generator_output["stop_reasons"][replacement_idx] - - if generator_output["rollout_logprobs"]: - generator_output["rollout_logprobs"][bad_idx] = generator_output["rollout_logprobs"][replacement_idx] + source_indices[bad_idx] = replacement_idx + if bad_indices: + generator_output = slice_generator_output(generator_output, source_indices) # Update UIDs accordingly replaced_uids = uids.copy() @@ -631,22 +623,7 @@ def get_bad_sample_replacements(good_uids: List[str], bad_uids: List[str]) -> Li def filter_generator_output(output: GeneratorOutput, kept_indices: List[int]) -> GeneratorOutput: """Filter GeneratorOutput based on kept indices.""" - filtered = { - "prompt_token_ids": [output["prompt_token_ids"][i] for i in kept_indices], - "response_ids": [output["response_ids"][i] for i in kept_indices], - "rewards": [output["rewards"][i] for i in kept_indices], - "loss_masks": [output["loss_masks"][i] for i in kept_indices], - "stop_reasons": None, - "rollout_metrics": output.get("rollout_metrics"), - "rollout_logprobs": ( - [output["rollout_logprobs"][i] for i in kept_indices] if output["rollout_logprobs"] else None - ), - } - - if output.get("stop_reasons"): - filtered["stop_reasons"] = [output["stop_reasons"][i] for i in kept_indices] - - return filtered + return slice_generator_output(output, kept_indices) def zero_variance_filter( @@ -725,6 +702,7 @@ def validate_generator_output(num_prompts: int, generator_output: GeneratorOutpu "stop_reasons", "trajectory_ids", "rollout_expert_indices", + "rollout_sample_support", "is_last_step", "pixel_values", "image_grid_thw", @@ -754,6 +732,11 @@ def validate_generator_output(num_prompts: int, generator_output: GeneratorOutpu f"Response ids and rollout logprobs must have the same length, " f"for sample {i} got {len(response_ids)} and {len(generator_output['rollout_logprobs'][i])}" ) + if generator_output.get("rollout_sample_support") is not None: + assert len(response_ids) == len(generator_output["rollout_sample_support"][i]), ( + "Response ids and rollout sample support must have the same length, " + f"for sample {i} got {len(response_ids)} and {len(generator_output['rollout_sample_support'][i])}" + ) # loss masks should be non-zero for at least one element for trainer if np.concatenate(generator_output["loss_masks"]).sum() == 0: diff --git a/skyrl/utils/routed_experts.py b/skyrl/utils/routed_experts.py new file mode 100644 index 0000000000..b31cdaccda --- /dev/null +++ b/skyrl/utils/routed_experts.py @@ -0,0 +1,102 @@ +from collections.abc import Sequence +from typing import TypeAlias + +import numpy as np +import torch + +from skyrl.utils.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") + + 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): + raise TypeError("routed expert indices must be a NumPy array") + if routed_experts.ndim != 3 or not np.issubdtype(routed_experts.dtype, np.integer): + raise ValueError( + "routed expert indices must be an integer [tokens, layers, topk] array, " + f"got shape {routed_experts.shape} and dtype {routed_experts.dtype}" + ) + if int(routed_experts.min(initial=0)) < 0: + raise ValueError("routed expert indices must be non-negative") + + max_expert_id = int(routed_experts.max(initial=0)) + if max_expert_id < 2**8: + dtype = np.dtype(np.uint8) + elif max_expert_id < 2**15: + dtype = np.dtype(np.int16) + elif max_expert_id < 2**31: + dtype = np.dtype(np.int32) + else: + raise ValueError(f"routed expert index exceeds signed int32: {max_expert_id}") + + compact = np.asarray(routed_experts, dtype=dtype, order="C") + if not compact.flags.writeable: + compact = compact.copy(order="C") + return compact + + +def make_replay_padding_indices( + shape: tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device | str | int | None = None, +) -> torch.Tensor: + """Return dummy routes with ``topk`` distinct experts in every row.""" + if not shape or shape[-1] < 1: + raise ValueError(f"Replay route padding requires a positive topk dimension, got {shape}") + padding_row = torch.arange(shape[-1], dtype=dtype, device=device) + return padding_row.expand(shape).clone() diff --git a/skyrl/utils/token_metadata.py b/skyrl/utils/token_metadata.py new file mode 100644 index 0000000000..58985035af --- /dev/null +++ b/skyrl/utils/token_metadata.py @@ -0,0 +1,253 @@ +"""Token-aligned metadata layout transforms shared by training features.""" + +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, +) + + +def _new_metadata_tensor( + source: torch.Tensor, + shape: tuple[int, ...], + padding_value: torch.Tensor | bool | int, +) -> torch.Tensor: + output = torch.empty(shape, dtype=source.dtype, device=source.device) + output[...] = padding_value + return output + + +@dataclass(frozen=True) +class TokenMetadataLayout: + """One shared description of Megatron's token padding and CP sharding.""" + + attention_mask: torch.Tensor + sequence_lengths: list[int] + aligned_sequence_length: int + padded_sequence_lengths: list[int] | None = None + # Retained to reconstruct CP-sharded packed outputs in canonical batch order. + cu_seqlens_padded: torch.Tensor | None = None + context_parallel_size: int = 1 + context_parallel_rank: int = 0 + + +def build_token_metadata_layout( + attention_mask: torch.Tensor, + device: torch.device, + *, + packed: bool, + fp8_enabled: bool, +) -> TokenMetadataLayout: + """Compute the shared layout once for all replayed token metadata.""" + import megatron.core.parallel_state as mpu + + aligned_attention_mask = attention_mask.to(device=device, dtype=torch.bool) + sequence_lengths_tensor = aligned_attention_mask.sum(dim=1, dtype=torch.int32) + sequence_lengths = sequence_lengths_tensor.tolist() + tp_size = mpu.get_tensor_model_parallel_world_size() + + if not packed: + align_size = get_unpacked_seq_align_size(tp_size, fp8_enabled=fp8_enabled) + max_sequence_length = max(sequence_lengths) + aligned_sequence_length = max_sequence_length + (-max_sequence_length % align_size) + return TokenMetadataLayout( + attention_mask=aligned_attention_mask, + sequence_lengths=sequence_lengths, + aligned_sequence_length=aligned_sequence_length, + ) + + cp_size = mpu.get_context_parallel_world_size() + align_size = get_packed_seq_align_size(tp_size, cp_size, fp8_enabled=fp8_enabled) + padded_sequence_lengths_tensor = sequence_lengths_tensor + (-sequence_lengths_tensor % align_size) + padded_sequence_lengths = padded_sequence_lengths_tensor.tolist() + cu_seqlens_padded = torch.cat( + ( + torch.zeros(1, dtype=torch.int32, device=device), + padded_sequence_lengths_tensor.cumsum(dim=0), + ) + ) + return TokenMetadataLayout( + attention_mask=aligned_attention_mask, + sequence_lengths=sequence_lengths, + aligned_sequence_length=sum(padded_sequence_lengths), + padded_sequence_lengths=padded_sequence_lengths, + cu_seqlens_padded=cu_seqlens_padded, + context_parallel_size=cp_size, + context_parallel_rank=mpu.get_context_parallel_rank() if cp_size > 1 else 0, + ) + + +def align_token_metadata( + metadata: torch.Tensor, + layout: TokenMetadataLayout, + padding_value: torch.Tensor | bool | int, + *, + next_token: bool = False, +) -> torch.Tensor: + """Apply padding, optional next-token shifting, and CP sharding.""" + if metadata.device != layout.attention_mask.device: + raise ValueError("Token-aligned metadata and attention_mask must be on the same device") + if metadata.shape[:2] != layout.attention_mask.shape: + raise ValueError( + f"Token-aligned metadata shape {metadata.shape[:2]} does not match " + f"attention_mask shape {layout.attention_mask.shape}" + ) + + if layout.padded_sequence_lengths is None: + if next_token: + raise ValueError("next-token metadata alignment is only used for packed sequences") + aligned = _new_metadata_tensor( + metadata, + (metadata.shape[0], layout.aligned_sequence_length, *metadata.shape[2:]), + padding_value, + ) + for row_index, sequence_length in enumerate(layout.sequence_lengths): + aligned[row_index, :sequence_length] = metadata[row_index, layout.attention_mask[row_index]] + return aligned + + packed = _new_metadata_tensor( + metadata, + (layout.aligned_sequence_length, *metadata.shape[2:]), + padding_value, + ) + offset = 0 + for row_index, (sequence_length, padded_length) in enumerate( + zip(layout.sequence_lengths, layout.padded_sequence_lengths, strict=True) + ): + packed[offset : offset + sequence_length] = metadata[row_index, layout.attention_mask[row_index]] + # Match Megatron's [seq0, pad0, seq1, pad1, ...] microbatch layout. + offset += padded_length + + if next_token: + # Each packed logit predicts the next token within its own padded sequence. + shifted = _new_metadata_tensor(metadata, packed.shape, padding_value) + offset = 0 + for padded_length in layout.padded_sequence_lengths: + shifted[offset : offset + padded_length - 1] = packed[offset + 1 : offset + padded_length] + offset += padded_length + packed = shifted + + if layout.context_parallel_size > 1: + out = _new_metadata_tensor( + metadata, + (packed.shape[0] // layout.context_parallel_size, *packed.shape[1:]), + padding_value, + ) + src_offset = 0 + dst_offset = 0 + for padded_length in layout.padded_sequence_lengths: + # CP uses matching front/back chunks of each padded sequence. + length_per_cp = padded_length // layout.context_parallel_size + half = length_per_cp // 2 + front_start = src_offset + half * layout.context_parallel_rank + back_start = src_offset + padded_length - half * (layout.context_parallel_rank + 1) + out[dst_offset : dst_offset + half] = packed[front_start : front_start + half] + out[dst_offset + half : dst_offset + length_per_cp] = packed[back_start : back_start + half] + src_offset += padded_length + dst_offset += length_per_cp + packed = out + + return packed.unsqueeze(0) + + +def scatter_packed_token_values_to_batch( + model_values: torch.Tensor, + layout: TokenMetadataLayout, + padding_value: bool | int, +) -> torch.Tensor: + """Scatter packed model outputs into canonical ``[batch, seq_len - 1]`` positions.""" + if layout.padded_sequence_lengths is None or layout.cu_seqlens_padded is None: + raise ValueError("Scattering packed token values requires a packed metadata layout") + if model_values.ndim != 2 or model_values.shape[0] != 1: + raise ValueError(f"Expected packed model values with shape [1, tokens], got {model_values.shape}") + + values = model_values.squeeze(0) + if layout.context_parallel_size > 1: + import megatron.core.parallel_state as mpu + + from skyrl.backends.skyrl_train.distributed.megatron.model_utils import ( + allgather_cp_sharded_packed_tensor, + ) + + values = allgather_cp_sharded_packed_tensor( + values, + layout.cu_seqlens_padded, + mpu.get_context_parallel_group(), + ) + + from skyrl.backends.skyrl_train.distributed.megatron.model_utils import ( + _packed_sequence_indices, + ) + + _, _, sequence_indices, sequence_offsets, _ = _packed_sequence_indices( + layout.cu_seqlens_padded, + values.shape[0], + values.device, + ) + valid_counts = torch.tensor(layout.sequence_lengths, dtype=torch.long, device=values.device) - 1 + packed_mask = sequence_offsets < valid_counts[sequence_indices] + + attention_mask = layout.attention_mask + token_ordinals = attention_mask.to(torch.long).cumsum(dim=1) + output_mask = attention_mask[:, :-1] & ( + token_ordinals[:, :-1] < torch.tensor(layout.sequence_lengths, device=values.device).unsqueeze(1) + ) + batch_values = _new_metadata_tensor( + model_values, + (attention_mask.shape[0], attention_mask.shape[1] - 1), + padding_value, + ) + batch_values[output_mask] = values[packed_mask] + return batch_values + + +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/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py b/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py index 2e5db6b265..dbb53642a9 100644 --- a/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py +++ b/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_router_replay.py @@ -17,7 +17,10 @@ ) from skyrl.backends.skyrl_train.training_batch import TrainingInputBatch from skyrl.train.config import SamplingParams, SkyRLTrainConfig -from skyrl.train.dataset.preprocess import convert_prompts_responses_to_batch_tensors +from skyrl.train.dataset.preprocess import ( + convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, +) from skyrl.train.generators.base import GeneratorInput from skyrl.train.generators.skyrl_gym_generator import SkyRLGymGenerator from skyrl.train.utils.utils import validate_cfg @@ -201,6 +204,7 @@ async def test_logprobs(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs): ) assert rii_tensor is not None + router_padding_mask = make_router_padding_mask(attention_mask, [len(sample) for sample in indices]) num_actions = response_mask.shape[1] batch_size = sequences.shape[0] training_input = TrainingInputBatch( @@ -216,6 +220,7 @@ async def test_logprobs(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs): else torch.zeros((batch_size, num_actions), dtype=torch.float32) ), "rollout_expert_indices": rii_tensor, + "router_padding_mask": router_padding_mask, "action_log_probs": torch.zeros((batch_size, num_actions), dtype=torch.float32), "base_action_log_probs": torch.zeros((batch_size, num_actions), dtype=torch.float32), "advantages": torch.zeros((batch_size, num_actions), dtype=torch.float32), @@ -333,10 +338,15 @@ def test_forward_backward(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs MOONLIGHT_NUM_LAYERS = 27 MOONLIGHT_TOPK = 6 MOONLIGHT_NUM_EXPERTS = 64 - rollout_expert_indices = torch.randint( - 0, MOONLIGHT_NUM_EXPERTS, (batch_size, seq_len, MOONLIGHT_NUM_LAYERS, MOONLIGHT_TOPK), dtype=torch.int32 + route_start = torch.randint( + 0, + MOONLIGHT_NUM_EXPERTS, + (batch_size, seq_len, MOONLIGHT_NUM_LAYERS, 1), + dtype=torch.int32, ) - rollout_expert_indices[attention_mask == 0] = 0 + route_offsets = torch.arange(MOONLIGHT_TOPK, dtype=torch.int32) + rollout_expert_indices = (route_start + route_offsets) % MOONLIGHT_NUM_EXPERTS + rollout_expert_indices[attention_mask == 0] = route_offsets gen = torch.Generator().manual_seed(42) training_input = TrainingInputBatch( @@ -348,6 +358,7 @@ def test_forward_backward(ray_init_fixture, tp, pp, cp, ep, etp, extra_tf_kwargs "loss_mask": loss_mask_t, "rollout_logprobs": -torch.rand((batch_size, num_actions), generator=gen) * 2.0, "rollout_expert_indices": rollout_expert_indices, + "router_padding_mask": ~attention_mask.bool(), "action_log_probs": -torch.rand((batch_size, num_actions), generator=gen) * 2.0, "base_action_log_probs": -torch.rand((batch_size, num_actions), generator=gen) * 2.0, "advantages": torch.randn((batch_size, num_actions), generator=gen), diff --git a/tests/backends/skyrl_train/inference_servers/test_build_vllm_cli_args.py b/tests/backends/skyrl_train/inference_servers/test_build_vllm_cli_args.py index ad3372e06f..b76166787e 100644 --- a/tests/backends/skyrl_train/inference_servers/test_build_vllm_cli_args.py +++ b/tests/backends/skyrl_train/inference_servers/test_build_vllm_cli_args.py @@ -40,6 +40,21 @@ def test_build_vllm_cli_args_succeeds_on_gpu_less_host(monkeypatch): # tests/backends/skyrl_train/mtp/test_build_vllm_cli_args_mtp.py +@pytest.mark.vllm +def test_sample_support_uses_processed_top_k_logprobs(): + cfg = SkyRLTrainConfig.from_cli_overrides( + [ + "generator.inference_engine.enable_return_sample_support_set=true", + "generator.sampling_params.top_k=8", + ] + ) + + args = build_vllm_cli_args(cfg) + + assert args.max_logprobs == 8 + assert args.logprobs_mode == "processed_logprobs" + + def test_resolve_policy_model_name_uses_served_model_name(): cfg = SkyRLTrainConfig() cfg.trainer.policy.model.path = "base-model" diff --git a/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py b/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py index 65158253ba..712f8f777c 100644 --- a/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py +++ b/tests/backends/skyrl_train/inference_servers/test_remote_inference_client.py @@ -8,6 +8,7 @@ import aiohttp import httpx +import numpy as np import pytest import pytest_asyncio import uvicorn @@ -19,6 +20,10 @@ SKYRL_LORA_ADAPTER_NAME, PauseMode, RemoteInferenceClient, + RemoteInferenceGenerator, +) +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + pack_routed_experts, ) from skyrl.backends.skyrl_train.inference_servers.setup import ( build_new_inference_client, @@ -31,6 +36,7 @@ def create_mock_vllm_server(server_id: int) -> FastAPI: 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 @@ -56,6 +62,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 { @@ -92,6 +102,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) @@ -113,6 +124,9 @@ async def generate(request: Request): for i in range(num_choices) ] } + if request.url.path == "/skyrl/v1/generate": + routes = np.arange(12).reshape(3, 2, 2) + response["choices"][0]["routed_experts"] = pack_routed_experts(routes) features = body.get("features") app.state.last_generate_features = features @@ -417,8 +431,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._generator is None class TestDataPlane: @@ -450,6 +463,81 @@ async def test_generate_with_session_id(self, client): result = await client.generate(input_batch) assert len(result["responses"]) == 1 + @pytest.mark.asyncio + async def test_generate_decodes_packed_routed_experts(self, mock_servers): + client = RemoteInferenceClient( + proxy_url=mock_servers["proxy_url"], + server_urls=mock_servers["server_urls"], + data_parallel_size=1, + enable_return_routed_experts=True, + ) + try: + result = await client.generate({"prompt_token_ids": [[1, 2, 3]], "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_external_generator_requests_sample_support(self, monkeypatch): + generator = RemoteInferenceGenerator(proxy_url="http://unused") + captured = {} + + async def fake_post(url, json, headers): + captured.update(url=url, json=json, headers=headers) + return { + "choices": [ + { + "token_ids": [7], + "finish_reason": "stop", + "logprobs": {"content": [{"logprob": -0.1}]}, + "rollout_sample_support": [[7, 8]], + } + ] + } + + monkeypatch.setattr(generator, "_post", fake_post) + result = await generator.generate( + prompt_token_ids=[1, 2], + sampling_params={}, + session_id=None, + model="default", + return_sample_support=True, + ) + + assert captured["url"].endswith("/skyrl/v1/generate") + assert captured["json"]["return_sample_support"] is True + assert result.sample_support == [[7, 8]] + + @pytest.mark.asyncio + async def test_generate_rejects_list_routed_experts(self, monkeypatch): + client = RemoteInferenceClient( + proxy_url="http://unused", + server_urls=["http://unused"], + data_parallel_size=1, + enable_return_routed_experts=True, + ) + + async def return_list_routes(*args, **kwargs): + return { + "choices": [ + { + "token_ids": [1], + "finish_reason": "stop", + "routed_experts": [[[0, 1]]], + } + ] + } + + monkeypatch.setattr(client._get_generator(), "_post", return_list_routes) + with pytest.raises(ValueError, match="must return packed"): + await client._generate_single([1], {}, None, "model") + @pytest.mark.asyncio async def test_chat_completion(self, client): """Test chat completion method.""" @@ -928,8 +1016,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._generator is None or client._generator._session is None async def _get_lora_registries(server_urls: List[str]) -> List[Dict[str, str]]: diff --git a/tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py b/tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py new file mode 100644 index 0000000000..92655372e0 --- /dev/null +++ b/tests/backends/skyrl_train/inference_servers/test_routed_experts_wire.py @@ -0,0 +1,92 @@ +import base64 + +import numpy as np +import pytest + +from skyrl.backends.skyrl_train.inference_servers.routed_experts_wire import ( + decode_packed_routed_experts, + pack_routed_experts, +) +from skyrl.utils.routed_experts import compact_routed_expert_indices + + +@pytest.mark.parametrize( + "routes,expected_dtype", + [ + (np.arange(12).reshape(3, 2, 2), "uint8"), + (np.array([[[2**8 - 1]]]), "uint8"), + (np.array([[[0, 2**8]]]), "int16"), + (np.array([[[0, 2**15 - 1]]]), "int16"), + (np.array([[[0, 2**15]]]), "int32"), + (np.array([[[0, 2**31 - 1]]], dtype=np.int64), "int32"), + (np.empty((0, 2, 2), dtype=np.int64), "uint8"), + (np.arange(24).reshape(6, 2, 2)[::2], "uint8"), + ], +) +def test_packed_routed_experts_round_trip(routes, expected_dtype): + payload = pack_routed_experts(routes) + decoded = decode_packed_routed_experts(payload) + + assert payload["dtype"] == expected_dtype + assert decoded.dtype.name == expected_dtype + assert decoded.flags.c_contiguous + assert np.array_equal(decoded, routes) + + +def test_packed_routed_experts_uses_raw_base64(): + assert pack_routed_experts(np.array([[[1, 2, 3]]]))["data"] == "AQID" + + +@pytest.mark.parametrize( + "routes", + [np.array([1, 2]), np.array([[[-1]]]), np.array([[[2**31]]], dtype=np.uint64)], +) +def test_pack_rejects_invalid_routes(routes): + with pytest.raises(ValueError): + pack_routed_experts(routes) + + +def test_pack_rejects_nested_lists(): + with pytest.raises(TypeError, match="NumPy array"): + pack_routed_experts([[[1, 2]]]) + + +def test_compaction_makes_read_only_arrays_writable(): + routes = np.arange(12, dtype=np.uint8).reshape(3, 2, 2) + routes.flags.writeable = False + + compact = compact_routed_expert_indices(routes) + + assert compact.dtype == np.uint8 + assert compact.flags.c_contiguous + assert compact.flags.writeable + + +def test_decode_rejects_incorrect_byte_count(): + with pytest.raises(ValueError, match="bytes"): + decode_packed_routed_experts({"data": "AQ==", "shape": [2, 1, 1], "dtype": "uint8"}) + + +@pytest.mark.parametrize( + "payload", + [ + {"data": "AQ==", "shape": [1, 1, 1], "dtype": "uint16"}, + {"data": "!", "shape": [1, 1, 1], "dtype": "uint8"}, + {"data": "AQ==", "shape": [True, 1, 1], "dtype": "uint8"}, + ], +) +def test_decode_rejects_malformed_payloads(payload): + with pytest.raises(ValueError): + decode_packed_routed_experts(payload) + + +def test_decode_rejects_noncanonical_dtype(): + routes = np.array([[[300]]], dtype=np.int32) + payload = { + "data": base64.b64encode(routes.tobytes()).decode("ascii"), + "shape": [1, 1, 1], + "dtype": "int32", + } + + with pytest.raises(ValueError, match="non-canonical dtype"): + decode_packed_routed_experts(payload) diff --git a/tests/backends/skyrl_train/inference_servers/test_vllm_sample_support.py b/tests/backends/skyrl_train/inference_servers/test_vllm_sample_support.py new file mode 100644 index 0000000000..761b543525 --- /dev/null +++ b/tests/backends/skyrl_train/inference_servers/test_vllm_sample_support.py @@ -0,0 +1,95 @@ +from types import SimpleNamespace + +import pytest + +pytest.importorskip("vllm") + +from skyrl.backends.skyrl_train.inference_servers.vllm_server_actor import ( + _sample_support_from_flat_logprobs, +) + +pytestmark = pytest.mark.vllm + + +def test_flat_logprobs_extracts_sampled_scores_and_support_rows(): + flat_logprobs = SimpleNamespace( + token_ids=[7, 7, 8, 9, 4, 3, 4, 5], + logprobs=[-0.1, -0.1, -0.2, -0.3, -0.4, -0.2, -0.4, -0.6], + ) + + sampled, support = _sample_support_from_flat_logprobs(flat_logprobs, top_k=3) + + assert sampled == [{"logprob": -0.1}, {"logprob": -0.4}] + assert support == [[7, 8, 9], [3, 4, 5]] + + +def test_flat_logprobs_replaces_top_p_masked_candidates(): + flat_logprobs = SimpleNamespace( + token_ids=[7, 7, 8, 9], + logprobs=[-0.1, -0.1, -0.2, float("-inf")], + ) + + _, support = _sample_support_from_flat_logprobs(flat_logprobs, top_k=3) + + assert support == [[7, 8, -1]] + + +def test_flat_logprobs_repairs_sampled_token_absent_from_support(): + # Three rows, top_k=3 (row_width=4): + # Row A: sampled id (100) absent from a fully-valid support row -> repair. + # Row B: sampled id (7) already present -> unchanged. + # Row C: sampled id (5) absent from a support row that has trailing -1 padding. + top_k = 3 + flat_logprobs = SimpleNamespace( + token_ids=[100, 8, 9, 10, 7, 7, 8, 9, 5, 6, 7, 8], + logprobs=[ + -0.1, + -0.2, + -0.3, + -0.4, # row A: all valid + -0.1, + -0.1, + -0.2, + -0.3, # row B: all valid + -0.4, + -0.5, + -0.6, + float("-inf"), # row C: last col filtered -> padding + ], + ) + + _, support = _sample_support_from_flat_logprobs(flat_logprobs, top_k=top_k) + sampled_ids = [100, 7, 5] + + # (b) each row keeps width == top_k + assert all(len(row) == top_k for row in support) + + # (a) every row's support now contains its sampled id + for sampled_id, row in zip(sampled_ids, support): + assert sampled_id in row + + # (c) the sampled id appears exactly once per repaired row (no duplicate) + assert support[0].count(100) == 1 + assert support[2].count(5) == 1 + + # (d) trailing -1 padding preserved on the padded row + assert support[2][-1] == -1 + + # (e) the unaffected row (sampled already present) is unchanged + assert support[1] == [7, 8, 9] + + # Concrete expected repair: weakest (trailing) valid member overwritten. + assert support[0] == [8, 9, 100] + assert support[2] == [6, 5, -1] + + +def test_flat_logprobs_top_k_one_repairs_single_support_column(): + # top_k == 1 (row_width == 2): a single support column that must hold the sampled id. + flat_logprobs = SimpleNamespace( + token_ids=[42, 9], + logprobs=[-0.1, -0.2], + ) + + _, support = _sample_support_from_flat_logprobs(flat_logprobs, top_k=1) + + assert support == [[42]] diff --git a/tests/backends/skyrl_train/test_token_based_batching_utils.py b/tests/backends/skyrl_train/test_token_based_batching_utils.py index 26a2b3c561..dd9d4710bc 100644 --- a/tests/backends/skyrl_train/test_token_based_batching_utils.py +++ b/tests/backends/skyrl_train/test_token_based_batching_utils.py @@ -199,6 +199,20 @@ def test_padding_microbatch_matches_seq_len(self): # Padding rows must not contribute to the loss. assert padding["loss_mask"].sum().item() == 0 + def test_padding_microbatch_uses_unique_dummy_routes(self): + batch = self._make_batch([4, 4], num_actions=2) + batch["rollout_expert_indices"] = torch.full((2, 4, 2, 3), 7, dtype=torch.int16) + batch["router_padding_mask"] = torch.zeros((2, 4), dtype=torch.bool) + batch["sample_support_ids"] = torch.full((2, 4, 8), 7, dtype=torch.int32) + iterator = TokenBasedBatchIterator(batch, max_tokens_per_microbatch=8) + + padding = iterator._create_padding_microbatch() + + expected = torch.tensor([0, 1, 2], dtype=torch.int16).expand_as(padding["rollout_expert_indices"]) + assert torch.equal(padding["rollout_expert_indices"], expected) + assert torch.all(padding["router_padding_mask"]) + assert torch.all(padding["sample_support_ids"] == -1) + def test_multimodal_tensorlist_microbatching(self): """Token-based microbatching must gather TensorList fields (multi-modal pixel_values / image_grid_thw) via the same index gather used for regular tensors.""" diff --git a/tests/backends/skyrl_train/test_train_batch.py b/tests/backends/skyrl_train/test_train_batch.py index b105750957..dd6546040d 100644 --- a/tests/backends/skyrl_train/test_train_batch.py +++ b/tests/backends/skyrl_train/test_train_batch.py @@ -551,6 +551,8 @@ def test_tensor_batch_none_tensor_list(): "rewards", "rollout_logprobs", "rollout_expert_indices", + "router_padding_mask", + "sample_support_ids", "pixel_values", "image_grid_thw", } @@ -576,6 +578,8 @@ def _make_full_training_batch(batch_size: int = 4, seq_len: int = 5) -> Training "rewards": torch.randn(batch_size, seq_len), "rollout_logprobs": torch.randn(batch_size, seq_len), "rollout_expert_indices": torch.randint(0, 8, (batch_size, seq_len, 2, 3), dtype=torch.long), + "router_padding_mask": torch.zeros((batch_size, seq_len), dtype=torch.bool), + "sample_support_ids": torch.randint(0, 100, (batch_size, seq_len, 4), dtype=torch.int32), "pixel_values": TensorList([torch.randn(i + 1, 3) for i in range(batch_size)]), # batch_size * (i + 1) * 3 "image_grid_thw": TensorList([torch.tensor([[1, 2, 3]]) for _ in range(batch_size)]), # batch_size * 1 * 3 } @@ -636,7 +640,22 @@ def test_pad_batch_all_fields(): # Regular tensor fields (not loss_mask, not TensorList): original rows untouched, # padding rows are copies of row 0. - regular_tensor_keys = EXPECTED_TRAINING_INPUT_FIELDS - {"loss_mask", "pixel_values", "image_grid_thw"} + assert torch.equal(padded["router_padding_mask"][:batch_size], batch["router_padding_mask"]) + assert torch.all(padded["router_padding_mask"][batch_size:]) + assert torch.equal(padded["rollout_expert_indices"][:batch_size], batch["rollout_expert_indices"]) + expected_routes = torch.tensor([0, 1, 2]).expand_as(padded["rollout_expert_indices"][batch_size:]) + assert torch.equal(padded["rollout_expert_indices"][batch_size:], expected_routes) + assert torch.equal(padded["sample_support_ids"][:batch_size], batch["sample_support_ids"]) + assert torch.all(padded["sample_support_ids"][batch_size:] == -1) + + regular_tensor_keys = EXPECTED_TRAINING_INPUT_FIELDS - { + "loss_mask", + "rollout_expert_indices", + "router_padding_mask", + "sample_support_ids", + "pixel_values", + "image_grid_thw", + } for key in regular_tensor_keys: assert torch.equal(padded[key][:batch_size], batch[key]), f"Original rows changed for {key!r}" for i in range(batch_size, batch_size + pad_size): diff --git a/tests/backends/skyrl_train/utils/test_replay_utils.py b/tests/backends/skyrl_train/utils/test_replay_utils.py new file mode 100644 index 0000000000..4ca819517e --- /dev/null +++ b/tests/backends/skyrl_train/utils/test_replay_utils.py @@ -0,0 +1,229 @@ +import inspect +import sys +import types +from types import SimpleNamespace + +import pytest +import torch + +from skyrl.backends.skyrl_train.utils import replay_utils +from skyrl.utils.routed_experts import make_replay_padding_indices +from skyrl.utils.token_metadata import build_token_metadata_layout + + +@pytest.fixture +def parallel_state(monkeypatch): + try: + import megatron.core.parallel_state as mpu + except ModuleNotFoundError: + megatron = types.ModuleType("megatron") + core = types.ModuleType("megatron.core") + mpu = types.ModuleType("megatron.core.parallel_state") + megatron.core = core + core.parallel_state = mpu + monkeypatch.setitem(sys.modules, "megatron", megatron) + monkeypatch.setitem(sys.modules, "megatron.core", core) + monkeypatch.setitem(sys.modules, "megatron.core.parallel_state", mpu) + + monkeypatch.setattr(mpu, "get_tensor_model_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_rank", lambda: 0, raising=False) + return mpu + + +def test_patch_topk_router_expert_bias_excludes_padding(monkeypatch): + router_module = types.ModuleType("megatron.core.transformer.moe.router") + + class TopKRouter: + def __init__(self): + self.local_tokens_per_expert = torch.zeros(3, dtype=torch.int64) + + def _apply_expert_bias(self, routing_map, padding_mask=None): + if padding_mask is not None: + routing_map = routing_map & (~padding_mask) + self.local_tokens_per_expert += routing_map.sum(dim=0) + + router_module.TopKRouter = TopKRouter + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router", router_module) + + replay_utils.patch_topk_router_expert_bias_padding_mask() + router = TopKRouter() + router._apply_expert_bias( + torch.tensor([[1, 0, 1], [0, 1, 1]], dtype=torch.bool), + torch.tensor([False, True]), + ) + + assert torch.equal(router.local_tokens_per_expert, torch.tensor([1, 0, 1])) + + +@pytest.mark.parametrize("dtype", [torch.uint8, torch.int16, torch.int32]) +def test_replay_padding_indices_are_unique(dtype): + padding = make_replay_padding_indices((2, 3, 4, 3), dtype=dtype) + + assert padding.shape == (2, 3, 4, 3) + assert torch.equal(padding, torch.tensor([0, 1, 2], dtype=dtype).expand_as(padding)) + + +def test_replay_has_no_dispatcher_specific_patch(): + assert "TokenDispatcher" not in inspect.getsource(replay_utils) + + +@pytest.mark.parametrize("route_dtype", [torch.uint8, torch.int16, torch.int32]) +def test_setup_replay_installs_indices_and_returns_model_mask(monkeypatch, parallel_state, route_dtype): + router_replay_module = types.ModuleType("megatron.core.transformer.moe.router_replay") + + class RouterReplay: + global_router_replay_instances = [object()] + replay_data = None + action = None + + @classmethod + def set_replay_data(cls, replay_data): + cls.replay_data = replay_data + + @classmethod + def set_global_router_replay_action(cls, action): + cls.action = action + + class RouterReplayAction: + REPLAY_FORWARD = "replay_forward" + + router_replay_module.RouterReplay = RouterReplay + router_replay_module.RouterReplayAction = RouterReplayAction + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router_replay", router_replay_module) + monkeypatch.setattr(replay_utils, "_get_current_pp_stage_layer_range", lambda model_config: (1, 1)) + monkeypatch.setattr( + replay_utils, + "scatter_router_padding_mask_for_model", + lambda mask, model, model_config: mask, + ) + apply_layout = replay_utils.align_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]) + return apply_layout(metadata, layout, padding_value) + + monkeypatch.setattr(replay_utils, "align_token_metadata", record_routed_layer_count) + + routes = torch.tensor( + [ + [ + [[0, 1], [0, 1], [0, 1]], + [[10, 11], [1, 2], [20, 21]], + [[12, 13], [3, 4], [22, 23]], + [[14, 15], [5, 6], [24, 25]], + ] + ], + dtype=route_dtype, + ) + attention_mask = torch.tensor([[0, 1, 1, 1]]) + router_padding_mask = torch.tensor([[1, 0, 0, 1]], dtype=torch.bool) + metadata_layout = build_token_metadata_layout( + attention_mask, + routes.device, + packed=False, + fp8_enabled=False, + ) + + model_kwargs = replay_utils.setup_per_microbatch_replay_forward( + routes, + router_padding_mask, + attention_mask, + model=object(), + model_config=SimpleNamespace(fp8=None), + metadata_layout=metadata_layout, + ) + + assert RouterReplay.replay_data[0].tolist() == [[1, 2], [3, 4], [5, 6]] + assert RouterReplay.replay_data[0].dtype == torch.int32 + assert RouterReplay.action == RouterReplayAction.REPLAY_FORWARD + assert model_kwargs["padding_mask"].tolist() == [[False, False, True]] + assert routed_layer_counts == [1] + + +@pytest.mark.parametrize( + ("model_kind", "pre_process", "expected"), + [ + ("gpt", True, [[False, False, True, True]]), + ("gpt", False, [[True, True]]), + ("hybrid", True, [[True, True]]), + ], +) +def test_sequence_parallel_mask_layout(monkeypatch, model_kind, pre_process, expected): + hybrid_model = types.ModuleType("megatron.core.models.hybrid.hybrid_model") + tensor_parallel = types.ModuleType("megatron.core.tensor_parallel") + utils = types.ModuleType("megatron.core.utils") + + class HybridModel: + def __init__(self): + self.pre_process = pre_process + + class GPTModel: + def __init__(self): + self.pre_process = pre_process + + hybrid_model.HybridModel = HybridModel + tensor_parallel.scatter_to_sequence_parallel_region = lambda value: value.chunk(2, dim=0)[1] + utils.unwrap_model = lambda model: model + monkeypatch.setitem(sys.modules, "megatron.core.models.hybrid.hybrid_model", hybrid_model) + monkeypatch.setitem(sys.modules, "megatron.core.tensor_parallel", tensor_parallel) + monkeypatch.setitem(sys.modules, "megatron.core.utils", utils) + + mask = torch.tensor([[0, 0, 1, 1]], dtype=torch.bool) + model = HybridModel() if model_kind == "hybrid" else GPTModel() + scattered = replay_utils.scatter_router_padding_mask_for_model( + mask, + model, + SimpleNamespace(sequence_parallel=True), + ) + + assert scattered.tolist() == expected + + +@pytest.fixture +def router_replay_module(monkeypatch): + module = types.ModuleType("megatron.core.transformer.moe.router_replay") + router = SimpleNamespace(replay_backward_list=[], action=None) + + class RouterReplay: + global_router_replay_instances = [router] + + @classmethod + def clear_global_indices(cls): + for instance in cls.global_router_replay_instances: + instance.replay_backward_list = [] + + @classmethod + def clear_global_router_replay_action(cls): + for instance in cls.global_router_replay_instances: + instance.action = None + + module.RouterReplay = RouterReplay + monkeypatch.setitem(sys.modules, "megatron.core.transformer.moe.router_replay", module) + return router + + +def test_router_replay_schedule_clears_stale_forward_only_fifo(router_replay_module): + router_replay_module.replay_backward_list = ["stale-forward-only"] + + with replay_utils.router_replay_schedule(enabled=True): + assert router_replay_module.replay_backward_list == [] + router_replay_module.replay_backward_list.extend(["microbatch-0", "microbatch-1"]) + assert router_replay_module.replay_backward_list.pop(0) == "microbatch-0" + assert router_replay_module.replay_backward_list.pop(0) == "microbatch-1" + + assert router_replay_module.replay_backward_list == [] + assert router_replay_module.action is None + + +def test_router_replay_schedule_clears_after_exception(router_replay_module): + with pytest.raises(RuntimeError, match="schedule failed"): + with replay_utils.router_replay_schedule(enabled=True): + router_replay_module.replay_backward_list.append("partially-consumed-schedule") + router_replay_module.action = "replay-backward" + raise RuntimeError("schedule failed") + + assert router_replay_module.replay_backward_list == [] + assert router_replay_module.action is None diff --git a/tests/backends/skyrl_train/utils/test_sample_support_replay.py b/tests/backends/skyrl_train/utils/test_sample_support_replay.py new file mode 100644 index 0000000000..752beaddbd --- /dev/null +++ b/tests/backends/skyrl_train/utils/test_sample_support_replay.py @@ -0,0 +1,234 @@ +import sys +import types + +import pytest +import torch + +from skyrl.backends.skyrl_train.utils.sample_support_replay import ( + sample_support_logprobs, + synthetic_eos_logprobs, +) +from skyrl.utils.token_metadata import TokenMetadataLayout + + +def _reference(logits, sampled_ids, support_ids): + outputs = [] + for row_logits, sampled_id, support in zip( + logits.reshape(-1, logits.shape[-1]), + sampled_ids.reshape(-1), + support_ids.reshape(-1, support_ids.shape[-1]), + strict=True, + ): + members = support[support >= 0].long() + outputs.append( + row_logits.new_zeros(()) + if members.numel() == 0 + else row_logits[sampled_id] - torch.logsumexp(row_logits[members], dim=0) + ) + return torch.stack(outputs).reshape(sampled_ids.shape) + + +def test_support_logprobs_match_reference_values_and_gradients(): + logits = torch.randn(2, 3, 11, dtype=torch.float64, requires_grad=True) + sampled_ids = torch.tensor([[2, 5, 1], [8, 3, 7]]) + support_ids = torch.tensor( + [ + [[2, 4, 6, -1], [5, -1, -1, -1], [-1, -1, -1, -1]], + [[8, 0, 9, 4], [3, 2, -1, -1], [7, 1, 5, -1]], + ], + dtype=torch.int32, + ) + + actual, valid = sample_support_logprobs( + logits, + sampled_ids, + support_ids, + vocab_start_index=0, + vocab_end_index=logits.shape[-1], + tp_group=None, + ) + expected = _reference(logits, sampled_ids, support_ids) + + assert valid.tolist() == [[True, True, False], [True, True, True]] + torch.testing.assert_close(actual, expected) + actual.sum().backward() + actual_grad = logits.grad.clone() + + reference_logits = logits.detach().clone().requires_grad_(True) + _reference(reference_logits, sampled_ids, support_ids).sum().backward() + torch.testing.assert_close(actual_grad, reference_logits.grad) + + +def test_fused_selected_projection_matches_explicit_logits_with_pair_chunking(): + temperature = 0.7 + hidden = torch.randn(2, 3, 5, dtype=torch.float64, requires_grad=True) + weight = torch.randn(9, 5, dtype=torch.float64, requires_grad=True) + sampled_ids = torch.tensor([[1, 4, 7], [2, 5, 8]]) + support_ids = torch.tensor( + [ + [[1, 0, 3], [4, 6, -1], [7, -1, -1]], + [[2, 1, 8], [5, 4, -1], [8, 0, 6]], + ], + dtype=torch.int32, + ) + + fused, _ = sample_support_logprobs( + hidden, + sampled_ids, + support_ids, + vocab_start_index=0, + vocab_end_index=weight.shape[0], + tp_group=None, + lm_head_weight=weight, + temperature=temperature, + chunk_size=4, + ) + fused.sum().backward() + fused_hidden_grad = hidden.grad.clone() + fused_weight_grad = weight.grad.clone() + + explicit_hidden = hidden.detach().clone().requires_grad_(True) + explicit_weight = weight.detach().clone().requires_grad_(True) + explicit_logits = (explicit_hidden @ explicit_weight.T) / temperature + explicit, _ = sample_support_logprobs( + explicit_logits, + sampled_ids, + support_ids, + vocab_start_index=0, + vocab_end_index=weight.shape[0], + tp_group=None, + ) + explicit.sum().backward() + + torch.testing.assert_close(fused, explicit, check_dtype=False) + torch.testing.assert_close(fused_hidden_grad, explicit_hidden.grad, rtol=1e-5, atol=1e-6) + torch.testing.assert_close(fused_weight_grad, explicit_weight.grad, rtol=1e-5, atol=1e-6) + + +def test_support_ids_must_be_int32(): + with pytest.raises(ValueError, match="int32"): + sample_support_logprobs( + torch.randn(1, 5), + torch.tensor([1]), + torch.tensor([[1, 2]], dtype=torch.int64), + vocab_start_index=0, + vocab_end_index=5, + tp_group=None, + ) + + +def _install_fake_distributed_logprob(monkeypatch, calls): + model_utils = types.ModuleType("skyrl.backends.skyrl_train.distributed.megatron.model_utils") + + class DistributedLogprob: + @staticmethod + def apply(source, targets, *args): + calls.append(source.shape) + return source.gather(-1, targets.unsqueeze(-1)).squeeze(-1) + + model_utils.DistributedLogprob = DistributedLogprob + monkeypatch.setitem(sys.modules, model_utils.__name__, model_utils) + return model_utils + + +def test_synthetic_eos_uses_one_fixed_slot_per_unpacked_trajectory(monkeypatch): + calls = [] + _install_fake_distributed_logprob(monkeypatch, calls) + logits = torch.arange(3 * 4 * 5, dtype=torch.float64).reshape(3, 4, 5).requires_grad_(True) + sampled_ids = torch.tensor([[0, 1, 2, 3], [1, 2, 3, 4], [2, 3, 4, 0]]) + synthetic_eos_mask = torch.tensor( + [[False, False, True, False], [False, False, False, False], [False, True, False, False]] + ) + + actual = synthetic_eos_logprobs( + logits, + sampled_ids, + synthetic_eos_mask, + vocab_start_index=0, + vocab_end_index=5, + tp_group=object(), + inference_only=False, + ) + + expected = torch.zeros_like(actual) + expected[0, 2] = logits.detach()[0, 2, 2] + expected[2, 1] = logits.detach()[2, 1, 3] + torch.testing.assert_close(actual, expected) + assert calls == [torch.Size([1, 3, 5])] + + actual.sum().backward() + expected_grad = torch.zeros_like(logits) + expected_grad[0, 2, 2] = 1 + expected_grad[2, 1, 3] = 1 + torch.testing.assert_close(logits.grad, expected_grad) + + +def test_synthetic_eos_uses_packed_cp_trajectory_segments(monkeypatch): + calls = [] + _install_fake_distributed_logprob(monkeypatch, calls) + logits = torch.arange(4 * 5, dtype=torch.float64).reshape(1, 4, 5).requires_grad_(True) + sampled_ids = torch.tensor([[0, 1, 2, 3]]) + synthetic_eos_mask = torch.tensor([[False, True, False, True]]) + layout = TokenMetadataLayout( + attention_mask=torch.ones((2, 3), dtype=torch.bool), + sequence_lengths=[3, 3], + aligned_sequence_length=8, + padded_sequence_lengths=[4, 4], + cu_seqlens_padded=torch.tensor([0, 4, 8], dtype=torch.int32), + context_parallel_size=2, + context_parallel_rank=0, + ) + + actual = synthetic_eos_logprobs( + logits, + sampled_ids, + synthetic_eos_mask, + vocab_start_index=0, + vocab_end_index=5, + tp_group=object(), + inference_only=False, + metadata_layout=layout, + ) + + expected = torch.zeros_like(actual) + expected[0, 1] = logits.detach()[0, 1, 1] + expected[0, 3] = logits.detach()[0, 3, 3] + torch.testing.assert_close(actual, expected) + assert calls == [torch.Size([1, 2, 5])] + + +def test_synthetic_eos_fused_projection_keeps_capacity_and_chunk_bound(monkeypatch): + calls = [] + model_utils = _install_fake_distributed_logprob(monkeypatch, calls) + + def fused_apply(backend, hidden, weight, targets, start, end, chunk_size, group, inference_only): + calls.append((hidden.shape, chunk_size)) + return hidden[..., 0] + + model_utils._fused_lm_head_logprob_apply = fused_apply + hidden = torch.arange(3 * 4 * 2, dtype=torch.float64).reshape(3, 4, 2).requires_grad_(True) + sampled_ids = torch.zeros((3, 4), dtype=torch.long) + synthetic_eos_mask = torch.tensor( + [[False, False, False, False], [False, True, False, False], [False, False, False, False]] + ) + + actual = synthetic_eos_logprobs( + hidden, + sampled_ids, + synthetic_eos_mask, + vocab_start_index=0, + vocab_end_index=5, + tp_group=object(), + inference_only=False, + lm_head_weight=torch.ones((5, 2), dtype=torch.float64), + chunk_size=2, + ) + + expected = torch.zeros_like(actual) + expected[1, 1] = hidden.detach()[1, 1, 0] + torch.testing.assert_close(actual, expected) + assert calls == [(torch.Size([1, 3, 2]), 2)] + actual.sum().backward() + expected_grad = torch.zeros_like(hidden) + expected_grad[1, 1, 0] = 1 + torch.testing.assert_close(hidden.grad, expected_grad) diff --git a/tests/train/dataset/test_preprocess.py b/tests/train/dataset/test_preprocess.py index 1df1d002a1..a3e4a121d6 100644 --- a/tests/train/dataset/test_preprocess.py +++ b/tests/train/dataset/test_preprocess.py @@ -4,11 +4,13 @@ from unittest.mock import MagicMock +import numpy as np import pytest import torch from skyrl.train.dataset.preprocess import ( convert_prompts_responses_to_batch_tensors, + make_router_padding_mask, ) @@ -56,6 +58,119 @@ def fake_tokenizer_decode_list(ids, **kwargs): return mock_tokenizer +def test_router_padding_mask_marks_left_padding_and_uncaptured_suffix(): + attention_mask = torch.tensor([[0, 1, 1, 1], [1, 1, 1, 1]]) + + mask = make_router_padding_mask(attention_mask, [2, 4]) + + assert mask.tolist() == [[True, False, False, True], [False, False, False, False]] + + +def test_routed_expert_tensor_uses_unique_dummy_routes(tokenizer): + routes = [ + np.asarray( + [ + [[2, 3], [4, 5]], + [[6, 7], [0, 1]], + ], + dtype=np.uint8, + ), + np.asarray( + [ + [[1, 2], [3, 4]], + [[5, 6], [7, 0]], + [[2, 4], [6, 7]], + ], + dtype=np.uint8, + ), + ] + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10], [20]], + responses=[[11, 12], [21, 22]], + rewards=[[0.0, 0.0], [0.0, 0.0]], + loss_masks=[[1, 1], [1, 1]], + rollout_expert_indices=routes, + ) + + assert routed.shape == (2, 3, 2, 2) + assert routed.dtype == torch.uint8 + assert routed[0, 2].tolist() == [[0, 1], [0, 1]] + + +@pytest.mark.parametrize( + ("max_expert_id", "source_dtype", "expected_dtype"), + [(2**8, np.int16, torch.int16), (2**15, np.int32, torch.int32)], +) +def test_routed_expert_tensor_promotes_mixed_batch_dtype( + tokenizer, + max_expert_id, + source_dtype, + expected_dtype, +): + routes = [ + np.asarray([[[1, 2]]], dtype=np.uint8), + np.asarray([[[max_expert_id, max_expert_id + 1]]], dtype=source_dtype), + ] + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10], [20]], + responses=[[11], [21]], + rewards=[[0.0], [0.0]], + loss_masks=[[1], [1]], + rollout_expert_indices=routes, + ) + + assert routed.dtype == expected_dtype + assert routed[1, 0].tolist() == [[max_expert_id, max_expert_id + 1]] + + +def test_routed_expert_tensor_accepts_read_only_arrays(tokenizer): + routes = np.asarray([[[1, 2]], [[3, 4]]], dtype=np.uint8) + routes.flags.writeable = False + + *_, routed = convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[routes], + ) + + assert routed.dtype == torch.uint8 + assert routed.tolist() == [[[[1, 2]], [[3, 4]]]] + + +def test_routed_expert_tensor_rejects_nested_lists(tokenizer): + with pytest.raises(TypeError, match="NumPy arrays"): + convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[[[[1, 2]], [[3, 4]]]], + ) + + +@pytest.mark.parametrize("dtype", [np.uint16, np.int64]) +def test_routed_expert_tensor_rejects_unsupported_dtypes(tokenizer, dtype): + routes = np.asarray([[[1, 2]], [[3, 4]]], dtype=dtype) + + with pytest.raises(TypeError, match="Unsupported routed expert dtype"): + convert_prompts_responses_to_batch_tensors( + tokenizer, + prompts=[[10]], + responses=[[11]], + rewards=[[0.0]], + loss_masks=[[1]], + rollout_expert_indices=[routes], + ) + + def test_convert_prompts_responses_to_batch_tensors_exact(tokenizer): """ Test with inputs of exact lengths. @@ -290,8 +405,8 @@ def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): topk = 2 # rollout_expert_indices[i] has shape [prompt_len_i + response_len_i, num_layers, topk] # Sample 0: 5 tokens, sample 1: 6 tokens - rei_0 = [[[1, 2]] * num_layers for _ in range(5)] # 5 tokens - rei_1 = [[[3, 4]] * num_layers for _ in range(6)] # 6 tokens + rei_0 = np.asarray([[[1, 2]] * num_layers for _ in range(5)], dtype=np.uint8) # 5 tokens + rei_1 = np.asarray([[[3, 4]] * num_layers for _ in range(6)], dtype=np.uint8) # 6 tokens seq, attn, action, rew, lm, lp, rei_tensor = convert_prompts_responses_to_batch_tensors( tokenizer, @@ -306,20 +421,21 @@ def test_rollout_expert_indices_shape_padding_and_alignment(tokenizer): # Shape: [batch=2, max_total=6, layers=2, topk=2] assert rei_tensor.shape == (2, 6, num_layers, topk) - # Sample 0 has total=5, so 1 left-pad position → first position should be zeros - assert rei_tensor[0, 0].tolist() == [[0, 0]] * num_layers # padding + dummy_routes = [[0, 1]] * num_layers + # Sample 0 has total=5, so the first position uses unique dummy routes. + assert rei_tensor[0, 0].tolist() == dummy_routes assert rei_tensor[0, 1].tolist() == [[1, 2]] * num_layers # first real token # Sample 1 has total=6, no padding assert rei_tensor[1, 0].tolist() == [[3, 4]] * num_layers # first real token - # Non-zero positions in rei_tensor align exactly with attention_mask==1 + # Dummy positions in rei_tensor align exactly with attention_mask==0. for i in range(2): for pos in range(6): if attn[i, pos] == 0: - assert rei_tensor[i, pos].tolist() == [[0, 0]] * num_layers + assert rei_tensor[i, pos].tolist() == dummy_routes else: - assert rei_tensor[i, pos].tolist() != [[0, 0]] * num_layers + assert rei_tensor[i, pos].tolist() != dummy_routes def test_rollout_expert_indices_none_when_not_provided(tokenizer): diff --git a/tests/train/dataset/test_sample_support_preprocess.py b/tests/train/dataset/test_sample_support_preprocess.py new file mode 100644 index 0000000000..7a97a9a0a2 --- /dev/null +++ b/tests/train/dataset/test_sample_support_preprocess.py @@ -0,0 +1,69 @@ +import pytest +import torch + +from skyrl.train.dataset.preprocess import build_dense_sample_support + + +def test_sample_support_is_response_aligned_int32(): + support = build_dense_sample_support( + [[[10, 11, -1], [12, -1, -1]], [[20, 21, 22]]], + [[10, 12], [20]], + [[1, 1], [1]], + sequence_length=5, + top_k=3, + eos_token_id=2, + ) + + assert support is not None + assert support.dtype == torch.int32 + assert support[0].tolist() == [[-1, -1, -1]] * 3 + [[10, 11, -1], [12, -1, -1]] + assert support[1].tolist() == [[-1, -1, -1]] * 4 + [[20, 21, 22]] + + +def test_empty_loss_masked_rows_and_loss_bearing_synthetic_eos_are_preserved(): + support = build_dense_sample_support( + [[[7, 8], [], []]], + [[7, 9, 2]], + [[1, 0, 1]], + sequence_length=4, + top_k=2, + eos_token_id=2, + ) + + assert support.tolist() == [[[-1, -1], [7, 8], [-1, -1], [-1, -1]]] + + +def test_empty_loss_bearing_non_eos_is_rejected(): + with pytest.raises(ValueError, match="loss-bearing non-EOS"): + build_dense_sample_support( + [[[]]], + [[3]], + [[1]], + sequence_length=2, + top_k=2, + eos_token_id=2, + ) + + +def test_multiple_loss_bearing_unsupported_eos_are_rejected(): + with pytest.raises(ValueError, match="more than one loss-bearing unsupported token"): + build_dense_sample_support( + [[[], []]], + [[2, 2]], + [[1, 1]], + sequence_length=2, + top_k=2, + eos_token_id=2, + ) + + +def test_loss_bearing_sampled_token_must_be_in_support(): + with pytest.raises(ValueError, match="sampled token 3 is missing"): + build_dense_sample_support( + [[[2, 4]]], + [[3]], + [[1]], + sequence_length=2, + top_k=2, + eos_token_id=2, + ) 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_generator_output_utils.py b/tests/train/generators/test_generator_output_utils.py index 7546858bb8..2d500d10ac 100644 --- a/tests/train/generators/test_generator_output_utils.py +++ b/tests/train/generators/test_generator_output_utils.py @@ -29,6 +29,7 @@ def test_generator_output_concatenation(): "rollout_metrics", "rollout_logprobs", "rollout_expert_indices", + "rollout_sample_support", # optional but present in the signature "trajectory_ids", "trajectory_generation_times", @@ -169,6 +170,7 @@ def test_case1_response_only_assistant(self): "rollout_logprobs": [[-0.5], [-0.3, -0.4]], "trajectory_ids": [tid, tid], "rollout_expert_indices": None, + "rollout_sample_support": [[[20, 21]], [[40, 42], [41, 43]]], "is_last_step": [False, True], } @@ -182,6 +184,7 @@ def test_case1_response_only_assistant(self): assert merged["loss_masks"] == [[1, 0, 1, 1]] # logprobs: A1=-0.5, O2=0.0, A2_tok1=-0.3, A2_tok2=-0.4 assert merged["rollout_logprobs"] == [[-0.5, 0.0, -0.3, -0.4]] + assert merged["rollout_sample_support"] == [[[20, 21], [], [40, 42], [41, 43]]] # rewards: A1=1.0, O2=0.0, A2_tok1=0.0, A2_tok2=5.0 assert merged["rewards"] == [[1.0, 0.0, 0.0, 5.0]] assert merged["stop_reasons"] == ["eos"] @@ -653,7 +656,7 @@ def test_asserts_no_expert_indices(self): "rollout_metrics": None, "rollout_logprobs": None, "trajectory_ids": [tid], - "rollout_expert_indices": [[[[1, 2]]]], + "rollout_expert_indices": [np.asarray([[[1, 2]]], dtype=np.uint8)], "is_last_step": [True], } with pytest.raises(AssertionError, match="rollout_expert_indices not supported"): diff --git a/tests/train/generators/test_skyrl_gym_generator.py b/tests/train/generators/test_skyrl_gym_generator.py index c29afbe920..1b1ac59ad3 100644 --- a/tests/train/generators/test_skyrl_gym_generator.py +++ b/tests/train/generators/test_skyrl_gym_generator.py @@ -5,6 +5,7 @@ from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock, patch +import numpy as np import pytest from skyrl.train.config import ChatTemplateConfig, GeneratorConfig @@ -14,7 +15,7 @@ GeneratorInput, GeneratorOutput, ) -from skyrl.train.generators.skyrl_gym_generator import SkyRLGymGenerator +from skyrl.train.generators.skyrl_gym_generator import SkyRLGymGenerator, TurnOutput from skyrl_gym.envs.base_text_env import BaseTextEnv, BaseTextEnvStepOutput # Mock constants, where 4 is the eos token id @@ -22,6 +23,25 @@ MOCK_TOKENIZER_ENCODED_IDS = [1, 2, 3, 4] +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], + reward=1.0, + rollout_sample_support=np.array([[10, 100], [11, 101]], dtype=np.int32), + added_eos=True, + ) + + np.testing.assert_array_equal( + output.get_turn_rollout_sample_support(), + np.array([[10, 100], [11, 101], [-1, -1], [-1, -1], [-1, -1]], dtype=np.int32), + ) + assert output.get_turn_loss_mask() == [1, 1, 0, 0, 0] + + # TODO (erictang000): clean up the mocking for tests in this file @pytest.fixture def mock_tokenizer(): @@ -376,7 +396,7 @@ def mock_generate(_, model=None): # No EOS: just add it expected_response_ids = mock_llm_output_ids + [mock_tokenizer.eos_token_id] - expected_loss_mask = [1] * (len(expected_response_ids)) + expected_loss_mask = [1] * len(expected_response_ids) if logprobs_setting is not None: assert output.rollout_logprobs is not None @@ -395,6 +415,73 @@ def mock_generate(_, model=None): assert output.stop_reason == "stop" +@pytest.mark.asyncio +@patch("skyrl_gym.make") +async def test_agent_loop_uses_incremental_replay_metadata_traces( + 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 + generator_cfg.inference_engine.enable_return_sample_support_set = True + generator_cfg.sampling_params.top_k = 2 + 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 = [] + generation_index = 0 + + def generate(input_batch, model=None): + nonlocal generation_index + 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 + sample_support = [[10, 100 + generation_index], [11, 110 + generation_index]] + generation_index += 1 + return { + "responses": ["mocked output"], + "response_ids": [output_ids], + "stop_reasons": ["stop"], + "rollout_expert_indices": [routes], + "rollout_sample_support": [sample_support], + } + + 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 = [] + + output = 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] + assert output.rollout_sample_support[:2] == [[10, 100], [11, 110]] + assert output.rollout_sample_support[-2:] == [[10, 101], [11, 111]] + assert all(row == [-1, -1] for row in output.rollout_sample_support[2:-2]) + + @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 6b9c58552c..4c9e0c2b05 100644 --- a/tests/train/test_config.py +++ b/tests/train/test_config.py @@ -128,6 +128,55 @@ def test_cli_overrides_empty_args(): assert cfg.trainer.seed == 42 +def test_sample_support_replay_requires_capture_and_megatron(): + with pytest.raises(ValueError, match="enable_sample_support_replay requires"): + SkyRLTrainConfig.from_cli_overrides(["trainer.algorithm.enable_sample_support_replay=true"]) + + with pytest.raises(ValueError, match="requires trainer.strategy=megatron"): + SkyRLTrainConfig.from_cli_overrides( + [ + "trainer.algorithm.enable_sample_support_replay=true", + "generator.inference_engine.enable_return_sample_support_set=true", + "generator.sampling_params.top_k=8", + ] + ) + + +@pytest.mark.parametrize( + ("override", "message"), + [ + ("generator.sampling_params.temperature=0", "temperature > 0"), + ("generator.sampling_params.top_k=1", "top_k > 1"), + ("generator.sampling_params.repetition_penalty=1.1", "repetition_penalty=1.0"), + ("generator.sampling_params.additional_kwargs.foo=bar", "additional_kwargs"), + ], +) +def test_sample_support_capture_rejects_unsupported_sampling_modifiers(override, message): + with pytest.raises(ValueError, match=message): + SkyRLTrainConfig.from_cli_overrides( + [ + "generator.inference_engine.enable_return_sample_support_set=true", + "generator.sampling_params.top_k=8", + override, + ] + ) + + +def test_sample_support_replay_accepts_top_k_top_p_and_min_p(): + cfg = SkyRLTrainConfig.from_cli_overrides( + [ + "trainer.strategy=megatron", + "trainer.algorithm.enable_sample_support_replay=true", + "generator.inference_engine.enable_return_sample_support_set=true", + "generator.sampling_params.top_k=8", + "generator.sampling_params.top_p=0.9", + "generator.sampling_params.min_p=0.05", + ] + ) + + assert cfg.trainer.algorithm.enable_sample_support_replay + + def test_cli_overrides_plus_prefix_rejected(): with pytest.raises(ValueError, match="The '\\+' prefix"): SkyRLTrainConfig.from_cli_overrides(["+new_field=value"]) diff --git a/tests/train/test_trainer.py b/tests/train/test_trainer.py index 800fc40767..b59208c6d5 100644 --- a/tests/train/test_trainer.py +++ b/tests/train/test_trainer.py @@ -114,6 +114,37 @@ def _get_test_data(trainer: RayPPOTrainer): return data +def test_fwd_logprobs_preserves_sample_support_and_loss_mask(dummy_config): + dummy_config.trainer.critic.model.path = None + trainer = object.__new__(RayPPOTrainer) + trainer.cfg = dummy_config + trainer.ref_model = None + trainer.dispatch = MagicMock() + trainer.all_metrics = {} + trainer._skip_policy_forward = MagicMock(return_value=False) + seen = {} + + def execute_forward_pass(model, batch, **kwargs): + seen.update(batch) + return torch.zeros((2, 3)) + + trainer._execute_forward_pass = execute_forward_pass + batch = TrainingInputBatch( + { + "sequences": torch.ones((2, 4), dtype=torch.long), + "attention_mask": torch.ones((2, 4), dtype=torch.long), + "loss_mask": torch.ones((2, 3)), + "sample_support_ids": torch.tensor([[[1, -1]] * 4, [[2, -1]] * 4], dtype=torch.int32), + } + ) + batch.metadata = {"response_length": 3} + + trainer.fwd_logprobs_values_reward(batch) + + assert torch.equal(seen["sample_support_ids"], batch["sample_support_ids"]) + assert torch.equal(seen["loss_mask"], batch["loss_mask"]) + + def test_calculate_kl_create_experience_batched(dummy_config): trainer = RayPPOTrainer( cfg=dummy_config, diff --git a/tests/train/test_trainer_utils.py b/tests/train/test_trainer_utils.py index 6a70e8ca66..2c80650574 100644 --- a/tests/train/test_trainer_utils.py +++ b/tests/train/test_trainer_utils.py @@ -10,6 +10,7 @@ from typing import Union from unittest.mock import Mock, mock_open, patch +import numpy as np import pytest import ray @@ -393,6 +394,8 @@ def test_handle_replace_sampling_sufficient_good_samples(): "stop_reasons": ["length"] * 6, "rollout_metrics": None, "rollout_logprobs": [[0.1, 0.2], [0.3, 0.4], [0.5, 0.25], [0.15, 0.25], [0.1, 0.2], [0.3, 0.4]], + "rollout_expert_indices": [np.asarray([[[i, i + 1]]], dtype=np.uint8) for i in range(6)], + "rollout_sample_support": [[[i, i + 10], [i + 1, i + 11]] for i in range(6)], } uids = ["uid1", "uid1", "uid2", "uid2", "uid3", "uid3"] # 2 samples per prompt sampling_config = {"n_samples_per_prompt": 2, "min_replace_ratio": 0.3} @@ -408,6 +411,17 @@ def test_handle_replace_sampling_sufficient_good_samples(): assert len(result_output["rewards"]) == 6 assert len(result_output["rollout_logprobs"]) == 6 assert len(result_uids) == 6 + route_by_response = { + tuple(response): routes + for response, routes in zip(generator_output["response_ids"], generator_output["rollout_expert_indices"]) + } + for response, routes in zip(result_output["response_ids"], result_output["rollout_expert_indices"]): + assert np.array_equal(routes, route_by_response[tuple(response)]) + support_by_response = dict( + zip(map(tuple, generator_output["response_ids"]), generator_output["rollout_sample_support"]) + ) + for response, support in zip(result_output["response_ids"], result_output["rollout_sample_support"]): + assert support == support_by_response[tuple(response)] # Check that bad uid2 samples were replaced with good samples uid2_indices = [i for i, uid in enumerate(result_uids) if uid == "uid2"] @@ -636,6 +650,7 @@ def test_handle_filter_sampling_single_sample_per_prompt(): def test_filter_generator_output(): """Test the filter_generator_output utility function.""" + routes = [np.asarray([[[i, i + 1]]], dtype=np.uint8) for i in range(3)] generator_output = { "prompt_token_ids": [[1, 2], [3, 4], [5, 6]], "response_ids": [[7, 8], [9, 10], [11, 12]], @@ -644,6 +659,8 @@ def test_filter_generator_output(): "stop_reasons": ["length", "length", "stop"], "rollout_metrics": {"metric": "value"}, "rollout_logprobs": [[0.16, 0.4], [0.1, 0.2], [0.3, 0.4]], + "rollout_expert_indices": routes, + "rollout_sample_support": [[[7, 70], [8, 80]], [[9, 90], [10, 100]], [[11, 110], [12, 120]]], } kept_indices = [0, 2] # Keep first and third samples @@ -656,6 +673,9 @@ def test_filter_generator_output(): assert filtered["stop_reasons"] == ["length", "stop"] assert filtered["rollout_metrics"] == {"metric": "value"} assert filtered["rollout_logprobs"] == [[0.16, 0.4], [0.3, 0.4]] + assert filtered["rollout_expert_indices"][0] is routes[0] + assert filtered["rollout_expert_indices"][1] is routes[2] + assert filtered["rollout_sample_support"] == [[[7, 70], [8, 80]], [[11, 110], [12, 120]]] def test_zero_variance_filter_mixed_groups(): diff --git a/tests/utils/test_token_metadata.py b/tests/utils/test_token_metadata.py new file mode 100644 index 0000000000..fd58df384b --- /dev/null +++ b/tests/utils/test_token_metadata.py @@ -0,0 +1,151 @@ +import sys +import types + +import numpy as np +import pytest +import torch + +from skyrl.utils import token_metadata +from skyrl.utils.routed_experts import RoutedExpertTrace +from skyrl.utils.token_metadata import TokenMetadataTrace + + +@pytest.fixture +def parallel_state(monkeypatch): + try: + import megatron.core.parallel_state as mpu + except ModuleNotFoundError: + megatron = types.ModuleType("megatron") + core = types.ModuleType("megatron.core") + mpu = types.ModuleType("megatron.core.parallel_state") + megatron.core = core + core.parallel_state = mpu + monkeypatch.setitem(sys.modules, "megatron", megatron) + monkeypatch.setitem(sys.modules, "megatron.core", core) + monkeypatch.setitem(sys.modules, "megatron.core.parallel_state", mpu) + + monkeypatch.setattr(mpu, "get_tensor_model_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_world_size", lambda: 1, raising=False) + monkeypatch.setattr(mpu, "get_context_parallel_rank", lambda: 0, raising=False) + return mpu + + +def test_microbatch_rows_share_one_packed_layout(monkeypatch, parallel_state): + monkeypatch.setattr(token_metadata, "get_packed_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1], [0, 0, 1, 1]]) + routes = torch.tensor( + [ + [[[0, 1]], [[10, 11]], [[12, 13]], [[14, 15]]], + [[[0, 1]], [[0, 1]], [[20, 21]], [[22, 23]]], + ], + dtype=torch.int16, + ) + router_mask = torch.tensor([[1, 0, 0, 1], [1, 1, 0, 0]], dtype=torch.bool) + + layout = token_metadata.build_token_metadata_layout( + attention_mask, + routes.device, + packed=True, + fp8_enabled=False, + ) + packed_routes = token_metadata.align_token_metadata( + routes, + layout, + torch.tensor([0, 1], dtype=routes.dtype), + ) + packed_mask = token_metadata.align_token_metadata(router_mask, layout, True) + + assert packed_routes[0, :, 0].tolist() == [ + [10, 11], + [12, 13], + [14, 15], + [0, 1], + [20, 21], + [22, 23], + [0, 1], + [0, 1], + ] + assert packed_mask.tolist() == [[False, False, True, True, False, False, True, True]] + assert layout.cu_seqlens_padded.tolist() == [0, 4, 8] + + +def test_packed_layout_aligns_next_token_metadata_and_scatters_rows(monkeypatch, parallel_state): + monkeypatch.setattr(token_metadata, "get_packed_seq_align_size", lambda *args, **kwargs: 4) + attention_mask = torch.tensor([[0, 1, 1, 1], [0, 0, 1, 1]]) + metadata = torch.tensor([[0, 10, 11, 12], [0, 0, 20, 21]], dtype=torch.int32) + layout = token_metadata.build_token_metadata_layout( + attention_mask, + metadata.device, + packed=True, + fp8_enabled=False, + ) + + aligned = token_metadata.align_token_metadata(metadata, layout, -1, next_token=True) + batch_values = token_metadata.scatter_packed_token_values_to_batch( + torch.arange(1, 9, dtype=torch.float32).unsqueeze(0), + layout, + 0, + ) + + assert aligned.tolist() == [[11, 12, -1, -1, 21, -1, -1, -1]] + assert batch_values.tolist() == [[0.0, 1.0, 2.0], [0.0, 0.0, 5.0]] + + +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]]) diff --git a/uv.lock b/uv.lock index 8fba78c710..01c249a8fb 100644 --- a/uv.lock +++ b/uv.lock @@ -8380,8 +8380,10 @@ fsdp = [ { name = "ninja" }, { name = "nixl", marker = "sys_platform == 'linux' or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-miniswe') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-tpu') or (extra == 'extra-5-skyrl-miniswe' and extra == 'extra-5-skyrl-tpu')" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8440,8 +8442,10 @@ megatron = [ { name = "nixl", marker = "sys_platform == 'linux'" }, { name = "nvidia-modelopt", marker = "sys_platform == 'linux'" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8481,8 +8485,10 @@ miniswe = [ { name = "ninja" }, { name = "nixl", marker = "sys_platform == 'linux' or (extra == 'extra-5-skyrl-fsdp' and extra == 'extra-5-skyrl-jax')" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8516,8 +8522,10 @@ skyrl-train = [ { name = "ninja" }, { name = "nixl", marker = "sys_platform == 'linux' or (extra == 'extra-5-skyrl-fsdp' and extra == 'extra-5-skyrl-jax') or (extra == 'extra-5-skyrl-fsdp' and extra == 'extra-5-skyrl-megatron') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-megatron') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-miniswe') or (extra == 'extra-5-skyrl-gpu' and extra == 'extra-5-skyrl-tpu') or (extra == 'extra-5-skyrl-jax' and extra == 'extra-5-skyrl-megatron') or (extra == 'extra-5-skyrl-megatron' and extra == 'extra-5-skyrl-miniswe') or (extra == 'extra-5-skyrl-megatron' and extra == 'extra-5-skyrl-tpu') or (extra == 'extra-5-skyrl-miniswe' and extra == 'extra-5-skyrl-tpu')" }, { name = "omegaconf" }, + { name = "orjson" }, { name = "peft" }, { name = "polars" }, + { name = "pybase64" }, { name = "pybind11" }, { name = "ray" }, { name = "s3fs" }, @@ -8609,12 +8617,14 @@ requires-dist = [ { name = "nvidia-modelopt", marker = "sys_platform == 'linux' and extra == 'megatron'" }, { name = "omegaconf", marker = "extra == 'skyrl-train'" }, { name = "optax", marker = "extra == 'jax'", specifier = ">=0.2.5" }, + { name = "orjson", marker = "extra == 'skyrl-train'", specifier = ">=3.11.9" }, { name = "peft", specifier = "==0.18.1" }, { name = "peft", marker = "extra == 'skyrl-train'", specifier = "==0.18.1" }, { name = "pillow", specifier = ">=11.3.0" }, { name = "polars", marker = "extra == 'skyrl-train'" }, { name = "pre-commit", marker = "extra == 'dev'" }, { name = "psycopg2-binary", marker = "extra == 'tinker'" }, + { name = "pybase64", marker = "extra == 'skyrl-train'", specifier = ">=1.4.2" }, { name = "pybind11", marker = "extra == 'skyrl-train'" }, { name = "pymdown-extensions", marker = "extra == 'dev'", specifier = ">=10.7" }, { name = "pytest", marker = "extra == 'dev'" },