From 27696242fa49b36f05ebf7d10b1e66c77bdc79bd Mon Sep 17 00:00:00 2001 From: YeonwooSung Date: Sat, 12 Sep 2026 17:03:03 +0900 Subject: [PATCH 1/5] [checkpoint] Warn when initial load is skipped on resume Resuming from checkpoint.folder is the fault-tolerance path and still wins over initial_load_in_hf / initial_load_path. Tell the user those options were ignored. Addresses #1900. --- docs/checkpoint.md | 2 +- tests/unit_tests/cpu/test_checkpoint.py | 53 ++++++++++++++++++++-- torchtitan/components/checkpointer/base.py | 15 ++++++ 3 files changed, 65 insertions(+), 5 deletions(-) diff --git a/docs/checkpoint.md b/docs/checkpoint.md index f5c4b7f8c35..7cfe093d510 100644 --- a/docs/checkpoint.md +++ b/docs/checkpoint.md @@ -69,7 +69,7 @@ NGPU=1 ./run_train.sh --module --config --checkpoint ### HuggingFace `torchtitan` offers two ways to work with Hugging Face models: either by directly saving and loading a Hugging Face checkpoint during training, or by using an example conversion script to directly reformat the model weights on cpu. -1. You can directly save huggingface model weights during training by using the `--checkpoint.last_save_in_hf` and `--checkpoint.last_save_model_only` options together. To directly load a `torchtitan` training session from a huggingface safetensors file, enable `--checkpoint.initial_load_in_hf`, and set either `--hf_assets_path` or `--checkpoint.initial_load_path` to the directory containing the huggingface checkpoint. `--checkpoint.initial_load_path` overrides `--hf_assets_path` if both are set. +1. You can directly save huggingface model weights during training by using the `--checkpoint.last_save_in_hf` and `--checkpoint.last_save_model_only` options together. To directly load a `torchtitan` training session from a huggingface safetensors file, enable `--checkpoint.initial_load_in_hf`, and set either `--hf_assets_path` or `--checkpoint.initial_load_path` to the directory containing the huggingface checkpoint. `--checkpoint.initial_load_path` overrides `--hf_assets_path` if both are set. If `checkpoint.folder` already contains a valid checkpoint, training resumes from that folder and ignores `initial_load_in_hf` / `initial_load_path` (fault-tolerance restart). The first run (empty folder) uses the initial load. 2. To directly reformat the weights without the need to run a training loop, run the corresponding conversion script. The naming scheme is `torchtitan`-centric, e.g. convert_from_hf means convert hf->tt. diff --git a/tests/unit_tests/cpu/test_checkpoint.py b/tests/unit_tests/cpu/test_checkpoint.py index 2cbad8751b3..c03206a14cc 100644 --- a/tests/unit_tests/cpu/test_checkpoint.py +++ b/tests/unit_tests/cpu/test_checkpoint.py @@ -482,15 +482,15 @@ def test_initial_load_path_used_when_folder_has_no_valid_checkpoints( self.assertTrue(res) manager.close() - @mock.patch("torchtitan.components.checkpointer.dcp.logger") + @mock.patch("torchtitan.components.checkpointer.base.logger") @mock.patch("torch.distributed.get_rank", return_value=0) @mock.patch.object(dist_checkpoint, "load") def test_initial_load_path_ignored_when_folder_has_valid_checkpoints( self, mock_load, mock_rank, mock_logger ): # Resuming from checkpoint.folder is the fault-tolerance path: all - # initial_* options are silently ignored so a job can keep the same - # arguments across automatic restarts. + # initial_* options are ignored so a job can keep the same arguments + # across automatic restarts. Warn so users notice the skip. initial_load_path = os.path.join(self.base_temp_dir, "initial", "step-100") os.makedirs(initial_load_path, exist_ok=True) ckpt_folder = os.path.join(self.test_folder, "checkpoints") @@ -519,7 +519,52 @@ def test_initial_load_path_ignored_when_folder_has_valid_checkpoints( mock_load.assert_called_once() _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) - mock_logger.warning.assert_not_called() + mock_logger.warning.assert_called() + warning_text = " ".join( + str(arg) for call in mock_logger.warning.call_args_list for arg in call.args + ) + self.assertRegex(warning_text, r"initial_load|ignored") + self.assertRegex(warning_text, r"folder|step") + manager.close() + + @mock.patch("torchtitan.components.checkpointer.base.logger") + @mock.patch("torch.distributed.get_rank", return_value=0) + @mock.patch.object(dist_checkpoint, "load") + def test_initial_load_in_hf_ignored_when_folder_has_valid_checkpoints( + self, mock_load, mock_rank, mock_logger + ): + ckpt_folder = os.path.join(self.test_folder, "checkpoints") + step_dir = os.path.join(ckpt_folder, "step-5") + os.makedirs(step_dir, exist_ok=True) + open(os.path.join(step_dir, ".metadata"), "w").close() + + cfg = self.trainer_config.checkpoint + cfg.folder = "checkpoints" + cfg.initial_load_in_hf = True + cfg.initial_load_model_only = True + manager = CheckpointManager( + dataloader=self.data_loader, + model_parts=self.model_parts, + optimizers=self.optimizers, + lr_schedulers=self.lr_schedulers, + states=self.states, + config=self.trainer_config.checkpoint, + sd_adapter=None, + base_folder=self.trainer_config.dump_folder, + ) + + res = manager.load(step=-1) + + self.assertTrue(res) + mock_load.assert_called_once() + _, kwargs = mock_load.call_args + self.assertEqual(kwargs.get("checkpoint_id"), step_dir) + mock_logger.warning.assert_called() + warning_text = " ".join( + str(arg) for call in mock_logger.warning.call_args_list for arg in call.args + ) + self.assertRegex(warning_text, r"initial_load|ignored") + self.assertRegex(warning_text, r"folder|step") manager.close() @mock.patch("torch.distributed.get_rank", return_value=0) diff --git a/torchtitan/components/checkpointer/base.py b/torchtitan/components/checkpointer/base.py index 5e3dc880a5f..e310e8e70f9 100644 --- a/torchtitan/components/checkpointer/base.py +++ b/torchtitan/components/checkpointer/base.py @@ -305,6 +305,21 @@ def load(self, step: int = -1) -> bool: raise FileNotFoundError( f"--checkpoint.load_step={step} not found at {checkpoint_id}" ) + # Fault-tolerance restart: an existing folder checkpoint wins + # over initial_* so the same job args can be reused. + if ( + self.initial_load_path + or self.initial_load_in_hf + or self.initial_load_in_hf_quantized + ): + logger.warning( + "Resuming from checkpoint.folder %s at step %s " + "(fault-tolerance restart); ignoring " + "initial_load_path / initial_load_in_hf / " + "initial_load_in_hf_quantized.", + self.folder, + step, + ) logger.info("Loading the checkpoint from %s.", checkpoint_id) begin = time.monotonic() From 47231c215ae422a33dc4cba512acd4c67a6c6d55 Mon Sep 17 00:00:00 2001 From: YeonwooSung Date: Sat, 12 Sep 2026 17:18:02 +0900 Subject: [PATCH 2/5] [checkpoint] Lock initial-load warning to the resume path Empty-folder loads must not warn. Resume warnings must render the actual step, not the sentinel -1. --- tests/unit_tests/cpu/test_checkpoint.py | 22 +++++++++++----------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/tests/unit_tests/cpu/test_checkpoint.py b/tests/unit_tests/cpu/test_checkpoint.py index c03206a14cc..a73276d2e3c 100644 --- a/tests/unit_tests/cpu/test_checkpoint.py +++ b/tests/unit_tests/cpu/test_checkpoint.py @@ -452,10 +452,11 @@ def test_load_finds_latest_and_calls_dcp_load(self, mock_load, mock_rank): self.assertTrue(res) manager.close() + @mock.patch("torchtitan.components.checkpointer.base.logger") @mock.patch("torch.distributed.get_rank", return_value=0) @mock.patch.object(dist_checkpoint, "load") def test_initial_load_path_used_when_folder_has_no_valid_checkpoints( - self, mock_load, mock_rank + self, mock_load, mock_rank, mock_logger ): initial_load_path = os.path.join(self.base_temp_dir, "initial", "step-100") os.makedirs(initial_load_path, exist_ok=True) @@ -480,6 +481,7 @@ def test_initial_load_path_used_when_folder_has_no_valid_checkpoints( _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), initial_load_path) self.assertTrue(res) + mock_logger.warning.assert_not_called() manager.close() @mock.patch("torchtitan.components.checkpointer.base.logger") @@ -520,11 +522,10 @@ def test_initial_load_path_ignored_when_folder_has_valid_checkpoints( _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) mock_logger.warning.assert_called() - warning_text = " ".join( - str(arg) for call in mock_logger.warning.call_args_list for arg in call.args - ) - self.assertRegex(warning_text, r"initial_load|ignored") - self.assertRegex(warning_text, r"folder|step") + fmt, *args = mock_logger.warning.call_args.args + rendered = fmt % tuple(args) + self.assertIn("initial_load", rendered) + self.assertIn("step 5", rendered) manager.close() @mock.patch("torchtitan.components.checkpointer.base.logger") @@ -560,11 +561,10 @@ def test_initial_load_in_hf_ignored_when_folder_has_valid_checkpoints( _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) mock_logger.warning.assert_called() - warning_text = " ".join( - str(arg) for call in mock_logger.warning.call_args_list for arg in call.args - ) - self.assertRegex(warning_text, r"initial_load|ignored") - self.assertRegex(warning_text, r"folder|step") + fmt, *args = mock_logger.warning.call_args.args + rendered = fmt % tuple(args) + self.assertIn("initial_load", rendered) + self.assertIn("step 5", rendered) manager.close() @mock.patch("torch.distributed.get_rank", return_value=0) From 4c6eb6838d2218fb6ab6297e122a1adbbee862b5 Mon Sep 17 00:00:00 2001 From: YeonwooSung Date: Tue, 15 Sep 2026 11:30:01 +0900 Subject: [PATCH 3/5] [checkpoint] Log the initial-load skip at info, not warning Resuming from checkpoint.folder is the fault-tolerance restart path, and it runs on every automatic restart of a long job. That is expected behavior, not a misconfiguration, so a warning on each restart is noise. The existing warnings in this file cover settings that can never take effect (load_only with enable_first_step_checkpoint, initial_load_model_only without initial_load_path); this one is different in kind. The message itself is unchanged, so the skip is still visible in the log, which is what #1900 needed. Tests now assert the skip message is emitted exactly once at info level and never as a warning. They match on the rendered message rather than on the last call, because the load path emits other info lines. --- tests/unit_tests/cpu/test_checkpoint.py | 51 +++++++++++++++++----- torchtitan/components/checkpointer/base.py | 5 ++- 2 files changed, 42 insertions(+), 14 deletions(-) diff --git a/tests/unit_tests/cpu/test_checkpoint.py b/tests/unit_tests/cpu/test_checkpoint.py index a73276d2e3c..f8518da5706 100644 --- a/tests/unit_tests/cpu/test_checkpoint.py +++ b/tests/unit_tests/cpu/test_checkpoint.py @@ -43,6 +43,13 @@ from torchtitan.quantization._fsdp_tensor import _ShardedFSDPTensor +def rendered_log_calls(mock_log_method) -> list[str]: + """Render every %-style call made to a mocked logger method.""" + return [ + fmt % tuple(args) for fmt, *args in (c.args for c in mock_log_method.mock_calls) + ] + + class FakeOptimizersContainer: """A fake OptimizersContainer that returns fake state dicts.""" @@ -481,7 +488,11 @@ def test_initial_load_path_used_when_folder_has_no_valid_checkpoints( _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), initial_load_path) self.assertTrue(res) - mock_logger.warning.assert_not_called() + # The initial load actually happened, so there is no skip to report. + for level in (mock_logger.info, mock_logger.warning): + self.assertFalse( + [msg for msg in rendered_log_calls(level) if "initial_load" in msg] + ) manager.close() @mock.patch("torchtitan.components.checkpointer.base.logger") @@ -492,7 +503,7 @@ def test_initial_load_path_ignored_when_folder_has_valid_checkpoints( ): # Resuming from checkpoint.folder is the fault-tolerance path: all # initial_* options are ignored so a job can keep the same arguments - # across automatic restarts. Warn so users notice the skip. + # across automatic restarts. Log it so users can see the skip. initial_load_path = os.path.join(self.base_temp_dir, "initial", "step-100") os.makedirs(initial_load_path, exist_ok=True) ckpt_folder = os.path.join(self.test_folder, "checkpoints") @@ -521,11 +532,19 @@ def test_initial_load_path_ignored_when_folder_has_valid_checkpoints( mock_load.assert_called_once() _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) - mock_logger.warning.assert_called() - fmt, *args = mock_logger.warning.call_args.args - rendered = fmt % tuple(args) - self.assertIn("initial_load", rendered) - self.assertIn("step 5", rendered) + skip_messages = [ + msg for msg in rendered_log_calls(mock_logger.info) if "initial_load" in msg + ] + self.assertEqual(len(skip_messages), 1) + self.assertIn("step 5", skip_messages[0]) + # The skip is the normal fault-tolerance restart, so it must not warn. + self.assertFalse( + [ + msg + for msg in rendered_log_calls(mock_logger.warning) + if "initial_load" in msg + ] + ) manager.close() @mock.patch("torchtitan.components.checkpointer.base.logger") @@ -560,11 +579,19 @@ def test_initial_load_in_hf_ignored_when_folder_has_valid_checkpoints( mock_load.assert_called_once() _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) - mock_logger.warning.assert_called() - fmt, *args = mock_logger.warning.call_args.args - rendered = fmt % tuple(args) - self.assertIn("initial_load", rendered) - self.assertIn("step 5", rendered) + skip_messages = [ + msg for msg in rendered_log_calls(mock_logger.info) if "initial_load" in msg + ] + self.assertEqual(len(skip_messages), 1) + self.assertIn("step 5", skip_messages[0]) + # The skip is the normal fault-tolerance restart, so it must not warn. + self.assertFalse( + [ + msg + for msg in rendered_log_calls(mock_logger.warning) + if "initial_load" in msg + ] + ) manager.close() @mock.patch("torch.distributed.get_rank", return_value=0) diff --git a/torchtitan/components/checkpointer/base.py b/torchtitan/components/checkpointer/base.py index e310e8e70f9..c67356ddf9a 100644 --- a/torchtitan/components/checkpointer/base.py +++ b/torchtitan/components/checkpointer/base.py @@ -306,13 +306,14 @@ def load(self, step: int = -1) -> bool: f"--checkpoint.load_step={step} not found at {checkpoint_id}" ) # Fault-tolerance restart: an existing folder checkpoint wins - # over initial_* so the same job args can be reused. + # over initial_* so the same job args can be reused. This is + # the normal restart path, so it is logged, not warned about. if ( self.initial_load_path or self.initial_load_in_hf or self.initial_load_in_hf_quantized ): - logger.warning( + logger.info( "Resuming from checkpoint.folder %s at step %s " "(fault-tolerance restart); ignoring " "initial_load_path / initial_load_in_hf / " From 92adda799361fb120e0e820ac2b60e3e7cdf847a Mon Sep 17 00:00:00 2001 From: YeonwooSung Date: Wed, 16 Sep 2026 11:32:20 +0900 Subject: [PATCH 4/5] [checkpoint] Match the skip log by phrase, not by "initial_load" The tests name the checkpoint folder after the running test, so the folder path contains "initial_load" and the substring filter also matched "Loading the checkpoint from ." That made the positive tests see two skip messages instead of one. Filter on "ignoring initial_load_path" instead. The phrase has spaces, so a path cannot match it. Moving the filter into the helper also drops the duplicated comprehensions at the call sites. --- tests/unit_tests/cpu/test_checkpoint.py | 43 +++++++++---------------- 1 file changed, 16 insertions(+), 27 deletions(-) diff --git a/tests/unit_tests/cpu/test_checkpoint.py b/tests/unit_tests/cpu/test_checkpoint.py index f8518da5706..e10dc1e0c5b 100644 --- a/tests/unit_tests/cpu/test_checkpoint.py +++ b/tests/unit_tests/cpu/test_checkpoint.py @@ -43,11 +43,18 @@ from torchtitan.quantization._fsdp_tensor import _ShardedFSDPTensor -def rendered_log_calls(mock_log_method) -> list[str]: - """Render every %-style call made to a mocked logger method.""" - return [ +# These tests name the checkpoint folder after the running test, so a bare +# "initial_load" also matches log lines that only echo that path back. Match on +# a phrase from the message instead; spaces keep it from matching a path. +INITIAL_LOAD_SKIP_MARKER = "ignoring initial_load_path" + + +def initial_load_skip_logs(mock_log_method) -> list[str]: + """Render the %-style calls to a mocked logger and keep the skip messages.""" + rendered = ( fmt % tuple(args) for fmt, *args in (c.args for c in mock_log_method.mock_calls) - ] + ) + return [msg for msg in rendered if INITIAL_LOAD_SKIP_MARKER in msg] class FakeOptimizersContainer: @@ -490,9 +497,7 @@ def test_initial_load_path_used_when_folder_has_no_valid_checkpoints( self.assertTrue(res) # The initial load actually happened, so there is no skip to report. for level in (mock_logger.info, mock_logger.warning): - self.assertFalse( - [msg for msg in rendered_log_calls(level) if "initial_load" in msg] - ) + self.assertFalse(initial_load_skip_logs(level)) manager.close() @mock.patch("torchtitan.components.checkpointer.base.logger") @@ -532,19 +537,11 @@ def test_initial_load_path_ignored_when_folder_has_valid_checkpoints( mock_load.assert_called_once() _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) - skip_messages = [ - msg for msg in rendered_log_calls(mock_logger.info) if "initial_load" in msg - ] + skip_messages = initial_load_skip_logs(mock_logger.info) self.assertEqual(len(skip_messages), 1) self.assertIn("step 5", skip_messages[0]) # The skip is the normal fault-tolerance restart, so it must not warn. - self.assertFalse( - [ - msg - for msg in rendered_log_calls(mock_logger.warning) - if "initial_load" in msg - ] - ) + self.assertFalse(initial_load_skip_logs(mock_logger.warning)) manager.close() @mock.patch("torchtitan.components.checkpointer.base.logger") @@ -579,19 +576,11 @@ def test_initial_load_in_hf_ignored_when_folder_has_valid_checkpoints( mock_load.assert_called_once() _, kwargs = mock_load.call_args self.assertEqual(kwargs.get("checkpoint_id"), step_dir) - skip_messages = [ - msg for msg in rendered_log_calls(mock_logger.info) if "initial_load" in msg - ] + skip_messages = initial_load_skip_logs(mock_logger.info) self.assertEqual(len(skip_messages), 1) self.assertIn("step 5", skip_messages[0]) # The skip is the normal fault-tolerance restart, so it must not warn. - self.assertFalse( - [ - msg - for msg in rendered_log_calls(mock_logger.warning) - if "initial_load" in msg - ] - ) + self.assertFalse(initial_load_skip_logs(mock_logger.warning)) manager.close() @mock.patch("torch.distributed.get_rank", return_value=0) From a99039ebf113f4432e4d4c5ac4ff37f7e63ca1f2 Mon Sep 17 00:00:00 2001 From: YeonwooSung Date: Wed, 16 Sep 2026 12:19:34 +0900 Subject: [PATCH 5/5] [checkpoint] Give the tracing fixture the initial_load_* attributes TestBaseCheckpointManagerTracing builds the manager as Mock(spec=BaseCheckpointManager). The spec only covers class attributes, and initial_load_path / initial_load_in_hf / initial_load_in_hf_quantized are set in __init__, so reading them from load() raised AttributeError. Set all three on the fixture. The values keep the new branch untaken, so the trace-span ordering and no_grad assertions are unaffected. --- tests/unit_tests/cpu/test_checkpoint.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tests/unit_tests/cpu/test_checkpoint.py b/tests/unit_tests/cpu/test_checkpoint.py index e10dc1e0c5b..6995762b128 100644 --- a/tests/unit_tests/cpu/test_checkpoint.py +++ b/tests/unit_tests/cpu/test_checkpoint.py @@ -1364,6 +1364,10 @@ def _manager(self, *, enable: bool = True): manager.enable = enable manager._save.return_value = True manager.folder = "/checkpoint" + # Set by __init__, so spec= does not cover them, but load() reads them. + manager.initial_load_path = None + manager.initial_load_in_hf = False + manager.initial_load_in_hf_quantized = False manager._storage = mock.Mock() manager._storage.isdir.return_value = True manager._create_checkpoint_id.return_value = "/checkpoint/step-10"