Skip to content
Open
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
15 changes: 13 additions & 2 deletions torchtitan/components/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from __future__ import annotations

import enum
import os
import queue
import re
import threading
Expand All @@ -25,7 +26,7 @@

from torch.distributed.checkpoint import HuggingFaceStorageWriter
from torch.distributed.checkpoint._consolidate_hf_safetensors import consolidate_safetensors_files_on_every_rank
from torch.distributed.checkpoint._fsspec_filesystem import FsspecReader, FsspecWriter
from torch.distributed.checkpoint._fsspec_filesystem import FileSystem as FsspecFileSystem, FsspecReader, FsspecWriter
from torch.distributed.checkpoint.staging import DefaultStager, StagingOptions
from torch.distributed.checkpoint.state_dict_saver import AsyncCheckpointerType, AsyncSaveResponse
from torch.distributed.checkpoint.stateful import Stateful
Expand All @@ -50,6 +51,16 @@
CHECKPOINT_UPLOAD_TIMEOUT_SECONDS = 600.0


class _CheckpointWriter(FsspecWriter):
def reset(self, checkpoint_id: str | os.PathLike | None = None) -> None:
super().reset()
if checkpoint_id:
# FsspecWriter.reset() otherwise recreates the filesystem without its timeout.
self.path = cast(FsspecFileSystem, self.fs).init_path(
checkpoint_id, timeout=CHECKPOINT_UPLOAD_TIMEOUT_SECONDS
)


class AsyncMode(str, enum.Enum):
DISABLED = "disabled"
ASYNC = "async"
Expand Down Expand Up @@ -565,7 +576,7 @@ def dcp_save(
# `consolidate_safetensors_files_on_every_rank` is used later to manage
# the multi-file merging process.
else:
storage_writer = FsspecWriter(checkpoint_id, timeout=CHECKPOINT_UPLOAD_TIMEOUT_SECONDS)
storage_writer = _CheckpointWriter(checkpoint_id, timeout=CHECKPOINT_UPLOAD_TIMEOUT_SECONDS)

# Execution Dispatch
checkpoint_save_id = None if to_hf else checkpoint_id # for HF the storage_writer handles the path
Expand Down
Loading