Fix stateful dataloader checkpointing across processes - #4165
Open
jrsmartin wants to merge 2 commits into
Open
Conversation
jrsmartin
marked this pull request as ready for review
August 14, 2026 20:59
Author
|
@SunMarc would you mind reviewing this PR when you have a chance - thank you! |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
What does this PR do?
This saves stateful dataloader checkpoints per process during distributed training. Stateful dataloader can hold process-specific cursor, dataset, worker, and RNG state. Currently, every process currently writes its state to the same
dl_state_dict.binfilename in a shared checkpoint directory, so the final file therefore contains whichever process wrote last, and other processes can restore the wrong data position.Distributed checkpoints now use rank-qualified filenames such as
dl_state_dict_rank0.bin. Loading prefers the current process's file and falls back to the legacy unqualified filename, preserving compatibility with existing checkpoints. Single-process checkpoint filenames remain unchanged.Related context
I didn't find an existing issue or forum thread for this exact shared-filename collision but #3080 discusses a separate multi-worker/prefetch issue and a comment there describes manually saving dataloader state on every rank. Please note that this PR doesn't address the multi-worker issue in #3080.
Tests
torchdatadependency is unavailable.make qualitypytest -q tests/test_state_checkpointing.py -k "not map_location"— 18 passedpytest -q tests/test_accelerator.py -k stateful_dataloader— 17 passedBefore submitting
ssh_portsupport #3038Who can review?
Anyone in the community is free to review the PR once the tests have passed.