diff --git a/src/maxtext/training_engine/checkpointing.py b/src/maxtext/training_engine/checkpointing.py index eb35d23f43..26ea0e0e85 100644 --- a/src/maxtext/training_engine/checkpointing.py +++ b/src/maxtext/training_engine/checkpointing.py @@ -21,9 +21,10 @@ from absl import logging from flax import nnx import jax +from maxtext.common import checkpoint_context from maxtext.configs import pyconfig from maxtext.training_engine import abstract_engine -import orbax.checkpoint as ocp +from orbax.checkpoint import v1 as ocp @dataclasses.dataclass @@ -53,27 +54,34 @@ def __init__( checkpoint_dir: The root directory for saving checkpoints. config: The training configuration. """ - self._checkpoint_manager: ocp.CheckpointManager | None = None + self._checkpointer: ocp.training.Checkpointer | None = None + # The v0 manager took async-ness as an option; the v1 Checkpointer chooses per call, + # so remember it and dispatch in save_checkpoint. + self._use_async = bool(config.async_checkpointing) if checkpoint_dir: - self._checkpoint_manager = ocp.CheckpointManager( - directory=checkpoint_dir, - options=ocp.CheckpointManagerOptions( - save_interval_steps=config.checkpoint_period, - max_to_keep=config.max_num_checkpoints_to_keep, - enable_async_checkpointing=config.async_checkpointing, + preservation_policy = None + if config.max_num_checkpoints_to_keep is not None: + preservation_policy = checkpoint_context.build_preservation_policy(max_to_keep=config.max_num_checkpoints_to_keep) + self._checkpointer = ocp.training.Checkpointer( + checkpoint_dir, + context=checkpoint_context.build_context(), + save_decision_policy=checkpoint_context.build_save_decision_policy( + save_interval_steps=config.checkpoint_period ), + preservation_policy=preservation_policy, ) def get_latest_step(self) -> int | None: """Returns the latest checkpoint step.""" - if self._checkpoint_manager: - return self._checkpoint_manager.latest_step() + if self._checkpointer: + latest = self._checkpointer.latest + return latest.step if latest is not None else None return None def wait_until_finished(self) -> None: """Waits for any ongoing async checkpoint saves to finish.""" - if self._checkpoint_manager: - self._checkpoint_manager.wait_until_finished() + if self._checkpointer: + self._checkpointer.wait() def get_saved_micro_step_count(self, step: int) -> int: """Returns how far into `step` the checkpoint already on disk got. @@ -85,10 +93,10 @@ def get_saved_micro_step_count(self, step: int) -> int: 0 if that checkpoint covers a complete step, otherwise the number of micro-batches accumulated into it. """ - if self._checkpoint_manager is None: + if self._checkpointer is None: return 0 try: - metadata = self._checkpoint_manager.metadata(step) + metadata = self._checkpointer.checkpointables_metadata(step) except Exception as e: # pylint: disable=broad-except logging.warning("Could not read metadata for step %d, treating it as complete: %s", step, e) return 0 @@ -126,9 +134,17 @@ def _delete_saved_step(self, step: int) -> None: Args: step: The step to delete. """ + if self._checkpointer is None: + return logging.info("Deleting intra-step checkpoint at step %d so a more complete one can replace it.", step) - self._checkpoint_manager.wait_until_finished() - self._checkpoint_manager.delete(step) + self._checkpointer.wait() + # The v1 Checkpointer exposes no public delete yet; use the + # underlying v0 manager's delete (the same call Orbax makes internally). + self._checkpointer._manager.delete(step) # pylint: disable=protected-access + # The v1 `checkpoints`/`latest` properties cache their step listing and are + # unaware of deletes made below them; refresh so the replacement save isn't + # misjudged. + self._checkpointer.reload() def save_checkpoint( self, @@ -148,7 +164,7 @@ def save_checkpoint( Returns: Whether the checkpoint was saved. """ - if self._checkpoint_manager is None: + if self._checkpointer is None: logging.info("Checkpointing is disabled, skipping save.") return False @@ -174,49 +190,31 @@ def save_checkpoint( params = nnx.state(checkpoint_state.model) jax.block_until_ready(params) - model_cp_args = ocp.args.PyTreeSave( - item=params, - save_args=jax.tree.map(lambda _: ocp.SaveArgs(), params), - ) - save_args = {"model_params": model_cp_args} + checkpointables = {"model_params": params} if checkpoint_state.optimizer: optimizer_state = nnx.state(checkpoint_state.optimizer, nnx.optimizer.OptState) jax.block_until_ready(optimizer_state) - optimizer_cp_args = ocp.args.PyTreeSave( - item=optimizer_state, - save_args=jax.tree.map(lambda _: ocp.SaveArgs(), optimizer_state), - ) - save_args["optimizer_state"] = optimizer_cp_args + checkpointables["optimizer_state"] = optimizer_state if checkpoint_state.accumulated_metrics is not None: jax.block_until_ready(checkpoint_state.accumulated_metrics) - metrics_cp_args = ocp.args.PyTreeSave( - item=checkpoint_state.accumulated_metrics, - save_args=jax.tree.map( - lambda _: ocp.SaveArgs(), - checkpoint_state.accumulated_metrics, - ), - ) - save_args["accumulated_metrics"] = metrics_cp_args + checkpointables["accumulated_metrics"] = checkpoint_state.accumulated_metrics if checkpoint_state.accumulated_grads: jax.block_until_ready(checkpoint_state.accumulated_grads) - grads_cp_args = ocp.args.PyTreeSave( - item=checkpoint_state.accumulated_grads, - save_args=jax.tree.map( - lambda _: ocp.SaveArgs(), - checkpoint_state.accumulated_grads, - ), - ) - save_args["accumulated_grads"] = grads_cp_args + checkpointables["accumulated_grads"] = checkpoint_state.accumulated_grads - return self._checkpoint_manager.save( - step=step, - args=ocp.args.Composite(**save_args), - custom_metadata=custom_metadata, - **kwargs, - ) + try: + if self._use_async: + response = self._checkpointer.save_checkpointables_async( + step, checkpointables, custom_metadata=custom_metadata, **kwargs + ) + return response is not None + return self._checkpointer.save_checkpointables(step, checkpointables, custom_metadata=custom_metadata, **kwargs) + except FileExistsError as e: # ocp.training StepAlreadyExistsError subclasses FileExistsError + logging.info("Checkpoint for step %d already exists, skipping save. (%s)", step, e) + return False def restore_checkpoint( self, @@ -232,7 +230,7 @@ def restore_checkpoint( Returns: A tuple of (step, checkpoint_state, custom metadata). """ - if self._checkpoint_manager is None: + if self._checkpointer is None: logging.info("Checkpointing is disabled, skipping restore.") return None, checkpoint_state, None @@ -242,43 +240,26 @@ def restore_checkpoint( logging.info("No checkpoint found, skipping restore.") return None, checkpoint_state, None - metadata = self._checkpoint_manager.metadata(step) - restore_args: dict[str, Any] = {} + metadata = self._checkpointer.checkpointables_metadata(step) + saved_checkpointables = metadata.metadata if isinstance(metadata.metadata, Mapping) else {} - abstract_params = nnx.state(checkpoint_state.model) - restore_args["model_params"] = ocp.args.PyTreeRestore( - item=abstract_params, - restore_args=ocp.checkpoint_utils.construct_restore_args(target=abstract_params), - ) + abstract_checkpointables: dict[str, Any] = {"model_params": nnx.state(checkpoint_state.model)} - if checkpoint_state.optimizer is not None and "optimizer_state" in metadata.item_metadata: - optimizer_state = nnx.state(checkpoint_state.optimizer, nnx.optimizer.OptState) - restore_args["optimizer_state"] = ocp.args.PyTreeRestore( - item=optimizer_state, - restore_args=ocp.checkpoint_utils.construct_restore_args( - target=nnx.state(checkpoint_state.optimizer, nnx.optimizer.OptState) - ), - ) + if checkpoint_state.optimizer is not None and "optimizer_state" in saved_checkpointables: + abstract_checkpointables["optimizer_state"] = nnx.state(checkpoint_state.optimizer, nnx.optimizer.OptState) - if "accumulated_metrics" in metadata.item_metadata: - restore_args["accumulated_metrics"] = ocp.args.PyTreeRestore() + if "accumulated_metrics" in saved_checkpointables: + abstract_checkpointables["accumulated_metrics"] = saved_checkpointables["accumulated_metrics"] - if "accumulated_grads" in metadata.item_metadata: - accumulated_grads_target = nnx.state(checkpoint_state.model, nnx.Param) - restore_args["accumulated_grads"] = ocp.args.PyTreeRestore( - item=accumulated_grads_target, - restore_args=ocp.checkpoint_utils.construct_restore_args(target=accumulated_grads_target), - ) + if "accumulated_grads" in saved_checkpointables: + abstract_checkpointables["accumulated_grads"] = nnx.state(checkpoint_state.model, nnx.Param) custom_metadata = None if metadata and hasattr(metadata, "custom_metadata"): custom_metadata = metadata.custom_metadata try: - restored_items = self._checkpoint_manager.restore( - step=step, - args=ocp.args.Composite(**restore_args), - ) + restored_items = self._checkpointer.load_checkpointables(step, abstract_checkpointables) except Exception as e: # pylint: disable=broad-except logging.exception("Failed to restore checkpoint: %s", e) return None, None, None @@ -296,5 +277,5 @@ def restore_checkpoint( def close(self) -> None: """Closes the checkpoint manager.""" - if self._checkpoint_manager: - self._checkpoint_manager.close() + if self._checkpointer: + self._checkpointer.close() diff --git a/tests/post_training/unit/maxtext_engine_test.py b/tests/post_training/unit/maxtext_engine_test.py index f8348cdee9..e199849c1b 100644 --- a/tests/post_training/unit/maxtext_engine_test.py +++ b/tests/post_training/unit/maxtext_engine_test.py @@ -32,7 +32,7 @@ from tests.utils.test_helpers import get_test_config_path import numpy as np import optax -import orbax.checkpoint as ocp +from orbax.checkpoint import v1 as ocp import pytest from tunix.experimental.common import datatypes from tunix.experimental.train import abstract_trainer @@ -116,16 +116,16 @@ def setup_config(self, enable_checkpointing: bool = False, **kwargs): def _mock_orbax_manager(self, engine, latest_step=None): """Installs a mock Orbax manager and returns it.""" mock_orbax_mgr = mock.MagicMock() - mock_orbax_mgr.latest_step.return_value = latest_step - mock_orbax_mgr.save.return_value = True - engine._checkpoint_manager._checkpoint_manager = mock_orbax_mgr + mock_orbax_mgr.latest = None if latest_step is None else mock.MagicMock(step=latest_step) + mock_orbax_mgr.save_checkpointables.return_value = True + engine._checkpoint_manager._checkpointer = mock_orbax_mgr return mock_orbax_mgr def _mock_saved_micro_step_count(self, mock_orbax_mgr, micro_step_count): """Makes the mocked Orbax manager report how far into its step a saved checkpoint got.""" saved_metadata = mock.MagicMock() saved_metadata.custom_metadata = {"micro_step_count": micro_step_count} - mock_orbax_mgr.metadata.return_value = saved_metadata + mock_orbax_mgr.checkpointables_metadata.return_value = saved_metadata def test_raises_type_error_for_non_pyconfig(self): invalid_config = abstract_engine.TrainingConfig() @@ -162,19 +162,23 @@ def test_max_text_trainer_instantiation_with_pyconfig(self): self.assertIsNone(t._accumulated_grads) self.assertEqual(t.train_step, 2) - @mock.patch("orbax.checkpoint.CheckpointManager") + @mock.patch.object(maxtext_engine.checkpointing.ocp.training, "Checkpointer") def test_max_text_trainer_checkpoint_manager_init(self, mock_create_mgr): mock_config = self.setup_config(enable_checkpointing=True) _ = maxtext_engine.MaxTextTrainingEngine(mock_config) - mock_create_mgr.assert_called_once_with( - directory=mock_config.checkpoint_dir, - options=ocp.CheckpointManagerOptions( - save_interval_steps=mock_config.checkpoint_period, - max_to_keep=mock_config.max_num_checkpoints_to_keep, - enable_async_checkpointing=mock_config.async_checkpointing, - ), + mock_create_mgr.assert_called_once() + args, kwargs = mock_create_mgr.call_args + self.assertEqual(str(args[0]), mock_config.checkpoint_dir) + self.assertEqual( + kwargs["save_decision_policy"], + ocp.training.save_decision_policies.FixedIntervalPolicy(interval=mock_config.checkpoint_period), ) + self.assertEqual( + kwargs["preservation_policy"], + ocp.training.preservation_policies.LatestN(mock_config.max_num_checkpoints_to_keep), + ) + self.assertIn("context", kwargs) def test_save_checkpoint_called_after_update(self): mock_config = self.setup_config(enable_checkpointing=True) @@ -186,18 +190,14 @@ def test_save_checkpoint_called_after_update(self): t.save_checkpoint(metadata=dummy_metadata) # Verify orbax save was called - mock_orbax_mgr.save.assert_called_once() - call_kwargs = mock_orbax_mgr.save.call_args.kwargs - self.assertEqual(call_kwargs["custom_metadata"]["micro_step_count"], 0) - self.assertEqual(call_kwargs["custom_metadata"]["additional_metadata"], dummy_metadata) - args_dict = ( - dict(call_kwargs["args"].items()) - if hasattr(call_kwargs["args"], "items") and callable(call_kwargs["args"].items) - else call_kwargs["args"].__dict__ - ) - self.assertIn("model_params", args_dict) - self.assertIn("accumulated_metrics", args_dict) - self.assertNotIn("accumulated_grads", args_dict) + mock_orbax_mgr.save_checkpointables.assert_called_once() + call = mock_orbax_mgr.save_checkpointables.call_args + self.assertEqual(call.kwargs["custom_metadata"]["micro_step_count"], 0) + self.assertEqual(call.kwargs["custom_metadata"]["additional_metadata"], dummy_metadata) + checkpointables = call.args[1] + self.assertIn("model_params", checkpointables) + self.assertIn("accumulated_metrics", checkpointables) + self.assertNotIn("accumulated_grads", checkpointables) def test_save_checkpoint_skips_if_already_saved(self): mock_config = self.setup_config(enable_checkpointing=True) @@ -206,8 +206,8 @@ def test_save_checkpoint_skips_if_already_saved(self): mock_orbax_mgr = self._mock_orbax_manager(t, latest_step=10) t.save_checkpoint(metadata={"step": 10}) - mock_orbax_mgr.save.assert_not_called() - mock_orbax_mgr.delete.assert_not_called() + mock_orbax_mgr.save_checkpointables.assert_not_called() + mock_orbax_mgr._manager.delete.assert_not_called() # pylint: disable=protected-access def test_save_checkpoint_overwrites_intra_step_checkpoint_at_same_step(self): t = maxtext_engine.MaxTextTrainingEngine(self.setup_config(enable_checkpointing=True)) @@ -220,14 +220,14 @@ def test_save_checkpoint_overwrites_intra_step_checkpoint_at_same_step(self): t.save_checkpoint(metadata=None) - mock_orbax_mgr.wait_until_finished.assert_called() - mock_orbax_mgr.delete.assert_called_once_with(10) - mock_orbax_mgr.save.assert_called_once() - call_kwargs = mock_orbax_mgr.save.call_args.kwargs - self.assertEqual(call_kwargs["step"], 10) + mock_orbax_mgr.wait.assert_called() + mock_orbax_mgr._manager.delete.assert_called_once_with(10) # pylint: disable=protected-access + mock_orbax_mgr.save_checkpointables.assert_called_once() + call = mock_orbax_mgr.save_checkpointables.call_args + self.assertEqual(call.args[0], 10) # Orbax's save-interval policy would otherwise decline a step it has already saved. - self.assertTrue(call_kwargs["force"]) - self.assertEqual(call_kwargs["custom_metadata"]["micro_step_count"], 0) + self.assertTrue(call.kwargs["force"]) + self.assertEqual(call.kwargs["custom_metadata"]["micro_step_count"], 0) def test_save_checkpoint_overwrites_less_complete_intra_step_checkpoint(self): t = maxtext_engine.MaxTextTrainingEngine(self.setup_config(enable_checkpointing=True)) @@ -240,8 +240,8 @@ def test_save_checkpoint_overwrites_less_complete_intra_step_checkpoint(self): t.save_checkpoint(metadata=None) - mock_orbax_mgr.delete.assert_called_once_with(10) - self.assertEqual(mock_orbax_mgr.save.call_args.kwargs["custom_metadata"]["micro_step_count"], 3) + mock_orbax_mgr._manager.delete.assert_called_once_with(10) # pylint: disable=protected-access + self.assertEqual(mock_orbax_mgr.save_checkpointables.call_args.kwargs["custom_metadata"]["micro_step_count"], 3) def test_save_checkpoint_never_overwrites_a_complete_step(self): t = maxtext_engine.MaxTextTrainingEngine(self.setup_config(enable_checkpointing=True)) @@ -253,8 +253,8 @@ def test_save_checkpoint_never_overwrites_a_complete_step(self): t.save_checkpoint(metadata=None) - mock_orbax_mgr.save.assert_not_called() - mock_orbax_mgr.delete.assert_not_called() + mock_orbax_mgr.save_checkpointables.assert_not_called() + mock_orbax_mgr._manager.delete.assert_not_called() # pylint: disable=protected-access def test_update_supersedes_intra_step_checkpoint_it_resumed_from(self): t = maxtext_engine.MaxTextTrainingEngine(self.setup_config(enable_checkpointing=True)) @@ -277,10 +277,10 @@ def test_update_supersedes_intra_step_checkpoint_it_resumed_from(self): self.assertEqual(t.update(), 5) # The completed step must land on disk now, not a checkpoint period later. - mock_orbax_mgr.delete.assert_called_once_with(5) - call_kwargs = mock_orbax_mgr.save.call_args.kwargs - self.assertEqual(call_kwargs["step"], 5) - self.assertEqual(call_kwargs["custom_metadata"]["micro_step_count"], 0) + mock_orbax_mgr._manager.delete.assert_called_once_with(5) # pylint: disable=protected-access + call = mock_orbax_mgr.save_checkpointables.call_args + self.assertEqual(call.args[0], 5) + self.assertEqual(call.kwargs["custom_metadata"]["micro_step_count"], 0) self.assertFalse(t._resumed_mid_step) def test_update_does_not_checkpoint_when_not_resumed_mid_step(self): @@ -298,7 +298,7 @@ def test_update_does_not_checkpoint_when_not_resumed_mid_step(self): t.fwd_bwd(payload) t.update() - mock_orbax_mgr.save.assert_not_called() + mock_orbax_mgr.save_checkpointables.assert_not_called() def test_save_checkpoint_drains_inflight_throttler(self): mock_config = self.setup_config(enable_checkpointing=True) @@ -313,7 +313,7 @@ def test_save_checkpoint_drains_inflight_throttler(self): t.save_checkpoint(metadata={"step": 10}) # Checkpoint should be saved and throttler queue should be drained. - mock_orbax_mgr.save.assert_called_once() + mock_orbax_mgr.save_checkpointables.assert_called_once() self.assertTrue(t._throttler._inflight_queue.empty()) def test_save_checkpoint_called_after_fwd_bwd_before_update(self): @@ -329,18 +329,14 @@ def test_save_checkpoint_called_after_fwd_bwd_before_update(self): t.save_checkpoint(metadata=dummy_metadata) # Verify orbax save was called - mock_orbax_mgr.save.assert_called_once() - call_kwargs = mock_orbax_mgr.save.call_args.kwargs - self.assertEqual(call_kwargs["custom_metadata"]["micro_step_count"], 1) - self.assertEqual(call_kwargs["custom_metadata"]["additional_metadata"], dummy_metadata) - args_dict = ( - dict(call_kwargs["args"].items()) - if hasattr(call_kwargs["args"], "items") and callable(call_kwargs["args"].items) - else call_kwargs["args"].__dict__ - ) - self.assertIn("model_params", args_dict) - self.assertIn("accumulated_metrics", args_dict) - self.assertIn("accumulated_grads", args_dict) + mock_orbax_mgr.save_checkpointables.assert_called_once() + call = mock_orbax_mgr.save_checkpointables.call_args + self.assertEqual(call.kwargs["custom_metadata"]["micro_step_count"], 1) + self.assertEqual(call.kwargs["custom_metadata"]["additional_metadata"], dummy_metadata) + checkpointables = call.args[1] + self.assertIn("model_params", checkpointables) + self.assertIn("accumulated_metrics", checkpointables) + self.assertIn("accumulated_grads", checkpointables) def test_close_writes_final_checkpoint(self): t = maxtext_engine.MaxTextTrainingEngine(self.setup_config(enable_checkpointing=True)) @@ -350,9 +346,9 @@ def test_close_writes_final_checkpoint(self): t.close() - call_kwargs = mock_orbax_mgr.save.call_args.kwargs - self.assertEqual(call_kwargs["step"], 4) - self.assertTrue(call_kwargs["force"]) + call = mock_orbax_mgr.save_checkpointables.call_args + self.assertEqual(call.args[0], 4) + self.assertTrue(call.kwargs["force"]) def test_close_saves_an_incomplete_step(self): t = maxtext_engine.MaxTextTrainingEngine(self.setup_config(enable_checkpointing=True)) @@ -366,7 +362,7 @@ def test_close_saves_an_incomplete_step(self): t.close() - self.assertEqual(mock_orbax_mgr.save.call_args.kwargs["step"], 5) + self.assertEqual(mock_orbax_mgr.save_checkpointables.call_args.args[0], 5) def test_restore_checkpoint_no_checkpoint_returns_defaults(self): mock_config = self.setup_config(enable_checkpointing=True) @@ -382,39 +378,39 @@ def test_restore_checkpoint_restores_ckpt_metadata(self): t = maxtext_engine.MaxTextTrainingEngine(mock_config) mock_orbax_mgr = self._mock_orbax_manager(t, latest_step=10) - # Mock metadata with item_metadata and custom_metadata attributes + # Mock metadata with checkpointables metadata and custom_metadata attributes dummy_metadata = mock.MagicMock() mock_metadata = mock.MagicMock() - mock_metadata.item_metadata = {"model_params": {}, "optimizer_state": {}} + mock_metadata.metadata = {"model_params": {}, "optimizer_state": {}} mock_metadata.custom_metadata = {"additional_metadata": dummy_metadata} - mock_orbax_mgr.metadata.return_value = mock_metadata + mock_orbax_mgr.checkpointables_metadata.return_value = mock_metadata # Return dummy model and optimizer state from orbax restore dummy_model = DummyNNXModel() dummy_opt = nnx.Optimizer(dummy_model, optax.sgd(0.01), wrt=nnx.Param) dummy_opt_state = nnx.state(dummy_opt, nnx.optimizer.OptState) - mock_orbax_mgr.restore.return_value = { + mock_orbax_mgr.load_checkpointables.return_value = { "model_params": nnx.state(dummy_model), "optimizer_state": dummy_opt_state, } - t._checkpoint_manager._checkpoint_manager = mock_orbax_mgr + t._checkpoint_manager._checkpointer = mock_orbax_mgr restored_metadata = t.restore_checkpoint(step=10) self.assertEqual(t.train_step, 10) self.assertEqual(restored_metadata, dummy_metadata) - mock_orbax_mgr.restore.assert_called_once() + mock_orbax_mgr.load_checkpointables.assert_called_once() def test_restore_intra_step_checkpoint(self): mock_config = self.setup_config(enable_checkpointing=True) t = maxtext_engine.MaxTextTrainingEngine(mock_config) mock_orbax_mgr = self._mock_orbax_manager(t, latest_step=5) - # Mock metadata with item_metadata and custom_metadata attributes + # Mock metadata with checkpointables metadata and custom_metadata attributes dummy_metadata = mock.MagicMock() mock_metadata = mock.MagicMock() - mock_metadata.item_metadata = {"model_params": {}, "optimizer_state": {}} + mock_metadata.metadata = {"model_params": {}, "optimizer_state": {}} mock_metadata.custom_metadata = {"micro_step_count": 2, "additional_metadata": dummy_metadata} - mock_orbax_mgr.metadata.return_value = mock_metadata + mock_orbax_mgr.checkpointables_metadata.return_value = mock_metadata metrics_buf = abstract_engine.MetricsBuffer(id=5, mode="train") # `weighted_metrics` is a plain dict; pylint cannot resolve that through @@ -428,13 +424,13 @@ def test_restore_intra_step_checkpoint(self): dummy_model = DummyNNXModel() dummy_opt = nnx.Optimizer(dummy_model, optax.sgd(0.01), wrt=nnx.Param) dummy_opt_state = nnx.state(dummy_opt, nnx.optimizer.OptState) - mock_orbax_mgr.restore.return_value = { + mock_orbax_mgr.load_checkpointables.return_value = { "model_params": nnx.state(dummy_model), "optimizer_state": dummy_opt_state, "accumulated_metrics": [metrics_buf], "accumulated_grads": dummy_grads, } - t._checkpoint_manager._checkpoint_manager = mock_orbax_mgr + t._checkpoint_manager._checkpointer = mock_orbax_mgr _ = t.restore_checkpoint(step=5) self.assertEqual(t._micro_step_count, 2)