Skip to content
Merged
Show file tree
Hide file tree
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
43 changes: 40 additions & 3 deletions checkpoint_engine/ps.py
Original file line number Diff line number Diff line change
Expand Up @@ -273,6 +273,8 @@ def __init__(
is_master=self._rank == 0,
)
self._store_counter = 0
self._store_barrier_counter = 0
self._store_barrier_trash: list[str] = []

def _get_memory_pool(self, checkpoint_name: str) -> list[MemoryBuffer]:
if checkpoint_name == self._current_shared_memory_pool_user:
Expand Down Expand Up @@ -555,16 +557,51 @@ def store_based_barrier(self, timeout: timedelta = timedelta(minutes=5)) -> None
allowing all ranks to synchronize regardless of which process group
they belong to.

Args:
store: The TCPStore instance to use for synchronization.
``_store_based_barrier`` is a one-shot initialization barrier: its
completion key remains in the store. Use a new group name for every
call so this method remains reusable with the shared root store.
"""
self._store_barrier_counter += 1
Comment thread
weixiao-huang marked this conversation as resolved.
group_name = f"parameter_server_barrier-{self._store_barrier_counter}"
torch.distributed.distributed_c10d._store_based_barrier(
rank=self._rank,
store=self._store,
group_name="parameter_server_barrier",
group_name=group_name,
rendezvous_count=self._world_size,
timeout=timeout,
)
Comment thread
weixiao-huang marked this conversation as resolved.
self._delete_stale_store_barrier_keys(group_name)

def _delete_stale_store_barrier_keys(self, group_name: str) -> None:
"""Delete keys left by barrier generations that no rank can still use.

Every generation leaves a counter key and a ``last_worker`` key in the
shared root store. Rank 0 retains the two latest generations. Once it
completes generation N + 2, every rank has entered that generation and
therefore has already returned from generation N, so N's keys are safe
to delete.
"""
if self._rank != 0:
return

self._store_barrier_trash.append(group_name)
if len(self._store_barrier_trash) <= 2:
return

stale_group_name = self._store_barrier_trash.pop(0)
prefix = torch.distributed.distributed_c10d.STORE_BASED_BARRIER_PREFIX
store_key = f"{prefix}:{stale_group_name}"
for key in (store_key, f"{store_key}:last_worker"):
self._delete_stale_store_key(key)

def _delete_stale_store_key(self, key: str) -> None:
try:
deleted = self._store.delete_key(key)
except RuntimeError as e:
logger.warning(f"[rank{self._rank}] failed to delete stale barrier key {key}: {e}")
else:
if not deleted:
logger.warning(f"[rank{self._rank}] stale barrier key {key} is already absent")

def update(
self,
Expand Down
62 changes: 62 additions & 0 deletions tests/test_store_barrier.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
from concurrent.futures import ThreadPoolExecutor
from datetime import timedelta
from unittest.mock import Mock, patch

import torch.distributed as dist

from checkpoint_engine.ps import ParameterServer


def test_store_based_barrier_uses_unique_group_name() -> None:
ps = ParameterServer.__new__(ParameterServer)
ps._rank = 0
ps._world_size = 2
ps._store = Mock()
ps._store_barrier_counter = 0
ps._store_barrier_trash = []
timeout = timedelta(seconds=5)

target = "torch.distributed.distributed_c10d._store_based_barrier"
with patch(target) as barrier:
ps.store_based_barrier(timeout)
ps.store_based_barrier(timeout)

assert [call.kwargs["group_name"] for call in barrier.call_args_list] == [
"parameter_server_barrier-1",
"parameter_server_barrier-2",
]
assert all(call.kwargs["store"] is ps._store for call in barrier.call_args_list)
assert all(call.kwargs["timeout"] == timeout for call in barrier.call_args_list)


def test_store_based_barrier_is_reusable_with_shared_tcp_store() -> None:
timeout = timedelta(seconds=5)
server_store = dist.TCPStore("127.0.0.1", 0, 2, True, timeout=timeout, wait_for_workers=False)
client_store = dist.TCPStore("127.0.0.1", server_store.port, 2, False, timeout=timeout)

parameter_servers = []
for rank, store in enumerate((server_store, client_store)):
ps = ParameterServer.__new__(ParameterServer)
ps._rank = rank
ps._world_size = 2
ps._store = store
ps._store_barrier_counter = 0
ps._store_barrier_trash = []
parameter_servers.append(ps)

with ThreadPoolExecutor(max_workers=2) as executor:
for _ in range(10):
futures = [executor.submit(ps.store_based_barrier, timeout) for ps in parameter_servers]
for future in futures:
future.result()

prefix = dist.distributed_c10d.STORE_BASED_BARRIER_PREFIX
for generation in range(1, 9):
store_key = f"{prefix}:parameter_server_barrier-{generation}"
assert not server_store.check([store_key])
assert not server_store.check([f"{store_key}:last_worker"])

for generation in (9, 10):
store_key = f"{prefix}:parameter_server_barrier-{generation}"
assert server_store.add(store_key, 0) == 2
assert server_store.get(f"{store_key}:last_worker") == b"1"
Loading