From 48975a369fb6313442be70d41d4c83fc268dc120 Mon Sep 17 00:00:00 2001 From: Luke Baumann Date: Thu, 1 Oct 2026 15:19:26 -0700 Subject: [PATCH] Default wait_for_slices to wait for all available slices. PiperOrigin-RevId: 991934395 --- pathwaysutils/elastic/elastic.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/pathwaysutils/elastic/elastic.py b/pathwaysutils/elastic/elastic.py index 1333f16..1ef4bf8 100644 --- a/pathwaysutils/elastic/elastic.py +++ b/pathwaysutils/elastic/elastic.py @@ -171,7 +171,7 @@ def get_active_slice_indices( def wait_for_slices( - slice_count: int, + slice_count: int | None = None, poll_interval: float | int = 10, timeout: float | int | None = None, slice_to_devices: Mapping[int, Sequence[jax.Device]] | None = None, @@ -180,7 +180,8 @@ def wait_for_slices( """Waits until after at least `slice_count` slices become active. Args: - slice_count: The number of slices to wait for. + slice_count: The number of slices to wait for. If None, defaults to the + total number of slices. poll_interval: The minimum number of seconds to wait between availability checks. If the check takes longer than this, the next check will start immediately after the current check completes. Defaults to 10 seconds. @@ -201,6 +202,10 @@ def wait_for_slices( _logger.debug("slice_to_devices is None. Getting from jax.devices().") slice_to_devices = get_slice_to_devices(jax.devices()) + if slice_count is None: + _logger.debug("slice_count is None. Using len(slice_to_devices).") + slice_count = len(slice_to_devices) + _logger.info( "Waiting for %s slices. Poll interval: %s, Timeout: %s", slice_count,