Skip to content
Open
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
32 changes: 32 additions & 0 deletions src/zeroband/comms.py
Original file line number Diff line number Diff line change
Expand Up @@ -402,6 +402,11 @@ def maybe_reinit_global_pg(self, admit_joiners: bool = False) -> bool:
# no op if disabled
return

if self.live_recovery.is_recovery_in_progress():
self._logger.info("Deferring global_pg reinit until live recovery completes")
self.live_recovery.wait_until_recovery_done()
return False

time_start = time.perf_counter()
self._logger.debug("[%s] Resolving world", self.world_info.global_unique_id)
if self._global_leader:
Expand Down Expand Up @@ -583,6 +588,8 @@ def format_grid(grid):


class LiveRecovery:
RECOVERY_POLL_INTERVAL = 0.05

def __init__(self, store: dist.Store):
self.logger = get_logger()
self.world_info = get_world_info()
Expand All @@ -593,6 +600,31 @@ def __init__(self, store: dist.Store):
def reset(self):
self.store.set(f"rank_{self.world_info.global_rank}", "null")

def _recovery_status(self) -> str:
try:
return self.store.get("status").decode("utf-8")
except Exception:
return "idle"

def begin_recovery(self, src_rank: int, dest_rank: int) -> None:
"""Mark cluster-wide live recovery in progress so peers defer PG reinit."""
self.store.set("status", f"in_progress:{src_rank}:{dest_rank}")

def end_recovery(self) -> None:
self.store.set("status", "idle")

def is_recovery_in_progress(self) -> bool:
return self._recovery_status().startswith("in_progress")

def wait_until_recovery_done(self, timeout: float = 3600.0) -> None:
if not self.is_recovery_in_progress():
return
start = time.perf_counter()
while self.is_recovery_in_progress():
if time.perf_counter() - start > timeout:
raise TimeoutError("Timed out waiting for live checkpoint recovery")
time.sleep(self.RECOVERY_POLL_INTERVAL)

def should_send_ckpt_to(self) -> int | None:
"""use this function to check if someone is awaiting for a live ckpt"""
data = self.store.get(f"rank_{self.world_info.global_rank}").decode("utf-8")
Expand Down
57 changes: 34 additions & 23 deletions src/zeroband/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,41 +240,52 @@ def train(config: Config):

if config.diloco is not None:
assert diloco is not None
# this is a patch for now to allow live recovery worker to not affect the all reduce at all
elastic_device_mesh.live_recovery.wait_until_recovery_done()

if not need_live_recovery:
elastic_device_mesh.maybe_reinit_global_pg(admit_joiners=True)

maybe_dest_rank = elastic_device_mesh.live_recovery.should_send_ckpt_to()
if maybe_dest_rank is not None:
src_rank = elastic_device_mesh.world_info.global_rank
logger.info(f"Start live recovery to rank {maybe_dest_rank}")
ckpt_manager.send_ckpt_to_peer(elastic_device_mesh.global_pg, maybe_dest_rank, blocking=True)

elastic_device_mesh.live_recovery.begin_recovery(src_rank, maybe_dest_rank)
try:
ckpt_manager.send_ckpt_to_peer(
elastic_device_mesh.global_pg, maybe_dest_rank, blocking=True
)
finally:
elastic_device_mesh.live_recovery.end_recovery()
elastic_device_mesh.live_recovery.reset()
else:
## receiving
time_start_live_recovery = time.perf_counter()
logger.info(f"Start live recovery from rank {config.ckpt.live_recovery_rank_src}")

## we create grad buffer and opts stats mamnually, the value will be overwritten by the ckpt but we need the DTensor to be correctly init before loading it

diloco.outer_optimizer.step() # need to step to init the DTensor stats

ckpt_manager.recv_ckpt_from_peer(elastic_device_mesh.global_pg)

log_hash_training_state(
config,
model,
inner_optimizer,
diloco,
metric_logger,
step=training_progress.step,
id="live_reco_recv",
)
need_live_recovery = False
src_rank = config.ckpt.live_recovery_rank_src
dest_rank = elastic_device_mesh.world_info.global_rank
logger.info(f"Start live recovery from rank {src_rank}")
elastic_device_mesh.live_recovery.begin_recovery(src_rank, dest_rank)
try:
## we create grad buffer and opts stats mamnually, the value will be overwritten by the ckpt but we need the DTensor to be correctly init before loading it

diloco.outer_optimizer.step() # need to step to init the DTensor stats

ckpt_manager.recv_ckpt_from_peer(elastic_device_mesh.global_pg)

log_hash_training_state(
config,
model,
inner_optimizer,
diloco,
metric_logger,
step=training_progress.step,
id="live_reco_recv",
)
need_live_recovery = False

if config.ckpt.remote_data_load:
ckpt_manager.remote_data_load()
if config.ckpt.remote_data_load:
ckpt_manager.remote_data_load()
finally:
elastic_device_mesh.live_recovery.end_recovery()

logger.info("live recovery done in %f", time.perf_counter() - time_start_live_recovery)

Expand Down
61 changes: 61 additions & 0 deletions tests/test_live_recovery.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
"""Unit tests for LiveRecovery coordination fence."""

from __future__ import annotations

import time
from unittest.mock import MagicMock

import pytest

from zeroband.comms import LiveRecovery


class _FakeStore:
def __init__(self) -> None:
self._data: dict[str, bytes] = {}

def get(self, key: str) -> bytes:
return self._data.get(key, b"idle")

def set(self, key: str, value: str | bytes) -> None:
if isinstance(value, str):
value = value.encode("utf-8")
self._data[key] = value


def _make_live_recovery() -> LiveRecovery:
lr = LiveRecovery.__new__(LiveRecovery)
lr.logger = MagicMock()
lr.world_info = MagicMock(global_rank=0)
lr.store = _FakeStore()
return lr


def test_live_recovery_fence_blocks_until_idle():
lr = _make_live_recovery()
assert not lr.is_recovery_in_progress()

lr.begin_recovery(src_rank=1, dest_rank=3)
assert lr.is_recovery_in_progress()

done = {"value": False}

def waiter() -> None:
lr.wait_until_recovery_done(timeout=2.0)
done["value"] = True

import threading

thread = threading.Thread(target=waiter)
thread.start()
time.sleep(0.1)
assert not done["value"]

lr.end_recovery()
thread.join(timeout=2.0)
assert done["value"]


def test_live_recovery_wait_is_noop_when_idle():
lr = _make_live_recovery()
lr.wait_until_recovery_done(timeout=0.1)