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
Original file line number Diff line number Diff line change
Expand Up @@ -196,21 +196,37 @@ def _create_replicator_file(
data_parallelism: int,
node_rank: int,
peer_ranks: List[int],
backup_interval_minutes: int,
backup_interval_minutes: Optional[int],
backup_interval_steps: Optional[int],
):
"""Creates a replicator file."""
_validate_replicator_ranks(
num_nodes=num_nodes, node_rank=node_rank, peer_ranks=peer_ranks
)
if (backup_interval_minutes is None) == (backup_interval_steps is None):
raise ValueError(
'Exactly one of backup_interval_minutes or backup_interval_steps '
'must be specified.'
)
if backup_interval_minutes is not None and backup_interval_minutes <= 0:
raise ValueError('backup_interval_minutes must be > 0.')
if backup_interval_steps is not None and backup_interval_steps <= 0:
raise ValueError('backup_interval_steps must be > 0.')

temp_file = epath.Path(file_path) / _TEMP_REPLICATOR_FILE_NAME
replicator_file = epath.Path(file_path) / _REPLICATOR_FILE
backup_interval_yaml = (
f'backup-interval-minutes: {backup_interval_minutes}'
if backup_interval_minutes is not None
else f'backup-interval-steps: {backup_interval_steps}'
)
replicator_yaml = f"""job-name: {run_name}
framework: orbax
assume-data-parallelism: {data_parallelism}
node-rank: {node_rank}
nodes: {num_nodes}
peer-ranks: {peer_ranks}
backup-interval-minutes: {backup_interval_minutes}"""
{backup_interval_yaml}"""
final_yaml = '\n'.join(
line.strip() for line in replicator_yaml.split('\n')
)
Expand All @@ -226,7 +242,8 @@ def _create_replicator_file(

def _initialize_mtc_colocated(
local_checkpoint_directory: epath.Path,
backup_interval_minutes: int,
backup_interval_minutes: Optional[int],
backup_interval_steps: Optional[int],
num_slices: int,
run_name: str,
data_parallelism: int,
Expand All @@ -236,9 +253,12 @@ def _initialize_mtc_colocated(
"""Initializes multi-tier checkpointing with a colocated Python sidecar on all workers.

Args:
local_checkpoint_directory: The local checkpoint directory on the
worker's filesystem.
backup_interval_minutes: The backup interval in minutes.
local_checkpoint_directory: The local checkpoint directory on the worker's
filesystem.
backup_interval_minutes: The backup interval in minutes. Exactly one of
`backup_interval_minutes` or `backup_interval_steps` must be specified.
backup_interval_steps: The backup interval in steps. Exactly one of
`backup_interval_minutes` or `backup_interval_steps` must be specified.
num_slices: The number of slices.
run_name: The run name.
data_parallelism: The data parallelism.
Expand Down Expand Up @@ -351,6 +371,7 @@ def _remaining_timeout_seconds() -> int:
node_rank=node_rank,
peer_ranks=peer_ranks,
backup_interval_minutes=backup_interval_minutes,
backup_interval_steps=backup_interval_steps,
)
_wait_for_replicator_file_to_disappear(
loc_dir,
Expand Down Expand Up @@ -416,7 +437,8 @@ def _initialize_jax_from_mtc(
def initialize_multi_tier_checkpointing(
local_checkpoint_directory: epath.Path,
*,
backup_interval_minutes: int = 30,
backup_interval_minutes: Optional[int] = None,
backup_interval_steps: Optional[int] = None,
num_slices: Optional[int] = None,
run_name: Optional[str] = None,
data_parallelism: Optional[int] = None,
Expand All @@ -430,19 +452,28 @@ def initialize_multi_tier_checkpointing(
Args:
local_checkpoint_directory: The local checkpoint directory.
backup_interval_minutes: The backup interval for the replicator service, in
minutes.
minutes. Exactly one of `backup_interval_minutes` or
`backup_interval_steps` must be specified.
backup_interval_steps: The backup interval for the replicator service, in
steps. Exactly one of `backup_interval_minutes` or `backup_interval_steps`
must be specified.
num_slices: The number of slices.
run_name: The name of the run.
data_parallelism: Number of identical pipelines in job, should be
equal to ICI data parallelism * DCN data parallelism. If not provided, it
will be inferred from the number of slices.
data_parallelism: Number of identical pipelines in job, should be equal to
ICI data parallelism * DCN data parallelism. If not provided, it will be
inferred from the number of slices.
jax_initialization_timeout_seconds: The timeout for JAX initialization.
use_mtc_process_ids: Use the MTC rank server to calculate process ids.
use_colocated_python: Whether to use Colocated Python for initialization.
devices: Optional JAX devices for Colocated Python initialization. This is
useful when the caller has already filtered controller-visible devices,
such as after an elastic restart.
"""
# Preserve previous default behavior where backup_interval_minutes defaulted
# to 30.
if backup_interval_minutes is None and backup_interval_steps is None:
backup_interval_minutes = 30

run_name = run_name if run_name else os.environ.get('JOBSET_NAME')
if not run_name:
raise ValueError(
Expand Down Expand Up @@ -473,6 +504,7 @@ def _resolve_parallelism_args():
_initialize_mtc_colocated(
local_checkpoint_directory=local_checkpoint_directory,
backup_interval_minutes=backup_interval_minutes,
backup_interval_steps=backup_interval_steps,
num_slices=num_slices, # pyrefly: ignore[bad-argument-type]
run_name=run_name,
data_parallelism=data_parallelism, # pyrefly: ignore[bad-argument-type]
Expand Down Expand Up @@ -565,6 +597,7 @@ def _resolve_parallelism_args():
node_rank=node_rank,
peer_ranks=peer_ranks,
backup_interval_minutes=backup_interval_minutes,
backup_interval_steps=backup_interval_steps,
)
_wait_for_replicator_file_to_disappear(
local_checkpoint_directory,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,7 @@ def test_create_replicator_file(self):
node_rank=0,
peer_ranks=[1],
backup_interval_minutes=10,
backup_interval_steps=None,
)
expected_replicator_data = {
"job-name": "test-run",
Expand All @@ -115,6 +116,71 @@ def test_create_replicator_file(self):
replicator_data = dict(yaml.safe_load(replicator_file.read_text()))
self.assertDictEqual(replicator_data, expected_replicator_data)

def test_create_replicator_file_steps(self):
tmp_dir = self.create_tempdir().full_path
epath.Path(tmp_dir).mkdir(parents=True, exist_ok=True)
replicator_file = epath.Path(tmp_dir) / initialization._REPLICATOR_FILE
self.assertFalse(replicator_file.exists())
initialization._create_replicator_file(
epath.Path(tmp_dir),
run_name="test-run",
num_nodes=2,
data_parallelism=1,
node_rank=0,
peer_ranks=[1],
backup_interval_minutes=None,
backup_interval_steps=100,
)
expected_replicator_data = {
"job-name": "test-run",
"framework": "orbax",
"assume-data-parallelism": 1,
"node-rank": 0,
"nodes": 2,
"peer-ranks": [1],
"backup-interval-steps": 100,
}

self.assertTrue(replicator_file.exists())
replicator_data = dict(yaml.safe_load(replicator_file.read_text()))
self.assertDictEqual(replicator_data, expected_replicator_data)

def test_create_replicator_file_rejects_both_intervals_set(self):
tmp_dir = self.create_tempdir().full_path
epath.Path(tmp_dir).mkdir(parents=True, exist_ok=True)
with self.assertRaisesRegex(
ValueError,
"Exactly one of backup_interval_minutes or backup_interval_steps",
):
initialization._create_replicator_file(
epath.Path(tmp_dir),
run_name="test-run",
num_nodes=2,
data_parallelism=1,
node_rank=0,
peer_ranks=[1],
backup_interval_minutes=10,
backup_interval_steps=100,
)

def test_create_replicator_file_rejects_neither_interval_set(self):
tmp_dir = self.create_tempdir().full_path
epath.Path(tmp_dir).mkdir(parents=True, exist_ok=True)
with self.assertRaisesRegex(
ValueError,
"Exactly one of backup_interval_minutes or backup_interval_steps",
):
initialization._create_replicator_file(
epath.Path(tmp_dir),
run_name="test-run",
num_nodes=2,
data_parallelism=1,
node_rank=0,
peer_ranks=[1],
backup_interval_minutes=None,
backup_interval_steps=None,
)

def test_create_replicator_file_rejects_invalid_node_rank(self):
tmp_dir = self.create_tempdir().full_path
epath.Path(tmp_dir).mkdir(parents=True, exist_ok=True)
Expand All @@ -127,6 +193,7 @@ def test_create_replicator_file_rejects_invalid_node_rank(self):
node_rank=-1,
peer_ranks=[1],
backup_interval_minutes=10,
backup_interval_steps=None,
)

def test_create_replicator_file_rejects_invalid_peer_rank(self):
Expand All @@ -141,6 +208,7 @@ def test_create_replicator_file_rejects_invalid_peer_rank(self):
node_rank=0,
peer_ranks=[2],
backup_interval_minutes=10,
backup_interval_steps=None,
)

def test_block_and_process_restore_dir_success(self):
Expand Down Expand Up @@ -494,13 +562,15 @@ def test_initialize_multi_tier_checkpointing_colocated_success(
data_parallelism=1,
use_colocated_python=True,
backup_interval_minutes=15,
backup_interval_steps=None,
devices=None,
)

# Verify colocated Python path is taken
mock_init_mtc_colocated.assert_called_once_with(
local_checkpoint_directory=tmp_dir_path,
backup_interval_minutes=15,
backup_interval_steps=None,
num_slices=1,
run_name="test-colocated-run",
data_parallelism=1,
Expand Down Expand Up @@ -536,6 +606,7 @@ def test_initialize_multi_tier_checkpointing_colocated_uses_devices(
mock_init_mtc_colocated.assert_called_once_with(
local_checkpoint_directory=tmp_dir_path,
backup_interval_minutes=30,
backup_interval_steps=None,
num_slices=1,
run_name="test-colocated-run",
data_parallelism=1,
Expand Down Expand Up @@ -579,6 +650,7 @@ def test_initialize_multi_tier_checkpointing_infers_defaults_when_none(
mock_init_mtc_colocated.assert_called_once_with(
local_checkpoint_directory=tmp_dir_path,
backup_interval_minutes=30,
backup_interval_steps=None,
num_slices=8,
run_name="test-colocated-run",
data_parallelism=8,
Expand Down Expand Up @@ -607,6 +679,7 @@ def test_initialize_multi_tier_checkpointing_infers_defaults_when_zero_or_negati
mock_init_mtc_colocated.assert_called_once_with(
local_checkpoint_directory=tmp_dir_path,
backup_interval_minutes=30,
backup_interval_steps=None,
num_slices=8,
run_name="test-colocated-run",
data_parallelism=8,
Expand Down Expand Up @@ -635,6 +708,7 @@ def test_initialize_multi_tier_checkpointing_infers_data_parallelism_from_num_sl
mock_init_mtc_colocated.assert_called_once_with(
local_checkpoint_directory=tmp_dir_path,
backup_interval_minutes=30,
backup_interval_steps=None,
num_slices=2,
run_name="test-colocated-run",
data_parallelism=2,
Expand Down Expand Up @@ -767,6 +841,7 @@ def specialize(self, *, out_specs_fn):
initialization._initialize_mtc_colocated(
local_checkpoint_directory=epath.Path("/tmp/mtc"),
backup_interval_minutes=15,
backup_interval_steps=None,
num_slices=2,
run_name="test-run",
data_parallelism=1,
Expand Down
Loading