diff --git a/tests/unit_tests/test_torch_checkpointing.py b/tests/unit_tests/test_torch_checkpointing.py index 835aeda5dc..ecaef51743 100644 --- a/tests/unit_tests/test_torch_checkpointing.py +++ b/tests/unit_tests/test_torch_checkpointing.py @@ -14,6 +14,7 @@ from concurrent.futures import Future from contextlib import nullcontext from pathlib import Path +from typing import Any from unittest import mock import torch @@ -32,6 +33,7 @@ SyncCheckpointSaverConfig, ) from torch_checkpointing.default_resharder import DefaultResharder +from torch_checkpointing.hf.resharder import HFSafetensorsDTensorResharder from torch_checkpointing.logging_utils import checkpoint_logging_context from torch_checkpointing.schema import ItemSpec from torch_checkpointing.storage.filesystem import LocalFileSystemStorageConfig @@ -54,7 +56,7 @@ def __init__(self) -> None: self.closed = False self.lock_calls = 0 self.load_calls = [] - self.load_result = None + self.load_result: Any = None self.prewarm_calls = [] self.save_calls = [] self.save_result = Future() @@ -90,13 +92,22 @@ def load_state_dict(self, state_dict) -> None: class _StateDictAdapter: - def __init__(self) -> None: + def __init__(self, hf_assets_path: str | None = None) -> None: self.fqn_to_index_mapping = {"hf_weight": 1} + self.hf_assets_path = hf_assets_path + self.from_hf_calls = [] self.to_hf_calls = [] + self.to_hf_results = [] def to_hf(self, state_dict): self.to_hf_calls.append(state_dict) - return {"hf_weight": state_dict["weight"]} + result = {"hf_weight": state_dict["weight"]} + self.to_hf_results.append(result) + return result + + def from_hf(self, state_dict): + self.from_hf_calls.append(state_dict) + return {"weight": state_dict["hf_weight"]} class TorchCheckpointingManagerTest(unittest.TestCase): @@ -110,6 +121,7 @@ def _build_manager( model_parts=None, optimizers=None, states=None, + sd_adapter=None, ) -> tuple[TorchCheckpointingManager, _BackendManager]: if backend_config is None: backend_config = _default_backend_config() @@ -132,7 +144,7 @@ def _build_manager( optimizers=optimizers or _Stateful("optimizer"), lr_schedulers=_Stateful("scheduler"), states=states or {"train_state": _Stateful("train")}, - sd_adapter=None, + sd_adapter=sd_adapter, base_folder=base_folder, storage_config=storage_config, ) @@ -639,7 +651,11 @@ def test_a_finished_hf_export_is_a_valid_checkpoint(self) -> None: manager = TorchCheckpointingManager.__new__(TorchCheckpointingManager) manager._storage = mock.Mock(spec=CheckpointStorage) - for marker in ("metadata.pkl", "model.safetensors.index.json"): + for marker in ( + "metadata.pkl", + "model.safetensors.index.json", + "model.safetensors", + ): with self.subTest(marker=marker): manager._storage.isfile.side_effect = ( lambda path, marker=marker: path.endswith(marker) @@ -819,19 +835,213 @@ def test_hf_final_save_converts_and_consolidates_before_commit( ) manager.close() + def test_hf_load_uses_temporary_model_only_manager_and_restores_model( + self, + ) -> None: + with tempfile.TemporaryDirectory() as base_folder: + checkpoint_id = os.path.join(base_folder, "hf_checkpoint") + os.makedirs(checkpoint_id) + with open( + os.path.join(checkpoint_id, "model.safetensors.index.json"), + "w", + ): + pass + model = nn.Linear(2, 2, bias=False) + adapter = _StateDictAdapter(hf_assets_path=checkpoint_id) + config = TorchCheckpointingManager.Config( + enable=True, + keep_latest_k=0, + initial_load_model_only=True, + initial_load_in_hf=True, + ) + backend_config = _default_backend_config() + backend_manager = _BackendManager() + hf_manager = _BackendManager() + expected_weight = torch.full_like(model.weight, 3) + hf_manager.load_result = { + MODEL: {"hf_weight": expected_weight}, + } + + with ( + mock.patch.object( + manager_module, + "_default_backend_config", + return_value=backend_config, + ), + mock.patch.object( + BackendCheckpointManager.Config, + "build", + autospec=True, + side_effect=[backend_manager, hf_manager], + ) as build, + ): + manager = config.build( + dataloader=None, + model_parts=[model], + optimizers=_Stateful("optimizer"), + lr_schedulers=_Stateful("scheduler"), + states={"train_state": _Stateful("train")}, + sd_adapter=adapter, + base_folder=base_folder, + ) + + self.assertTrue(manager.load()) + + self.assertEqual([], backend_manager.load_calls) + self.assertEqual(1, len(hf_manager.load_calls)) + loaded_id, into, kwargs = hf_manager.load_calls[0] + self.assertEqual(checkpoint_id, loaded_id) + self.assertIs(adapter.to_hf_results[0], into[MODEL]) + self.assertEqual({"strict": True}, kwargs) + self.assertEqual(1, len(adapter.from_hf_calls)) + self.assertIs( + expected_weight, + adapter.from_hf_calls[0]["hf_weight"], + ) + torch.testing.assert_close(model.weight, expected_weight) + + hf_config = build.call_args_list[1].args[0] + self.assertIsInstance( + manager._manager_config.save, + AsyncCheckpointSaverConfig, + ) + self.assertIsInstance(hf_config.save, SyncCheckpointSaverConfig) + self.assertIsNone(hf_config.save.writer_config.barrier_config) + self.assertEqual({MODEL}, set(hf_config.items)) + self.assertIsNone(hf_config.default) + self.assertIsInstance( + hf_config.items[MODEL].resharder, + HFSafetensorsDTensorResharder, + ) + self.assertEqual( + manager._manager_config.items[MODEL].requires_copy, + hf_config.items[MODEL].requires_copy, + ) + self.assertEqual( + manager._manager_config.items[MODEL].layout, + hf_config.items[MODEL].layout, + ) + self.assertEqual( + manager._manager_config.items[MODEL].required, + hf_config.items[MODEL].required, + ) + self.assertIsInstance( + manager._manager_config.items[MODEL].resharder, + DefaultResharder, + ) + self.assertNotIsInstance( + manager._manager_config.items[MODEL].resharder, + HFSafetensorsDTensorResharder, + ) + self.assertIs( + manager._manager_config.storage_config, + hf_config.storage_config, + ) + self.assertTrue(hf_manager.closed) + self.assertFalse(backend_manager.closed) + manager.close() + self.assertTrue(backend_manager.closed) + + def test_hf_load_closes_temporary_manager_when_load_fails(self) -> None: + with tempfile.TemporaryDirectory() as base_folder: + checkpoint_id = os.path.join(base_folder, "hf_checkpoint") + os.makedirs(checkpoint_id) + with open(os.path.join(checkpoint_id, "model.safetensors"), "wb"): + pass + adapter = _StateDictAdapter() + config = TorchCheckpointingManager.Config( + enable=True, + keep_latest_k=0, + initial_load_path=checkpoint_id, + initial_load_model_only=True, + initial_load_in_hf=True, + load_only=True, + ) + backend_manager = _BackendManager() + hf_manager = _BackendManager() + hf_manager.load = mock.Mock(side_effect=RuntimeError("load failed")) + + with ( + mock.patch.object( + manager_module, + "_default_backend_config", + return_value=_default_backend_config(), + ), + mock.patch.object( + BackendCheckpointManager.Config, + "build", + autospec=True, + side_effect=[backend_manager, hf_manager], + ), + ): + manager = config.build( + dataloader=None, + model_parts=[nn.Linear(2, 2)], + optimizers=_Stateful("optimizer"), + lr_schedulers=_Stateful("scheduler"), + states={"train_state": _Stateful("train")}, + sd_adapter=adapter, + base_folder=base_folder, + ) + + with self.assertRaisesRegex(RuntimeError, "load failed"): + manager.load() + + self.assertTrue(hf_manager.closed) + self.assertEqual([], adapter.from_hf_calls) + manager.close() + + def test_hf_load_rejects_quantized_checkpoint(self) -> None: + with tempfile.TemporaryDirectory() as base_folder: + checkpoint_id = os.path.join(base_folder, "hf_checkpoint") + os.makedirs(checkpoint_id) + with open(os.path.join(checkpoint_id, "model.safetensors"), "wb"): + pass + config = TorchCheckpointingManager.Config( + enable=True, + keep_latest_k=0, + initial_load_path=checkpoint_id, + initial_load_model_only=True, + initial_load_in_hf=True, + initial_load_in_hf_quantized=True, + load_only=True, + ) + manager, backend_manager = self._build_manager( + config, + base_folder=base_folder, + sd_adapter=_StateDictAdapter(), + ) + + with self.assertRaisesRegex(ValueError, "quantized"): + manager.load() + + self.assertEqual([], backend_manager.load_calls) + manager.close() + def test_native_load_restores_model_and_optimizer(self) -> None: with tempfile.TemporaryDirectory() as base_folder: checkpoint_id = os.path.join(base_folder, "checkpoint", "step-5") os.makedirs(checkpoint_id) with open(os.path.join(checkpoint_id, "metadata.pkl"), "wb"): pass + hf_export_id = os.path.join(base_folder, "checkpoint", "step-8") + os.makedirs(hf_export_id) + with open(os.path.join(hf_export_id, "model.safetensors"), "wb"): + pass + hf_checkpoint_id = os.path.join(base_folder, "hf_checkpoint") + os.makedirs(hf_checkpoint_id) + with open(os.path.join(hf_checkpoint_id, "model.safetensors"), "wb"): + pass model = nn.Linear(2, 2, bias=False) optimizer = _Stateful("optimizer") + adapter = _StateDictAdapter() config = TorchCheckpointingManager.Config( enable=True, folder="checkpoint", keep_latest_k=0, - initial_load_model_only=False, + initial_load_path=hf_checkpoint_id, + initial_load_model_only=True, + initial_load_in_hf=True, load_only=True, ) manager, backend_manager = self._build_manager( @@ -839,6 +1049,7 @@ def test_native_load_restores_model_and_optimizer(self) -> None: base_folder=base_folder, model_parts=[model], optimizers=optimizer, + sd_adapter=adapter, ) expected_weight = torch.full_like(model.weight, 3) backend_manager.load_result = { @@ -846,13 +1057,14 @@ def test_native_load_restores_model_and_optimizer(self) -> None: OPTIMIZER: {"value": "restored"}, } - self.assertTrue(manager.load(step=5)) + self.assertTrue(manager.load()) torch.testing.assert_close(model.weight, expected_weight) self.assertEqual("restored", optimizer.value) self.assertEqual(checkpoint_id, backend_manager.load_calls[0][0]) self.assertEqual(set(manager.states), set(backend_manager.load_calls[0][1])) self.assertEqual({"strict": True}, backend_manager.load_calls[0][2]) + self.assertEqual([], adapter.to_hf_calls) manager.close() def test_native_load_requires_every_requested_key(self) -> None: diff --git a/torchtitan/components/checkpointer/torch_checkpointing.py b/torchtitan/components/checkpointer/torch_checkpointing.py index 07ab30779f..fcb08f429f 100644 --- a/torchtitan/components/checkpointer/torch_checkpointing.py +++ b/torchtitan/components/checkpointer/torch_checkpointing.py @@ -39,6 +39,7 @@ METADATA_FILE_NAME as TORCH_CHECKPOINTING_METADATA_FILE_NAME, ) from torch_checkpointing.hf.consolidation import consolidate_hf_safetensors_checkpoint +from torch_checkpointing.hf.resharder import HFSafetensorsDTensorResharder from torch_checkpointing.logging_utils import checkpoint_logging_context from torch_checkpointing.schema import ItemSpec from torch_checkpointing.staging import CheckpointStagerConfig @@ -73,6 +74,7 @@ # Index the HF consolidation writes at the root of a final export; the # backend names it after the checkpoint item it consolidated. _HF_INDEX_FILE_NAME = f"{MODEL}.safetensors.index.json" +_HF_SINGLE_FILE_NAME = f"{MODEL}.safetensors" # Logger the backend emits its checkpoint events and metrics on. _BACKEND_LOGGER_NAME = "torch_checkpointing" @@ -380,10 +382,11 @@ def __del__(self) -> None: @sl.log_trace_span("checkpoint_load") @torch.no_grad() def _load(self, step: int = -1) -> bool: + from_hf = False has_checkpoint_folder = self._storage.isdir(self.folder) load_step = -1 if has_checkpoint_folder: - load_step = self._find_load_step() if step == -1 else step + load_step = self._find_native_load_step() if step == -1 else step if step != -1 and not has_checkpoint_folder: raise FileNotFoundError( f"--checkpoint.load_step={step} not found because " @@ -391,16 +394,32 @@ def _load(self, step: int = -1) -> bool: ) if load_step == -1: - if self.initial_load_in_hf: - raise ValueError( - "TorchCheckpointingManager does not yet support loading " - "Hugging Face checkpoints." - ) - if not self.initial_load_path: + from_hf = self.initial_load_in_hf + if from_hf: + if self.initial_load_in_hf_quantized: + raise ValueError( + "TorchCheckpointingManager does not support loading " + "quantized Hugging Face checkpoints." + ) + if self.sd_adapter is None: + raise ValueError( + "checkpoint.initial_load_in_hf is True, but sd_adapter " + "is not provided." + ) + checkpoint_id = self.initial_load_path or self.sd_adapter.hf_assets_path + if not checkpoint_id: + raise ValueError( + "checkpoint.initial_load_in_hf requires either " + "checkpoint.initial_load_path or model.hf_assets_path." + ) + model_only = True + elif not self.initial_load_path: logger.info("No checkpoint was provided, this is a fresh start.") return False - checkpoint_id = self.initial_load_path - model_only = self.initial_load_model_only + else: + checkpoint_id = self.initial_load_path + model_only = self.initial_load_model_only + if not self._storage.isdir(checkpoint_id): raise ValueError( f"Checkpoint.initial_load_path is invalid: {checkpoint_id}" @@ -415,9 +434,14 @@ def _load(self, step: int = -1) -> bool: f"--checkpoint.load_step={step} not found at {checkpoint_id}" ) - if not self._is_valid_checkpoint(checkpoint_id): + is_valid_checkpoint = ( + self._is_hf_checkpoint(checkpoint_id) + if from_hf + else self._is_native_checkpoint(checkpoint_id) + ) + if not is_valid_checkpoint: raise ValueError( - f"Checkpoint {checkpoint_id!r} is not a native " + f"Checkpoint {checkpoint_id!r} is not a supported " "torch_checkpointing checkpoint." ) logger.info("Loading the checkpoint from %s.", checkpoint_id) @@ -428,12 +452,36 @@ def _load(self, step: int = -1) -> bool: # values and resume from a model that is not the one that was saved. # exclude_from_loading is applied by _states_to_load, so anything still # in `states` here is genuinely required. - loaded = self._manager.load( - checkpoint_id, - into=_stateful_to_state_dict(states), - strict=True, - ) - _restore_state_dict(states, loaded) + if from_hf: + assert self.sd_adapter is not None + hf_state = self.sd_adapter.to_hf(_stateful_to_state_dict(states)[MODEL]) + model_spec = replace( + self._manager_config.items[MODEL], + resharder=HFSafetensorsDTensorResharder(), + ) + hf_config = replace( + _with_sync_save(self._manager_config, use_barrier=False), + items={MODEL: model_spec}, + default=None, + ) + hf_manager = hf_config.build() + try: + loaded = hf_manager.load( + checkpoint_id, + into={MODEL: hf_state}, + strict=True, + ) + finally: + hf_manager.close() + native_state = self.sd_adapter.from_hf(loaded[MODEL]) + _restore_state_dict(states, {MODEL: native_state}) + else: + loaded = self._manager.load( + checkpoint_id, + into=_stateful_to_state_dict(states), + strict=True, + ) + _restore_state_dict(states, loaded) GarbageCollection.collect("GC collection for checkpoint loading.") logger.info( "Finished loading the checkpoint in %.2f seconds.", @@ -481,6 +529,30 @@ def _parse_step(self, filename: str) -> tuple[int, bool] | None: return None return int(match.group("step")), bool(match.group("tmp")) + def _find_native_load_step(self) -> int: + valid_steps = [] + for filename in self._storage.listdir(self.folder): + parsed = self._parse_step(filename) + if parsed is None: + continue + step, is_staging = parsed + if is_staging: + continue + checkpoint_id = filesystem.join(self.folder, filename) + if self._is_native_checkpoint(checkpoint_id): + valid_steps.append(step) + return max(valid_steps) if valid_steps else -1 + + def _is_native_checkpoint(self, checkpoint_id: str) -> bool: + return self._storage.isfile( + filesystem.join(checkpoint_id, TORCH_CHECKPOINTING_METADATA_FILE_NAME) + ) + + def _is_hf_checkpoint(self, checkpoint_id: str) -> bool: + return self._storage.isfile( + filesystem.join(checkpoint_id, _HF_INDEX_FILE_NAME) + ) or self._storage.isfile(filesystem.join(checkpoint_id, _HF_SINGLE_FILE_NAME)) + def _is_valid_checkpoint(self, checkpoint_id: str) -> bool: # Either shape this manager publishes. A resumable checkpoint has the # backend's metadata at its root. A final HF export does not: its @@ -488,9 +560,9 @@ def _is_valid_checkpoint(self, checkpoint_id: str) -> bool: # were written to, and the root holds the consolidated HF files. Probing # only for the former would classify a finished export as abandoned and # let the next run's retention delete it. - return self._storage.isfile( - filesystem.join(checkpoint_id, TORCH_CHECKPOINTING_METADATA_FILE_NAME) - ) or self._storage.isfile(filesystem.join(checkpoint_id, _HF_INDEX_FILE_NAME)) + return self._is_native_checkpoint(checkpoint_id) or self._is_hf_checkpoint( + checkpoint_id + ) def _maybe_wait_for_staging(self) -> None: # Acquiring the backend lock is what blocks until staging for the last