diff --git a/docs/guides/checkpointing_solutions/multi_tier_checkpointing.md b/docs/guides/checkpointing_solutions/multi_tier_checkpointing.md index 3af6fb880c..59ea384233 100644 --- a/docs/guides/checkpointing_solutions/multi_tier_checkpointing.md +++ b/docs/guides/checkpointing_solutions/multi_tier_checkpointing.md @@ -132,7 +132,7 @@ 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 | | :------------------------------------------------- | :--------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | :-------- | :------ | @@ -140,7 +140,8 @@ This configuration manages a `multi-tiered checkpointing` system designed for bo | `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 diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 227c7e36f1..a5c1ea714d 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -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. @@ -1358,4 +1370,3 @@ elastic_backup_kind: "snapshot" elastic_timeout_seconds: 300 elastic_max_retries: 10 elastic_min_slice_count: -1 - diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 4865557b5a..77117390b9 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -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.", @@ -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: diff --git a/src/maxtext/utils/max_utils.py b/src/maxtext/utils/max_utils.py index e6aebce059..e2b5870000 100644 --- a/src/maxtext/utils/max_utils.py +++ b/src/maxtext/utils/max_utils.py @@ -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, @@ -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"], diff --git a/tests/unit/max_utils_test.py b/tests/unit/max_utils_test.py index 097d85d491..04da4c6428 100644 --- a/tests/unit/max_utils_test.py +++ b/tests/unit/max_utils_test.py @@ -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 @@ -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} @@ -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, @@ -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"], @@ -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"], @@ -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, @@ -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) @@ -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()