From 7e0f0b0fed9b5cd2185a89905b4f8c6c5c3dcc8e Mon Sep 17 00:00:00 2001 From: Sam Stoelinga Date: Thu, 20 Aug 2026 21:42:38 -0700 Subject: [PATCH] Add wait_for_devices_ready utility to pathwaysutils. Workloads connecting to Pathways or initializing JAX may need to ensure all target devices are placed, warm, and ready to accept computations before starting training or inference. PiperOrigin-RevId: 968248469 --- pathwaysutils/__init__.py | 1 + pathwaysutils/_initialize.py | 47 +++++++++++++++++++++ pathwaysutils/test/initialize_test.py | 60 +++++++++++++++++++++++++++ 3 files changed, 108 insertions(+) diff --git a/pathwaysutils/__init__.py b/pathwaysutils/__init__.py index 02cb847..73ca030 100644 --- a/pathwaysutils/__init__.py +++ b/pathwaysutils/__init__.py @@ -17,6 +17,7 @@ initialize: Callable[[], None] = _initialize.initialize is_pathways_backend_used: Callable[[], bool] = _initialize.is_pathways_backend_used +wait_for_devices_ready = _initialize.wait_for_devices_ready del _initialize diff --git a/pathwaysutils/_initialize.py b/pathwaysutils/_initialize.py index 27476c9..bd22d06 100644 --- a/pathwaysutils/_initialize.py +++ b/pathwaysutils/_initialize.py @@ -13,6 +13,8 @@ # limitations under the License. """Initialization functions for Pathways-on-Cloud utilities.""" +from collections.abc import Sequence +import concurrent.futures import datetime import logging import os @@ -106,3 +108,48 @@ def initialize() -> None: _logger.debug( "Did not detect Pathways-on-Cloud backend. No changes applied." ) + + +def wait_for_devices_ready( + devices: Sequence[jax.Device] | None = None, + timeout: float | int | None = None, +) -> None: + """Waits for the given devices to be ready and available for computation. + + Args: + devices: The sequence of JAX devices to wait for. If None, defaults to all + available devices via `jax.devices()`. + timeout: The maximum number of seconds to wait. If None, there is no timeout + (waits indefinitely). + + Raises: + TimeoutError: If the timeout is reached before the devices become ready. + """ + if devices is None: + devices = jax.devices() + + if not devices: + return + + _logger.info( + "Waiting for %d devices to be ready (timeout=%s).", len(devices), timeout + ) + fn = lambda x: x + 1 + results = [jax.jit(fn, device=d)(0) for d in devices] + if timeout is None: + jax.block_until_ready(results) + else: + executor = concurrent.futures.ThreadPoolExecutor(max_workers=1) + future = executor.submit(jax.block_until_ready, results) + try: + future.result(timeout=timeout) + except concurrent.futures.TimeoutError as e: + executor.shutdown(wait=False, cancel_futures=True) + raise TimeoutError( + f"Timed out waiting for {len(devices)} devices to be ready after" + f" {timeout} seconds." + ) from e + else: + executor.shutdown(wait=True) + + _logger.info("All %d devices are ready.", len(devices)) diff --git a/pathwaysutils/test/initialize_test.py b/pathwaysutils/test/initialize_test.py index 31eaa0f..6431ca0 100644 --- a/pathwaysutils/test/initialize_test.py +++ b/pathwaysutils/test/initialize_test.py @@ -13,10 +13,13 @@ # limitations under the License. import os +import time +from unittest import mock from absl.testing import absltest from absl.testing import parameterized import jax +import pathwaysutils from pathwaysutils import _initialize @@ -99,6 +102,63 @@ def test_persistence_enabled(self): del os.environ["ENABLE_PATHWAYS_PERSISTENCE"] self.assertFalse(_initialize._is_persistence_enabled()) + def test_wait_for_devices_ready_default(self): + # Should execute without errors on default devices. + pathwaysutils.wait_for_devices_ready() + + def test_wait_for_devices_ready_explicit_devices(self): + devices = jax.devices()[:1] + pathwaysutils.wait_for_devices_ready(devices) + + def test_wait_for_devices_ready_empty(self): + pathwaysutils.wait_for_devices_ready([]) + + def test_wait_for_devices_ready_calls_jit_and_block_until_ready(self): + mock_dev1 = mock.create_autospec(jax.Device, instance=True) + mock_dev2 = mock.create_autospec(jax.Device, instance=True) + mock_devices = [mock_dev1, mock_dev2] + + mock_jit_fn = mock.MagicMock(return_value="result") + mock_jit = self.enter_context( + mock.patch.object(jax, "jit", return_value=mock_jit_fn) + ) + mock_block = self.enter_context(mock.patch.object(jax, "block_until_ready")) + + pathwaysutils.wait_for_devices_ready(mock_devices) + + self.assertEqual(mock_jit.call_count, 2) + mock_jit.assert_any_call(mock.ANY, device=mock_dev1) + mock_jit.assert_any_call(mock.ANY, device=mock_dev2) + self.assertIs( + mock_jit.call_args_list[0][0][0], mock_jit.call_args_list[1][0][0] + ) + mock_block.assert_called_once_with(["result", "result"]) + + def test_wait_for_devices_ready_logs(self): + with self.assertLogs(_initialize._logger, level="INFO") as logs: + pathwaysutils.wait_for_devices_ready(jax.devices()[:1], timeout=10) + self.assertLen(logs.output, 2) + self.assertIn( + "Waiting for 1 devices to be ready (timeout=10).", logs.output[0] + ) + self.assertIn("All 1 devices are ready.", logs.output[1]) + + def test_wait_for_devices_ready_with_timeout_success(self): + pathwaysutils.wait_for_devices_ready(timeout=60) + + def test_wait_for_devices_ready_with_timeout_exceeded(self): + def slow_block_until_ready(results): + time.sleep(1) + return results + + self.enter_context( + mock.patch.object( + jax, "block_until_ready", side_effect=slow_block_until_ready + ) + ) + with self.assertRaises(TimeoutError): + pathwaysutils.wait_for_devices_ready(timeout=0.01) + if __name__ == "__main__": absltest.main()