Skip to content
Open
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
34 changes: 18 additions & 16 deletions pathwaysutils/elastic/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,11 +116,12 @@ def __init__(self, devices: Sequence[jax.Device] | None = None) -> None:

self.all_slice_indices = frozenset(self.slice_to_devices.keys())

self._active_slice_indices = frozenset()
self.available_inactive_slices = frozenset()

self.active_slice_indices = elastic.get_active_slice_indices(
slice_to_devices=self.slice_to_devices
)
self.inactive_slice_indices = self.all_slice_indices - self.active_slice_indices
self.available_inactive_slices = frozenset()

self._stop_event = None
self._monitor_thread = None
Expand Down Expand Up @@ -180,22 +181,26 @@ def default_device(self) -> jax.Device:
except StopIteration as error:
raise ValueError("No active slices") from error

@property
def active_slice_indices(self) -> frozenset[int]:
"""The indices of active slices."""
return self._active_slice_indices

@active_slice_indices.setter
def active_slice_indices(self, value: set[int] | frozenset[int]) -> None:
self._active_slice_indices = frozenset(value)
self.inactive_slice_indices = (
self.all_slice_indices - self._active_slice_indices
)
self.available_inactive_slices = frozenset(
self.available_inactive_slices - self._active_slice_indices
)

@property
def active_slice_count(self) -> int:
"""The number of active slices."""
return len(self.active_slice_indices)

@property
def new_slice_event(self) -> threading.Event:
"""Deprecated compatibility property for un-updated MaxText code.

TODO: b/527183831 - Remove this property once MaxText CL 2 is submitted.
"""
event = threading.Event()
if self.available_inactive_slices:
event.set()
return event

def scale_by_active_slices(self, x: int | float) -> int | float:
"""Scale x by the number of active slices."""
if isinstance(x, int):
Expand Down Expand Up @@ -372,9 +377,6 @@ def attempt_execution(attempt: int) -> Any:
poll_interval=poll_interval,
timeout=timeout,
)
self.inactive_slice_indices = (
self.all_slice_indices - self.active_slice_indices
)
# Reset available_inactive_slices at attempt start since
# active_slice_indices has just been updated by wait_for_slices.
self.available_inactive_slices = frozenset()
Expand Down
Loading