From 5cdbca3c325dca0413fb990e3b629fbe9a5d3dd8 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 18 Sep 2026 23:18:13 -0700 Subject: [PATCH] Preserve checkpoint upload timeout across writer resets --- torchtitan/components/checkpoint.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/torchtitan/components/checkpoint.py b/torchtitan/components/checkpoint.py index ab7780ddfb9..b4b08339b9d 100644 --- a/torchtitan/components/checkpoint.py +++ b/torchtitan/components/checkpoint.py @@ -7,6 +7,7 @@ from __future__ import annotations import enum +import os import queue import re import threading @@ -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 @@ -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" @@ -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