Skip to content

Commit 5bad1a2

Browse files
Fixed bug in DataLoader where sharding remainder was always dropped, ignoring ShardOptions.drop_remainder=False.
PiperOrigin-RevId: 934593946
1 parent 941d86c commit 5bad1a2

3 files changed

Lines changed: 58 additions & 1 deletion

File tree

CHANGELOG.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ changes. Best viewed [here](https://google-grain.readthedocs.io/en/latest/change
1212
* Deprecations:
1313

1414
* Bug fixes:
15+
* Fixed bug in DataLoader where sharding remainder was dropped even when ShardOptions.drop_remainder=False.
1516

1617
## Grain 0.2.18 (June 17, 2026)
1718

grain/_src/python/data_loader.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,15 @@ def __init__(
125125
super().__init__(dataset.MapDataset.source(data_source))
126126
self._sampler = sampler
127127
self._shard_options = shard_options
128-
self.length = self._sampler_size() // self._shard_options.shard_count
128+
sampler_size = self._sampler_size()
129+
shard_count = self._shard_options.shard_count
130+
if self._shard_options.drop_remainder:
131+
self.length = sampler_size // shard_count
132+
else:
133+
remainder = sampler_size % shard_count
134+
self.length = sampler_size // shard_count + (
135+
1 if self._shard_options.shard_index < remainder else 0
136+
)
129137

130138
def _sampler_size(self) -> int:
131139
"""Returns the length of the sampler."""

grain/_src/python/data_loader_test.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -825,6 +825,54 @@ def test_start_prefetch(
825825
data_loader_iterator.start_prefetch()
826826
self.assertEqual(list(data_loader_iterator), list(range(16)))
827827

828+
def test_data_loader_sharding_drop_remainder(self):
829+
range_data_source = RangeDataSource(start=0, stop=10, step=1)
830+
831+
# Shard 0, drop_remainder=False
832+
sampler = samplers.SequentialSampler(
833+
num_records=len(range_data_source),
834+
shard_options=sharding.ShardOptions(
835+
shard_index=0, shard_count=3, drop_remainder=False
836+
),
837+
)
838+
dataloader = data_loader_lib.DataLoader(
839+
data_source=range_data_source,
840+
sampler=sampler,
841+
)
842+
actual = list(dataloader)
843+
self.assertEqual(actual, [0, 3, 6, 9])
844+
self.assertLen(actual, 4)
845+
846+
# Shard 0, drop_remainder=True
847+
sampler = samplers.SequentialSampler(
848+
num_records=len(range_data_source),
849+
shard_options=sharding.ShardOptions(
850+
shard_index=0, shard_count=3, drop_remainder=True
851+
),
852+
)
853+
dataloader = data_loader_lib.DataLoader(
854+
data_source=range_data_source,
855+
sampler=sampler,
856+
)
857+
actual = list(dataloader)
858+
self.assertEqual(actual, [0, 3, 6])
859+
self.assertLen(actual, 3)
860+
861+
# Shard 1, drop_remainder=False
862+
sampler = samplers.SequentialSampler(
863+
num_records=len(range_data_source),
864+
shard_options=sharding.ShardOptions(
865+
shard_index=1, shard_count=3, drop_remainder=False
866+
),
867+
)
868+
dataloader = data_loader_lib.DataLoader(
869+
data_source=range_data_source,
870+
sampler=sampler,
871+
)
872+
actual = list(dataloader)
873+
self.assertEqual(actual, [1, 4, 7])
874+
self.assertLen(actual, 3)
875+
828876

829877
class PyGrainDatasetIteratorTest(absltest.TestCase):
830878

0 commit comments

Comments
 (0)