From 3b59c5117ee3f69020b934fabb18eb11bc35c42e Mon Sep 17 00:00:00 2001 From: Jacky Fang Date: Tue, 18 Aug 2026 08:12:18 +0000 Subject: [PATCH] fix(checkpointing): raise RuntimeError on fatal checkpointing errors --- src/maxtext/common/checkpointing.py | 4 +-- tests/unit/checkpointing_test.py | 32 +++++++++++++++++++ tests/unit/train_state_nnx_checkpoint_test.py | 15 +++++++++ 3 files changed, 49 insertions(+), 2 deletions(-) diff --git a/src/maxtext/common/checkpointing.py b/src/maxtext/common/checkpointing.py index 854c1b3968..7191190f2b 100644 --- a/src/maxtext/common/checkpointing.py +++ b/src/maxtext/common/checkpointing.py @@ -893,8 +893,8 @@ def maybe_save_checkpoint(checkpoint_manager, state, config, data_iterator, step return def _checkpoint_error_handler(err): - """Handles checkpointing errors, when not in an elastic context.""" - raise exceptions.StopTraining(f"Checkpointing failed. {str(err)}") from err + """Handles checkpointing errors.""" + raise RuntimeError(f"Checkpointing failed. {str(err)}") from err with checkpoint_exception_guard(config, checkpoint_manager, _checkpoint_error_handler): checkpoint_saved = save_checkpoint( diff --git a/tests/unit/checkpointing_test.py b/tests/unit/checkpointing_test.py index 4ef2ab30b7..9440d64821 100644 --- a/tests/unit/checkpointing_test.py +++ b/tests/unit/checkpointing_test.py @@ -480,5 +480,37 @@ async def await_creation(self): self.assertEqual(iterators_restore_v1[1].state, 20) +class CheckpointErrorHandlerTest(parameterized.TestCase): + """Tests for checkpoint error handling in maybe_save_checkpoint.""" + + def setUp(self): + super().setUp() + self.mock_manager = mock.MagicMock() + self.mock_manager.latest_step.return_value = None + self.mock_manager.reached_preemption.return_value = False + self.state = mock.Mock() + + def test_error_handler_raises_runtime_error(self): + """Unexpected checkpointing errors should raise RuntimeError with original error chained.""" + config = mock.Mock() + config.checkpoint_period = 1 + config.pure_nnx = True + config.enable_diloco = False + config.async_checkpointing = False + config.enable_continuous_checkpointing = False + config.enable_emergency_checkpoint = False + config.enable_multi_tier_checkpointing = False + config.local_checkpoint_period = 0 + config.enable_autocheckpoint = False + config.elastic_enabled = False + + original_error = RuntimeError("GCS failure") + with mock.patch.object(checkpointing, "save_checkpoint", side_effect=original_error): + with self.assertRaises(RuntimeError) as cm: + checkpointing.maybe_save_checkpoint(self.mock_manager, self.state, config, data_iterator=None, step=1) + self.assertIn("Checkpointing failed. GCS failure", str(cm.exception)) + self.assertIs(cm.exception.__cause__, original_error) + + if __name__ == "__main__": absltest.main() diff --git a/tests/unit/train_state_nnx_checkpoint_test.py b/tests/unit/train_state_nnx_checkpoint_test.py index 172bd09a72..301b1098c6 100644 --- a/tests/unit/train_state_nnx_checkpoint_test.py +++ b/tests/unit/train_state_nnx_checkpoint_test.py @@ -624,6 +624,21 @@ def test_maybe_save_checkpoint_checks_scale_up_after_unsaved_dispatch(self): save_checkpoint_mock.assert_called_once() mock_maybe_scale_up.assert_called_once_with(config, mgr) + def test_maybe_save_checkpoint_error_handler_raises_runtime_error(self): + """Fatal checkpoint errors raise RuntimeError with the original cause attached.""" + state = mock.Mock() + config = self._config(checkpoint_period=1) + mgr = mock.MagicMock() + mgr.latest_step.return_value = None + mgr.reached_preemption.return_value = False + + original_error = RuntimeError("Disk I/O error") + with mock.patch.object(checkpointing, "save_checkpoint", side_effect=original_error): + with self.assertRaises(RuntimeError) as cm: + checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=5) + self.assertIn("Checkpointing failed. Disk I/O error", str(cm.exception)) + self.assertIs(cm.exception.__cause__, original_error) + class TestLinenCheckpointFormatConverters(unittest.TestCase): """to_linen_checkpoint_dict / from_linen_checkpoint_dict (NNX <-> Linen on-disk layout)."""