diff --git a/tests/experimental/rollout/rollout_test.py b/tests/experimental/rollout/rollout_test.py index accb3d688..0219a1d1f 100644 --- a/tests/experimental/rollout/rollout_test.py +++ b/tests/experimental/rollout/rollout_test.py @@ -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.""" diff --git a/tests/experimental/rollout/vllm_sampler_adapter_test.py b/tests/experimental/rollout/vllm_sampler_adapter_test.py new file mode 100644 index 000000000..9c46023ce --- /dev/null +++ b/tests/experimental/rollout/vllm_sampler_adapter_test.py @@ -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() diff --git a/tests/experimental/weight_sync/raiden_synchronizer_test.py b/tests/experimental/weight_sync/raiden_synchronizer_test.py index 156d8c533..d94c01537 100644 --- a/tests/experimental/weight_sync/raiden_synchronizer_test.py +++ b/tests/experimental/weight_sync/raiden_synchronizer_test.py @@ -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() \ No newline at end of file diff --git a/tunix/experimental/examples/math_gsm8k_dist/run_rollout_node.py b/tunix/experimental/examples/math_gsm8k_dist/run_rollout_node.py index f8492fc99..26af29cf1 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_rollout_node.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_rollout_node.py @@ -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", diff --git a/tunix/experimental/rollout/__init__.py b/tunix/experimental/rollout/__init__.py index 0e16324f8..4677a8107 100644 --- a/tunix/experimental/rollout/__init__.py +++ b/tunix/experimental/rollout/__init__.py @@ -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 diff --git a/tunix/experimental/rollout/manager.py b/tunix/experimental/rollout/manager.py index 9ab81353e..1aedc560b 100644 --- a/tunix/experimental/rollout/manager.py +++ b/tunix/experimental/rollout/manager.py @@ -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 diff --git a/tunix/experimental/rollout/vllm_sampler_adapter.py b/tunix/experimental/rollout/vllm_sampler_adapter.py new file mode 100644 index 000000000..5639e9796 --- /dev/null +++ b/tunix/experimental/rollout/vllm_sampler_adapter.py @@ -0,0 +1,401 @@ +# 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. + +"""vLLM Sampler adapter implementing Tunix WeightSyncDestination via Raiden.""" + +from __future__ import annotations + +import asyncio +import importlib +import logging +import os +from typing import Any, List, Mapping, Sequence + +import numpy as np +from tunix.experimental.weight_sync import weight_sync +from tunix.experimental.weight_sync import weight_sync_coordinator +from tunix.experimental.rollout import sampler as base_sampler_lib + +Sampler = base_sampler_lib.Sampler +logger = logging.getLogger(__name__) + + +def _get_rl_vllm_sampler_cls(): + """Lazy import of tpu_inference.rl.RLVllmSampler. + + Resolved through importlib so static analyzers do not try to follow the + tpu-inference dependency, which is not available in every environment. + """ + try: + return getattr(importlib.import_module("tpu_inference.rl"), "RLVllmSampler") + except (ImportError, AttributeError) as e: + raise ImportError( + "tpu_inference.rl.RLVllmSampler is not available. Please ensure" + " tpu-inference is installed." + ) from e + + +# Hooks RLVllmSampler must expose for Raiden weight sync; verified once at +# init rather than on every call. +_REQUIRED_RAIDEN_METHODS = ( + "bind_raiden_sync", + "get_raiden_metadata", + "pre_weight_sync", + "raiden_h2d", + "post_weight_sync", +) + + +def _format_sampling_response(r: Any) -> base_sampler_lib.SamplingResponse: + """Formats raw sampler output into a standardized SamplingResponse.""" + if isinstance(r, base_sampler_lib.SamplingResponse): + return r + tok_ids = getattr(r, "token_ids", np.zeros(0, dtype=np.int32)) + if not isinstance(tok_ids, np.ndarray): + tok_ids = np.array(tok_ids, dtype=np.int32) + prompt_ids = getattr(r, "prompt_token_ids", None) + if prompt_ids is None: + prompt_ids = np.zeros(0, dtype=np.int32) + elif not isinstance(prompt_ids, np.ndarray): + prompt_ids = np.array(prompt_ids, dtype=np.int32) + lps = getattr(r, "logprobs", None) + if lps is not None and not isinstance(lps, np.ndarray): + lps = np.array(lps, dtype=np.float32) + return base_sampler_lib.SamplingResponse( + request_id=getattr(r, "request_id", ""), + text=getattr(r, "text", ""), + token_ids=tok_ids, + logprobs=lps, + prompt_token_ids=prompt_ids, + finish_reason=getattr(r, "finish_reason", "stop"), + routed_experts=getattr(r, "routed_experts", None), + error=getattr(r, "error", None), + ) + + +class VllmSamplerAdapter(Sampler, weight_sync.WeightSyncDestination): + """Sampler adapter wrapping tpu-inference RLVllmSampler with full Raiden weight sync.""" + + def __init__( + self, + server_id: str = "vllm-rollout-0", + engine_args: Any = None, + model_name: str = "", + sampler_instance: Any = None, + worker_index: int = 0, + parallelism: int = 4, + weight_sync_mode: weight_sync.WeightSyncMode | str | None = None, + **kwargs, + ): + self.server_id = server_id + self.engine_args = engine_args + self.model_name = model_name or (engine_args.model if engine_args else "") + self.sampler = sampler_instance + self.worker_index = worker_index + self._parallelism = parallelism + + # Defaults to RAIDEN when unspecified: RLVllmSampler drives weight sync + # through its own native Raiden hooks, so callers that construct the + # adapter directly keep the historical always-Raiden behaviour. + if isinstance(weight_sync_mode, weight_sync.WeightSyncMode): + self.weight_sync_mode = weight_sync_mode + elif isinstance(weight_sync_mode, str): + self.weight_sync_mode = weight_sync.WeightSyncMode(weight_sync_mode) + else: + self.weight_sync_mode = weight_sync.WeightSyncMode.RAIDEN + self.enable_raiden = ( + self.weight_sync_mode == weight_sync.WeightSyncMode.RAIDEN + ) + if not self.enable_raiden: + logger.info( + "VllmSamplerAdapter [%s] weight_sync_mode=%s; Raiden weight sync is" + " disabled.", + self.server_id, + self.weight_sync_mode.value, + ) + + self._tracker = weight_sync_coordinator.WorkerRoundTracker() + self._sync_lock = asyncio.Lock() + self._policy_version = 0 + self._kv_cache_freed = False + + if self.sampler is None and self.engine_args is not None: + sampler_cls = _get_rl_vllm_sampler_cls() + self.sampler = sampler_cls(engine_args=self.engine_args) + self._verify_sampler_protocol() + + def initialize(self) -> None: + """Initializes RLVllmSampler if not already initialized.""" + if self.sampler is None: + if self.engine_args is None and self.model_name: + from vllm.engine.arg_utils import AsyncEngineArgs # pylint: disable=g-import-not-at-top + self.engine_args = AsyncEngineArgs(model=self.model_name) + if self.engine_args is not None: + sampler_cls = _get_rl_vllm_sampler_cls() + self.sampler = sampler_cls(engine_args=self.engine_args) + if self.sampler is None: + raise RuntimeError( + f"VllmSamplerAdapter [{self.server_id}] requires valid" + " engine_args or model_name." + ) + self._verify_sampler_protocol() + + def _verify_sampler_protocol(self) -> None: + """Fails fast when a Raiden-enabled sampler lacks the required hooks.""" + if not self.enable_raiden or self.sampler is None: + return + missing = [ + name + for name in _REQUIRED_RAIDEN_METHODS + if not hasattr(self.sampler, name) + ] + if missing: + raise RuntimeError( + f"VllmSamplerAdapter [{self.server_id}] sampler" + f" {type(self.sampler).__name__} is missing required Raiden" + f" methods: {', '.join(missing)}." + ) + + def _require_sampler(self) -> Any: + """Returns the sampler, failing if it has not been initialized.""" + if self.sampler is None: + raise RuntimeError( + f"VllmSamplerAdapter [{self.server_id}] is not initialized." + ) + return self.sampler + + # --------------------------------------------------------------------------- + # Lifecycle & Inference Methods + # --------------------------------------------------------------------------- + + async def start(self, **kwargs) -> Any: + """Starts the underlying sampler engine.""" + return await self._require_sampler().start(**kwargs) + + async def stop(self, **kwargs) -> Any: + """Stops the underlying sampler engine.""" + return await self._require_sampler().stop(**kwargs) + + async def pause(self, **kwargs) -> Any: + """Pauses inference processing on this worker slice.""" + return await self._require_sampler().pause(**kwargs) + + async def resume(self, **kwargs) -> Any: + """Resumes inference processing on this worker slice.""" + return await self._require_sampler().resume(**kwargs) + + async def get_mesh(self, **kwargs) -> Any: + """Returns the underlying device mesh topology.""" + return await self._require_sampler().get_mesh(**kwargs) + + async def sample( + self, + sampling_requests: ( + base_sampler_lib.SamplingRequest + | Sequence[base_sampler_lib.SamplingRequest] + | Any + ), + **kwargs, + ) -> ( + base_sampler_lib.SamplingResponse + | List[base_sampler_lib.SamplingResponse] + | Any + ): + """Generates completions using underlying tpu-inference RLVllmSampler.""" + if sampling_requests is None: + raise ValueError("sampling_requests cannot be None.") + + is_sequence = isinstance(sampling_requests, (list, tuple)) + raw_responses = await self._require_sampler().sample( + sampling_requests, **kwargs + ) + + if isinstance(raw_responses, (list, tuple)): + formatted = [_format_sampling_response(r) for r in raw_responses] + if is_sequence: + return formatted + return formatted[0] if formatted else base_sampler_lib.SamplingResponse() + + return _format_sampling_response(raw_responses) + + # --------------------------------------------------------------------------- + # WeightSyncDestination Protocol Implementation + # --------------------------------------------------------------------------- + + async def bind_weight_sync( + self, + sync_request: base_sampler_lib.WeightSyncRequest | Any = None, + **kwargs: Any, + ) -> Any: + """Idempotent transport binding called while the worker is STILL SERVING.""" + del sync_request, kwargs + if not self.enable_raiden: + return None + return await self._require_sampler().bind_raiden_sync( + worker_index=self.worker_index, parallelism=self._parallelism + ) + + async def get_weight_sync_metadata( + self, + **kwargs: Any, + ) -> Sequence[weight_sync.WorkUnitMetadata] | Any: + """Returns transport metadata with Raiden endpoints and TensorMetadata.""" + del kwargs + if not self.enable_raiden: + raise NotImplementedError( + f"VllmSamplerAdapter [{self.server_id}] does not support" + " get_weight_sync_metadata when Raiden is disabled" + f" (weight_sync_mode={self.weight_sync_mode.value})." + ) + meta = await self._require_sampler().get_raiden_metadata() + return [weight_sync.WorkUnitMetadata.from_dict(m) for m in meta or []] + + async def pre_weight_sync(self, sync_request: Any = None, **kwargs: Any) -> Any: + """Quiesces intake, drains pending requests, resets prefix cache, and drops KV cache.""" + if not self.enable_raiden: + return True + sampler = self._require_sampler() + async with self._sync_lock: + if not self._tracker.admit(sync_request, "prepared"): + return True + + logger.info("Executing pre_weight_sync for server_id=%s", self.server_id) + + # delegate to RLVllmSampler's native pause + clear + free-kv-cache + await sampler.pre_weight_sync(free_kv_cache=True) + self._kv_cache_freed = True + + self._tracker.complete(sync_request, "prepared") + return True + + async def weight_sync(self, sync_request: Any = None, **kwargs: Any) -> Any: + """Flushes/awaits H2D transfers and refreshes state_leaves.""" + if not self.enable_raiden: + # RLVllmSampler owns its weight buffers, so there is no host-side + # update_params fallback equivalent to the in-process adapter's. + raise RuntimeError( + f"VllmSamplerAdapter [{self.server_id}] supports Raiden weight sync" + " only; no fallback path exists" + f" (weight_sync_mode={self.weight_sync_mode.value})." + ) + sampler = self._require_sampler() + async with self._sync_lock: + if not self._tracker.admit(sync_request, "h2d_done"): + return True + + logger.info("Executing weight_sync barrier on Raiden synchronizers...") + checksums = await sampler.raiden_h2d() + if checksums: + logger.info("Destination weights checksums: %s", checksums) + + if hasattr(sampler, "refresh_model_state_leaves"): + result = sampler.refresh_model_state_leaves() + if asyncio.iscoroutine(result): + await result + + self._tracker.complete(sync_request, "h2d_done") + return True + + async def post_weight_sync(self, sync_request: Any = None, **kwargs: Any) -> Any: + """Reinitializes KV cache, restores request intake, and bumps active policy version.""" + if not self.enable_raiden: + return True + sampler = self._require_sampler() + async with self._sync_lock: + if not self._tracker.admit(sync_request, "committed"): + return True + + logger.info("Executing post_weight_sync: restoring serving state...") + # delegate to RLVllmSampler's native reinitialize-kv-cache + resume + if self._kv_cache_freed: + await sampler.post_weight_sync(sync_request) + self._kv_cache_freed = False + else: + await self.resume() + + version = getattr(sync_request, "policy_version", None) + if version is not None: + self._policy_version = version + else: + self._policy_version += 1 + + if os.environ.get("VERIFY_WEIGHTS", "").lower() == "true" and hasattr( + sampler, "raiden_metrics" + ): + logger.info("Raiden transfer metrics: %s", await sampler.raiden_metrics()) + + self._tracker.complete(sync_request, "committed") + return self._policy_version + + async def abort_weight_sync(self, sync_request: Any = None, **kwargs: Any) -> Any: + """Safely rolls back to serving previous policy version without publishing staging.""" + if not self.enable_raiden: + return True + sampler = self._require_sampler() + async with self._sync_lock: + if not self._tracker.admit(sync_request, "aborted"): + return False + + logger.warning( + "Aborting weight sync round: rolling back to policy_version=%d", + self._policy_version, + ) + # RLVllmSampler has no dedicated abort path; post_weight_sync does + # the same recovery (reinitialize KV cache + resume). + if self._kv_cache_freed: + await sampler.post_weight_sync(sync_request) + self._kv_cache_freed = False + else: + await self.resume() + + self._tracker.complete(sync_request, "aborted") + return True + + async def get_weight_sync_status(self) -> Mapping[str, Any]: + """Reports worker-side round status for coordinator recovery checks.""" + return self._tracker.report() + + async def get_transfer_status(self, req_id: Any, **kwargs) -> Any: + """Queries status of an ongoing weight transfer or KV-cache migration.""" + sampler = self._require_sampler() + if hasattr(sampler, "get_transfer_status"): + return await sampler.get_transfer_status(req_id, **kwargs) + return "UNKNOWN" + + async def get_load_info(self, **kwargs) -> base_sampler_lib.LoadInfo: + """Returns load information from the underlying engine.""" + info = await self._require_sampler().get_load_info(**kwargs) + return base_sampler_lib.LoadInfo( + num_requests_waiting=getattr(info, "num_requests_waiting", 0), + num_requests_running=getattr(info, "num_requests_running", 0), + kv_cache_usage_perc=getattr(info, "kv_cache_usage_perc", 0.0), + ) + + async def migrate_kv_cache( + self, + source_server_id: str, + target_server_id: str, + token_ids: List[int], + **kwargs, + ) -> bool: + """Triggers KV-cache transfer across TPU slices.""" + sampler = self._require_sampler() + if hasattr(sampler, "migrate_kv_cache"): + return await sampler.migrate_kv_cache( + route_key=kwargs.get("route_key", ""), + source_server_id=source_server_id, + target_server_id=target_server_id, + token_ids=token_ids, + ) + return False diff --git a/tunix/experimental/weight_sync/raiden_synchronizer.py b/tunix/experimental/weight_sync/raiden_synchronizer.py index ebe89c51a..ea5763a72 100644 --- a/tunix/experimental/weight_sync/raiden_synchronizer.py +++ b/tunix/experimental/weight_sync/raiden_synchronizer.py @@ -25,6 +25,23 @@ import jax.numpy as jnp from tunix.experimental.weight_sync import weight_sync + +def _log_rss(tag: str) -> None: + """Logs process peak RSS (GB) -- pinpoints which bind() stage spikes host + memory, since ru_maxrss is a high-water mark that only grows. + + Off by default; enable with --v=1. `resource` is imported lazily because it + is Unix-only. + """ + if not logging.vlog_is_on(1): + return + import resource # pylint: disable=g-import-not-at-top + + rss_gb = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1e6 + logging.vlog( + 1, "raiden bind rss checkpoint [%s]: %.1f GB (peak)", tag, rss_gb + ) + _ws_lib: Any = None try: from tpu_sync.api.jax import weight_synchronizer as _ws_lib # pytype: disable=import-error pylint: disable=g-import-not-at-top @@ -51,16 +68,33 @@ def local_ip() -> str: def to_host_cpu_state(state: Any) -> Any: - """Pulls arrays to client host memory; proxy arrays cannot bind directly.""" - cpu = jax.local_devices(backend="cpu")[0] + """Pulls arrays to client host memory; proxy arrays cannot bind directly. - def pull(leaf): + TODO: drop the explicit source-buffer deletion below once the FFI path is + enabled -- FFI releases the transit buffers itself. + """ + cpu = jax.local_devices(backend="cpu")[0] + leaves, treedef = jax.tree_util.tree_flatten(state) + del state + new_leaves = [] + for i in range(len(leaves)): + leaf = leaves[i] + leaves[i] = None arr = getattr(leaf, "value", leaf) if hasattr(arr, "shape") and hasattr(arr, "dtype"): - return jax.device_put(jax.device_get(arr), cpu) - return leaf - - return jax.tree_util.tree_map(pull, state) + new_leaves.append(jax.device_put(jax.device_get(arr), cpu)) + # Release the Pathways proxy transit buffer immediately; deleting is far + # cheaper than the gc.collect() sweep this replaces. + if hasattr(arr, "delete"): + try: + arr.delete() + except Exception: # pylint: disable=broad-except + pass + else: + new_leaves.append(leaf) + del leaf, arr + del leaves + return jax.tree_util.tree_unflatten(treedef, new_leaves) def flatten_weights(state: Any) -> Tuple[List[str], List[Any]]: @@ -75,22 +109,33 @@ def flatten_weights(state: Any) -> Tuple[List[str], List[Any]]: def _bindable(arr: Any) -> bool: - """True if the native layer can bind this leaf.""" + """True if the native layer can bind this leaf. + + Binding an unsupported leaf (e.g. RNG keys) can SIGSEGV, so only + floating-point, rank>=1, fully TPU- or CPU-resident arrays qualify. CPU is + allowed because host_stage deliberately copies proxy-backed (Pathways) + arrays to host CPU memory before bind() gets here -- rejecting "not TPU" + would drop every leaf it just staged. + """ try: + if not hasattr(arr, "shape") or not hasattr(arr, "dtype"): + return False + if arr.ndim < 1: + return False + if not jnp.issubdtype(arr.dtype, jnp.floating): + return False devices = arr.devices() - except AttributeError: + if not devices: + return False + return all(getattr(d, "platform", "?") in ("tpu", "cpu") for d in devices) + except Exception: return False - on_local_hw = all( - getattr(d, "platform", "?") in ("cpu", "tpu") for d in devices - ) - return on_local_hw and jnp.issubdtype(arr.dtype, jnp.number) def _filter_bindable( names: List[str], arrays: List[Any] ) -> Tuple[List[str], List[Any]]: - """Drops leaves the native layer cannot bind; binding them is undefined - behavior (observed: random RuntimeError or SIGSEGV on RNG-key arrays).""" + """Drops leaves _bindable rejects.""" logging.vlog( 1, "raiden bind census: %s", @@ -103,6 +148,12 @@ def _filter_bindable( dropped = [] for name, arr in zip(names, arrays): if _bindable(arr): + if hasattr(arr, "block_until_ready"): + try: + # binding an in-flight buffer is part of what SIGSEGVs + arr.block_until_ready() + except Exception: + pass keep_names.append(name) keep_arrays.append(arr) else: @@ -184,14 +235,18 @@ def active(self) -> bool: return self._sync is not None def bind(self, state: Any) -> None: - """Binds this host's weights, or rebinds them after a training step. - - With host_stage the arrays are copied to local CPU memory first; arrays - backed by the pathways proxy cannot bind in place. - """ + """Binds or rebinds weights to the Raiden transport.""" + _log_rss("bind:start") + # Clear previous buffers before staging to avoid holding duplicate weight + # copies in host memory during rebinds. + self.names = [] + self.arrays = [] if self._host_stage: state = to_host_cpu_state(state) + _log_rss("bind:after_host_stage") self.names, self.arrays = _filter_bindable(*flatten_weights(state)) + del state + _log_rss("bind:after_flatten") if _ws_lib is None: return if self._sync is None: @@ -206,8 +261,10 @@ def bind(self, state: Any) -> None: bind_ip=None, auto_h2d=self._auto_h2d, ) + _log_rss("bind:after_native_construct") else: self._sync.bind_weights(self.arrays) + _log_rss("bind:after_native_rebind") def _require_sync(self, op: str) -> Any: if self._sync is None: @@ -225,6 +282,11 @@ def h2d(self) -> None: raise RuntimeError(f"{self.job_name}: bind() must run before h2d()") if self._sync is not None: self._sync.h2d() + jax.block_until_ready(self.arrays) + + def release_host_arrays(self) -> None: + """Drops host-staged array references to reclaim memory between sync rounds.""" + self.arrays = [] def metrics(self) -> dict: return self._sync.get_metrics() if self._sync else {} diff --git a/tunix/experimental/weight_sync/weight_sync.py b/tunix/experimental/weight_sync/weight_sync.py index e32890a10..d3d178bf1 100644 --- a/tunix/experimental/weight_sync/weight_sync.py +++ b/tunix/experimental/weight_sync/weight_sync.py @@ -187,6 +187,79 @@ class WorkUnitMetadata: variables: tuple[TensorMetadata, ...] = () mesh_axes: Optional[tuple[str, ...]] = None + @classmethod + def from_dict(cls, d: Any) -> WorkUnitMetadata: + """Reconstructs WorkUnitMetadata from a dictionary or returns metadata directly.""" + if isinstance(d, cls): + return d + if not isinstance(d, dict): + raise TypeError(f"Expected WorkUnitMetadata or dict, got {type(d)}") + + unit_raw = d.get("unit") + if isinstance(unit_raw, dict): + unit = WorkUnitId(**unit_raw) + elif isinstance(unit_raw, WorkUnitId): + unit = unit_raw + else: + unit = WorkUnitId(job_name=str(unit_raw or "destination")) + + variables_raw = d.get("variables", ()) + variables = [] + for v in variables_raw: + if isinstance(v, TensorMetadata): + variables.append(v) + elif isinstance(v, dict): + variables.append( + TensorMetadata( + name=v["name"], + shape=tuple(v["shape"]), + mesh_shape=tuple(v["mesh_shape"]), + layout=tuple(v["layout"]), + item_size=int(v["item_size"]), + layer_idx=int(v.get("layer_idx", 0)), + sharding_spec=tuple(v.get("sharding_spec", ())), + ) + ) + elif hasattr(v, "name"): + variables.append( + TensorMetadata( + name=v.name, + shape=tuple(v.shape), + mesh_shape=tuple(v.mesh_shape), + layout=tuple(v.layout), + item_size=int(v.item_size), + layer_idx=int(getattr(v, "layer_idx", 0)), + sharding_spec=tuple(getattr(v, "sharding_spec", ())), + ) + ) + + return cls( + unit=unit, + shards=tuple(d.get("shards", ())), + control_plane_rpc_address=str(d.get("control_plane_rpc_address", "")), + global_shape=( + tuple(d["global_shape"]) + if d.get("global_shape") is not None + else None + ), + mesh_shape=( + tuple(d["mesh_shape"]) if d.get("mesh_shape") is not None else None + ), + layout=tuple(d["layout"]) if d.get("layout") is not None else None, + item_size=( + int(d["item_size"]) if d.get("item_size") is not None else None + ), + variables=tuple(variables), + mesh_axes=( + tuple(d["mesh_axes"]) if d.get("mesh_axes") is not None else None + ), + ) + + +def dict_to_metadata(d: Any) -> WorkUnitMetadata: + """Reconstructs WorkUnitMetadata from a dictionary (delegates to WorkUnitMetadata.from_dict).""" + return WorkUnitMetadata.from_dict(d) + @dataclasses.dataclass(frozen=True) class TransferResult: diff --git a/tunix/experimental/weight_sync/weight_sync_coordinator.py b/tunix/experimental/weight_sync/weight_sync_coordinator.py index 38894d41b..2532a93a6 100644 --- a/tunix/experimental/weight_sync/weight_sync_coordinator.py +++ b/tunix/experimental/weight_sync/weight_sync_coordinator.py @@ -991,8 +991,16 @@ async def record_workers(final_error: str = "") -> None: " quiesced; no rollback needed" ) from e - src_metadata = [m for per_source in src_meta_lists for m in per_source] - dst_metadata = [m for per_dest in dst_meta_lists for m in per_dest] + src_metadata = [ + weight_sync.dict_to_metadata(m) + for per_source in src_meta_lists + for m in per_source + ] + dst_metadata = [ + weight_sync.dict_to_metadata(m) + for per_dest in dst_meta_lists + for m in per_dest + ] if not src_metadata or not dst_metadata: failures.append( f"metadata: {len(src_metadata)} source, {len(dst_metadata)}" diff --git a/tunix/experimental/worker/rollout_worker.py b/tunix/experimental/worker/rollout_worker.py index 856a5cd6a..824806503 100644 --- a/tunix/experimental/worker/rollout_worker.py +++ b/tunix/experimental/worker/rollout_worker.py @@ -347,7 +347,7 @@ def _sampling_to_rollout_response( if prompt_token_arr.size == 0: raise RuntimeError( "Sampler response is missing prompt_token_ids for " - f"{request.request_id}." + f"request_id={request.request_id!r} traj_id={request.traj_id!r}." ) metadata = dict(request.metadata or {}) metadata.setdefault("text", text)