Skip to content

Restore shuffle order when resuming a stateful DataLoader - #4203

Closed
YeonwooSung wants to merge 5 commits into
huggingface:mainfrom
YeonwooSung:fix/stateful-dataloader-shuffle-resume
Closed

YeonwooSung wants to merge 5 commits into
huggingface:mainfrom
YeonwooSung:fix/stateful-dataloader-shuffle-resume

Conversation

@YeonwooSung

Copy link
Copy Markdown

What does this PR do?

use_stateful_dataloader=True restored how far the loader had walked, but not the shuffle permutation. After a completed epoch, a fresh loader rebuilt RandomSampler at 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:

  • BatchSamplerShard now yields a stateful iterator. StatefulDataLoader stores 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.
  • SeedableRandomSampler epoch / initial_seed are restored.
  • DataLoaderAdapter.state_dict also records _accelerate_iteration after a completed pass so set_epoch is applied on load.

A new loader + load_state_dict (or Accelerator.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, using prepare_data_loader(..., num_processes=2) (no GPU):

  • Full epoch 0, checkpoint after 3 batches of epoch 1, resume into a new loader — ranks 0 and 1, with and without use_seedable_sampler
  • Mid-first-epoch resume, same matrix
  • Rank 0 and rank 1 see different orders
  • Resume does not depend on re-seeding torch
  • RandomSampler.generator is None
  • save_accelerator_state / load_accelerator_state

Existing stateful DataLoader tests still pass.

Who can review?

@SunMarc

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.
@GoldenStain

Copy link
Copy Markdown

Additional Findings & Verification on Stateful Dataloader Shuffle Order

Thanks @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 itertools.islice(iter(index_sampler), consumed). This draws a fresh permutation from the global RNG at restore time:

  • In a minimal single-process test (same seed, no intermediate RNG consumers), the RNG draw sequence coincidentally matches the original run's epoch-start draw.
  • In real multi-epoch training (model initialization, a second evaluation dataloader, transform/augmentation RNG draws), the RNG consumption diverged wildly, and the restored permutation was completely scrambled.

2. Epoch-boundary checkpoints and auto-incrementing samplers

When a checkpoint is saved on the final microbatch / after a completed epoch (where torchdata records _iterator_finished=True), torchdata's restore logic first fast-forwards through the finished epoch (one extra full sampler cycle) before beginning the new epoch.

  • Any sampler whose __iter__ or iterator auto-increments its epoch (e.g. SeedableRandomSampler doing self.epoch += 1) gets shifted by an extra epoch permutation.
  • A robust order-restoring sampler must have a set_epoch(epoch) that only assigns and never auto-increments implicitly during iteration/restore.

3. SeedableRandomSampler epoch-0 swap timing

In Accelerate 1.13.0, DataLoaderShard constructor / adapter initialization materialized an initial iterator before dataloader.set_sampler(sampler) replaced the inner sampler.

4. Summary of Required Test Bars

To ensure complete robustness, any shuffle restoration fix should verify all four scenarios under DDP:

  1. Mid-epoch-0 checkpoint resume (clean process)
  2. Mid-epoch-N ($N \ge 1$) checkpoint resume
  3. Epoch-boundary / completed-epoch checkpoint resume ($N \to N+1$)
  4. Resumed stream matches the uninterrupted reference stream bit-for-bit across all ranks without requiring the user to manually re-seed torch before recreating the loader.

(Note: We also observed a separate data partition defect when num_workers > 0 under DDP due to prepare-time iterator prefetch before cross-rank RNG sync, which is tracked separately in #4204).

__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.
@YeonwooSung

Copy link
Copy Markdown
Author

Thanks for the detailed restore matrix — this is exactly the extra cycle we were missing.

Sampler contract. SeedableRandomSampler.__iter__ used to call set_epoch(self.epoch + 1) after a full pass. Combined with torchdata replaying a finished epoch when _iterator_finished=True, that shifted the next permutation by one extra epoch. set_epoch is now assignment-only; DataLoaderShard already calls it at the start of each pass. Completed-epoch snapshots persist iteration + 1 explicitly so restore does not depend on that implicit increment.

Tests. Resume is now checked against a never-interrupted reference stream after extra RNG consumption (a second shuffled eval loader, torch.randn, torch.randperm) and without the caller re-seeding torch. The four bars run for both ranks under prepare_data_loader(..., num_processes=2), with and without use_seedable_sampler:

  1. Mid-epoch-0
  2. Mid-epoch-N (N = 1)
  3. Epoch-boundary / completed-epoch (N → N+1)
  4. Bit-for-bit match of the concatenated resumed stream against the reference

The epoch-0 base._iterator = None swap after set_sampler is unchanged. The num_workers > 0 partition issue remains out of scope here (tracked in #4204).

@SunMarc SunMarc closed this Sep 7, 2026
@YeonwooSung
YeonwooSung deleted the fix/stateful-dataloader-shuffle-resume branch September 8, 2026 00:09
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

use_stateful_dataloader under multi-process training restores the cursor, not the shuffle order — sampler permutation is never serialized

3 participants