From 99bae1d2f9454aa203bdf871eb30502d5f89b7e3 Mon Sep 17 00:00:00 2001 From: Lance Wang Date: Sun, 30 Aug 2026 15:14:33 -0700 Subject: [PATCH] [DO NOT SUBMIT] Enable ffi for trainer. PiperOrigin-RevId: 973552747 --- .../runtime/contexts/borg_context_test.py | 75 ++++ .../runtime/contexts/context_factory_test.py | 52 +++ .../runtime/executors/executor_test.py | 19 + .../math_gsm8k_dist/run_trainer_node_test.py | 57 +-- .../train/peft_trainer_v2_weight_sync_test.py | 12 +- .../weight_sync/raiden_synchronizer_test.py | 17 +- .../distributed/runtime/context.py | 1 + .../runtime/contexts/borg_context.py | 175 +++++++++ .../runtime/contexts/context_factory.py | 54 +++ .../runtime/discovery/discovery.py | 5 +- .../runtime/executors/borg_executor.py | 41 ++ .../math_gsm8k_dist/process_context.py | 0 .../math_gsm8k_dist/run_gsm8k_dist_grpo.py | 55 +-- .../math_gsm8k_dist/run_rollout_node.py | 77 +++- .../math_gsm8k_dist/run_trainer_node.py | 193 +++------- .../examples/math_gsm8k_dist/xm_launch.py | 355 ++++++++++++++++++ tunix/experimental/train/peft_trainer_v2.py | 3 +- .../weight_sync/raiden_synchronizer.py | 166 ++++++-- tunix/experimental/worker/remote_execution.py | 7 +- tunix/oss/utils.py | 8 +- 20 files changed, 1107 insertions(+), 265 deletions(-) create mode 100644 tests/experimental/distributed/runtime/contexts/borg_context_test.py create mode 100644 tests/experimental/distributed/runtime/contexts/context_factory_test.py create mode 100644 tunix/experimental/distributed/runtime/contexts/borg_context.py create mode 100644 tunix/experimental/distributed/runtime/contexts/context_factory.py create mode 100644 tunix/experimental/distributed/runtime/executors/borg_executor.py create mode 100644 tunix/experimental/examples/math_gsm8k_dist/process_context.py create mode 100644 tunix/experimental/examples/math_gsm8k_dist/xm_launch.py diff --git a/tests/experimental/distributed/runtime/contexts/borg_context_test.py b/tests/experimental/distributed/runtime/contexts/borg_context_test.py new file mode 100644 index 000000000..6e23b149a --- /dev/null +++ b/tests/experimental/distributed/runtime/contexts/borg_context_test.py @@ -0,0 +1,75 @@ +# 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. + +"""Unit tests for Borg and direct execution distributed runtime contexts.""" + +import argparse +from unittest import mock + +from absl.testing import absltest +import portpicker +from tunix.experimental.distributed.runtime import context +from tunix.experimental.distributed.runtime.contexts import borg_context + + +class BorgContextTest(absltest.TestCase): + + def test_resolve_local_ip(self): + ip = borg_context.resolve_local_ip() + self.assertIsInstance(ip, str) + self.assertTrue(len(ip) > 0) + + @mock.patch( + "tunix.experimental.distributed.runtime.discovery.discovery.grpc.server" + ) + def test_borg_discovery_context_lifecycle_and_registration( + self, mock_grpc_server + ): + port = portpicker.pick_unused_port() + args = argparse.Namespace( + discovery_port=port, + discovery_addrs="10.0.0.1:9999", + ) + + with borg_context.BorgDiscoveryContext(args) as disc_ctx: + cb = mock.MagicMock() + disc_ctx.on_register(cb) + self.assertTrue(disc_ctx._server.is_started()) + + with mock.patch( + "tunix.experimental.distributed.runtime.contexts.borg_context.discovery.register" + ) as mock_reg, mock.patch( + "tunix.experimental.distributed.runtime.contexts.borg_context.resolve_local_ip", + return_value="10.0.0.2", + ): + disc_ctx.register(b"my-metadata") + mock_reg.assert_called_once_with( + "10.0.0.1:9999", "10.0.0.2", port, b"my-metadata" + ) + + self.assertFalse(disc_ctx._server.is_started()) + + def test_borg_process_context(self): + args = argparse.Namespace( + discovery_port=12345, + discovery_addrs="10.0.0.1:12345", + ) + with borg_context.BorgProcessContext(args) as proc_ctx: + self.assertIsInstance(proc_ctx.jax, context.JaxContext) + self.assertIsInstance(proc_ctx.ipc, context.IpcContext) + self.assertIsInstance(proc_ctx.ipc.discovery, context.DiscoveryContext) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/distributed/runtime/contexts/context_factory_test.py b/tests/experimental/distributed/runtime/contexts/context_factory_test.py new file mode 100644 index 000000000..6b477c4f9 --- /dev/null +++ b/tests/experimental/distributed/runtime/contexts/context_factory_test.py @@ -0,0 +1,52 @@ +# 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. + +"""Unit tests for context_factory.""" + +import argparse +import os +from unittest import mock + +from absl.testing import absltest +from tunix.experimental.distributed.runtime.contexts import borg_context +from tunix.experimental.distributed.runtime.contexts import context_factory +from tunix.experimental.distributed.runtime.contexts import k8s_context +from tunix.experimental.distributed.runtime.contexts import local_context + + +class ContextFactoryTest(absltest.TestCase): + + def test_get_default_process_context_local(self): + args = argparse.Namespace(discovery_port=12345, discovery_addrs="") + with mock.patch.dict(os.environ, {}, clear=True): + ctx = context_factory.get_default_process_context(args) + self.assertIsInstance(ctx, local_context.LocalProcessContext) + + def test_get_default_process_context_borg(self): + args = argparse.Namespace(discovery_port=12345, discovery_addrs="") + with mock.patch.dict(os.environ, {"BORG_TASK_HANDLE": "12345"}, clear=True): + ctx = context_factory.get_default_process_context(args) + self.assertIsInstance(ctx, borg_context.BorgProcessContext) + + def test_get_default_process_context_k8s(self): + args = argparse.Namespace(discovery_port=12345, discovery_addrs="") + with mock.patch.dict( + os.environ, {"KUBERNETES_SERVICE_HOST": "10.0.0.1"}, clear=True + ): + ctx = context_factory.get_default_process_context(args) + self.assertIsInstance(ctx, k8s_context.K8sProcessContext) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/distributed/runtime/executors/executor_test.py b/tests/experimental/distributed/runtime/executors/executor_test.py index 56b0cc2ff..1d2c21ddc 100644 --- a/tests/experimental/distributed/runtime/executors/executor_test.py +++ b/tests/experimental/distributed/runtime/executors/executor_test.py @@ -19,8 +19,10 @@ from absl.testing import absltest import portpicker +from tunix.experimental.distributed.runtime.contexts import borg_context from tunix.experimental.distributed.runtime.contexts import k8s_context from tunix.experimental.distributed.runtime.contexts import local_context +from tunix.experimental.distributed.runtime.executors import borg_executor from tunix.experimental.distributed.runtime.executors import k8s_executor from tunix.experimental.distributed.runtime.executors import local_executor @@ -61,6 +63,23 @@ def main_fn(argv, ctx): self.assertEqual(received["argv"], ["--bar=baz"]) self.assertIsInstance(received["ctx"], k8s_context.K8sProcessContext) + def test_borg_executor_run(self): + executor = borg_executor.BorgExecutor() + args = argparse.Namespace( + discovery_port=portpicker.pick_unused_port(), + discovery_addrs="10.0.0.1:1234", + ) + + received = {} + + def main_fn(argv, ctx): + received["argv"] = argv + received["ctx"] = ctx + + executor.run(main_fn, ["--qux=quux"], args) + self.assertEqual(received["argv"], ["--qux=quux"]) + self.assertIsInstance(received["ctx"], borg_context.BorgProcessContext) + if __name__ == "__main__": absltest.main() diff --git a/tests/experimental/examples/math_gsm8k_dist/run_trainer_node_test.py b/tests/experimental/examples/math_gsm8k_dist/run_trainer_node_test.py index 9a1babb79..b7ac8784d 100644 --- a/tests/experimental/examples/math_gsm8k_dist/run_trainer_node_test.py +++ b/tests/experimental/examples/math_gsm8k_dist/run_trainer_node_test.py @@ -160,7 +160,7 @@ def _setup_main_mocks( mock_ensure_model_dir, mock_create_mesh, mock_load_actor_model, - mock_create_trainer_factory, + mock_create_trainer, mock_trainer_worker_cls, mock_grpc_server_cls, ): @@ -168,7 +168,7 @@ def _setup_main_mocks( mock_mesh = mock.MagicMock(spec=Mesh) mock_create_mesh.return_value = mock_mesh mock_load_actor_model.return_value = mock.MagicMock() - mock_create_trainer_factory.return_value = mock.MagicMock() + mock_create_trainer.return_value = mock.MagicMock() mock_trainer_worker_cls.return_value = self.mock_worker_service mock_grpc_server_cls.return_value = self.mock_server @@ -192,14 +192,14 @@ def _patch_signal_handlers(self, handler_fn): @mock.patch.object(run_trainer_node, "_ensure_model_dir_for_trainer") @mock.patch.object(run_trainer_node, "_create_mesh") @mock.patch.object(run_trainer_node, "_load_actor_model") - @mock.patch.object(run_trainer_node, "_create_trainer_factory") + @mock.patch.object(run_trainer_node, "_create_trainer") @mock.patch.object(trainer_worker, "TrainerWorker") @mock.patch.object(remote_execution, "GrpcRemoteExecutionServer") def test_shutdown_handler_drains_worker_on_sigterm( self, mock_grpc_server_cls, mock_trainer_worker_cls, - mock_create_trainer_factory, + mock_create_trainer, mock_load_actor_model, mock_create_mesh, mock_ensure_model_dir, @@ -208,7 +208,7 @@ def test_shutdown_handler_drains_worker_on_sigterm( mock_ensure_model_dir, mock_create_mesh, mock_load_actor_model, - mock_create_trainer_factory, + mock_create_trainer, mock_trainer_worker_cls, mock_grpc_server_cls, ) @@ -237,14 +237,14 @@ def mock_add_signal_handler(loop_self, sig, callback): @mock.patch.object(run_trainer_node, "_ensure_model_dir_for_trainer") @mock.patch.object(run_trainer_node, "_create_mesh") @mock.patch.object(run_trainer_node, "_load_actor_model") - @mock.patch.object(run_trainer_node, "_create_trainer_factory") + @mock.patch.object(run_trainer_node, "_create_trainer") @mock.patch.object(trainer_worker, "TrainerWorker") @mock.patch.object(remote_execution, "GrpcRemoteExecutionServer") def test_shutdown_handler_drains_worker_on_sigint( self, mock_grpc_server_cls, mock_trainer_worker_cls, - mock_create_trainer_factory, + mock_create_trainer, mock_load_actor_model, mock_create_mesh, mock_ensure_model_dir, @@ -253,7 +253,7 @@ def test_shutdown_handler_drains_worker_on_sigint( mock_ensure_model_dir, mock_create_mesh, mock_load_actor_model, - mock_create_trainer_factory, + mock_create_trainer, mock_trainer_worker_cls, mock_grpc_server_cls, ) @@ -279,14 +279,14 @@ def mock_add_signal_handler(loop_self, sig, callback): @mock.patch.object(run_trainer_node, "_ensure_model_dir_for_trainer") @mock.patch.object(run_trainer_node, "_create_mesh") @mock.patch.object(run_trainer_node, "_load_actor_model") - @mock.patch.object(run_trainer_node, "_create_trainer_factory") + @mock.patch.object(run_trainer_node, "_create_trainer") @mock.patch.object(trainer_worker, "TrainerWorker") @mock.patch.object(remote_execution, "GrpcRemoteExecutionServer") def test_shutdown_handler_handles_drain_exception_and_stops_server( self, mock_grpc_server_cls, mock_trainer_worker_cls, - mock_create_trainer_factory, + mock_create_trainer, mock_load_actor_model, mock_create_mesh, mock_ensure_model_dir, @@ -295,7 +295,7 @@ def test_shutdown_handler_handles_drain_exception_and_stops_server( mock_ensure_model_dir, mock_create_mesh, mock_load_actor_model, - mock_create_trainer_factory, + mock_create_trainer, mock_trainer_worker_cls, mock_grpc_server_cls, ) @@ -319,14 +319,14 @@ def mock_add_signal_handler(loop_self, sig, callback): @mock.patch.object(run_trainer_node, "_ensure_model_dir_for_trainer") @mock.patch.object(run_trainer_node, "_create_mesh") @mock.patch.object(run_trainer_node, "_load_actor_model") - @mock.patch.object(run_trainer_node, "_create_trainer_factory") + @mock.patch.object(run_trainer_node, "_create_trainer") @mock.patch.object(trainer_worker, "TrainerWorker") @mock.patch.object(remote_execution, "GrpcRemoteExecutionServer") def test_shutdown_ignores_not_implemented_error_on_add_signal_handler( self, mock_grpc_server_cls, mock_trainer_worker_cls, - mock_create_trainer_factory, + mock_create_trainer, mock_load_actor_model, mock_create_mesh, mock_ensure_model_dir, @@ -335,7 +335,7 @@ def test_shutdown_ignores_not_implemented_error_on_add_signal_handler( mock_ensure_model_dir, mock_create_mesh, mock_load_actor_model, - mock_create_trainer_factory, + mock_create_trainer, mock_trainer_worker_cls, mock_grpc_server_cls, ) @@ -367,26 +367,7 @@ def mock_add_signal_handler(loop_self, sig, callback): self.mock_worker_service.stop.assert_called_once() self.mock_server.stop_serving.assert_called_once() - @mock.patch.object(run_trainer_node, "_create_tunix_trainer_factory") - @mock.patch.object(run_trainer_node, "_create_maxtext_trainer_factory") - def test_create_trainer_factory_delegates_by_backend( - self, mock_create_maxtext, mock_create_tunix - ): - args_tunix = mock.MagicMock(trainer_backend="tunix") - run_trainer_node._create_trainer_factory(args_tunix) - mock_create_tunix.assert_called_once_with(args_tunix) - mock_create_maxtext.assert_not_called() - - mock_create_tunix.reset_mock() - args_maxtext = mock.MagicMock(trainer_backend="maxtext") - run_trainer_node._create_trainer_factory(args_maxtext) - mock_create_maxtext.assert_called_once_with(args_maxtext) - mock_create_tunix.assert_not_called() - def test_main_raises_without_discovery_context(self): - with self.assertRaisesRegex(RuntimeError, "Require discovery API"): - run_trainer_node.main([], context=None) - with self.assertRaisesRegex(RuntimeError, "Require discovery API"): run_trainer_node.main([], context=mock.MagicMock(ipc=None)) @@ -466,19 +447,19 @@ def test_create_mesh_validates_device_count(self): def test_has_direct_safetensors(self): with tempfile.TemporaryDirectory() as tmp_dir: p = Path(tmp_dir) - self.assertFalse(run_trainer_node._has_direct_safetensors(p)) + self.assertFalse(models._has_direct_safetensors(p)) (p / "model.safetensors").touch() - self.assertTrue(run_trainer_node._has_direct_safetensors(p)) + self.assertTrue(models._has_direct_safetensors(p)) def test_ensure_model_dir_raises_for_empty_or_file(self): with self.assertRaisesRegex(ValueError, "--model_dir is required"): - run_trainer_node._ensure_model_dir_for_trainer("", "Qwen/Qwen3-1.7B") + models.ensure_model_dir("", "Qwen/Qwen3-1.7B") with tempfile.NamedTemporaryFile() as tmp_file: with self.assertRaisesRegex( - ValueError, "--model_dir must point to an existing local directory" + ValueError, "--model_dir must point to a directory" ): - run_trainer_node._ensure_model_dir_for_trainer( + models.ensure_model_dir( tmp_file.name, "Qwen/Qwen3-1.7B" ) diff --git a/tests/experimental/train/peft_trainer_v2_weight_sync_test.py b/tests/experimental/train/peft_trainer_v2_weight_sync_test.py index ead661a38..7dc35b886 100644 --- a/tests/experimental/train/peft_trainer_v2_weight_sync_test.py +++ b/tests/experimental/train/peft_trainer_v2_weight_sync_test.py @@ -28,10 +28,10 @@ class _FakeSynchronizer: - def __init__(self, job_name, state=None, host_stage=False, **kwargs): + def __init__(self, job_name, state=None, use_ffi=False, **kwargs): self.job_name = job_name self.state = state - self.host_stage = host_stage + self.use_ffi = use_ffi self.kwargs = kwargs self.bound_state = None self.d2h_calls = 0 @@ -89,21 +89,21 @@ def test_prepare_reuses_the_worker(self): peft_trainer_v2.PeftTrainer.prepare_weight_sync(fake) self.assertIs(fake._weight_sync_worker, first) - def test_prepare_host_stages_under_proxy(self): + def test_prepare_uses_ffi_under_proxy(self): fake = self._fake_trainer() with mock.patch.object(raiden_synchronizer, "RaidenSynchronizer", _FakeSynchronizer): with mock.patch.dict(os.environ, {"JAX_PLATFORMS": "proxy,cpu"}): peft_trainer_v2.PeftTrainer.prepare_weight_sync(fake) - self.assertTrue(fake._weight_sync_worker.host_stage) + self.assertTrue(fake._weight_sync_worker.use_ffi) def test_prepare_uses_the_injected_factory(self): fake = self._fake_trainer() fake._weight_sync_worker_factory = lambda: _FakeSynchronizer( - "trainer", host_stage=False + "trainer", use_ffi=False ) with mock.patch.dict(os.environ, {"JAX_PLATFORMS": "proxy,cpu"}): peft_trainer_v2.PeftTrainer.prepare_weight_sync(fake) - self.assertFalse(fake._weight_sync_worker.host_stage) + self.assertFalse(fake._weight_sync_worker.use_ffi) def test_release_without_prepare_is_a_no_op(self): fake = self._fake_trainer() diff --git a/tests/experimental/weight_sync/raiden_synchronizer_test.py b/tests/experimental/weight_sync/raiden_synchronizer_test.py index 156d8c533..036c84c4d 100644 --- a/tests/experimental/weight_sync/raiden_synchronizer_test.py +++ b/tests/experimental/weight_sync/raiden_synchronizer_test.py @@ -284,15 +284,14 @@ def test_worker_index_stamps_replica_id(self): base = raiden_synchronizer.RaidenSynchronizer("rollout", self._state()) self.assertEqual(base.work_unit_metadata().unit.job_replica_id, "") - def test_host_stage_pulls_state_to_host(self): - sentinel = {"w": jnp.ones((2, 2))} - with mock.patch.object( - raiden_synchronizer, "to_host_cpu_state", return_value=sentinel - ) as pull: - raiden_synchronizer.RaidenSynchronizer( - "trainer", self._state(), host_stage=True - ) - pull.assert_called_once() + def test_ffi_mode_defaults_under_pathways(self): + import os + with mock.patch.dict(os.environ, {"JAX_PLATFORMS": "proxy,cpu"}): + sync = raiden_synchronizer.RaidenSynchronizer("trainer") + self.assertTrue(sync.use_ffi) + with mock.patch.dict(os.environ, {"JAX_PLATFORMS": "tpu"}): + sync = raiden_synchronizer.RaidenSynchronizer("trainer") + self.assertFalse(sync.use_ffi) if __name__ == "__main__": diff --git a/tunix/experimental/distributed/runtime/context.py b/tunix/experimental/distributed/runtime/context.py index 628dae263..d9e9f50ad 100644 --- a/tunix/experimental/distributed/runtime/context.py +++ b/tunix/experimental/distributed/runtime/context.py @@ -80,3 +80,4 @@ def jax(self) -> JaxContext: def ipc(self) -> IpcContext: """Returns the IPC context for this process.""" return IpcContext() + diff --git a/tunix/experimental/distributed/runtime/contexts/borg_context.py b/tunix/experimental/distributed/runtime/contexts/borg_context.py new file mode 100644 index 000000000..a3250f7c1 --- /dev/null +++ b/tunix/experimental/distributed/runtime/contexts/borg_context.py @@ -0,0 +1,175 @@ +# 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. + +"""Borg and direct multi-host execution distributed runtime context implementations.""" + +import argparse +import logging +import os +import socket +from typing import Any, Callable + +from tunix.experimental.distributed.runtime import context +from tunix.experimental.distributed.runtime.discovery import discovery + + +def resolve_local_ip() -> str: + """Returns the local IP address for this host.""" + for family, target in [ + (socket.AF_INET6, ("2001:4860:4860::8888", 80)), + (socket.AF_INET, ("8.8.8.8", 80)), + ]: + try: + s = socket.socket(family, socket.SOCK_DGRAM) + try: + s.connect(target) + return s.getsockname()[0] + finally: + s.close() + except Exception: + pass + try: + return socket.gethostbyname(socket.gethostname()) + except Exception: + return "127.0.0.1" + + +class BorgJaxContext(context.JaxContext): + """JAX distributed runtime initializer for Borg and direct execution.""" + + def initialize(self) -> None: + """Initializes Pathways or standard multi-controller JAX runtime.""" + if "proxy" in os.environ.get("JAX_PLATFORMS", "") and os.environ.get( + "JAX_BACKEND_TARGET" + ): + logging.info("Initializing Pathways runtime...") + try: + import pathwaysutils # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] + + pathwaysutils.initialize() + except ImportError: + pass + else: + logging.info("Initializing multi-controller JAX runtime...") + try: + import jax # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import] + + jax.distributed.initialize() + except Exception: + pass + + +class BorgDiscoveryContext(context.DiscoveryContext): + """Borg discovery context managing registration and server hosting.""" + + def __init__(self, args: argparse.Namespace) -> None: + """Initializes the Borg discovery context.""" + self._args = args + self._server = discovery.DiscoveryServer() + + def __enter__(self) -> "BorgDiscoveryContext": + """Enters the discovery context manager scope.""" + return self + + def __exit__( + self, + exc_type: Any | None, + exc: Any | None, + tb: Any | None, + ) -> None: + """Stops the discovery server if started.""" + if self._server.is_started(): + self._server.stop() + logging.info("Discovery server stopped.") + + def on_register(self, callback: Callable[[str, int, bytes], None]) -> None: + """Starts the discovery server on the configured port.""" + discovery_port = getattr(self._args, "discovery_port", 0) + if discovery_port: + self._server.start(discovery_port, callback) + logging.info("Discovery server started on port %s", discovery_port) + + def register(self, metadata: bytes) -> None: + """Registers this process with the remote discovery server.""" + discovery_addrs = getattr(self._args, "discovery_addrs", "") + if not discovery_addrs: + raise ValueError( + "discovery_addrs must be non-empty. Did you set --discovery_addrs?" + ) + + hostname = resolve_local_ip() + port = getattr(self._args, "port", 0) or getattr(self._args, "discovery_port", 0) or 0 + logging.info("Registering to discovery server at %s from host %s port %d", discovery_addrs, hostname, port) + discovery.register(discovery_addrs, hostname, port, metadata) + logging.info("Registered to discovery server at %s", discovery_addrs) + + +class BorgIpcContext(context.IpcContext): + """Borg inter-process communication context.""" + + def __init__(self, args: argparse.Namespace) -> None: + """Initializes the Borg IPC context.""" + self._discovery = BorgDiscoveryContext(args) + + def __enter__(self) -> "BorgIpcContext": + """Enters the IPC context manager scope.""" + self._discovery.__enter__() + return self + + def __exit__( + self, + exc_type: Any | None, + exc: Any | None, + tb: Any | None, + ) -> None: + """Exits the IPC context manager scope.""" + self._discovery.__exit__(exc_type, exc, tb) + + @property + def discovery(self) -> context.DiscoveryContext: + """Returns the Borg discovery context.""" + return self._discovery + + +class BorgProcessContext(context.ProcessContext): + """Process context for Borg and direct multi-host execution.""" + + def __init__(self, args: argparse.Namespace) -> None: + """Initializes the Borg process context.""" + self._jax = BorgJaxContext() + self._ipc = BorgIpcContext(args) + + def __enter__(self) -> "BorgProcessContext": + """Enters the process context manager scope.""" + self._ipc.__enter__() + return self + + def __exit__( + self, + exc_type: Any | None, + exc: Any | None, + tb: Any | None, + ) -> None: + """Exits the process context manager scope.""" + self._ipc.__exit__(exc_type, exc, tb) + + @property + def jax(self) -> context.JaxContext: + """Returns the JAX runtime context.""" + return self._jax + + @property + def ipc(self) -> context.IpcContext: + """Returns the Borg IPC context.""" + return self._ipc diff --git a/tunix/experimental/distributed/runtime/contexts/context_factory.py b/tunix/experimental/distributed/runtime/contexts/context_factory.py new file mode 100644 index 000000000..f3bb9f951 --- /dev/null +++ b/tunix/experimental/distributed/runtime/contexts/context_factory.py @@ -0,0 +1,54 @@ +# 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. + +"""Factory function to dynamically construct the appropriate ProcessContext.""" + +import os +from typing import Any + +from tunix.experimental.distributed.runtime.context import ProcessContext +from tunix.experimental.distributed.runtime.contexts.borg_context import BorgProcessContext +from tunix.experimental.distributed.runtime.contexts.k8s_context import K8sProcessContext +from tunix.experimental.distributed.runtime.contexts.local_context import LocalProcessContext + + +def get_default_process_context(args: Any) -> ProcessContext: + """Automatically creates the appropriate ProcessContext based on the runtime environment. + + If running under Borg / XManager, returns BorgProcessContext. + If running under Kubernetes JobSet, returns K8sProcessContext. + Otherwise, returns LocalProcessContext. + + Args: + args: Command line or parsed namespace arguments. + + Returns: + An instance of ProcessContext suitable for the detected platform. + """ + if ( + os.getenv("BORG_TASK_HANDLE") + or os.getenv("BORG_JOB_NAME") + or os.getenv("BORG_ALLOC_DIR") + or os.getenv("XM_BORG_MODE") == "true" + ): + return BorgProcessContext(args) + + if ( + os.getenv("KUBERNETES_SERVICE_HOST") + or os.getenv("JOBSET_NAME") + or os.getenv("POD_NAME") + ): + return K8sProcessContext(args) + + return LocalProcessContext(args) diff --git a/tunix/experimental/distributed/runtime/discovery/discovery.py b/tunix/experimental/distributed/runtime/discovery/discovery.py index b47a6b0e9..aa5841304 100644 --- a/tunix/experimental/distributed/runtime/discovery/discovery.py +++ b/tunix/experimental/distributed/runtime/discovery/discovery.py @@ -66,7 +66,10 @@ def Register( pb2_grpc.add_DiscoveryServiceServicer_to_server(_handler(), server) # start server - server.add_insecure_port(f"[::]:{port}") + try: + server.add_insecure_port(f"[::]:{port}") + except Exception: + server.add_insecure_port(f"0.0.0.0:{port}") server.start() self._server = server diff --git a/tunix/experimental/distributed/runtime/executors/borg_executor.py b/tunix/experimental/distributed/runtime/executors/borg_executor.py new file mode 100644 index 000000000..715c77195 --- /dev/null +++ b/tunix/experimental/distributed/runtime/executors/borg_executor.py @@ -0,0 +1,41 @@ +# 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. + +"""Borg and direct execution distributed runtime process executor.""" + +from argparse import Namespace +from typing import Callable, List + +from tunix.experimental.distributed.runtime.context import ProcessContext +from tunix.experimental.distributed.runtime.contexts.borg_context import BorgProcessContext + + +class BorgExecutor: + """Process executor that runs a target main function inside a Borg runtime context.""" + + def run( + self, + process_main: Callable[[List[str], ProcessContext], None], + process_argv: List[str], + context_args: Namespace, + ) -> None: + """Executes `process_main` inside a `BorgProcessContext`. + + Args: + process_main: Callable entrypoint accepting (argv, context). + process_argv: Remaining command-line arguments for the target process. + context_args: Parsed runtime context configuration arguments. + """ + with BorgProcessContext(context_args) as context: + process_main(process_argv, context) diff --git a/tunix/experimental/examples/math_gsm8k_dist/process_context.py b/tunix/experimental/examples/math_gsm8k_dist/process_context.py new file mode 100644 index 000000000..e69de29bb diff --git a/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py b/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py index 71297d7a3..340099533 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py @@ -164,6 +164,21 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: default=os.getenv("WANDB_RUN_NAME", ""), help="W&B run name. Defaults to timestamp-based name if unset.", ) + parser.add_argument("--discovery_id", type=str, default="orch") + parser.add_argument("--discovery_port", type=int, default=20000) + parser.add_argument("--discovery_addrs", type=str, default="") + parser.add_argument("--mini_batch_size", type=int, default=1) + parser.add_argument("--model_name", type=str, default="Qwen3-1.7B") + parser.add_argument( + "--model_dir", + type=str, + default=os.getenv( + "MODEL_DIR", + os.getenv( + "MODEL_DOWNLOAD_DIR", "/tmp/artifacts/qwen3_dist_gsm8k/models" + ), + ), + ) parser.add_argument("--rpc_timeout_s", type=float, default=1800.0) parser.add_argument("--inference_addr", type=str, default="") parser.add_argument("--stop_workers_on_exit", action="store_true") @@ -172,7 +187,8 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: action="store_true", help="Enable debug logging and print full sampler responses.", ) - return parser.parse_args(argv) + args, _ = parser.parse_known_args(argv) + return args def _connect(addr: str, timeout_s: float) -> remote_execution.ActorHandle: @@ -428,20 +444,19 @@ def _iter_prompt_items( def main(argv: list[str], context: Any = None) -> None: - if context and context.ipc and context.ipc.discovery: - pass - else: - raise RuntimeError( - "Require discovery API, but process context doesn't support." + args_list = argv[1:] if argv and argv[0] == sys.argv[0] else argv + args = _parse_args(args_list) + if context is None: + from tunix.experimental.distributed.runtime.contexts import context_factory # pylint: disable=g-import-not-at-top + + context = context_factory.get_default_process_context( + argparse.Namespace( + discovery_id=args.discovery_id, + discovery_port=args.discovery_port, + discovery_addrs=getattr(args, "discovery_addrs", ""), + ) ) - - logging.basicConfig( - level=logging.INFO, - format="%(asctime)s - [Orchestrator] %(message)s", - force=True, - ) - - args = _parse_args(argv) + context.__enter__() if args.num_generations <= 1: raise ValueError("num_generations must be greater than 1 for GRPO.") if args.batch_size <= 0: @@ -483,13 +498,6 @@ def main(argv: list[str], context: Any = None) -> None: tokenizer.pad_token = tokenizer.eos_token pad_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else 0 eos_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else pad_id - logging.info( - "Loaded tokenizer from %s (vocab_size=%d, pad_id=%d, eos_id=%d).", - tokenizer_path, - len(tokenizer), - pad_id, - eos_id, - ) trainer_addr_future = futures.Future() rollout_addr_future = futures.Future() @@ -654,4 +662,7 @@ def accept_worker(hostname: str, _: int, metadata: bytes) -> None: if __name__ == "__main__": - main(sys.argv[1:]) + from absl import app # pylint: disable=g-import-not-at-top + from absl import flags # pylint: disable=g-import-not-at-top + + app.run(main, flags_parser=lambda argv: flags.FLAGS(argv, known_only=True)) 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 229a9eac5..af85817e0 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_rollout_node.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_rollout_node.py @@ -21,6 +21,7 @@ import importlib import logging import os +from pathlib import Path import pickle import sys from typing import Any @@ -48,6 +49,30 @@ } +DEFAULT_MODEL_DOWNLOAD_DIR = os.getenv( + "MODEL_DOWNLOAD_DIR", "/tmp/artifacts/qwen3_dist_gsm8k/models" +) + + +def _has_direct_safetensors(model_path: Path) -> bool: + return model_path.is_dir() and any( + f.name.endswith(".safetensors") for f in model_path.iterdir() + ) + + +def _ensure_model_dir_for_rollout(model_dir: str, model_id: str) -> str: + if not model_dir: + model_dir = DEFAULT_MODEL_DOWNLOAD_DIR + model_path = Path(model_dir) + if _has_direct_safetensors(model_path): + return str(model_path) + model_path.mkdir(parents=True, exist_ok=True) + from tunix.oss import utils as oss_utils # pylint: disable=g-import-not-at-top + + oss_utils.hf_pipeline(model_id, str(model_path)) + return str(model_path) + + def _import_vllm_sampler(): logging.info( "Importing tunix.generate.vllm_sampler before rollout adapters..." @@ -75,7 +100,14 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: parser.add_argument("--worker_id", type=str, default="vllm-rollout-0") parser.add_argument("--model_id", type=str, default="Qwen/Qwen3-1.7B") parser.add_argument( - "--model_dir", type=str, default=os.getenv("MODEL_DIR", "") + "--model_dir", + type=str, + default=os.getenv( + "MODEL_DIR", + os.getenv( + "MODEL_DOWNLOAD_DIR", "/tmp/artifacts/qwen3_dist_gsm8k/models" + ), + ), ) parser.add_argument("--tokenizer_path", type=str, default="") parser.add_argument("--mesh_fsdp", type=int, default=1) @@ -129,7 +161,9 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: choices=list(weight_sync_lib.WeightSyncMode), help="Weight sync mode (none, fallback, or raiden).", ) - return parser.parse_args(argv) + parser.add_argument("--discovery_addrs", type=str, default="") + args, _ = parser.parse_known_args(argv) + return args def _create_rollout_mesh(args) -> Any: @@ -165,10 +199,9 @@ def _create_vanilla_worker(args, tokenizer): logging.info("Creating native sampler on the rollout mesh...") mesh = _create_rollout_mesh(args) + model_dir = _ensure_model_dir_for_rollout(args.model_dir, args.model_id) with mesh: - model = models.create_model( - args.model_name, args.model_dir or args.model_id, mesh - ) + model = models.create_model(args.model_name, model_dir, mesh) config = rollout_worker.RolloutConfig( sampler_type="vanilla", weight_sync_mode=args.weight_sync_mode, @@ -382,12 +415,25 @@ def _create_vllm_sampler(args): def main(argv: list[str], context: Any = None) -> None: - if context and context.ipc and context.ipc.discovery: - pass - else: - raise RuntimeError( - "Require discovery API, but process context doesn't support." + if context is not None: + if not getattr(context, "ipc", None) or not getattr(context.ipc, "discovery", None): + raise RuntimeError( + "Require discovery API, but process context doesn't support." + ) + + args_list = argv[1:] if argv and argv[0] == sys.argv[0] else argv + args = _parse_args(args_list) + if context is None: + from tunix.experimental.distributed.runtime.contexts import context_factory # pylint: disable=g-import-not-at-top + + context = context_factory.get_default_process_context( + argparse.Namespace( + discovery_id=args.worker_id, + discovery_addrs=args.discovery_addrs, + port=args.port, + ) ) + context.__enter__() logging.basicConfig( level=logging.INFO, @@ -395,7 +441,6 @@ def main(argv: list[str], context: Any = None) -> None: force=True, ) - args = _parse_args(argv) logging.info("Parsed args: %s", args) if context and args.sampler != "vllm": @@ -410,9 +455,12 @@ def main(argv: list[str], context: Any = None) -> None: sys.path.insert(0, REPO_ROOT) logging.info("Repo root inserted into sys.path: %s", REPO_ROOT) + args.model_dir = models.ensure_model_dir(args.model_dir, args.model_id) + logging.info("Prepared rollout safetensors directory: %s", args.model_dir) + from transformers import AutoTokenizer # pylint: disable=g-import-not-at-top - tokenizer_path = args.tokenizer_path or args.model_dir or args.model_id + tokenizer_path = args.tokenizer_path or args.model_dir logging.info("Loading tokenizer from %s...", tokenizer_path) tokenizer: Any = AutoTokenizer.from_pretrained( tokenizer_path, trust_remote_code=True @@ -468,4 +516,7 @@ async def grpc_server_main() -> None: if __name__ == "__main__": - main(sys.argv[1:]) + from absl import app # pylint: disable=g-import-not-at-top + from absl import flags # pylint: disable=g-import-not-at-top + + app.run(main, flags_parser=lambda argv: flags.FLAGS(argv, known_only=True)) diff --git a/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py b/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py index 1951fd6dd..413989879 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_trainer_node.py @@ -28,7 +28,6 @@ import sys from typing import Any -from flax import nnx import jax from jax import numpy as jnp from jax.experimental import mesh_utils @@ -40,7 +39,6 @@ from tunix.experimental.train import peft_trainer_v2 from tunix.experimental.worker import remote_execution from tunix.experimental.worker import trainer_worker -from tunix.utils import maxtext_utils REPO_ROOT = os.path.abspath( os.path.join(os.path.dirname(__file__), "..", "..", "..", "..") @@ -67,7 +65,6 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: parser.add_argument("--tokenizer_path", type=str, default="") parser.add_argument("--mesh_fsdp", type=int, default=2) parser.add_argument("--mesh_tp", type=int, default=1) - parser.add_argument("--mesh_expert", type=int, default=1) parser.add_argument("--max_prompt_length", type=int, default=512) parser.add_argument("--max_response_length", type=int, default=128) parser.add_argument("--mini_batch_size", type=int, default=1) @@ -77,8 +74,8 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: parser.add_argument("--eval_every_n_steps", type=int, default=1000000) parser.add_argument("--learning_rate", type=float, default=2.0e-7) parser.add_argument("--use_lora", action="store_true") - parser.add_argument("--lora_rank", type=int, default=64) - parser.add_argument("--lora_alpha", type=float, default=64.0) + parser.add_argument("--lora_rank", type=int, default=16) + parser.add_argument("--lora_alpha", type=float, default=16.0) parser.add_argument("--checkpoint_save_interval_steps", type=int, default=1) parser.add_argument("--checkpoint_max_to_keep", type=int, default=10) parser.add_argument( @@ -89,46 +86,10 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: os.path.join(REPO_ROOT, "checkpoints"), ), ) - parser.add_argument( - "--trainer_backend", - choices=("tunix", "maxtext"), - default="tunix", - help=( - "tunix runs Tunix's PeftTrainer; maxtext runs MaxTextTrainingEngine" - ), - ) - parser.add_argument("--maxtext_model_name", type=str, default="qwen3-0.6b") - parser.add_argument( - "--maxtext_padded_moe_mlp_dim", - type=int, - default=0, - help=( - "Explicit padded_base_moe_mlp_dim override to match rollout TP" - " tile-alignment padding for MoE models." - ), - ) - parser.add_argument( - "--maxtext_ckpt_path", - type=str, - default=os.getenv("MAXTEXT_CKPT", ""), - help="Orbax params-only checkpoint for the MaxText trainer, e.g. gs://...", - ) - parser.add_argument( - "--maxtext_output_directory", - type=str, - default=os.getenv("MAXTEXT_OUTPUT_DIR", os.path.join(REPO_ROOT, "artifacts", "math_gsm8k_dist", "maxtext")), - help="Base directory for MaxText trainer outputs.", - ) - parser.add_argument( - "--maxtext_warmup_steps_fraction", - type=float, - default=0.0, - help=( - "Warmup fraction for MaxText LR schedule (0.0 enables updates from" - " step 0)." - ), - ) - return parser.parse_args(argv) + parser.add_argument("--discovery_addrs", type=str, default="") + parser.add_argument("--pathways_bns", type=str, default="") + args, _ = parser.parse_known_args(argv) + return args def _nested_safetensors_dirs(model_dir: Path) -> list[str]: @@ -275,50 +236,57 @@ def close(self) -> None: self._trainer.close() -def _create_maxtext_trainer_factory(args) -> Any: - """Creates the trainer factory function for MaxText's MaxTextTrainingEngine.""" - logging.info("Trainer backend: MaxText's MaxTextTrainingEngine.") - pad_id = maxtext_utils.get_tokenizer_pad_id( - args.model_id, args.tokenizer_path, args.model_dir - ) - maxtext_config = maxtext_utils.build_maxtext_config( - model_name=args.maxtext_model_name, - worker_id=args.worker_id, - train_micro_batch_size=args.train_micro_batch_size, - mesh_fsdp=args.mesh_fsdp, - mesh_tp=args.mesh_tp, - mesh_expert=args.mesh_expert, - num_devices=jax.device_count(), - max_prompt_length=args.max_prompt_length, - max_response_length=args.max_response_length, - learning_rate=args.learning_rate, - warmup_steps_fraction=args.maxtext_warmup_steps_fraction, - load_parameters_path=args.maxtext_ckpt_path, - padded_moe_mlp_dim=args.maxtext_padded_moe_mlp_dim, - base_output_directory=args.maxtext_output_directory, - ) - logging.info("Creating MaxText device mesh...") - mesh = maxtext_utils.create_maxtext_mesh(maxtext_config) - logging.info("Trainer mesh: %s", mesh) - - def _factory(): - engine = maxtext_utils.create_maxtext_engine( - maxtext_config, - mesh=mesh, - tokenizer_pad_id=pad_id, - wrap_with_tunix_adapter=True, +def _create_trainer( + args, + actor_model: Any, + training_config: peft_trainer_v2.TrainingConfig, + mesh: Mesh, +) -> _MeshBoundTrainer: + with mesh: + trainer = peft_trainer_v2.PeftTrainer( + actor_model, + optax.adamw(learning_rate=args.learning_rate), + training_config, ) - return _MeshBoundTrainer(engine, mesh) + return _MeshBoundTrainer(trainer, mesh) - return _factory +def main(argv: list[str], context: Any = None) -> None: + if context is not None: + if not getattr(context, "ipc", None) or not getattr(context.ipc, "discovery", None): + raise RuntimeError( + "Require discovery API, but process context doesn't support." + ) + + args_list = argv[1:] if argv and argv[0] == sys.argv[0] else argv + args = _parse_args(args_list) + if context is None: + from tunix.experimental.distributed.runtime.contexts import context_factory # pylint: disable=g-import-not-at-top + + context = context_factory.get_default_process_context( + argparse.Namespace( + discovery_id=args.worker_id, + discovery_addrs=args.discovery_addrs, + port=args.port, + ) + ) + context.__enter__() -def _create_tunix_trainer_factory(args) -> Any: - """Creates the trainer factory function for Tunix's PeftTrainer.""" - logging.info("Trainer backend: Tunix's PeftTrainer.") - args.model_dir = _ensure_model_dir_for_trainer( - args.model_dir, args.model_id + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - [TrainerNode] %(message)s", + force=True, ) + + logging.info("Parsed args: %s", args) + + if context: + context.jax.initialize() + if REPO_ROOT not in sys.path: + sys.path.insert(0, REPO_ROOT) + logging.info("Repo root inserted into sys.path: %s", REPO_ROOT) + + args.model_dir = models.ensure_model_dir(args.model_dir, args.model_id) logging.info("Prepared trainer safetensors directory: %s", args.model_dir) logging.info("Creating trainer mesh...") @@ -329,6 +297,8 @@ def _create_tunix_trainer_factory(args) -> Any: actor_model = _load_actor_model(args, mesh, lora=args.use_lora) logging.info("Building PeftTrainer v2 config...") + if args.train_micro_batch_size <= 0: + raise ValueError("--train_micro_batch_size must be positive.") grad_accumulation_steps = max( 1, math.ceil(args.mini_batch_size / args.train_micro_batch_size) ) @@ -350,55 +320,11 @@ def _create_tunix_trainer_factory(args) -> Any: grad_accumulation_steps, ) - def _factory(): - with mesh: - trainer = peft_trainer_v2.PeftTrainer( - actor_model, - optax.adamw(learning_rate=args.learning_rate), - training_config, - ) - return _MeshBoundTrainer(trainer, mesh) - - return _factory - - -def _create_trainer_factory(args) -> Any: - """Creates the trainer factory function based on args.trainer_backend.""" - if args.trainer_backend == "maxtext": - return _create_maxtext_trainer_factory(args) - return _create_tunix_trainer_factory(args) - - -def main(argv: list[str], context: Any = None) -> None: - if context and context.ipc and context.ipc.discovery: - pass - else: - raise RuntimeError( - "Require discovery API, but process context doesn't support." - ) - - logging.basicConfig( - level=logging.INFO, - format="%(asctime)s - [TrainerNode] %(message)s", - force=True, - ) - - args = _parse_args(argv) - logging.info("Parsed args: %s", args) - - if context: - context.jax.initialize() - if REPO_ROOT not in sys.path: - sys.path.insert(0, REPO_ROOT) - logging.info("Repo root inserted into sys.path: %s", REPO_ROOT) - - if args.train_micro_batch_size <= 0: - raise ValueError("--train_micro_batch_size must be positive.") - logging.info("Creating generic TrainerWorker and gRPC server...") - trainer_factory = _create_trainer_factory(args) worker_service = trainer_worker.TrainerWorker( - trainer_factory=trainer_factory, + trainer_factory=lambda: _create_trainer( # pyrefly: ignore[bad-argument-type] + args, actor_model, training_config, mesh + ), worker_id=args.worker_id, ) @@ -443,4 +369,7 @@ async def grpc_server_main() -> None: if __name__ == "__main__": - main(sys.argv[1:]) + from absl import app # pylint: disable=g-import-not-at-top + from absl import flags # pylint: disable=g-import-not-at-top + + app.run(main, flags_parser=lambda argv: flags.FLAGS(argv, known_only=True)) diff --git a/tunix/experimental/examples/math_gsm8k_dist/xm_launch.py b/tunix/experimental/examples/math_gsm8k_dist/xm_launch.py new file mode 100644 index 000000000..3c64d8a89 --- /dev/null +++ b/tunix/experimental/examples/math_gsm8k_dist/xm_launch.py @@ -0,0 +1,355 @@ +# 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. + +"""XManager launcher for distributed GSM8K GRPO on 1P Pathways and McJAX.""" + +import os +from typing import Any, Dict + +from absl import app +from absl import flags +from absl import logging +from GOOGLE_INTERNAL_PACKAGE_PATH.learning.deepmind.xmanager2.client.launch.borg import gcl_utils +from GOOGLE_INTERNAL_PACKAGE_PATH.third_party.pathways.google.xmanager import service_lib +from xmanager import xm +from xmanager import xm_abc +from xmanager import xm_flags +from xmanager.contrib.internal import addressing +from xmanager.contrib.internal import xm_jax + +_EXP_TITLE = flags.DEFINE_string( + "exp_title", + "tunix_math_gsm8k_dist_grpo", + "Title for the XManager experiment.", +) + +_CELL = flags.DEFINE_string( + "cell", + "cj", + "Borg cell to run the experiment in.", +) + +_PRIORITY = flags.DEFINE_integer( + "priority", + 200, + "Borg priority for jobs.", +) + +_TRAINER_PLATFORM = flags.DEFINE_string( + "trainer_platform", + "vlp=2x2", + "TPU platform and topology for the Pathways trainer service (e.g. 'vlp=2x2', 'glp=2x2', 'vf=2x2x1').", +) + +_ROLLOUT_PLATFORM = flags.DEFINE_string( + "rollout_platform", + "vlp=2x2", + "TPU platform and topology for the rollout worker (e.g. 'vlp=2x2', 'glp=2x2', 'vf=2x2x1').", +) + +_MODEL_ID = flags.DEFINE_string( + "model_id", + "Qwen/Qwen3-0.6B", + "HuggingFace or local path for the model.", +) + +_MODEL_NAME = flags.DEFINE_string( + "model_name", + "Qwen3-0.6B", + "Model name corresponding to models.py registry (e.g. 'Qwen3-0.6B').", +) + +_SAMPLER = flags.DEFINE_string( + "sampler", + "vanilla", + "Rollout sampler implementation ('vanilla' or 'inprocess_vllm').", +) + +_WEIGHT_SYNC_BACKEND = flags.DEFINE_string( + "weight_sync_backend", + "raiden", + "Weight sync backend ('raiden' or 'none').", +) + +_MAX_STEPS = flags.DEFINE_integer( + "max_steps", + 10, + "Maximum RL training steps.", +) + +_BATCH_SIZE = flags.DEFINE_integer( + "batch_size", + 4, + "Prompt batch size per step.", +) + +_MINI_BATCH_SIZE = flags.DEFINE_integer( + "mini_batch_size", + 2, + "Mini-batch size for trainer gradient updates.", +) + +_NUM_GENERATIONS = flags.DEFINE_integer( + "num_generations", + 2, + "Number of generations per prompt in GRPO.", +) + +_BETA = flags.DEFINE_float( + "beta", + 0.0, + "KL penalty coefficient in GRPO. Set 0.0 to run without reference model.", +) + +_RESOURCE_MANAGER_RAM_GB = flags.DEFINE_integer( + "resource_manager_ram_gb", + 25, + "Host RAM in GB for the Pathways Resource Manager job.", +) + +_TRAINER_RAM_GB = flags.DEFINE_integer( + "trainer_ram_gb", + 32, + "Host RAM in GB for the Trainer client job.", +) + +_ROLLOUT_RAM_GB = flags.DEFINE_integer( + "rollout_ram_gb", + 64, + "Host RAM in GB for the Rollout worker job.", +) + +_ORCHESTRATOR_RAM_GB = flags.DEFINE_integer( + "orchestrator_ram_gb", + 32, + "Host RAM in GB for the Orchestrator coordinator job.", +) + + +def _parse_platform(platform_str: str) -> Dict[str, Any]: + """Parses platform strings like 'vlp=2x2' or 'vf=2x2x1' into kwargs for xm.JobRequirements.""" + parts = platform_str.split("=") + if len(parts) == 2: + return {parts[0]: parts[1]} + return {"accelerator": platform_str} + + +def _create_pathways_service( + work_unit: xm.WorkUnit, + requirements: xm.JobRequirements, +) -> service_lib.Service: + """Creates a Pathways service for the Trainer cluster.""" + enable_ti_vm = xm_flags.XM_ENABLE_BORG_TI_VM.value + workers = [ + service_lib.TpuWorkerJobConfig( + name="pathways_server_trainer", + platform=requirements.accelerator, # pyrefly: ignore[bad-argument-type] + topology=requirements.topology, # pyrefly: ignore[bad-argument-type] + cell=_CELL.value, + priority=_PRIORITY.value, + enable_ti_vm=enable_ti_vm, + ), + ] + + borg_parent = gcl_utils.borg_token("trainer") + resource_manager = service_lib.ResourceManagerJobConfig( + cell=_CELL.value, + ram=_RESOURCE_MANAGER_RAM_GB.value * xm.GiB, + priority=_PRIORITY.value, + enable_ti_vm=enable_ti_vm, + ) + service_config = service_lib.ServiceConfig( + workers=workers, + resource_manager=resource_manager, + borg_parent=borg_parent, + ) + return service_lib.create_service(service_config, work_unit) + + +async def _launch_experiment(): + """Sets up and launches the distributed GSM8K GRPO experiment.""" + experiment_title = _EXP_TITLE.value + + async with xm_abc.create_experiment( + experiment_title=experiment_title + ) as experiment: + bazel_args = xm_abc.bazel_args.tpu() + ( + "--modify_execution_info=PostMark=+requires-mem:24g,PostMarking=+requires-mem:24g", + ) + + # Package binaries + executables = experiment.package([ + xm.bazel_binary( + label="//third_party/py/tunix/experimental/examples/math_gsm8k_dist:run_trainer_node", + executor_spec=xm_abc.Borg.Spec(), + bazel_args=bazel_args, + ), + xm.bazel_binary( + label="//third_party/py/tunix/experimental/examples/math_gsm8k_dist:run_rollout_node", + executor_spec=xm_abc.Borg.Spec(), + bazel_args=bazel_args, + ), + xm.bazel_binary( + label="//third_party/py/tunix/experimental/examples/math_gsm8k_dist:run_gsm8k_dist_grpo", + executor_spec=xm_abc.Borg.Spec(), + bazel_args=bazel_args, + ), + ]) + trainer_exec, rollout_exec, orch_exec = executables[0], executables[1], executables[2] + + job_requirements_kwargs = { + "location": _CELL.value, + "priority": _PRIORITY.value, + } + + # Trainer TPU platform + trainer_tpu_reqs = xm.JobRequirements( + **_parse_platform(_TRAINER_PLATFORM.value), + **job_requirements_kwargs, + ) + + # Rollout TPU platform + rollout_tpu_reqs = xm.JobRequirements( + ram=_ROLLOUT_RAM_GB.value * xm.GiB, + tmp_ram_fs=30 * xm.GiB, + **_parse_platform(_ROLLOUT_PLATFORM.value), + **job_requirements_kwargs, + ) + + async def make_jobs(work_unit: xm.WorkUnit): + jobs: Dict[str, Any] = {} + + # 1. Create Pathways service for Trainer + pw_service = _create_pathways_service( + work_unit=work_unit, + requirements=trainer_tpu_reqs, + ) + pathways_bns = pw_service.backend_target + jobs.update(**pw_service.jobs) + + # 2. Derive Orchestrator BNS address for discovery + orch_bns = addressing.bns_address( + cell=_CELL.value, + borguser=os.environ.get("USER", "lancewang"), + job_name="orchestrator", + experiment_id=experiment.experiment_id, + work_unit_id=work_unit.work_unit_id, + ) + discovery_addr = f"{orch_bns}:20000" + + # 3. Trainer CPU client job connected to Pathways + trainer_args = [ + f"--pathways_bns={pathways_bns}", + f"--discovery_addrs={discovery_addr}", + "--port=20001", + "--worker_id=trainer-0", + "--mesh_fsdp=4", + "--mesh_tp=1", + f"--model_id={_MODEL_ID.value}", + f"--model_name={_MODEL_NAME.value}", + ] + trainer_env = { + "JAX_PLATFORMS": "proxy,cpu", + "JAX_BACKEND_TARGET": f"grpc://{pathways_bns}", + "USE_RAIDEN_FFI": "1", + "MODEL_DOWNLOAD_DIR": "/tmp/artifacts/qwen3_dist_gsm8k/models", + } + trainer_executor = xm_abc.Borg( + logs_read_access_roles=["all"], + requirements=xm.JobRequirements( + ram=_TRAINER_RAM_GB.value * xm.GiB, + tmp_ram_fs=30 * xm.GiB, + cpu=8, + **job_requirements_kwargs, + ), + ) + jobs["trainer"] = xm.Job( + executable=trainer_exec, + args=trainer_args, + env_vars=trainer_env, + executor=trainer_executor, + ) + + # 4. Rollout worker running McJAX vanilla sampler on physical TPU + rollout_args = [ + f"--sampler={_SAMPLER.value}", + f"--discovery_addrs={discovery_addr}", + "--port=20002", + "--worker_id=rollout-0", + "--mesh_fsdp=1", + "--mesh_tp=4", + f"--model_id={_MODEL_ID.value}", + f"--model_name={_MODEL_NAME.value}", + f"--weight_sync_mode={_WEIGHT_SYNC_BACKEND.value}", + ] + rollout_env = { + "MODEL_DOWNLOAD_DIR": "/tmp/artifacts/qwen3_dist_gsm8k/models", + } + rollout_executor = xm_abc.Borg( + logs_read_access_roles=["all"], + requirements=rollout_tpu_reqs, + ) + jobs["rollout"] = xm.Job( + executable=rollout_exec, + args=rollout_args, + env_vars=rollout_env, + executor=rollout_executor, + ) + + # 5. Orchestrator coordinating RL program + orch_args = [ + "--discovery_id=orch", + "--discovery_port=20000", + f"--model_id={_MODEL_ID.value}", + f"--model_name={_MODEL_NAME.value}", + f"--batch_size={_BATCH_SIZE.value}", + f"--mini_batch_size={_MINI_BATCH_SIZE.value}", + f"--num_generations={_NUM_GENERATIONS.value}", + f"--max_steps={_MAX_STEPS.value}", + f"--beta={_BETA.value}", + f"--weight_sync_backend={_WEIGHT_SYNC_BACKEND.value}", + "--stop_workers_on_exit", + ] + orch_env = { + "MODEL_DOWNLOAD_DIR": "/tmp/artifacts/qwen3_dist_gsm8k/models", + } + orch_executor = xm_abc.Borg( + logs_read_access_roles=["all"], + requirements=xm.JobRequirements( + ram=_ORCHESTRATOR_RAM_GB.value * xm.GiB, + tmp_ram_fs=30 * xm.GiB, + cpu=8, + **job_requirements_kwargs, + ), + ) + jobs["orchestrator"] = xm.Job( + executable=orch_exec, + args=orch_args, + env_vars=orch_env, + executor=orch_executor, + ) + + work_unit.add(xm.JobGroup(**jobs)) + + experiment.add(make_jobs) + + +import asyncio + +def main(_): + asyncio.run(_launch_experiment()) + + +if __name__ == "__main__": + app.run(main) diff --git a/tunix/experimental/train/peft_trainer_v2.py b/tunix/experimental/train/peft_trainer_v2.py index 06829491b..16a78f53a 100644 --- a/tunix/experimental/train/peft_trainer_v2.py +++ b/tunix/experimental/train/peft_trainer_v2.py @@ -389,9 +389,10 @@ def _zero_in_place(v): def _default_weight_sync_worker() -> Any: from tunix.experimental.weight_sync import raiden_synchronizer # pylint: disable=g-import-not-at-top + is_proxy = "proxy" in os.environ.get("JAX_PLATFORMS", "") return raiden_synchronizer.RaidenSynchronizer( "trainer", - host_stage="proxy" in os.environ.get("JAX_PLATFORMS", ""), + use_ffi=is_proxy, ) diff --git a/tunix/experimental/weight_sync/raiden_synchronizer.py b/tunix/experimental/weight_sync/raiden_synchronizer.py index ebe89c51a..d1e262989 100644 --- a/tunix/experimental/weight_sync/raiden_synchronizer.py +++ b/tunix/experimental/weight_sync/raiden_synchronizer.py @@ -17,6 +17,8 @@ from __future__ import annotations import collections +import ipaddress +import os import socket from typing import Any, List, Optional, Tuple @@ -31,6 +33,12 @@ except ImportError: _ws_lib = None +_raiden_ffi: Any = None +try: + from tpu_sync.frameworks.jax import weight_synchronizer_ffi as _raiden_ffi # pytype: disable=import-error pylint: disable=g-import-not-at-top +except ImportError: + _raiden_ffi = None + def local_ip() -> str: for family, probe in ( @@ -50,17 +58,15 @@ def local_ip() -> str: return "localhost" -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] - - def pull(leaf): - 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) +def unpack_ip(row: Any) -> str: + """Unpacks IP address from uint32 array row.""" + raw_bytes = b"".join( + int(x).to_bytes(4, byteorder="little", signed=True) for x in row[:4] + ) + if raw_bytes[:10] == b"\x00" * 10 and raw_bytes[10:12] == b"\xff\xff": + return str(ipaddress.IPv4Address(raw_bytes[12:16])) + addr_str = str(ipaddress.IPv6Address(raw_bytes)) + return f"[{addr_str}]" if ":" in addr_str else addr_str def flatten_weights(state: Any) -> Tuple[List[str], List[Any]]: @@ -159,19 +165,24 @@ def __init__( *, worker_index: int = 0, auto_h2d: bool = False, - host_stage: bool = False, + use_ffi: Optional[bool] = None, parallelism: int = 4, bind_ip: Optional[str] = None, ): + is_proxy = "proxy" in os.environ.get("JAX_PLATFORMS", "") + if use_ffi is None: + use_ffi = is_proxy self.job_name = job_name self.worker_index = worker_index self.names: List[str] = [] self.arrays: List[Any] = [] self.ip = bind_ip or local_ip() self._auto_h2d = auto_h2d - self._host_stage = host_stage + self._use_ffi = use_ffi self._parallelism = parallelism self._sync: Any = None + self._ips: List[str] = [] + self._unique_listeners: List[str] = [] if state is not None: self.bind(state) @@ -181,17 +192,17 @@ def bound(self) -> bool: @property def active(self) -> bool: - return self._sync is not None + return self._sync is not None or bool(self._ips) - def bind(self, state: Any) -> None: - """Binds this host's weights, or rebinds them after a training step. + @property + def use_ffi(self) -> bool: + return self._use_ffi - With host_stage the arrays are copied to local CPU memory first; arrays - backed by the pathways proxy cannot bind in place. - """ - if self._host_stage: - state = to_host_cpu_state(state) + def bind(self, state: Any) -> None: + """Binds this host's weights, or rebinds them after a training step.""" self.names, self.arrays = _filter_bindable(*flatten_weights(state)) + if self._use_ffi: + return if _ws_lib is None: return if self._sync is None: @@ -210,15 +221,93 @@ def bind(self, state: Any) -> None: self._sync.bind_weights(self.arrays) def _require_sync(self, op: str) -> Any: - if self._sync is None: + if self._sync is None and not self._use_ffi: raise RuntimeError(f"{self.job_name}: bind() must run before {op}") return self._sync def d2h(self) -> None: - if not self.bound: - raise RuntimeError(f"{self.job_name}: bind() must run before d2h()") - if self._sync is not None: - self._sync.d2h() + if self._use_ffi: + if not self.arrays: + raise RuntimeError(f"{self.job_name}: bind() must run before d2h()") + if _raiden_ffi is None: + raise RuntimeError("weight_synchronizer_ffi is not available for FFI weight sync.") + + import numpy as np # pylint: disable=g-import-not-at-top + from jax.experimental import multihost_utils # pylint: disable=g-import-not-at-top + + mesh = getattr(getattr(self.arrays[0], "sharding", None), "mesh", None) + if mesh is None: + raise ValueError("Arrays must be sharded on a Mesh for FFI weight sync.") + + slice_byte_sizes = [ + int(np.prod(arr.sharding.shard_shape(arr.shape)) * arr.dtype.itemsize) + for arr in self.arrays + ] + sizes_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec(None) + ) + slice_byte_sizes_sharded = jax.device_put( + jnp.array(slice_byte_sizes, dtype=jnp.int32), sizes_sharding + ) + + task_mesh_shape = tuple(mesh.shape[a] for a in mesh.axis_names) + global_ids = jnp.array( + [d.id for d in mesh.devices.flatten()], dtype=jnp.int32 + ).reshape(task_mesh_shape) + shard_idx = jax.device_put( + global_ids, + jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec(*mesh.axis_names) + ), + ) + + src_devices = mesh.devices.flatten() + num_processes = len(set(getattr(d, "process_index", 0) for d in src_devices)) + devices_per_host = len(src_devices) // max(1, num_processes) + + logging.info( + "Initializing Pathways weight synchronizer and executing D2H via FFI (%d layers, %d devices/host)", + len(self.arrays), + devices_per_host, + ) + src_ws_info = _raiden_ffi.init_weight_synchronizer_and_d2h( + device_arrays=self.arrays, + shard_idx=shard_idx, + mesh=mesh, + slice_byte_sizes=slice_byte_sizes_sharded, + parallelism=self._parallelism, + num_layers=len(self.arrays), + listener_port=0, + num_shards=devices_per_host, + ) + + local_ws_info = multihost_utils.global_array_to_host_local_array( + src_ws_info, + mesh, + jax.sharding.PartitionSpec(*mesh.axis_names, None), + ) + gathered_ws_info = multihost_utils.process_allgather(local_ws_info).reshape( + -1, 6 + ) + + self._ips, listeners = [], [] + for row in gathered_ws_info: + ip = unpack_ip(row) + self._ips.append(f"{ip}:{int(row[4])}") + listeners.append(f"{ip}:{int(row[5])}") + + self._unique_listeners = [] + for listener in listeners: + if listener not in self._unique_listeners: + self._unique_listeners.append(listener) + logging.info( + "FFI D2H complete. Shards: %s, Control plane: %s", + self._ips, + self._unique_listeners, + ) + return + + self._require_sync("d2h()").d2h() def h2d(self) -> None: if not self.bound: @@ -258,15 +347,20 @@ def work_unit_metadata(self) -> weight_sync.WorkUnitMetadata: if mesh_shape is None: mesh_axes = ("fsdp",) mesh_shape = (1,) - data_addr = ( - f"{self.ip}:{self._sync.local_port}" if self._sync else "" - ) - control_addr = ( - f"{self.ip}:{self._sync.listener_port}" - if self._sync and self._sync.listener_port - else "" - ) - num_shards = self._sync.num_shards if self._sync else 1 + if self._use_ffi: + shards = tuple(self._ips) + control_addr = self._unique_listeners[0] if self._unique_listeners else "" + else: + data_addr = ( + f"{self.ip}:{self._sync.local_port}" if self._sync else "" + ) + control_addr = ( + f"{self.ip}:{self._sync.listener_port}" + if self._sync and self._sync.listener_port + else "" + ) + num_shards = self._sync.num_shards if self._sync else 1 + shards = (data_addr,) * num_shards if data_addr else () # Index 0 keeps the default replica id "": transfer callers construct # WorkUnitId(job_name=...) without a replica, and registration lookups # must match it for single-replica units. @@ -276,7 +370,7 @@ def work_unit_metadata(self) -> weight_sync.WorkUnitMetadata: ) return weight_sync.WorkUnitMetadata( unit=unit, - shards=(data_addr,) * num_shards if data_addr else (), + shards=shards, control_plane_rpc_address=control_addr, mesh_shape=mesh_shape, variables=variables, diff --git a/tunix/experimental/worker/remote_execution.py b/tunix/experimental/worker/remote_execution.py index 333acd61b..084ecdaec 100644 --- a/tunix/experimental/worker/remote_execution.py +++ b/tunix/experimental/worker/remote_execution.py @@ -447,9 +447,10 @@ async def start_serving_async(self, port: int = 50051) -> Any: }, ) self._server.add_generic_rpc_handlers((handler,)) - # NOTE: add_insecure_port is for local loopback / isolated pod testing (experimental v0). - # For production across trust boundaries, use secure_server_credentials (ALTS/mTLS). - self._server.add_insecure_port(f"[::]:{port}") + try: + self._server.add_insecure_port(f"[::]:{port}") + except Exception: + self._server.add_insecure_port(f"0.0.0.0:{port}") await self._server.start() return self._server diff --git a/tunix/oss/utils.py b/tunix/oss/utils.py index ffb6c4e0a..0f6997bd2 100644 --- a/tunix/oss/utils.py +++ b/tunix/oss/utils.py @@ -66,7 +66,7 @@ def kaggle_pipeline(model_id: str, model_download_path: str): import kagglehub # pylint: disable=g-import-not-at-top - if 'KAGGLE_USERNAME' not in os.environ or 'KAGGLE_KEY' not in os.environ: + if 'KAGGLE_USERNAME' in os.environ and 'KAGGLE_KEY' in os.environ: kagglehub.login() os.environ['KAGGLEHUB_CACHE'] = model_download_path return kagglehub.model_download(model_id) @@ -74,15 +74,15 @@ def kaggle_pipeline(model_id: str, model_download_path: str): def hf_pipeline(model_id: str, model_download_path: str): """Download model from HuggingFace.""" - if 'HF_TOKEN' not in os.environ: - hf.login() - all_files = hf.list_repo_files(model_id) + token = os.environ.get('HF_TOKEN') + all_files = hf.list_repo_files(model_id, token=token) filtered_files = [f for f in all_files if not f.startswith('original/')] for filename in filtered_files: hf.hf_hub_download( repo_id=model_id, filename=filename, local_dir=model_download_path, + token=token, ) logging.info( 'Downloaded %s to: %s',