Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
ee54593
refactor(inference): extract reusable RemoteGenerateClient
dyurk-lila Jul 16, 2026
d881ecd
refactor(generators): assemble routed-expert traces incrementally
dyurk-lila Jul 16, 2026
707724c
test(generators): drop obsolete turn route argument
dyurk-lila Jul 20, 2026
4bb8674
Update skyrl/backends/skyrl_train/utils/routed_experts.py
erictang000 Aug 5, 2026
31d6b3c
perf(r3): collate routes through a container-aware pool, and refuse r…
Aug 14, 2026
e080afe
style: tighten route collation comments and tests
dyurk-lila Aug 22, 2026
5dccdc0
refactor(wire): generic packed-ndarray codec and an N-field response-…
Aug 14, 2026
a604e5c
style: tighten wire codec comments and tests
dyurk-lila Aug 22, 2026
58ed80b
perf(r3): route rollout experts to the trainer packed with cu_seqlens
Aug 14, 2026
c1e71a0
style: tighten packed-route comments and tests
dyurk-lila Aug 22, 2026
646f87f
perf(batch): ship large per-token side channels to the trainer zero-copy
Aug 14, 2026
bafe553
style: tighten zero-copy transport comments and tests
dyurk-lila Aug 22, 2026
0caad2f
style: apply Black formatting
dyurk-lila Aug 22, 2026
1c3510a
feat(sample-support): capture the sampler's bounded support and trace…
Aug 14, 2026
9a8327d
style: tighten sample-support capture comments and tests
dyurk-lila Aug 22, 2026
098017b
perf(sample-support): carry the captured support to the trainer as on…
Aug 14, 2026
ec57e6f
style: tighten packed sample-support comments and tests
dyurk-lila Aug 22, 2026
b366738
feat(sample-support): score policy logprobs over the recorded support…
Aug 14, 2026
6be01c8
style: tighten Megatron sample-support comments and tests
dyurk-lila Aug 22, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/content/docs/algorithms/off_policy_correction.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -155,6 +155,7 @@ trainer:
```

To enable rollout router replay, set `generator.inference_engine.enable_return_routed_experts=True`, `trainer.policy.megatron_config.moe_enable_routing_replay=True`, and use the `mp` distributed_executor_backend for vLLM. Note that
R3 requires per-token routes that line up with the trained tokens, so `SkyRLGymGenerator` refuses it together with `generator.step_wise_trajectories`, `generator.use_conversation_multi_turn=False`, a custom `generator.chat_template`, and `generator.vision_language_generator` (each breaks that alignment; see the [step-wise training](../tutorials/step-wise-training) page for the step-wise reason). Routes are captured for train batches only. Note that

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

This sentence is quite long and lists several distinct points. For better readability, consider breaking it down into smaller sentences and using a list. For example:

R3 requires per-token routes that line up with the trained tokens. This alignment is broken by certain generator configurations. Therefore, `SkyRLGymGenerator` refuses to enable R3 when used with any of the following:

- `generator.step_wise_trajectories` (see the [step-wise training](../tutorials/step-wise-training) page for details)
- `generator.use_conversation_multi_turn=False`
- a custom `generator.chat_template`
- `generator.vision_language_generator`

R3 does induce additional training bias when mini-batching, since routing decisions are fixed for all mini-batches in a training batch. However, it has been shown to be important for stabilizing large-scale MoE training, particularly
in models adopting a DeepSeek-V3 like architecture (notably the GLM family) due to the use of sigmoid-based affinity scoring instead of softmax for top-k routing.

Expand Down
7 changes: 6 additions & 1 deletion docs/content/docs/tutorials/step-wise-training.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,10 @@ class GeneratorOutput(TypedDict):
rollout_metrics: Optional[Dict[str, Any]]
rollout_logprobs: Optional[List[List[float]]]
trajectory_ids: Optional[List[TrajectoryID]]
rollout_expert_indices: Optional[List[List[List[List[int]]]]]
trajectory_generation_times: Optional[List[float]]
trajectory_time_splits: Optional[Dict[str, List[float]]]
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]]
```
Expand All @@ -85,6 +88,8 @@ When `step_wise_trajectories=True`, some related fields:
| `is_last_step` | `List[bool]` | Marks the final step of each trajectory. Must have at least one `True`, and the last element must be `True`. |
| `trajectory_ids` | `List[TrajectoryID]` | Associates each step-sample with its parent trajectory. All steps of the same trajectory share the same `TrajectoryID`. |
| `rollout_logprobs` | `List[List[float]]` | Per-token logprobs from the inference engine, aligned with `response_ids`. Required for TIS. |
| `rollout_sample_support` | `List[List[List[int]]]` | Per generated token, the bounded top-k vocab IDs the sampler drew from, right-padded with `SAMPLE_SUPPORT_PADDING` (`-1`). One dense `[tokens, top_k]` block per row, aligned with `response_ids`. Every position must be present: an observation token or a synthetic (loop-appended) EOS carries an all-padding row rather than being absent, so the block stays rectangular and aligned. Enabled with `generator.inference_engine.enable_return_sample_support_set`, and requested per request (train batches only — eval's greedy `top_k=-1` params cannot satisfy the capture contract). |
| `rollout_expert_indices` | `Optional[List[RoutedExpertIndices]]` | Per token, the `[layers, topk]` MoE routes vLLM recorded, as one `[tokens, layers, topk]` integer array per trajectory (rollout router replay, R3). **Refused with `step_wise_trajectories=True`**: each step-wise row's prompt is the whole history so far while routes are recorded for that step's generated tokens only, so a step's routes would replay onto the first N prompt tokens of its row with no length mismatch to assert on, silently training against routing that does not match the rollout. |

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The explanation for why R3 is refused with step_wise_trajectories=True is a bit dense. For improved clarity, you could rephrase it. For example:

Refused with step_wise_trajectories=True: This is because in step-wise training, each row's prompt consists of the entire history, but routes are only recorded for the tokens generated in that step. This would cause a mismatch where the step's routes are replayed onto the initial prompt tokens of the training row. Without a length mismatch to trigger an error, this would lead to silently training against incorrect routing information.


### Concrete Example

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -471,7 +471,8 @@ def efficient_entropy_triton_kernel_epilogue(
global_logprobs_scalar = tl.sum(global_logprobs, axis=0)
tl.atomic_add(global_logprobs_scalar_ptr, global_logprobs_scalar)
elif reduction == 2:
global_logprobs_scalar = tl.sum(global_logprobs, axis=0) / num_tokens.to(tl.float32)
# Multiplication promotes both Triton tensors and specialized Python integers.
global_logprobs_scalar = tl.sum(global_logprobs, axis=0) / (num_tokens * 1.0)
tl.atomic_add(global_logprobs_scalar_ptr, global_logprobs_scalar)


Expand Down Expand Up @@ -593,7 +594,7 @@ def efficient_entropy_triton_epilogue_tp_update(
logprobs_scalar = tl.sum(logprobs, axis=0)
tl.atomic_add(logprobs_scalar_ptr, logprobs_scalar)
elif reduction == 2:
logprobs_scalar = tl.sum(logprobs, axis=0) / num_tokens.to(tl.float32)
logprobs_scalar = tl.sum(logprobs, axis=0) / (num_tokens * 1.0)
tl.atomic_add(logprobs_scalar_ptr, logprobs_scalar)


Expand Down Expand Up @@ -900,7 +901,7 @@ def efficient_entropy_backward_kernel_general_mainloop_MN(
d_logprobs = tl.load(d_logprobs_ptr)
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
else: # mean
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), (num_tokens * 1.0))
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
d_logprobs = -1 * d_logprobs

Expand Down Expand Up @@ -1061,7 +1062,7 @@ def efficient_entropy_backward_kernel_d_hidden(
d_logprobs = tl.load(d_logprobs_ptr)
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
else:
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), (num_tokens * 1.0))
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
d_logprobs = -1 * d_logprobs

Expand Down Expand Up @@ -1190,7 +1191,7 @@ def efficient_entropy_backward_kernel_d_weight(
d_logprobs = tl.load(d_logprobs_ptr)
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
else:
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), (num_tokens * 1.0))
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
d_logprobs = -1 * d_logprobs

Expand Down Expand Up @@ -1323,7 +1324,7 @@ def efficient_entropy_backward_kernel_general_d_logits(
d_logprobs = tl.load(d_logprobs_ptr)
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
else: # mean
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), (num_tokens * 1.0))
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
d_logprobs = -1 * d_logprobs

Expand Down Expand Up @@ -1473,7 +1474,7 @@ def efficient_entropy_backward_kernel_general_d_logits_split_N(
d_logprobs = tl.load(d_logprobs_ptr)
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
else:
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))
d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), (num_tokens * 1.0))
d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))
d_logprobs = -1 * d_logprobs
entropy_b = tl.load(entropy_b_ptr + offs_am * stride_entropy_b, mask=offs_am < num_tokens, other=0.0)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,8 @@
import torch
import torch.distributed as dist

from skyrl.backends.skyrl_train.utils.packed_tensor import lengths_from_offsets


@torch.no_grad()
def _compute_distributed_log_softmax(
Expand Down Expand Up @@ -924,7 +926,7 @@ def _packed_sequence_indices(
token_indices = torch.arange(total_tokens, device=device)
seq_indices = torch.searchsorted(cu_seqlens_padded[1:], token_indices, right=True)
seq_offsets = token_indices - cu_seqlens_padded[seq_indices]
seq_lens_padded = cu_seqlens_padded[1:] - cu_seqlens_padded[:-1]
seq_lens_padded = lengths_from_offsets(cu_seqlens_padded)
return cu_seqlens_padded, token_indices, seq_indices, seq_offsets, seq_lens_padded


Expand Down
158 changes: 150 additions & 8 deletions skyrl/backends/skyrl_train/distributed/megatron/token_metadata.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,16 @@
"""Token-aligned metadata layout transforms shared by training features."""

from collections.abc import Callable, Sequence
from dataclasses import dataclass

import numpy as np
import torch

from skyrl.backends.skyrl_train.distributed.megatron.packing_utils import (
get_packed_seq_align_size,
get_unpacked_seq_align_size,
)
from skyrl.backends.skyrl_train.utils.packed_tensor import PackedTensor

# Megatron is imported lazily inside functions so that non-Megatron backends can
# import this module for its layout dataclass and padding transforms.
Expand Down Expand Up @@ -99,34 +102,113 @@ def align_token_metadata(
f"attention_mask shape {layout.attention_mask.shape}"
)

return _align_token_rows(
lambda row_index: metadata[row_index, layout.attention_mask[row_index]],
metadata,
metadata.shape[2:],
layout,
padding_value,
next_token=next_token,
)


def align_packed_token_metadata(
metadata: PackedTensor,
layout: TokenMetadataLayout,
padding_value: torch.Tensor | bool | int,
*,
next_token: bool = False,
segment_starts: Sequence[int] | None = None,
) -> torch.Tensor:
"""Align metadata that already arrives packed as ``[sum(seqlen), *row_shape]``.

This relies on left padding, which makes each trajectory's real tokens contiguous.
Without ``segment_starts``, segments must match ``layout.sequence_lengths``;
otherwise each segment is placed at its specified real-token offset.
"""
if metadata.device != layout.attention_mask.device:
raise ValueError("Token-aligned metadata and attention_mask must be on the same device")
if len(metadata) != len(layout.sequence_lengths):
raise ValueError(
f"Packed metadata holds {len(metadata)} segments for {len(layout.sequence_lengths)} trajectories"
)
segment_lengths = metadata.sequence_lengths.tolist()
if segment_starts is None:
if segment_lengths != list(layout.sequence_lengths):
raise ValueError(
f"Packed metadata segments {segment_lengths} do not match "
f"trajectory lengths {list(layout.sequence_lengths)}"
)
else:
if len(segment_starts) != len(metadata):
raise ValueError(f"Got {len(segment_starts)} segment starts for {len(metadata)} segments")
for row_index, (start, length) in enumerate(zip(segment_starts, segment_lengths, strict=True)):
if start < 0 or start + length > layout.sequence_lengths[row_index]:
raise ValueError(
f"Segment {row_index} spans real tokens [{start}, {start + length}) of a "
f"{layout.sequence_lengths[row_index]}-token trajectory"
)

return _align_token_rows(
metadata.segment,
metadata.values,
metadata.row_shape,
layout,
padding_value,
next_token=next_token,
segment_starts=segment_starts,
)


def _align_token_rows(
rows_for: Callable[[int], torch.Tensor],
source: torch.Tensor,
row_shape: tuple[int, ...] | torch.Size,
layout: TokenMetadataLayout,
padding_value: torch.Tensor | bool | int,
*,
next_token: bool = False,
segment_starts: Sequence[int] | None = None,
) -> torch.Tensor:
"""Place each trajectory's real-token rows into Megatron's layout and CP-shard them.

``rows_for(row_index)`` yields one trajectory's rows. They land at the front of its
padded region unless ``segment_starts`` names a per-trajectory destination offset.
"""
if layout.padded_sequence_lengths is None:
if next_token:
raise ValueError("next-token metadata alignment is only used for packed sequences")
aligned = _new_metadata_tensor(
metadata,
(metadata.shape[0], layout.aligned_sequence_length, *metadata.shape[2:]),
source,
(len(layout.sequence_lengths), layout.aligned_sequence_length, *row_shape),
padding_value,
)
for row_index, sequence_length in enumerate(layout.sequence_lengths):
aligned[row_index, :sequence_length] = metadata[row_index, layout.attention_mask[row_index]]
rows = rows_for(row_index)
start = 0 if segment_starts is None else segment_starts[row_index]
end = sequence_length if segment_starts is None else start + rows.shape[0]
aligned[row_index, start:end] = rows
return aligned

packed = _new_metadata_tensor(
metadata,
(layout.aligned_sequence_length, *metadata.shape[2:]),
source,
(layout.aligned_sequence_length, *row_shape),
padding_value,
)
offset = 0
for row_index, (sequence_length, padded_length) in enumerate(
zip(layout.sequence_lengths, layout.padded_sequence_lengths, strict=True)
):
packed[offset : offset + sequence_length] = metadata[row_index, layout.attention_mask[row_index]]
rows = rows_for(row_index)
start = offset if segment_starts is None else offset + segment_starts[row_index]
end = offset + sequence_length if segment_starts is None else start + rows.shape[0]
packed[start:end] = rows
# Match Megatron's [seq0, pad0, seq1, pad1, ...] microbatch layout.
offset += padded_length

if next_token:
# Each packed logit predicts the next token within its own padded sequence.
shifted = _new_metadata_tensor(metadata, packed.shape, padding_value)
shifted = _new_metadata_tensor(source, packed.shape, padding_value)
offset = 0
for padded_length in layout.padded_sequence_lengths:
shifted[offset : offset + padded_length - 1] = packed[offset + 1 : offset + padded_length]
Expand All @@ -135,7 +217,7 @@ def align_token_metadata(

if layout.context_parallel_size > 1:
out = _new_metadata_tensor(
metadata,
source,
(packed.shape[0] // layout.context_parallel_size, *packed.shape[1:]),
padding_value,
)
Expand Down Expand Up @@ -205,3 +287,63 @@ def scatter_packed_token_values_to_batch(
)
batch_values[output_mask] = values[packed_mask]
return batch_values


class TokenMetadataTrace:
"""Accumulate arrays whose first dimension is aligned to tokens."""

def __init__(self) -> None:
self._chunks: list[np.ndarray] = []
self._schema: tuple[tuple[int, ...], np.dtype] | None = None
self._num_rows = 0
self._finalized = False

@property
def num_rows(self) -> int:
return self._num_rows

def append(self, rows: np.ndarray, *, expected_rows: int) -> None:
if self._finalized:
raise RuntimeError("token metadata trace is already finalized")
if isinstance(expected_rows, bool) or not isinstance(expected_rows, int) or expected_rows < 0:
raise ValueError(f"expected_rows must be a non-negative integer, got {expected_rows!r}")
if not isinstance(rows, np.ndarray):
raise TypeError("token metadata rows must be a NumPy array")
if rows.ndim < 1:
raise ValueError("token metadata must have a token-row dimension")
if rows.shape[0] != expected_rows:
raise ValueError(f"token metadata has {rows.shape[0]} rows, expected {expected_rows}")
if not rows.flags.c_contiguous:
raise ValueError("token metadata rows must be contiguous")

schema = (rows.shape[1:], rows.dtype)
if self._schema is None:
self._schema = schema
elif schema != self._schema:
raise ValueError(f"token metadata schema changed from {self._schema} to {schema}")

self._chunks.append(rows)
self._num_rows += expected_rows

def append_padding(self, count: int, *, fill: int = -1) -> None:
"""Append ``count`` rows of ``fill`` in the schema already established by ``append``."""
if isinstance(count, bool) or not isinstance(count, int) or count < 0:
raise ValueError(f"padding count must be a non-negative integer, got {count!r}")
if count == 0:
return
if self._schema is None:
raise ValueError("cannot pad token metadata before any rows are captured")

row_shape, dtype = self._schema
self.append(np.full((count, *row_shape), fill, dtype=dtype, order="C"), expected_rows=count)

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)
6 changes: 6 additions & 0 deletions skyrl/backends/skyrl_train/inference_servers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from typing import TYPE_CHECKING, Any, Dict, Hashable, List, Optional, Tuple, TypedDict

from skyrl.backends.skyrl_train.utils.routed_experts import RoutedExpertIndices
from skyrl.backends.skyrl_train.utils.sample_support import SampleSupport

if TYPE_CHECKING:
from skyrl.backends.skyrl_train.weight_sync import WeightUpdateRequest
Expand Down Expand Up @@ -34,6 +35,9 @@ 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]]
# Per-batch opt-in; the engine must enable sample-support capture at startup.
return_sample_support: Optional[bool]


class InferenceEngineOutput(TypedDict):
Expand All @@ -50,6 +54,8 @@ class InferenceEngineOutput(TypedDict):
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[RoutedExpertIndices]]
# One ``[generated_tokens, top_k]`` int32 support array per prompt.
rollout_sample_support: Optional[List[SampleSupport]]


class InferenceEngineInterface(ABC):
Expand Down
Loading
Loading