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
133 changes: 4 additions & 129 deletions tests/experimental/rollout/inprocess_vllm_sampler_adapter_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,16 +12,15 @@
# See the License for the specific language governing permissions and
# limitations under the License.

"""Tests for InprocessVllmSamplerAdapter with Tunix VllmSampler and Raiden delegate."""
"""Tests for InprocessVllmSamplerAdapter with Tunix VllmSampler."""

import asyncio
from unittest import mock
from absl.testing import absltest
import numpy as np

from tunix.experimental.rollout import inprocess_vllm_sampler_adapter
from tunix.experimental.rollout import sampler as base_sampler_lib
from tunix.experimental.weight_sync import raiden_weight_sync_delegate
from tunix.experimental.weight_sync import weight_sync
from tunix.generate import base_sampler


Expand All @@ -37,7 +36,6 @@ def setUp(self):
padded_prompt_tokens=np.array([[1, 2]], dtype=np.int32),
logprobs=None,
)
self.mock_vllm_sampler.mesh = "mock_mesh"
self.mock_vllm_lib = mock.MagicMock()
self.mock_vllm_lib.VllmSampler.return_value = self.mock_vllm_sampler

Expand All @@ -50,7 +48,6 @@ def setUp(self):

self.mock_tokenizer = mock.MagicMock()
self.mock_config = mock.MagicMock()
self.mock_config.enable_raiden = False

self.sampler_adapter = (
inprocess_vllm_sampler_adapter.InprocessVllmSamplerAdapter(
Expand All @@ -64,17 +61,6 @@ def tearDown(self):
self.patcher.stop()
super().tearDown()

def test_implements_sampler_protocol(self):
self.assertIsInstance(self.sampler_adapter, base_sampler_lib.Sampler)

def test_lifecycle_methods(self):
self.assertTrue(asyncio.run(self.sampler_adapter.start()))
self.assertTrue(asyncio.run(self.sampler_adapter.pause()))
self.assertTrue(asyncio.run(self.sampler_adapter.resume()))
self.assertEqual(asyncio.run(self.sampler_adapter.get_mesh()), "mock_mesh")
self.assertTrue(asyncio.run(self.sampler_adapter.stop()))
self.mock_vllm_sampler.stop.assert_called_once()

def test_single_sampling_request(self):
req = base_sampler_lib.SamplingRequest(
request_id="vllm_req_01",
Expand Down Expand Up @@ -121,123 +107,12 @@ def test_batch_sampling_requests(self):
np.testing.assert_array_equal(responses[0].prompt_token_ids, [1, 2])
np.testing.assert_array_equal(responses[1].prompt_token_ids, [3, 4])

def test_weight_sync_without_raiden_delegate(self):
def test_weight_sync(self):
mock_weights = {"layer1": "weights"}
req = base_sampler_lib.WeightSyncRequest(weights=mock_weights)
res = asyncio.run(self.sampler_adapter.weight_sync(sync_request=req))
res = asyncio.run(self.sampler_adapter.weight_sync(mock_weights))
self.assertTrue(res)
self.mock_vllm_sampler.update_params.assert_called_once_with(mock_weights)

# Missing sync_request should raise ValueError
with self.assertRaises(ValueError):
asyncio.run(self.sampler_adapter.weight_sync(sync_request=None))

# Missing weights in sync_request should raise ValueError
empty_req = base_sampler_lib.WeightSyncRequest()
with self.assertRaises(ValueError):
asyncio.run(self.sampler_adapter.weight_sync(sync_request=empty_req))

self.assertIsNone(asyncio.run(self.sampler_adapter.bind_weight_sync()))
self.assertTrue(asyncio.run(self.sampler_adapter.pre_weight_sync()))
self.assertTrue(asyncio.run(self.sampler_adapter.post_weight_sync()))
with self.assertRaises(NotImplementedError):
asyncio.run(self.sampler_adapter.get_weight_sync_metadata())

def test_weight_sync_with_raiden_delegate(self):
mock_delegate = mock.MagicMock(
spec=raiden_weight_sync_delegate.RaidenWeightSyncDelegate
)
mock_delegate.is_bounded.return_value = False
mock_delegate.bind_weight_sync = mock.AsyncMock(return_value=True)
mock_delegate.get_weight_sync_metadata = mock.AsyncMock(
return_value=[{"unit": "rollout"}]
)
mock_delegate.pre_weight_sync = mock.AsyncMock(return_value=True)
mock_delegate.weight_sync = mock.AsyncMock(return_value=5)
mock_delegate.post_weight_sync = mock.AsyncMock(return_value=True)

fake_transformer_state = {"param": "tensor"}
self.mock_vllm_sampler.transformer_state = fake_transformer_state

raiden_config = mock.MagicMock()
raiden_config.weight_sync_mode = weight_sync.WeightSyncMode.RAIDEN

raiden_adapter = inprocess_vllm_sampler_adapter.InprocessVllmSamplerAdapter(
server_id="vllm_raiden_slice",
tokenizer=self.mock_tokenizer,
config=raiden_config,
raiden_sync_delegate=mock_delegate,
)

sync_req = base_sampler_lib.WeightSyncRequest(policy_version=5)

# 1. bind_weight_sync
asyncio.run(raiden_adapter.bind_weight_sync(sync_req))
mock_delegate.bind_weight_sync.assert_awaited_once_with(
sync_request=sync_req, state=fake_transformer_state
)
mock_delegate.is_bounded.return_value = True

# 2. get_weight_sync_metadata
metadata = asyncio.run(raiden_adapter.get_weight_sync_metadata())
self.assertEqual(metadata, [{"unit": "rollout"}])

# 3. pre_weight_sync
self.assertTrue(asyncio.run(raiden_adapter.pre_weight_sync(sync_req)))
mock_delegate.pre_weight_sync.assert_awaited_once_with(
sync_request=sync_req
)

# 4. weight_sync
version = asyncio.run(raiden_adapter.weight_sync(sync_req))
self.assertEqual(version, 5)
mock_delegate.weight_sync.assert_awaited_once_with(sync_request=sync_req)

# 5. post_weight_sync
self.assertTrue(asyncio.run(raiden_adapter.post_weight_sync(sync_req)))
mock_delegate.post_weight_sync.assert_awaited_once_with(
sync_request=sync_req
)

def test_raiden_bind_without_transformer_state_raises(self):
mock_delegate = mock.MagicMock(
spec=raiden_weight_sync_delegate.RaidenWeightSyncDelegate
)
mock_delegate.is_bounded.return_value = False
# Ensure vllm_sampler has no transformer_state
if hasattr(self.mock_vllm_sampler, "transformer_state"):
del self.mock_vllm_sampler.transformer_state

raiden_config = mock.MagicMock()
raiden_config.weight_sync_mode = weight_sync.WeightSyncMode.RAIDEN

raiden_adapter = inprocess_vllm_sampler_adapter.InprocessVllmSamplerAdapter(
server_id="vllm_raiden_slice",
tokenizer=self.mock_tokenizer,
config=raiden_config,
raiden_sync_delegate=mock_delegate,
)

with self.assertRaisesRegex(RuntimeError, "transformer_state"):
asyncio.run(raiden_adapter.bind_weight_sync())

def test_other_sampler_methods(self):
self.assertEqual(
asyncio.run(self.sampler_adapter.get_transfer_status("req_1")),
"SUCCESS",
)
self.assertTrue(
asyncio.run(
self.sampler_adapter.migrate_kv_cache(
source_server_id="s1",
target_server_id="s2",
token_ids=[1, 2, 3],
)
)
)
load_info = asyncio.run(self.sampler_adapter.get_load_info())
self.assertIsInstance(load_info, base_sampler_lib.LoadInfo)

def test_uninitialized_sampler_raises(self):
uninit = inprocess_vllm_sampler_adapter.InprocessVllmSamplerAdapter(
server_id="empty"
Expand Down
48 changes: 1 addition & 47 deletions tests/experimental/rollout/manager_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,14 +13,11 @@
# limitations under the License.

import asyncio
import types
import unittest
from unittest import mock

from absl.testing import absltest
from tunix.experimental.rollout import manager as manager_lib
from tunix.experimental.rollout import sampler as sampler_lib
from tunix.experimental.weight_sync import weight_sync


class _FakeSampler(sampler_lib.Sampler):
Expand Down Expand Up @@ -94,7 +91,7 @@ async def test_post_reopens_admission(self):
await manager.pre_weight_sync()
await manager.post_weight_sync()
self.assertTrue(manager._traffic.is_admission_open())

async def test_reopen_admission_after_abort(self):
manager = self._manager()
await manager.pre_weight_sync()
Expand Down Expand Up @@ -140,48 +137,5 @@ async def test_drain_timeout_returns(self):
task.cancel()
manager._active_tasks.pop("t0", None)


class WeightSyncModeTest(absltest.TestCase):

def test_config_weight_sync_mode_raiden(self):
config = types.SimpleNamespace(
sampler_type="vanilla",
weight_sync_mode=weight_sync.WeightSyncMode.RAIDEN,
)
manager = manager_lib.RolloutManager(
config=config, tokenizer="mock", chat_parser="mock"
)
self.assertTrue(getattr(manager.sampler, "enable_raiden", False))
self.assertIsNotNone(getattr(manager.sampler, "raiden_sync_delegate", None))

def test_config_weight_sync_mode_fallback(self):
config = types.SimpleNamespace(
sampler_type="vanilla",
weight_sync_mode=weight_sync.WeightSyncMode.FALLBACK,
)
manager = manager_lib.RolloutManager(
config=config, tokenizer="mock", chat_parser="mock"
)
self.assertFalse(getattr(manager.sampler, "enable_raiden", False))
self.assertIsNone(getattr(manager.sampler, "raiden_sync_delegate", None))

@mock.patch(
"tunix.experimental.rollout.inprocess_vllm_sampler_adapter._get_vllm_sampler_cls"
)
def test_config_weight_sync_mode_inprocess_vllm_raiden(self, mock_get_vllm):
mock_lib = mock.MagicMock()
mock_lib.VllmSampler.return_value = mock.MagicMock()
mock_get_vllm.return_value = mock_lib
config = types.SimpleNamespace(
sampler_type="inprocess_vllm",
weight_sync_mode=weight_sync.WeightSyncMode.RAIDEN,
)
manager = manager_lib.RolloutManager(
config=config, tokenizer="mock", chat_parser="mock"
)
self.assertTrue(getattr(manager.sampler, "enable_raiden", False))
self.assertIsNotNone(getattr(manager.sampler, "raiden_sync_delegate", None))


if __name__ == "__main__":
absltest.main()
Loading
Loading