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..6995762b128 100644 --- a/tests/unit_tests/cpu/test_checkpoint.py +++ b/tests/unit_tests/cpu/test_checkpoint.py @@ -43,6 +43,20 @@ from torchtitan.quantization._fsdp_tensor import _ShardedFSDPTensor +# 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: """A fake OptimizersContainer that returns fake state dicts.""" @@ -452,10 +466,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,17 +495,20 @@ 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) + # The initial load actually happened, so there is no skip to report. + for level in (mock_logger.info, mock_logger.warning): + self.assertFalse(initial_load_skip_logs(level)) 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. 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") @@ -519,7 +537,50 @@ 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() + 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(initial_load_skip_logs(mock_logger.warning)) + 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) + 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(initial_load_skip_logs(mock_logger.warning)) manager.close() @mock.patch("torch.distributed.get_rank", return_value=0) @@ -1303,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" diff --git a/torchtitan/components/checkpointer/base.py b/torchtitan/components/checkpointer/base.py index 5e3dc880a5f..c67356ddf9a 100644 --- a/torchtitan/components/checkpointer/base.py +++ b/torchtitan/components/checkpointer/base.py @@ -305,6 +305,22 @@ 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. 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.info( + "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()