Repository navigation
Restore shuffle order when resuming a stateful DataLoader - #4203
YeonwooSung wants to merge 5 commits into
Conversation
StatefulDataLoader only saved the cursor (_num_yielded). BatchSamplerShard hid the inner RandomSampler, so a fresh loader after an epoch boundary restored position against a new permutation. Snapshot generator/epoch state on the shard iterator and restore DataLoaderShard.iteration. Fixes huggingface#4195
The last-batch snapshot still held the pre-epoch shuffle and omitted _accelerate_iteration at epoch 0, so a fresh loader replayed the epoch that just ended. Capture post-epoch sampler state with a zero cursor.
Replacing every `_sampler_iter_state` with `{yielded, shuffle_state}`
broke torchdata's `{samples_yielded}` schema, so a finished
DataLoaderShard(batch_size=...) raised KeyError on load_state_dict.
Gate the rewrite on BatchSamplerShard and leave torchdata's payload alone.
Additional Findings & Verification on Stateful Dataloader Shuffle OrderThanks @YeonwooSung for opening PR #4203! We have done extensive empirical probing across single-process, 2-process Gloo, and multi-GPU setups around exact shuffle continuation. Here are a few important edge cases and observations from our test matrix that may be helpful for review and test coverage in #4203: 1. How restore proceeds without sampler state (and why single-process minimal repros can pass by coincidence)In vanilla Accelerate + torchdata without sampler state persistence, torchdata falls back to
2. Epoch-boundary checkpoints and auto-incrementing samplersWhen a checkpoint is saved on the final microbatch / after a completed epoch (where torchdata records
3.
|
__iter__ used to call set_epoch(epoch + 1) after a full pass. When torchdata restores a completed-epoch checkpoint it can replay one extra sampler cycle, which then shifted the next permutation by an extra epoch. set_epoch now only assigns. DataLoaderShard already calls set_epoch at the start of each pass. Completed-epoch snapshots persist the next epoch explicitly. Tests compare resume against an uninterrupted reference stream after extra RNG consumption, covering mid-epoch-0, mid-epoch-N, and epoch-boundary checkpoints.
|
Thanks for the detailed restore matrix — this is exactly the extra cycle we were missing. Sampler contract. Tests. Resume is now checked against a never-interrupted reference stream after extra RNG consumption (a second shuffled eval loader,
The epoch-0 |
What does this PR do?
use_stateful_dataloader=Truerestored how far the loader had walked, but not the shuffle permutation. After a completed epoch, a fresh loader rebuiltRandomSamplerat epoch 0 and then skipped N batches of the wrong order. Mid-epoch resume on a new process had the same failure when the sampler's generator was not serialized.This persists the state needed to rebuild the current epoch's permutation:
BatchSamplerShardnow yields a stateful iterator.StatefulDataLoaderstores that as_sampler_iter_state(generator snapshot taken before the inner sampler consumes it, plus how many batches were yielded).RandomSampler(generator=None)gets a dedicated generator so the permutation is not tied to the caller's global RNG.SeedableRandomSamplerepoch /initial_seedare restored.DataLoaderAdapter.state_dictalso records_accelerate_iterationafter a completed pass soset_epochis applied on load.A new loader +
load_state_dict(orAccelerator.save_state/load_state) then sees the same remaining batches.Fixes #4195. Related: #4196 takes a similar iterator snapshot but does not cover the epoch-boundary case or
Accelerator.save_state.Tests
All under
@require_torchdata_stateful_dataloader, usingprepare_data_loader(..., num_processes=2)(no GPU):use_seedable_samplerRandomSampler.generator is Nonesave_accelerator_state/load_accelerator_stateExisting stateful DataLoader tests still pass.
Who can review?
@SunMarc