Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
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
146 changes: 138 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)}"
)
Comment on lines +137 to +141

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.

high

If layout.sequence_lengths is a PyTorch tensor, calling list(layout.sequence_lengths) will return a list of 0-D tensors. Comparing a list of Python ints (segment_lengths) to a list of 0-D tensors evaluates to False in Python, which will cause this validation to always fail and raise a ValueError incorrectly. Coerce layout.sequence_lengths to a list of Python ints using .tolist() if it is a tensor.

Suggested change
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)}"
)
layout_lengths = layout.sequence_lengths.tolist() if hasattr(layout.sequence_lengths, "tolist") else list(layout.sequence_lengths)
if segment_lengths != layout_lengths:
raise ValueError(
f"Packed metadata segments {segment_lengths} do not match "
f"trajectory lengths {layout_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,51 @@ def scatter_packed_token_values_to_batch(
)
batch_values[output_mask] = values[packed_mask]
return batch_values


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

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

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

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

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

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

def finalize(self, *, expected_rows: int) -> np.ndarray:
if self._finalized:
raise RuntimeError("token metadata trace is already finalized")
if self._num_rows != expected_rows:
raise ValueError(f"token metadata trace has {self._num_rows} rows, expected {expected_rows}")
if not self._chunks:
raise ValueError("token metadata trace has no chunks")

self._finalized = True
return self._chunks[0] if len(self._chunks) == 1 else np.concatenate(self._chunks, axis=0)
1 change: 1 addition & 0 deletions skyrl/backends/skyrl_train/inference_servers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ class InferenceEngineInput(TypedDict):
# Optional prefix-cache salt forwarded to vLLM as the request ``cache_salt`` so cache blocks are
# only shared between requests carrying the same salt. See ``GeneratorConfig.use_cache_salt``.
cache_salt: Optional[str]
routed_experts_prompt_starts: Optional[List[int]]


class InferenceEngineOutput(TypedDict):
Expand Down
Loading
Loading