Shuffling and Repeating Samples¶
Sample Order Is Part of the Curriculum. While the mixture specification determines
how many samples each Dataset contributes to each training epoch and how samples from
different Datasets are interleaved, the shuffle specification determines the order in
which the StaticMixtureWorkSource draws samples from each individual Dataset before
that interleaving occurs. Shuffling is critical because the samples that are adjacent to
one another within the shards of a Dataset often share common attributes (like source,
format, language, or quality) that likely do not reflect the overall representation of
those attributes in the overall data that we want to train on. Ideally, each local batch
of training data is made up of a mixture of samples that reflect the overall
distribution of the training data to ensure that the optimization process does not
over-index on one specific type of data, which can make the training process unstable or
lead to misleading loss curves. A random shuffle of the training data is one of the best
tools we have to help local batches better reflect the overall distribution of data from
the training set as a whole.
Of course, there is a tradeoff between the strength of the random shuffle we perform and
the throughput of our data loading pipeline: while a purely random global shuffle reduces
local correlation between samples caused by the original ordering during the training
epoch to a maximum, it comes at a cost for the I/O locality and efficiency of reading the
data, especially for data formats that are not designed for efficient random access to
individual records within a file or cloud setups. Zephon provides a number of settings on
the StaticMixtureWorkSource that allows you to control how samples are shuffled within
each training epoch so that you can maximize the performance of your data loading pipeline
while minimizing the impact of spurious correlations between samples on your training
loop.
Choosing A Shuffle Strategy. Our shuffling configuration is inspired by Mosaic
Streaming.
By default, the StaticMixtureWorkSource will shuffle the order of the individual
shards within a Dataset. You can turn this default behavior off by setting
shuffle_shards=False on the constructor. Shuffling the shards within each Dataset
every time it is processed during an epoch is a mechanism for getting a minimal amount
of randomness in your training loop without any meaningful cost in I/O efficiency,
because the WorkSource will be able to draw the samples from each shard in the
Dataset sequentially.
The next level of shuffling is shuffle_within_shard, which is disabled by default; it
randomizes the order in which samples from a single shard are read on each epoch. This is
a good option to enable when the adjacent samples within a shard have strongly correlated
properties that we would like to avoid training on, but it does impose an additional
compute cost when the format of our Dataset is not designed for efficient random access
of individual records (most notably Parquet).
The third level is the shuffle_block_size setting, which allows us to optionally shuffle
samples within a single Dataset across a window of samples that may cross shards (it
defaults to None, which disables this additional shuffling). If you don’t have a strong
sense of what a good shuffle window size would be for your dataset, setting
shuffle_block_size="auto" will have Zephon choose a shuffle block size that is equal to
eight times the number of samples in the largest shard in any Dataset included in the
mixture (or simply the size of the Dataset itself if it happens to be smaller than that
window size). This setting provides cross-shard mixing over several shards and is a solid
default until you have a compelling reason to set a different window size yourself.
Finally, setting the shuffle_block_size to the string "global" will apply a truly
global shuffle across the entire Dataset. Note that this requires us to materialize all
sample IDs of the entire Dataset upfront and shuffle them. At trillion-token scale
training, this might not be possible as just the sample identifiers can take hundreds of
gigabytes of DRAM, and this is slow. Our other shuffle algorithms avoid full
materialization of the identifiers up front.
A global shuffle is not as important as you might think it is: our friends at Mosaic Streaming have observed that the entropy of a shuffle across the shards is extremely close to an actual global shuffle.
Repeating Samples. In addition to determining the order in which the records from
each Dataset are processed, we also need to decide when to stop drawing samples from
each Dataset, i.e., when an epoch ends. By default, the StaticMixtureWorkSource
repeats datasets that finish a pass sooner under the requested mixture in order to
maintain the requested mixture until every dataset has reached the end of its first
pass. To continue drawing samples until every dataset has reached the end of at least
n passes, set the stop_after_passes argument to n instead of the default value of
1:
"""Keep drawing samples until every dataset has completed three passes."""
from zephon.io import Dataset
from zephon.work import MixtureSpec, StaticMixtureWorkSource
code = Dataset.from_path("code", "s3://my-bucket/corpora/code")
web = Dataset.from_path("web", "s3://my-bucket/corpora/web")
# Datasets that get there first keep repeating to hold the mixture while the
# others catch up, reshuffled on each pass unless reshuffle_on_repeat is off.
work_source = StaticMixtureWorkSource(
datasets=[code, web],
mixture=MixtureSpec({"code": 0.25, "web": 0.75}),
seed=42,
stop_after_passes=3,
)
Note that each time a Dataset is processed, its shuffle settings will be applied to it
with a deterministic per-epoch shuffle seed in order to generate a different ordering of
the samples for that pass; you can disable this behavior by setting the
reshuffle_on_repeat argument to False.
The stop_after_passes argument controls how many passes every dataset must complete
before the work source stops. Datasets that reach this target sooner continue repeating to
maintain the requested mixture while the remaining datasets catch up.
You can instead control what happens when an individual dataset runs out of samples using
exhausted_policy:
exhausted_policy="stop"stops the entire work source as soon as any dataset can no longer supply its requested allocation. No dataset repeats, but other datasets may still have unread samples. This means that setting bothexhausted_policy="stop"andstop_after_passesis an invalid configuration.exhausted_policy="repeat"restarts each dataset independently whenever it runs out. Without an additional stopping limit, the work source continues indefinitely. This is useful when your training loop controls termination through a fixed number of steps or tokens.
With exhausted_policy="repeat", you can also set max_repeats=N. This allows each
dataset its initial pass plus at most N additional passes. The entire WorkSource stops
when any dataset can no longer supply its requested allocation without exceeding that
limit. Other datasets may have completed fewer passes, or may still be on their first
pass.
stop_after_passes waits for every dataset to reach a minimum number of passes, while
max_repeats stops the source when any dataset would exceed its maximum number of
repeats. Just like with exhausted_policy="stop", stop_after_passes can also not be
combined with max_repeats.
When none of these settings is supplied, Zephon defaults to stop_after_passes=1 and
automatically repeats datasets as needed. Specifying exhausted_policy or max_repeats
replaces this default stopping behavior. Assuming you don’t set max_repeats:
What you supply |
Effective |
Effective |
When the source stops |
|---|---|---|---|
Nothing |
|
|
Once every dataset has completed at least one pass |
Only |
|
|
Never automatically; datasets repeat indefinitely |
Only |
|
|
When any dataset cannot supply its next allocation, no repeats |
Only |
|
|
Once every dataset has completed at least |
The exhausted_policy, max_repeats, and reshuffle_on_repeat settings can each be
configured per-Dataset to allow for more complex training plans:
"""Give each Dataset its own repetition behaviour."""
from zephon.io import Dataset
from zephon.work import MixtureSpec, StaticMixtureWorkSource
code = Dataset.from_path("code", "s3://my-bucket/corpora/code")
web = Dataset.from_path("web", "s3://my-bucket/corpora/web")
# The small code set may be seen up to three times over; the large web set is
# read once and stops the source when it runs out.
work_source = StaticMixtureWorkSource(
datasets=[code, web],
mixture=MixtureSpec({"code": 0.25, "web": 0.75}),
seed=42,
exhausted_policy={"code": "repeat", "web": "stop"},
max_repeats={"code": 2, "web": None},
reshuffle_on_repeat={"code": True, "web": False},
)