From abff6d22e5550cbf4d367fe981ac15fc3a5da26c Mon Sep 17 00:00:00 2001 From: Bozhen Peng <42286547+kiteretsu903@users.noreply.github.com> Date: Sat, 3 Oct 2026 19:56:04 -0700 Subject: [PATCH] perf(verl): reduce peak memory when decoding routed experts --- agentlightning/verl/rollout_adapter.py | 21 ++++++---- tests/verl/test_rollout_adapter.py | 58 +++++++++++++++++++++++++- 2 files changed, 69 insertions(+), 10 deletions(-) diff --git a/agentlightning/verl/rollout_adapter.py b/agentlightning/verl/rollout_adapter.py index 18f2e1cca..92b87565c 100644 --- a/agentlightning/verl/rollout_adapter.py +++ b/agentlightning/verl/rollout_adapter.py @@ -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) @@ -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 diff --git a/tests/verl/test_rollout_adapter.py b/tests/verl/test_rollout_adapter.py index 8d0843926..fd1684ed7 100644 --- a/tests/verl/test_rollout_adapter.py +++ b/tests/verl/test_rollout_adapter.py @@ -6,6 +6,7 @@ import io import json import sys +import weakref import zipfile from types import SimpleNamespace from typing import ClassVar @@ -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: @@ -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