Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
8 changes: 4 additions & 4 deletions tests/experimental/rollout/rollout_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,12 +54,12 @@ def tearDown(self):
self.service.stop()

def test_sampler_config_types(self):
"""Verifies sampler_type='vllm' raises NotImplementedError."""
"""Verifies unknown sampler_type raises ValueError."""
from tunix.experimental.rollout import manager as manager_lib # pylint: disable=g-import-not-at-top

config_vllm = worker.RolloutConfig(sampler_type="vllm")
with self.assertRaises(NotImplementedError):
manager_lib.RolloutManager(config=config_vllm)
config_invalid = worker.RolloutConfig(sampler_type="unknown")
with self.assertRaises(ValueError):
manager_lib.RolloutManager(config=config_invalid)

def test_single_trajectory_generation(self):
"""Verifies single multi-turn episode execution."""
Expand Down
283 changes: 283 additions & 0 deletions tests/experimental/rollout/vllm_sampler_adapter_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,283 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Tests for VllmSamplerAdapter with tpu-inference RLVllmSampler."""

import asyncio
from types import SimpleNamespace
from unittest import mock

from absl.testing import absltest
import numpy as np
from tunix.experimental.rollout import sampler as base_sampler_lib
from tunix.experimental.rollout import vllm_sampler_adapter


class VllmSamplerAdapterTest(absltest.TestCase):

def setUp(self):
super().setUp()
self.mock_sampler_instance = mock.AsyncMock()
self.sampler_adapter = vllm_sampler_adapter.VllmSamplerAdapter(
server_id="vllm_slice_01",
sampler_instance=self.mock_sampler_instance,
)

def test_single_sampling_request(self):
mock_response = SimpleNamespace(
request_id="req_01",
text="completion result",
token_ids=np.array([101, 102, 103], dtype=np.int32),
logprobs=np.array([-0.1, -0.2, -0.3], dtype=np.float32),
finish_reason="stop",
routed_experts=None,
error=None,
)
self.mock_sampler_instance.sample.return_value = mock_response

req = base_sampler_lib.SamplingRequest(
request_id="req_01",
prompt="Test prompt",
sampling_params=base_sampler_lib.SamplingParams(
max_tokens=16,
temperature=0.7,
return_logprobs=True,
),
)
response = asyncio.run(self.sampler_adapter.sample(req))
self.assertIsInstance(response, base_sampler_lib.SamplingResponse)
self.assertEqual(response.request_id, "req_01")
self.assertEqual(response.text, "completion result")
np.testing.assert_array_equal(response.token_ids, [101, 102, 103])
np.testing.assert_allclose(response.logprobs, [-0.1, -0.2, -0.3])
self.assertEqual(response.finish_reason, "stop")
self.mock_sampler_instance.sample.assert_called_once_with(req)

def test_batch_sampling_requests(self):
mock_responses = [
SimpleNamespace(
request_id="req_a",
text="completion a",
token_ids=np.array([10, 20], dtype=np.int32),
logprobs=np.array([-0.5, -0.6], dtype=np.float32),
finish_reason="stop",
routed_experts=None,
error=None,
),
SimpleNamespace(
request_id="req_b",
text="completion b",
token_ids=np.array([30, 40], dtype=np.int32),
logprobs=np.array([-0.7, -0.8], dtype=np.float32),
finish_reason="length",
routed_experts=None,
error=None,
),
]
self.mock_sampler_instance.sample.return_value = mock_responses

reqs = [
base_sampler_lib.SamplingRequest(
request_id="req_a",
prompt="Prompt A",
),
base_sampler_lib.SamplingRequest(
request_id="req_b",
prompt="Prompt B",
),
]
responses = asyncio.run(self.sampler_adapter.sample(reqs))
self.assertIsInstance(responses, list)
self.assertLen(responses, 2)
self.assertEqual(responses[0].request_id, "req_a")
self.assertEqual(responses[0].text, "completion a")
self.assertEqual(responses[1].request_id, "req_b")
self.assertEqual(responses[1].text, "completion b")
self.assertEqual(responses[1].finish_reason, "length")

def test_lifecycle_delegations(self):
self.mock_sampler_instance.start.return_value = None
self.mock_sampler_instance.stop.return_value = True
self.mock_sampler_instance.pause.return_value = True
self.mock_sampler_instance.resume.return_value = True
self.mock_sampler_instance.get_mesh.return_value = "mock_mesh"

asyncio.run(self.sampler_adapter.start())
self.mock_sampler_instance.start.assert_called_once()

asyncio.run(self.sampler_adapter.stop())
self.mock_sampler_instance.stop.assert_called_once()

asyncio.run(self.sampler_adapter.pause())
self.mock_sampler_instance.pause.assert_called_once()

asyncio.run(self.sampler_adapter.resume())
self.mock_sampler_instance.resume.assert_called_once()

mesh = asyncio.run(self.sampler_adapter.get_mesh())
self.assertEqual(mesh, "mock_mesh")

def test_weight_sync_delegations(self):
self.mock_sampler_instance.get_raiden_metadata.return_value = [
{
"unit": {
"job_name": "rollout",
"job_replica_id": "",
"data_name": "",
"data_replica_idx": 0,
},
"variables": (),
"mesh_shape": (1,),
"mesh_axes": ("fsdp",),
"data_address": "",
"control_plane_rpc_address": "",
}
]
self.mock_sampler_instance.raiden_h2d.return_value = {"tensor": 1.0}

meta = asyncio.run(self.sampler_adapter.get_weight_sync_metadata())
self.assertLen(meta, 1)

sync_req = base_sampler_lib.WeightSyncRequest(
policy_version=1, extra_config={"req_id": "r1", "uuid": 1}
)
res_pre = asyncio.run(self.sampler_adapter.pre_weight_sync(sync_req))
self.assertTrue(res_pre)
self.mock_sampler_instance.pre_weight_sync.assert_called_once_with(
free_kv_cache=True
)

res_sync = asyncio.run(self.sampler_adapter.weight_sync(sync_req))
self.assertTrue(res_sync)
self.mock_sampler_instance.raiden_h2d.assert_called_once()

res_post = asyncio.run(self.sampler_adapter.post_weight_sync(sync_req))
self.assertEqual(res_post, 1)
self.mock_sampler_instance.post_weight_sync.assert_called_once_with(sync_req)

status = asyncio.run(self.sampler_adapter.get_weight_sync_status())
self.assertEqual(status.get("phase"), "committed")

# Test abort path on next round (with higher uuid)
sync_req_2 = base_sampler_lib.WeightSyncRequest(
policy_version=2, extra_config={"req_id": "r2", "uuid": 2}
)
asyncio.run(self.sampler_adapter.pre_weight_sync(sync_req_2))
res_abort = asyncio.run(self.sampler_adapter.abort_weight_sync(sync_req_2))
self.assertTrue(res_abort)
status_2 = asyncio.run(self.sampler_adapter.get_weight_sync_status())
self.assertEqual(status_2.get("phase"), "aborted")

def test_get_load_info(self):
self.mock_sampler_instance.get_load_info.return_value = SimpleNamespace(
num_requests_waiting=3,
num_requests_running=1,
kv_cache_usage_perc=25.5,
)
load_info = asyncio.run(self.sampler_adapter.get_load_info())
self.assertIsInstance(load_info, base_sampler_lib.LoadInfo)
self.assertEqual(load_info.num_requests_waiting, 3)
self.assertEqual(load_info.num_requests_running, 1)
self.assertAlmostEqual(load_info.kv_cache_usage_perc, 25.5)

def test_migrate_kv_cache(self):
self.mock_sampler_instance.migrate_kv_cache.return_value = True
res = asyncio.run(
self.sampler_adapter.migrate_kv_cache(
source_server_id="src_0",
target_server_id="dst_0",
token_ids=[1, 2, 3],
route_key="key_1",
)
)
self.assertTrue(res)
self.mock_sampler_instance.migrate_kv_cache.assert_called_once_with(
route_key="key_1",
source_server_id="src_0",
target_server_id="dst_0",
token_ids=[1, 2, 3],
)

def test_uninitialized_raises(self):
uninit = vllm_sampler_adapter.VllmSamplerAdapter(server_id="empty")
with self.assertRaises(RuntimeError):
asyncio.run(uninit.sample(base_sampler_lib.SamplingRequest(prompt="hi")))
with self.assertRaises(RuntimeError):
asyncio.run(uninit.stop())
with self.assertRaises(RuntimeError):
asyncio.run(uninit.pause())
with self.assertRaises(RuntimeError):
asyncio.run(uninit.resume())
with self.assertRaises(RuntimeError):
asyncio.run(uninit.get_mesh())
with self.assertRaises(RuntimeError):
asyncio.run(uninit.get_weight_sync_metadata())
with self.assertRaises(RuntimeError):
asyncio.run(uninit.pre_weight_sync(None))
with self.assertRaises(RuntimeError):
asyncio.run(uninit.weight_sync(None))
with self.assertRaises(RuntimeError):
asyncio.run(uninit.post_weight_sync(None))
with self.assertRaises(RuntimeError):
asyncio.run(uninit.get_transfer_status("req"))
with self.assertRaises(RuntimeError):
asyncio.run(uninit.get_load_info())
with self.assertRaises(RuntimeError):
asyncio.run(
uninit.migrate_kv_cache(
source_server_id="s", target_server_id="t", token_ids=[1]
)
)

def test_sample_none_requests_raises(self):
with self.assertRaises(ValueError):
asyncio.run(self.sampler_adapter.sample(None))

def test_incomplete_sampler_fails_at_init(self):
"""Raiden-enabled adapters reject samplers missing required hooks."""
incomplete = mock.AsyncMock(spec=["start", "stop", "sample"])
with self.assertRaises(RuntimeError) as ctx:
vllm_sampler_adapter.VllmSamplerAdapter(
server_id="vllm_slice_01", sampler_instance=incomplete
)
message = str(ctx.exception)
for name in vllm_sampler_adapter._REQUIRED_RAIDEN_METHODS: # pylint: disable=protected-access
self.assertIn(name, message)

def test_incomplete_sampler_allowed_when_raiden_disabled(self):
"""The protocol check only applies when Raiden weight sync is enabled."""
incomplete = mock.AsyncMock(spec=["start", "stop", "sample"])
adapter = vllm_sampler_adapter.VllmSamplerAdapter(
server_id="vllm_slice_01",
sampler_instance=incomplete,
weight_sync_mode="none",
)
self.assertFalse(adapter.enable_raiden)

def test_weight_sync_apis_fail_when_uninitialized(self):
"""Weight sync entry points fail instead of silently initializing."""
uninit = vllm_sampler_adapter.VllmSamplerAdapter(server_id="vllm_slice_01")
for coro in (
uninit.bind_weight_sync(),
uninit.pre_weight_sync(None),
uninit.weight_sync(None),
uninit.post_weight_sync(None),
uninit.abort_weight_sync(None),
):
with self.assertRaises(RuntimeError):
asyncio.run(coro)


if __name__ == "__main__":
absltest.main()
8 changes: 8 additions & 0 deletions tests/experimental/weight_sync/raiden_synchronizer_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,6 +294,14 @@ def test_host_stage_pulls_state_to_host(self):
)
pull.assert_called_once()

def test_release_host_arrays(self):
sync = raiden_synchronizer.RaidenSynchronizer("trainer", self._state())
self.assertTrue(sync.bound)
self.assertNotEmpty(sync.arrays)
sync.release_host_arrays()
self.assertEmpty(sync.arrays)
self.assertTrue(sync.bound)


if __name__ == "__main__":
absltest.main()
Original file line number Diff line number Diff line change
Expand Up @@ -363,6 +363,7 @@ def _create_vllm_sampler(args):
server_id=args.worker_id,
engine_args=engine_args,
model_name=vllm_model,
weight_sync_mode=args.weight_sync_mode,
)
config = rollout_worker.RolloutConfig(
sampler_type="vllm",
Expand Down
1 change: 1 addition & 0 deletions tunix/experimental/rollout/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,5 +21,6 @@
from tunix.experimental.rollout.manager import RolloutManager
from tunix.experimental.rollout.sampler import Sampler
from tunix.experimental.rollout.vanilla_sampler_adapter import VanillaSamplerAdapter
from tunix.experimental.rollout.vllm_sampler_adapter import VllmSamplerAdapter
from tunix.experimental.trajectory.trajectory import Trajectory
from tunix.experimental.trajectory.trajectory import TrajectoryError
9 changes: 6 additions & 3 deletions tunix/experimental/rollout/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,9 +70,12 @@ def __init__(
)

if sampler_type == "vllm":
raise NotImplementedError(
"vLLM sampler is not implemented yet. Use 'inprocess_vllm' or"
" 'vanilla'."
from tunix.experimental.rollout import vllm_sampler_adapter # pylint: disable=g-import-not-at-top

sampler = vllm_sampler_adapter.VllmSamplerAdapter( # pyrefly: ignore[bad-instantiation]
server_id="vllm_sampler",
model_name=getattr(config, "rollout_vllm_model_version", ""),
weight_sync_mode=weight_sync_mode,
)
elif "inprocess_vllm" in sampler_type:
from tunix.experimental.rollout import inprocess_vllm_sampler_adapter # pylint: disable=g-import-not-at-top
Expand Down
Loading
Loading