Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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()
Original file line number Diff line number Diff line change
@@ -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()
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()
Original file line number Diff line number Diff line change
Expand Up @@ -160,15 +160,15 @@ 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,
):
mock_ensure_model_dir.return_value = "/tmp/mock_model_dir"
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

Expand All @@ -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,
Expand All @@ -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,
)
Expand Down Expand Up @@ -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,
Expand All @@ -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,
)
Expand All @@ -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,
Expand All @@ -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,
)
Expand All @@ -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,
Expand All @@ -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,
)
Expand Down Expand Up @@ -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))

Expand Down Expand Up @@ -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"
)

Expand Down
12 changes: 6 additions & 6 deletions tests/experimental/train/peft_trainer_v2_weight_sync_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down
Loading
Loading