diff --git a/src/dependencies/extra_deps/tpu_overrides.txt b/src/dependencies/extra_deps/tpu_overrides.txt index 647a1ca427..5f837af0a5 100644 --- a/src/dependencies/extra_deps/tpu_overrides.txt +++ b/src/dependencies/extra_deps/tpu_overrides.txt @@ -1,3 +1,4 @@ datasets>=4.8.5 fsspec==2026.2.0 gcsfs==2026.2.0 +orbax-checkpoint>=0.12.1 diff --git a/src/dependencies/requirements/generated_requirements/tpu-requirements.txt b/src/dependencies/requirements/generated_requirements/tpu-requirements.txt index 26ba4fcdda..42eb293874 100644 --- a/src/dependencies/requirements/generated_requirements/tpu-requirements.txt +++ b/src/dependencies/requirements/generated_requirements/tpu-requirements.txt @@ -4,12 +4,12 @@ absl-py>=2.4.0 aiofiles>=25.1.0 aiohappyeyeballs>=2.6.2 -aiohttp>=3.13.5 +aiohttp>=3.14.1 aiosignal>=1.4.0 annotated-doc>=0.0.4 annotated-types>=0.7.0 antlr4-python3-runtime>=4.9.3 -anyio>=4.13.0 +anyio>=4.14.1 aqtp>=0.9.0 array-record>=0.8.3 astroid>=4.0.4 @@ -22,29 +22,29 @@ certifi>=2026.2.25 cffi>=2.0.0 ; platform_python_implementation != 'PyPy' cfgv>=3.5.0 charset-normalizer>=3.4.7 -chex>=0.1.91 -click>=8.4.0 +chex>=0.1.92 +click>=8.4.2 cloud-accelerator-diagnostics>=0.1.1 cloudpickle>=3.1.2 clu>=0.0.12 colorama>=0.4.6 contourpy>=1.3.3 -cryptography>=48.0.0 +cryptography>=49.0.0 cycler>=0.12.1 -datasets>=4.8.5 +datasets>=5.0.0 decorator>=5.3.1 dill>=0.4.1 -distlib>=0.4.0 +distlib>=0.4.3 distro>=1.9.0 dm-tree>=0.1.10 docstring-parser>=0.18.0 -drjax>=0.1.4 +drjax>=0.2.0 editdistance>=0.8.1 einops>=0.8.2 einshape>=1.0 etils>=1.14.0 execnet>=2.1.2 -fastapi>=0.136.1 +fastapi>=0.138.0 filelock>=3.28.0 flatbuffers>=25.12.19 flax>=0.12.7 @@ -53,39 +53,39 @@ frozenlist>=1.8.0 fsspec>=2026.2.0 gast>=0.7.0 gcsfs>=2026.2.0 -google-api-core>=2.30.3 -google-api-python-client>=2.196.0 -google-auth>=2.53.0 +google-api-core>=2.31.0 +google-api-python-client>=2.197.0 +google-auth>=2.55.0 google-auth-httplib2>=0.4.0 google-auth-oauthlib>=1.4.0 -google-cloud-aiplatform>=1.153.1 -google-cloud-appengine-logging>=1.9.0 -google-cloud-audit-log>=0.5.0 -google-cloud-bigquery>=3.41.0 +google-cloud-aiplatform>=1.158.0 +google-cloud-appengine-logging>=1.10.0 +google-cloud-audit-log>=0.6.0 +google-cloud-bigquery>=3.42.1 google-cloud-core>=2.6.0 -google-cloud-logging>=3.15.0 -google-cloud-mldiagnostics>=1.0.2 -google-cloud-monitoring>=2.30.0 -google-cloud-resource-manager>=1.17.0 -google-cloud-storage>=3.10.1 -google-cloud-storage-control>=1.11.0 +google-cloud-logging>=3.16.0 +google-cloud-mldiagnostics>=1.0.3 +google-cloud-monitoring>=2.31.0 +google-cloud-resource-manager>=1.18.0 +google-cloud-storage>=3.12.0 +google-cloud-storage-control>=1.12.0 google-crc32c>=1.8.0 -google-genai>=2.4.0 +google-genai>=2.10.0 google-pasta>=0.2.0 -google-resumable-media>=2.9.0 +google-resumable-media>=2.10.0 googleapis-common-protos>=1.75.0 -grain>=0.2.16 +grain>=0.2.18 grpc-google-iam-v1>=0.14.4 grpcio>=1.80.0 grpcio-status>=1.80.0 gviz-api>=1.10.0 h11>=0.16.0 h5py>=3.14.0 -hf-xet>=1.5.0 ; platform_machine == 'AMD64' or platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64' +hf-xet>=1.5.1 ; platform_machine == 'AMD64' or platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64' httpcore>=1.0.9 httplib2>=0.31.2 httpx>=0.28.1 -huggingface-hub>=1.15.0 +huggingface-hub>=1.20.1 humanize>=4.15.0 hypothesis>=6.142.1 identify>=2.6.19 @@ -96,7 +96,7 @@ iniconfig>=2.3.0 isort>=8.0.1 jax>=0.10.0 jaxlib>=0.10.0 -jaxtyping>=0.3.9 +jaxtyping>=0.3.11 jinja2>=3.1.6 jsonlines>=4.0.0 keras>=3.14.0 @@ -114,9 +114,9 @@ mccabe>=0.7.0 mdurl>=0.1.2 ml-collections>=1.1.0 ml-dtypes>=0.5.4 -ml-goodput-measurement>=0.0.16 +ml-goodput-measurement>=0.2.0 mpmath>=1.3.0 -msgpack>=1.1.2 +msgpack>=1.2.1 msgspec>=0.21.1 multidict>=6.7.1 multiprocess>=0.70.19 @@ -129,24 +129,26 @@ nodeenv>=1.10.0 numpy>=2.0.2 numpy-typing-compat>=20251206.2.0 nvidia-cuda-cccl>=13.2.75 +nvidia-ml-py>=13.610.43 oauthlib>=3.3.1 -omegaconf>=2.3.0 -opentelemetry-api>=1.42.0 +omegaconf>=2.3.1 +opentelemetry-api>=1.43.0 opt-einsum>=3.4.0 optax>=0.2.8 optree>=0.19.0 optype>=0.17.0 -orbax-checkpoint>=0.11.39 +orbax-checkpoint>=0.12.1 packaging>=26.1 pandas>=3.0.3 parameterized>=0.9.0 pathspec>=1.1.1 -pathwaysutils>=0.1.8 +pathwaysutils>=0.1.9 pillow>=12.2.0 -platformdirs>=4.9.6 +platformdirs>=4.10.0 pluggy>=1.6.0 portpicker>=1.6.0 pre-commit>=4.6.0 +prometheus-client>=0.20.0 promise>=2.3 propcache>=0.5.2 proto-plus>=1.28.0 @@ -164,22 +166,24 @@ pyelftools>=0.32 pyglove>=0.4.5 pygments>=2.20.0 pyink>=25.12.0 -pylint>=4.0.5 +pylint>=4.0.6 +pynvml>=13.0.1 +pyopenssl>=26.3.0 pyparsing>=3.3.2 pyproject-hooks>=1.2.0 pytest>=8.4.2 pytest-xdist>=3.8.0 python-dateutil>=2.9.0.post0 -python-discovery>=1.3.1 +python-discovery>=1.4.2 pytokens>=0.4.1 pytype>=2024.10.11 pyyaml>=6.0.3 -qwix>=0.1.6 +qwix>=0.1.8 regex>=2026.5.9 requests>=2.33.1 requests-oauthlib>=2.0.0 rich>=15.0.0 -safetensors>=0.7.0 +safetensors>=0.8.0 scipy>=1.17.1 scipy-stubs>=1.17.1.4 sentencepiece>=0.2.1 @@ -191,7 +195,7 @@ simplejson>=4.1.1 six>=1.17.0 sniffio>=1.3.1 sortedcontainers>=2.4.0 -starlette>=1.0.0 +starlette>=1.3.1 sympy>=1.14.0 tabulate>=0.10.0 tenacity>=9.1.4 @@ -201,9 +205,9 @@ tensorboard-plugin-profile>=2.13.0 tensorboardx>=2.6.5 tensorflow>=2.20.0 tensorflow-datasets>=4.9.10 -tensorflow-metadata>=1.17.3 +tensorflow-metadata>=1.21.0 tensorflow-text>=2.20.1 -tensorstore>=0.1.82 +tensorstore>=0.1.84 termcolor>=3.3.0 tiktoken>=0.13.0 tokamax>=0.0.12 @@ -211,8 +215,8 @@ tokenizers>=0.22.2 toml>=0.10.2 tomlkit>=0.15.0 toolz>=1.1.0 -tqdm>=4.66.3 -transformers>=5.9.0 +tqdm>=4.68.3 +transformers>=5.12.1 treescope>=0.1.10 typeguard>=2.13.3 typer>=0.25.1 @@ -221,15 +225,15 @@ typing-inspection>=0.4.2 tzdata>=2026.2 ; sys_platform == 'emscripten' or sys_platform == 'win32' uritemplate>=4.2.0 urllib3>=2.6.3 -uvicorn>=0.47.0 +uvicorn>=0.49.0 uvloop>=0.22.1 -virtualenv>=21.3.3 +virtualenv>=21.5.1 wadler-lindig>=0.1.7 websockets>=16.0 werkzeug>=3.1.8 wheel>=0.46.3 -wrapt>=2.1.2 -xxhash>=3.7.0 +wrapt>=2.2.2 +xxhash>=3.7.1 yarl>=1.24.2 zipp>=3.23.1 zstandard>=0.25.0 diff --git a/src/maxtext/common/checkpointing.py b/src/maxtext/common/checkpointing.py index 73f475bb39..2f0d8c3a49 100644 --- a/src/maxtext/common/checkpointing.py +++ b/src/maxtext/common/checkpointing.py @@ -443,17 +443,6 @@ def create_orbax_checkpoint_manager( logger=orbax_logger, ) - # Use Colocated Python checkpointing optimization (Single Controller only). - if enable_single_controller and colocated_python_checkpointing: - max_logging.log("Registering colocated python array handler") - checkpointing_impl = ocp.pathways.CheckpointingImpl.from_options( - use_colocated_python=True, - ) - ocp.pathways.register_type_handlers( - use_single_replica_array_handler=enable_single_replica_ckpt_restoring, - checkpointing_impl=checkpointing_impl, - ) - max_logging.log("Checkpoint manager created!") return manager @@ -500,6 +489,7 @@ def create_orbax_emergency_replicator_checkpoint_manager( local_checkpoint_dir: str, save_interval_steps: int, global_mesh: jax.sharding.Mesh, + colocated_python_checkpointing: bool = False, ): """Returns an emergency replicator checkpoint manager.""" flags.FLAGS.experimental_orbax_use_distributed_process_id = True @@ -509,6 +499,7 @@ def create_orbax_emergency_replicator_checkpoint_manager( epath.Path(local_checkpoint_dir), options=emergency_replicator_checkpoint_manager.ReplicatorCheckpointManagerOptions( save_interval_steps=save_interval_steps, + use_colocated_python=colocated_python_checkpointing, ), global_mesh=global_mesh, ) @@ -838,9 +829,7 @@ def map_to_pspec(data): (EmergencyCheckpointManager, EmergencyReplicatorCheckpointManager), ): checkpoint_path = str(checkpoint_manager.directory / str(step) / "items") - with handle_checkpoint_mismatch( - "restore NNX checkpoint", checkpoint_path - ): + with handle_checkpoint_mismatch("restore NNX checkpoint", checkpoint_path): restored_nnx = _load_linen_checkpoint_into_nnx( checkpoint_path, abstract_unboxed_pre_state, @@ -876,9 +865,7 @@ def map_to_pspec(data): EmergencyReplicatorCheckpointManager, ), ): - restored = checkpoint_manager.restore( - step, args=Composite(state=checkpoint_args) - ).state + restored = checkpoint_manager.restore(step, args=Composite(state=checkpoint_args)).state _assert_no_shaped_dtype_struct(restored) return ( restored, @@ -906,9 +893,7 @@ def map_to_pspec(data): # Case 3: Default/Fallback case. # This case acts as a wildcard ('_') and matches if none of the preceding cases were met. case _: - restored = checkpoint_manager.restore( - step, args=Composite(items=checkpoint_args) - ) + restored = checkpoint_manager.restore(step, args=Composite(items=checkpoint_args)) _assert_no_shaped_dtype_struct(restored) return (restored, None) @@ -918,9 +903,7 @@ def map_to_pspec(data): else: params = abstract_unboxed_pre_state.params - with handle_checkpoint_mismatch( - "load parameters", load_parameters_from_path - ): + with handle_checkpoint_mismatch("load parameters", load_parameters_from_path): restored_params = load_params_from_path( load_parameters_from_path, params, @@ -932,9 +915,7 @@ def map_to_pspec(data): return None, restored_params elif load_full_state_from_path != "": max_logging.log(f"Loading full state from path: {load_full_state_from_path}") - with handle_checkpoint_mismatch( - "load full state", load_full_state_from_path - ): + with handle_checkpoint_mismatch("load full state", load_full_state_from_path): restored_state = _load_full_state_from_path( path=load_full_state_from_path, abstract_unboxed_pre_state=abstract_unboxed_pre_state, diff --git a/src/maxtext/utils/max_utils.py b/src/maxtext/utils/max_utils.py index 8fd0b303e3..4ee35a3844 100644 --- a/src/maxtext/utils/max_utils.py +++ b/src/maxtext/utils/max_utils.py @@ -246,6 +246,17 @@ def maybe_initialize_jax_distributed_system(raw_keys): return if raw_keys["enable_single_controller"]: max_logging.log("Skipping jax distributed system since its not needed for single controller.") + if raw_keys["enable_multi_tier_checkpointing"]: + max_logging.log("Initializing multi-tier checkpointing for single controller...") + initialize_multi_tier_checkpointing( + local_checkpoint_directory=raw_keys["local_checkpoint_directory"], + backup_interval_minutes=raw_keys["multi_tier_checkpointing_backup_interval_minutes"], + run_name=raw_keys["run_name"], + jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"], + data_parallelism=raw_keys["mtc_data_parallelism"], + num_slices=raw_keys["num_slices"], + use_colocated_python=True, + ) return if jax.distributed.is_initialized(): max_logging.log("Jax distributed system is already initialized.") @@ -290,6 +301,7 @@ def maybe_initialize_jax_distributed_system(raw_keys): run_name=raw_keys["run_name"], jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"], data_parallelism=raw_keys["mtc_data_parallelism"], + num_slices=raw_keys["num_slices"], ) max_logging.log("Jax distributed system initialized on TPUs for multi-tier checkpointing!") elif raw_keys["enable_checkpointing"] and raw_keys["compile_topology_num_slices"] == -1: diff --git a/src/maxtext/utils/train_utils.py b/src/maxtext/utils/train_utils.py index ca1b54314a..7a3e4854da 100644 --- a/src/maxtext/utils/train_utils.py +++ b/src/maxtext/utils/train_utils.py @@ -18,6 +18,7 @@ import subprocess import jax import functools +import orbax.checkpoint.pathways as ocp_pathways from functools import partial from flax import nnx @@ -55,6 +56,7 @@ def create_checkpoint_manager(config, mesh, init_state_fn): config.local_checkpoint_directory, config.local_checkpoint_period, mesh, + config.colocated_python_checkpointing, ) elif config.enable_emergency_checkpoint: abstract_state, _, _ = maxtext_utils.get_abstract_state(config, mesh, init_state_fn, is_training=True) @@ -97,6 +99,17 @@ def create_checkpoint_manager(config, mesh, init_state_fn): config.checkpoint_todelete_full_path, ) + # Use Colocated Python checkpointing dispatchers optimization (Single Controller only). + if checkpoint_manager is not None and config.enable_single_controller and config.colocated_python_checkpointing: + max_logging.log("Registering colocated python array handler") + checkpointing_impl = ocp_pathways.CheckpointingImpl.from_options( + use_colocated_python=True, + ) + ocp_pathways.register_type_handlers( + use_single_replica_array_handler=config.enable_single_replica_ckpt_restoring, + checkpointing_impl=checkpointing_impl, + ) + return checkpoint_manager diff --git a/tests/unit/max_utils_test.py b/tests/unit/max_utils_test.py index e108ab0d45..d00dbdd130 100644 --- a/tests/unit/max_utils_test.py +++ b/tests/unit/max_utils_test.py @@ -387,6 +387,7 @@ def _base_keys(self, **overrides): "multi_tier_checkpointing_backup_interval_minutes": 5, "run_name": "test_run", "mtc_data_parallelism": 1, + "num_slices": 2, "enable_checkpointing": True, "compile_topology_num_slices": -1, } @@ -458,6 +459,23 @@ def test_tpu_multi_tier_checkpointing(self, mock_mtc): 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"], + num_slices=self._base_keys()["num_slices"], + ) + + @mock.patch("maxtext.utils.max_utils.initialize_multi_tier_checkpointing") + @mock.patch("jax.distributed.initialize") + def test_single_controller_multi_tier_checkpointing_uses_colocated_python(self, mock_init, mock_mtc): + raw_keys = self._base_keys(enable_single_controller=True, enable_multi_tier_checkpointing=True) + max_utils.maybe_initialize_jax_distributed_system(raw_keys) + mock_init.assert_not_called() + 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"], + 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"], + num_slices=self._base_keys()["num_slices"], + use_colocated_python=True, ) @mock.patch("jax.distributed.initialize") diff --git a/tests/unit/train_compile_test.py b/tests/unit/train_compile_test.py index 1bbc9f2848..870358cdc0 100644 --- a/tests/unit/train_compile_test.py +++ b/tests/unit/train_compile_test.py @@ -28,10 +28,8 @@ from jax.experimental.compilation_cache import compilation_cache import pytest from tempfile import gettempdir, NamedTemporaryFile -import transformers -from maxtext.checkpoint_conversion.utils.hf_model_configs import DeepseekV32Config from maxtext.configs import pyconfig from maxtext.trainers.pre_train.train_compile import main as train_compile_main from tests.utils.test_helpers import get_test_config_path @@ -928,7 +926,6 @@ def test_mhc_integration(self): def test_engram_integration(self): """AOT test for Engram implementation""" compiled_trainstep_file = "/tmp/test_engram_integration" - transformers.AutoConfig.register("deepseek_v32", DeepseekV32Config) train_compile_main( ( "", diff --git a/tests/unit/train_state_nnx_checkpoint_test.py b/tests/unit/train_state_nnx_checkpoint_test.py index 60c0fc7300..edac15f784 100644 --- a/tests/unit/train_state_nnx_checkpoint_test.py +++ b/tests/unit/train_state_nnx_checkpoint_test.py @@ -65,6 +65,35 @@ def _replicate_for_orbax(pytree): return jax.tree.map(lambda x: jax.device_put(x, sharding) if isinstance(x, jax.Array) else x, pytree) +class TestEmergencyReplicatorCheckpointManager(unittest.TestCase): + """Tests for emergency replicator checkpoint manager construction.""" + + def test_colocated_python_option_is_forwarded(self): + checkpoint_manager = object() + mesh = object() + + with mock.patch.object( + checkpointing, + "EmergencyReplicatorCheckpointManager", + return_value=checkpoint_manager, + ) as manager_cls: + result = checkpointing.create_orbax_emergency_replicator_checkpoint_manager( + "/tmp/mtc", + save_interval_steps=10, + global_mesh=mesh, + colocated_python_checkpointing=True, + ) + + self.assertIs(result, checkpoint_manager) + manager_cls.assert_called_once() + args, kwargs = manager_cls.call_args + self.assertEqual(str(args[0]), "/tmp/mtc") + options = kwargs["options"] + self.assertEqual(options.save_interval_steps, 10) + self.assertTrue(options.use_colocated_python) + self.assertIs(kwargs["global_mesh"], mesh) + + @pytest.mark.cpu_only class TestTrainStateNNXCheckpoint(unittest.TestCase): """Class to test NNX checkpoint.""" @@ -460,9 +489,7 @@ def _nnx_pure(self): return { "model": { "decoder": {"norm": {"scale": jnp.ones((3,))}}, - "dropout": { - "rngs": {"default": {"key": jnp.ones((2,), dtype=jnp.uint32)}} - }, # NNX-only + "dropout": {"rngs": {"default": {"key": jnp.ones((2,), dtype=jnp.uint32)}}}, # NNX-only }, "optimizer": { "step": jnp.asarray(7, dtype=jnp.uint32), @@ -481,9 +508,7 @@ def test_to_linen_layout(self): linen = train_state_nnx.to_linen_checkpoint_dict(self._nnx_pure()) self.assertEqual(set(linen.keys()), {"params", "step", "opt_state"}) self.assertIn("params", linen["params"]) # params/params/ collection wrap - self.assertNotIn( - "dropout", linen["params"]["params"] - ) # NNX-only rngs/dropout stripped + self.assertNotIn("dropout", linen["params"]["params"]) # NNX-only rngs/dropout stripped self.assertEqual(linen["step"].dtype, jnp.int32) # Linen step is int32 # opt_state is a list with None for the EmptyState slot, mu/nu wrapped under params. self.assertIsInstance(linen["opt_state"], list) @@ -493,19 +518,11 @@ def test_to_linen_layout(self): def test_round_trip_preserves_values(self): nnx_pure = self._nnx_pure() - back = train_state_nnx.from_linen_checkpoint_dict( - train_state_nnx.to_linen_checkpoint_dict(nnx_pure) - ) + back = train_state_nnx.from_linen_checkpoint_dict(train_state_nnx.to_linen_checkpoint_dict(nnx_pure)) self.assertEqual(set(back.keys()), {"model", "optimizer"}) - self.assertEqual( - back["optimizer"]["step"].dtype, jnp.uint32 - ) # NNX step back to uint32 - self.assertEqual( - set(back["optimizer"]["opt_state"].keys()), {0, 2} - ) # int-keyed dict, EmptyState dropped - self.assertNotIn( - "params", back["optimizer"]["opt_state"][0]["mu"] - ) # mu/nu unwrapped + self.assertEqual(back["optimizer"]["step"].dtype, jnp.uint32) # NNX step back to uint32 + self.assertEqual(set(back["optimizer"]["opt_state"].keys()), {0, 2}) # int-keyed dict, EmptyState dropped + self.assertNotIn("params", back["optimizer"]["opt_state"][0]["mu"]) # mu/nu unwrapped self.assertTrue( jnp.array_equal( nnx_pure["model"]["decoder"]["norm"]["scale"], diff --git a/tests/unit/train_utils_test.py b/tests/unit/train_utils_test.py index 4e4a50ae98..e6436877e1 100644 --- a/tests/unit/train_utils_test.py +++ b/tests/unit/train_utils_test.py @@ -16,8 +16,11 @@ import unittest from dataclasses import dataclass +from types import SimpleNamespace +from unittest import mock from unittest.mock import MagicMock +from maxtext.utils import train_utils from maxtext.utils.train_utils import ( validate_train_config, create_training_optimizer, @@ -194,6 +197,53 @@ def test_sgd_optimizer_returns_tx(self): self.assertTrue(hasattr(tx, "init")) +class TestCreateCheckpointManager(unittest.TestCase): + """Tests for create_checkpoint_manager.""" + + def test_single_controller_mtc_registers_colocated_python_handlers(self): + config = SimpleNamespace( + enable_multi_tier_checkpointing=True, + enable_checkpoint_cloud_logger=False, + run_name="test_run", + local_checkpoint_directory="/tmp/mtc", + local_checkpoint_period=10, + colocated_python_checkpointing=True, + enable_single_controller=True, + enable_single_replica_ckpt_restoring=True, + ) + mesh = object() + checkpoint_manager = object() + checkpointing_impl = object() + + with ( + mock.patch.object( + train_utils.checkpointing, + "create_orbax_emergency_replicator_checkpoint_manager", + return_value=checkpoint_manager, + ) as create_manager, + mock.patch.object( + train_utils.ocp_pathways.CheckpointingImpl, + "from_options", + return_value=checkpointing_impl, + ) as from_options, + mock.patch.object(train_utils.ocp_pathways, "register_type_handlers") as register_type_handlers, + ): + result = train_utils.create_checkpoint_manager(config, mesh, init_state_fn=object()) + + self.assertIs(result, checkpoint_manager) + create_manager.assert_called_once_with( + config.local_checkpoint_directory, + config.local_checkpoint_period, + mesh, + config.colocated_python_checkpointing, + ) + from_options.assert_called_once_with(use_colocated_python=True) + register_type_handlers.assert_called_once_with( + use_single_replica_array_handler=config.enable_single_replica_ckpt_restoring, + checkpointing_impl=checkpointing_impl, + ) + + class TestValidateCompletedSteps(unittest.TestCase): """Tests for validate_completed_steps."""