Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/checkpoint.md
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,7 @@ NGPU=1 ./run_train.sh --module <module_name> --config <config_name> --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.

Expand Down
75 changes: 70 additions & 5 deletions tests/unit_tests/cpu/test_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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)
Expand All @@ -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")
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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"
Expand Down
16 changes: 16 additions & 0 deletions torchtitan/components/checkpointer/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading