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
21 changes: 12 additions & 9 deletions agentlightning/verl/rollout_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,15 +191,15 @@ def _build_routed_experts_batch(
max_response_length: int,
device: torch.device,
) -> torch.Tensor:
decoded = [np.load(io.BytesIO(base64.b64decode(payload)), allow_pickle=False) for payload, _, _, _ in rows]
batch = torch.zeros(
(len(rows), max_prompt_length + max_response_length, *decoded[0].shape[1:]),
dtype=torch.uint8,
device=device,
)
for index, (routes, (_, original_prompt_length, prompt_length, response_length)) in enumerate(
zip(decoded, rows, strict=True)
):
batch: torch.Tensor | None = None
for index, (payload, original_prompt_length, prompt_length, response_length) in enumerate(rows):
routes = np.load(io.BytesIO(base64.b64decode(payload)), allow_pickle=False)
if batch is None:
batch = torch.zeros(
(len(rows), max_prompt_length + max_response_length, *routes.shape[1:]),
dtype=torch.uint8,
device=device,
)
if len(routes) < original_prompt_length + response_length - 1:
raise RuntimeError("R3 routed_experts is shorter than its token sequence")
routes = torch.from_numpy(routes).to(device=device, dtype=torch.uint8)
Expand All @@ -210,6 +210,9 @@ def _build_routed_experts_batch(
batch[index, max_prompt_length : max_prompt_length + response_routes] = routes[
original_prompt_length : original_prompt_length + response_routes
]
# Release the NumPy-backed tensor before decoding the next row.
del routes
assert batch is not None
return batch


Expand Down
58 changes: 57 additions & 1 deletion tests/verl/test_rollout_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import io
import json
import sys
import weakref
import zipfile
from types import SimpleNamespace
from typing import ClassVar
Expand All @@ -16,11 +17,12 @@
pytest.importorskip("tensordict")
pytest.importorskip("verl")

import numpy as np
import torch

from agentlightning.verl import rollout_adapter as rollout_adapter_module
from agentlightning.verl.agl_rollout_manager import CompletedRollout, Triplet
from agentlightning.verl.rollout_adapter import RolloutAdapter
from agentlightning.verl.rollout_adapter import RolloutAdapter, _build_routed_experts_batch


class FakeTokenizer:
Expand All @@ -30,6 +32,60 @@ def decode(self, ids: list[int], skip_special_tokens: bool = True) -> str:
return " ".join(str(i) for i in ids)


def _encode_routes(routes: np.ndarray) -> str:
buffer = io.BytesIO()
np.save(buffer, routes, allow_pickle=False)
return base64.b64encode(buffer.getvalue()).decode()


@pytest.mark.parametrize("dtype", [np.uint8, np.int64])
def test_routed_experts_batch_preserves_padding_and_truncation(dtype: type[np.generic]) -> None:
rows = []
for tokens, original_prompt_length, prompt_length, response_length in [
([10, 11, 12, 13], 2, 2, 3),
([20, 21, 22, 23, 24, 25, 26], 6, 4, 2),
([30, 31, 32, 33, 34, 35, 36], 1, 1, 3),
]:
routes = (np.array(tokens)[:, None, None] * 4 + np.arange(4).reshape(1, 2, 2)).astype(dtype)
rows.append((_encode_routes(routes), original_prompt_length, prompt_length, response_length))

batch = _build_routed_experts_batch(rows, 4, 3, torch.device("cpu"))

expected = torch.zeros((3, 7, 2, 2), dtype=torch.uint8)
for row, tokens in enumerate([[0, 0, 10, 11, 12, 13, 0], [20, 21, 22, 23, 26, 0, 0], [0, 0, 0, 30, 31, 32, 33]]):
for position, token in enumerate(tokens):
if token:
expected[row, position] = torch.arange(4, dtype=torch.uint8).reshape(2, 2) + token * 4
torch.testing.assert_close(batch, expected, rtol=0, atol=0)


def test_routed_experts_batch_releases_decoded_rows(monkeypatch: pytest.MonkeyPatch) -> None:
payload = _encode_routes(np.arange(48, dtype=np.uint8).reshape(12, 2, 2))
decoded: list[weakref.ReferenceType[np.ndarray]] = []
load = np.load

def tracked_load(*args, **kwargs):
assert all(ref() is None for ref in decoded), "Previous decoded routes are still retained"
routes = load(*args, **kwargs)
decoded.append(weakref.ref(routes))
return routes

monkeypatch.setattr(rollout_adapter_module.np, "load", tracked_load)
batch = _build_routed_experts_batch([(payload, 4, 4, 3)] * 4, 4, 3, torch.device("cpu"))

assert len(decoded) == 4
assert all(ref() is None for ref in decoded)
expected = torch.arange(28, dtype=torch.uint8).reshape(7, 2, 2).expand(4, -1, -1, -1)
torch.testing.assert_close(batch, expected, rtol=0, atol=0)


def test_routed_experts_batch_rejects_short_route_sequence() -> None:
valid = _encode_routes(np.zeros((4, 2, 2), dtype=np.uint8))
short = _encode_routes(np.zeros((3, 2, 2), dtype=np.uint8))
with pytest.raises(RuntimeError, match="R3 routed_experts is shorter than its token sequence"):
_build_routed_experts_batch([(valid, 2, 2, 3), (short, 2, 2, 3)], 4, 3, torch.device("cpu"))


class FakeTable:
def __init__(self, columns: list[str]) -> None:
self.columns = columns
Expand Down