Skip to content
Merged
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
9 changes: 7 additions & 2 deletions pathwaysutils/elastic/elastic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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,
Expand Down
Loading