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,