Skip to content

Fix sampler handling in skip_first_batches - #4363

Open
Michael-RDev wants to merge 2 commits into
huggingface:mainfrom
Michael-RDev:fix/skip-first-batches-sampler-preservation
Open

Michael-RDev wants to merge 2 commits into
huggingface:mainfrom
Michael-RDev:fix/skip-first-batches-sampler-preservation

Conversation

@Michael-RDev

Copy link
Copy Markdown

What does this PR do?

I tried to use skip_first_batches() with batch_size=None and it failed even when skipping zero batches. The loader worked fine on its own, but this function assumed it had a batch sampler.

For samplers that already yield complete batches, moving the sampler into batch_sampler could also change what gets passed to the collate function and what the loader returns.

This change preserves the original sampler placement and batch_size=None across ordinary, sharded, and dispatched loaders. No new dependencies or public API changes.

Tests

Added regression tests covering unbatched samples, samplers that yield complete batches, custom ordering and collation, and repeated skipping. The affected cases failed before the fix and pass afterward.

Verified with CPU PyTorch:

  • python -m pytest tests/test_data_loader.py -q --tb=short — 52 passed.
  • python -m ruff check . — passed.
  • python -m ruff format --check . — passed.

Before submitting

  • This PR fixes a typo or improves the docs. — Not applicable.
  • Read the contributor guidelines.
  • Discussed/approved through a GitHub issue or forum. — Not yet.
  • Updated documentation. — No documentation changes; this fixes existing behavior.
  • Added necessary regression tests.

Who can review?

Anyone familiar with Accelerate's data loaders and batch skipping

- Preserved sampler placement and batch_size=NOne so skipping batches
  doesn't change collation or returned data structs.
- Added regression tests for ordinary, shareded, and dispatched loaders
- All 52 data-loader tests and lint check passes

@chrikrah chrikrah left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@Michael-RDev skip_first_batches raises TypeError: 'NoneType' object is not iterable on a batch_size=None loader. Your 49de500 returns a loader that iterates. Approving. Ten of the twelve new parameterisations fail with src/accelerate/data_loader.py restored to 01c73fb. The two survivors are DataLoaderShard with a wrapped BatchSampler, a path main already handled.

non-blocking for the logic, though Quality Check stays red. The comments at data_loader.py:1409 and 1438 want a space after #.

$ cat probe4363.py
import torch
from torch.utils.data import DataLoader
from accelerate.data_loader import skip_first_batches
ds = torch.arange(20).reshape(10, 2)
dl = DataLoader(ds, sampler=list(reversed(range(10))), batch_size=None)
print("original first item:", list(dl)[0].tolist())
new = skip_first_batches(dl, num_batches=2)
print("resumed batch_size:", new.batch_size, "batch_sampler:", new.batch_sampler)
print("resumed first item:", list(new)[0].tolist())

$ python probe4363.py   # 01c73fb, CPython 3.12.3, torch 2.14.1+cpu
  File ".../accelerate/data_loader.py", line 1346, in __iter__
    for index, samples in enumerate(self.batch_sampler):
TypeError: 'NoneType' object is not iterable

$ python probe4363.py   # 49de500
original first item: [18, 19]
resumed batch_size: None batch_sampler: None
resumed first item: [14, 15]

$ python -m pytest tests/test_data_loader.py -q
# 01c73fb: 28 passed, 12 skipped in 2.70s
# 49de500: 40 passed, 12 skipped, 36 subtests passed in 3.29s

$ python -m pytest tests/test_data_loader.py -k skip_first_batches -q   # 49de500 with data_loader.py restored to 01c73fb
40 failed, 8 passed, 34 deselected, 6 subtests passed in 4.77s
# the 8 are the 6 pre-existing tests plus _06_DataLoaderShard and _07_DataLoaderShard

$ make quality   # ruff 0.13.1, pinned at setup.py:19
# 01c73fb: All checks passed! / 186 files already formatted
# 49de500: E265 at src/accelerate/data_loader.py:1409:5 and :1438:9, Found 2 errors.
#          ruff format --check . alone: 1 file would be reformatted, 185 files already formatted

$ cat probe_autobatch.py
import torch
from torch.utils.data import DataLoader, BatchSampler
from accelerate.data_loader import skip_first_batches
bs = BatchSampler(range(12), batch_size=2, drop_last=False)
dl = DataLoader(torch.arange(24).reshape(12, 2), sampler=bs, batch_size=2)
print([tuple(b.shape) for b in skip_first_batches(dl, 1)])

$ python probe_autobatch.py
# 01c73fb: [(2, 2), (2, 2), (2, 2), (2, 2), (2, 2)]
# 49de500: [(2, 2, 2), (2, 2, 2), (1, 2, 2)]

# not run: the XLA and multi-process paths, CPU-only box with one process

Every new case passes batch_size=None, so nothing covers the auto-batching route. The probe above shows it moving. @Michael-RDev, would you add a case with sampler=BatchSampler(...) and a real batch_size, so a test pins the output you meant?

@Michael-RDev

Copy link
Copy Markdown
Author

Sounds good, I’ll add it soon. Thanks for the review :)

- Wrap the loader batch sampler when automatic batching is enabled so skipping counts complete loader batches and preserves drop_last. Also keep the sampler path for loaders with batch_size=None.

- Add regression coverage for ordinary, sharded, and dispatched loaders, including nested batches, custom collation, repeated skipping, and exhaustion.

@chrikrah chrikrah left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@Michael-RDev approving at 0ea0251. With a BatchSampler as sampler and a real batch_size, skipping one batch now drops one loader batch, and all 13 new tests pin it.

$ python probe_autobatch.py      # the probe from my first review, torch 2.14.1+cpu, CPython 3.12.3
[(2, 2, 2), (2, 2, 2)]          # 0ea0251; 49de500 gave [(2, 2, 2), (2, 2, 2), (1, 2, 2)]

$ python -m pytest tests/test_data_loader.py -q -rs
52 passed, 13 skipped, 72 subtests passed      # 0ea0251
27 passed, 13 skipped                          # 01c73fb, the base
SKIPPED [1] tests/test_data_loader.py:711: test requires the datasets library
# one skip more than my first review's 28/12 at 01c73fb: datasets is not installed in this venv

$ python -m pytest tests/test_data_loader.py -q -k "auto_batching or nested_batch"   # data_loader.py taken from 49de500
43 failed, 12 passed, 40 deselected, 42 subtests passed   # every one of the 13 new tests fails

$ ruff check . && ruff format --check .        # ruff 0.13.1, pinned at setup.py:19
All checks passed! / 186 files already formatted

@SunMarc, would you take this one for merge?

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants