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
Original file line number Diff line number Diff line change
Expand Up @@ -132,15 +132,16 @@ This configuration manages a `multi-tiered checkpointing` system designed for bo

- **Local checkpointing**: Saves checkpoints much more frequently to a fast, local directory on each host (i.e. a ramdisk). If a preemption or failure occurs, the job can restore from this recent local copy almost instantly, minimizing lost work without needing to download from slower persistent storage. This feature is enabled by setting `enable_checkpointing`, `enable_multi_tier_checkpointing`, `local_checkpoint_directory`, and a non-zero `local_checkpoint_period` flags.

- **Backup checkpointing**: These are checkpoints saved periodically to persistent storage(i.e. GCS bucket). They ensure that you can recover your training state even after a complete job failure(repair of all nodepools). From User's perspective all restoration is from local ramdisk, its replicator service responsibility to make the checkpointing available to local storage in case of job restart. The interval for backup can be enabled by setting a non-zero `multi_tier_checkpointing_backup_interval_minutes` flags.
- **Backup checkpointing**: These are checkpoints saved periodically to persistent storage (i.e. GCS bucket). They ensure that you can recover your training state even after a complete job failure (repair of all nodepools). From User's perspective all restoration is from local ramdisk, its replicator service responsibility to make the checkpoints available in local storage in case of job restart. The interval for backup can be enabled by setting a non-zero `multi_tier_checkpointing_backup_interval_minutes` or `multi_tier_checkpointing_backup_interval_steps` flags (but not both).

| Flag | Description | Type | Default |
| :------------------------------------------------- | :--------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | :-------- | :------ |
| `enable_checkpointing` | A master switch to enable (`True`) or disable (`False`) saving checkpoints during the training run. | `boolean` | `True` |
| `enable_multi_tier_checkpointing` | When set to (`True`), this flag enables the multi-tier checkpointing feature on maxtext level. | `boolean` | `False` |
| `local_checkpoint_directory` | The high-speed local filesystem path(i.e. ramdisk) where **Multi-tier checkpoints** are saved. Setting this path, along with a non-zero `local_checkpoint_period`, enables the Multi-tier Checkpointing feature. | `string` | `""` |
| `local_checkpoint_period` | The interval, in training steps, for how often a **Multi-tier checkpoint** is saved in local ramdisks. | `integer` | `0` |
| `multi_tier_checkpointing_backup_interval_minutes` | The interval, in minutes, for how often a **Multi-tier checkpoint** is saved to backup from local ramdisks. | `integer` | `0` |
| `multi_tier_checkpointing_backup_interval_minutes` | The interval, in minutes, for how often a **Multi-tier checkpoint** is saved to backup from local ramdisks. | `integer` | `None` |
| `multi_tier_checkpointing_backup_interval_steps` | The interval, in steps, for how often a **Multi-tier checkpoint** is saved to backup from local ramdisks. | `integer` | `None` |

### Workload creation using XPK

Expand Down
17 changes: 14 additions & 3 deletions src/maxtext/configs/base.yml
Original file line number Diff line number Diff line change
Expand Up @@ -481,9 +481,21 @@ base_output_directory: ""
# enable_multi_tier_checkpointing=true local_checkpoint_directory="/local" local_checkpoint_period=20 multi_tier_checkpointing_backup_interval_minutes=20
enable_multi_tier_checkpointing: false

# The interval to backup local checkpoints to the persistent storage(GCS bucket) in minutes.
# The interval to backup local checkpoints to the persistent storage
# (GCS bucket) in minutes.
# It should be a positive number when enabling multi-tier checkpointing.
multi_tier_checkpointing_backup_interval_minutes: 0
# Note: This parameter and `multi_tier_checkpointing_backup_interval_steps`
# are mutually exclusive. Exactly one must be specified when enabling
# multi-tier checkpointing.
multi_tier_checkpointing_backup_interval_minutes: null

# The interval to backup local checkpoints to the persistent storage
# (GCS bucket) in steps.
# It should be a positive number when enabling multi-tier checkpointing.
# Note: This parameter and `multi_tier_checkpointing_backup_interval_minutes`
# are mutually exclusive. Exactly one must be specified when enabling
# multi-tier checkpointing.
multi_tier_checkpointing_backup_interval_steps: null

# Number of identical pipelines in job, should be equal to ICI data parallelism * DCN data parallelism.
# It should be a positive number when enabling multi-tier checkpointing. If set to 0, it will be set to num of slices.
Expand Down Expand Up @@ -1358,4 +1370,3 @@ elastic_backup_kind: "snapshot"
elastic_timeout_seconds: 300
elastic_max_retries: 10
elastic_min_slice_count: -1

27 changes: 23 additions & 4 deletions src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -407,10 +407,14 @@ class EmergencyCheckpointing(BaseModel):
)
local_checkpoint_directory: PathStr = Field("", description="Local directory for emergency checkpoints.")
local_checkpoint_period: NonNegativeInt = Field(0, description="Frequency (in steps) for local emergency checkpoints.")
multi_tier_checkpointing_backup_interval_minutes: NonNegativeInt = Field(
0,
multi_tier_checkpointing_backup_interval_minutes: PositiveInt | None = Field(
None,
description="Interval in minutes to back up local checkpoints to persistent storage.",
)
multi_tier_checkpointing_backup_interval_steps: PositiveInt | None = Field(
None,
description="Interval in steps to back up local checkpoints to persistent storage.",
)
mtc_data_parallelism: int = Field(
0,
description="Number of identical pipelines in the job for multi-tier checkpointing. 0 defaults to num_slices.",
Expand Down Expand Up @@ -3483,8 +3487,23 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
raise ValueError("`local_checkpoint_directory` must be set for multi-tier checkpointing.")
if self.local_checkpoint_period <= 0:
raise ValueError("`local_checkpoint_period` must be > 0 for multi-tier checkpointing.")
if self.multi_tier_checkpointing_backup_interval_minutes <= 0:
raise ValueError("`multi_tier_checkpointing_backup_interval_minutes` must be > 0.")
if (self.multi_tier_checkpointing_backup_interval_minutes is None) == (
self.multi_tier_checkpointing_backup_interval_steps is None
):
raise ValueError(
"Exactly one of `multi_tier_checkpointing_backup_interval_minutes`"
" or `multi_tier_checkpointing_backup_interval_steps` must be"
" specified."
)
if (
self.multi_tier_checkpointing_backup_interval_steps is not None
and self.multi_tier_checkpointing_backup_interval_steps < self.local_checkpoint_period
):
raise ValueError(
"`multi_tier_checkpointing_backup_interval_steps`"
f" ({self.multi_tier_checkpointing_backup_interval_steps}) must be"
f" >= `local_checkpoint_period` ({self.local_checkpoint_period})."
)
if self.colocated_python_checkpointing and not self.enable_single_controller:
raise ValueError("`colocated_python_checkpointing` is only supported with `enable_single_controller` set to True.")
if self.enable_emergency_checkpoint:
Expand Down
2 changes: 2 additions & 0 deletions src/maxtext/utils/max_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,6 +252,7 @@ def maybe_initialize_jax_distributed_system(raw_keys):
initialize_multi_tier_checkpointing(
local_checkpoint_directory=raw_keys["local_checkpoint_directory"],
backup_interval_minutes=raw_keys["multi_tier_checkpointing_backup_interval_minutes"],
backup_interval_steps=raw_keys["multi_tier_checkpointing_backup_interval_steps"],
run_name=raw_keys["run_name"],
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
use_colocated_python=True,
Expand Down Expand Up @@ -298,6 +299,7 @@ def maybe_initialize_jax_distributed_system(raw_keys):
initialize_multi_tier_checkpointing(
local_checkpoint_directory=raw_keys["local_checkpoint_directory"],
backup_interval_minutes=raw_keys["multi_tier_checkpointing_backup_interval_minutes"],
backup_interval_steps=raw_keys["multi_tier_checkpointing_backup_interval_steps"],
run_name=raw_keys["run_name"],
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
data_parallelism=raw_keys["mtc_data_parallelism"],
Expand Down
64 changes: 63 additions & 1 deletion tests/unit/max_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,19 @@
# limitations under the License.

"""Tests for the common Max Utils"""

import os
import sys

try:
from maxtext.utils import max_utils as max_utils_module
from maxtext.utils import max_logging as max_logging_module

sys.modules["maxtext.utils"] = max_utils_module
sys.modules["maxtext.utils.max_utils"] = max_utils_module
sys.modules["maxtext.utils.max_logging"] = max_logging_module
except ImportError:
pass
import time
import unittest
from unittest import mock
Expand Down Expand Up @@ -289,6 +300,7 @@ def test_initialize_jax_for_gpu_invalid_devices(self, _mock_log, _mock_devices,
@mock.patch("maxtext.utils.max_logging.log")
def test_initialize_jax_for_gpu_no_devices(self, _mock_log, _mock_devices, mock_init, mock_config_update):
"""When coordinator env is set but neither CUDA_VISIBLE_DEVICES nor SLURM_STEP_GPUS is set, JAX uses all devices

(config) and init gets no local ids.
"""
raw_keys = {"jax_distributed_initialization_timeout": 300}
Expand Down Expand Up @@ -399,6 +411,7 @@ def _base_keys(self, **overrides):
"enable_multi_tier_checkpointing": False,
"local_checkpoint_directory": "/tmp/ckpt",
"multi_tier_checkpointing_backup_interval_minutes": 5,
"multi_tier_checkpointing_backup_interval_steps": None,
"run_name": "test_run",
"mtc_data_parallelism": 1,
"num_slices": 2,
Expand Down Expand Up @@ -470,6 +483,7 @@ def test_tpu_multi_tier_checkpointing(self, mock_mtc):
mock_mtc.assert_called_once_with(
local_checkpoint_directory=self._base_keys()["local_checkpoint_directory"],
backup_interval_minutes=self._base_keys()["multi_tier_checkpointing_backup_interval_minutes"],
backup_interval_steps=self._base_keys()["multi_tier_checkpointing_backup_interval_steps"],
run_name=self._base_keys()["run_name"],
jax_initialization_timeout_seconds=self._base_keys()["jax_distributed_initialization_timeout"],
data_parallelism=self._base_keys()["mtc_data_parallelism"],
Expand All @@ -485,6 +499,7 @@ def test_single_controller_multi_tier_checkpointing_uses_colocated_python(self,
mock_mtc.assert_called_once_with(
local_checkpoint_directory=self._base_keys()["local_checkpoint_directory"],
backup_interval_minutes=self._base_keys()["multi_tier_checkpointing_backup_interval_minutes"],
backup_interval_steps=self._base_keys()["multi_tier_checkpointing_backup_interval_steps"],
run_name=self._base_keys()["run_name"],
jax_initialization_timeout_seconds=self._base_keys()["jax_distributed_initialization_timeout"],
data_parallelism=self._base_keys()["mtc_data_parallelism"],
Expand Down Expand Up @@ -522,6 +537,7 @@ def test_single_controller_multi_tier_checkpointing_uses_elastic_utils_kwargs(
mock_mtc.assert_called_once_with(
local_checkpoint_directory=self._base_keys()["local_checkpoint_directory"],
backup_interval_minutes=self._base_keys()["multi_tier_checkpointing_backup_interval_minutes"],
backup_interval_steps=self._base_keys()["multi_tier_checkpointing_backup_interval_steps"],
run_name=self._base_keys()["run_name"],
jax_initialization_timeout_seconds=self._base_keys()["jax_distributed_initialization_timeout"],
data_parallelism=1,
Expand All @@ -530,6 +546,47 @@ def test_single_controller_multi_tier_checkpointing_uses_elastic_utils_kwargs(
devices=active_devices,
)

@mock.patch("maxtext.utils.max_utils.elastic_utils.single_controller_mtc_init_kwargs")
@mock.patch("maxtext.utils.max_utils.initialize_multi_tier_checkpointing")
@mock.patch("jax.distributed.initialize")
def test_single_controller_multi_tier_checkpointing_with_steps_override(
self, mock_init, mock_mtc, mock_mtc_init_kwargs
):
active_devices = (
mock.Mock(slice_index=0),
mock.Mock(slice_index=0),
)
mock_mtc_init_kwargs.return_value = {
"data_parallelism": 1,
"num_slices": 1,
"devices": active_devices,
}
raw_keys = self._base_keys(
enable_single_controller=True,
enable_multi_tier_checkpointing=True,
elastic_enabled=True,
mtc_data_parallelism=0,
num_slices=2,
multi_tier_checkpointing_backup_interval_minutes=None,
multi_tier_checkpointing_backup_interval_steps=100,
)

max_utils.maybe_initialize_jax_distributed_system(raw_keys)

mock_init.assert_not_called()
mock_mtc_init_kwargs.assert_called_once_with(raw_keys)
mock_mtc.assert_called_once_with(
local_checkpoint_directory=raw_keys["local_checkpoint_directory"],
backup_interval_minutes=None,
backup_interval_steps=100,
run_name=raw_keys["run_name"],
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
data_parallelism=1,
num_slices=1,
use_colocated_python=True,
devices=active_devices,
)

@mock.patch("jax.distributed.initialize")
def test_tpu_checkpointing_no_emergency_calls_jax_init(self, mock_init):
raw_keys = self._base_keys(enable_checkpointing=True, compile_topology_num_slices=-1)
Expand Down Expand Up @@ -657,4 +714,9 @@ def test_reorder_roundtrip(self):


if __name__ == "__main__":
unittest.main()
try:
from absl.testing import absltest

absltest.main()
except ImportError:
unittest.main()
Loading