A pretraining run reads trillions of tokens over weeks on hundreds or thousands of GPUs. The offline half of the data work, which covers crawling, filtering, deduplication and tokenization, decides which tokens exist. This article is about the online half: the component that turns a directory of tokenized shards into the exact sequence of batches each GPU sees, step after step, and that has to keep doing so through crashes, node replacements and a mid-run change of batch size.
The bandwidth involved is small. The hard requirements are determinism, exact resumption and auditability: naming the documents behind a loss spike, and restarting without repeating or skipping tokens. Both need one design decision: make every batch a pure function of a seed and a global sample position. We will build that function, wire it to data-parallel ranks, add weighted mixtures, and then use it to debug a loss spike.
The contract a training pipeline must keep
Write the contract down before writing code. A training data pipeline for a large model should guarantee five things:
- Coverage by plan. Each source contributes the fraction of tokens the mixture says, and repeats data only when its budget says it may.
- No duplication across ranks. In one optimizer step, no two data-parallel ranks train on the same sample.
- Determinism. Given the seed, index version and position, the batch is bit-identical on every machine and every restart.
- Exact resumption. After a restart, training continues at the first sample that was not consumed, even if the number of GPUs changed.
- Throughput with headroom. The GPU never waits on data. That is a separate topic covered in the dataloader bottleneck article.
A streaming design with a shuffle buffer breaks all five, because its state lives where you cannot checkpoint it. The fix is to compute every decision up front as an index, so the runtime only does lookups.
Token storage and the sample index
Start with storage. Offline tokenization should emit, per source, one flat array of token ids with an end-of-document token between documents, plus an offsets file that records where each document starts. Choose the dtype from the vocabulary: uint16 works only when the vocabulary has at most 65,536 entries. A 128,256-entry vocabulary such as Llama 3's needs uint32. Store raw little-endian arrays you can memory-map. Avoid per-document pickles and JSON lines on the hot path.
Next, define what a sample is. With sequence length L, the simplest packed scheme cuts the flat array into consecutive windows of L + 1 tokens. Consecutive windows overlap by one token, so the input is the first L tokens and the target is the same window shifted by one. Documents cross window boundaries, which wastes nothing. The cost is that attention can reach across a document boundary unless you also pass document-aware masks. The packing article in the transformer math series works through that trade. Here we only need the window to be addressable by an integer.
import numpy as np
class TokenSource:
"""One tokenized source: flat token array + document offsets, memory-mapped."""
def __init__(self, prefix, seq_len, dtype=np.uint32):
self.tokens = np.memmap(prefix + ".bin", dtype=dtype, mode="r")
self.doc_starts = np.load(prefix + ".idx.npy") # int64, ascending
self.L = seq_len
self.n_samples = (len(self.tokens) - 1) // seq_len # windows of L+1, stride L
def window(self, k):
s = k * self.L
return np.asarray(self.tokens[s : s + self.L + 1], dtype=np.int64)
def docs_in(self, k):
"""Document ids overlapping sample k, for audits and replays."""
s, e = k * self.L, k * self.L + self.L + 1
first = np.searchsorted(self.doc_starts, s, side="right") - 1
last = np.searchsorted(self.doc_starts, e - 1, side="right") - 1
return list(range(first, last + 1))
def epoch_perm(n, seed, source_id, epoch):
"""A fresh shuffle of a source's samples for each pass over it."""
rng = np.random.default_rng([seed, source_id, epoch])
return rng.permutation(n)The permutation is seeded by a tuple of seed, source and epoch, never by global RNG state. That is what makes it reproducible on any host in any order. For a source with billions of samples, memory-map the permutation or replace it with a keyed bijection such as a small Feistel network.
Mixtures: a deterministic blend
A pretraining mixture might be 60 percent web, 15 percent code, 10 percent books, 10 percent academic text and 5 percent math. Sampling a source at random per position meets those weights only on average, and short windows drift. A better blend is deterministic and error-minimising: at each global position, pick the source that is furthest behind its target share. Megatron-LM's blended dataset uses this idea. It keeps every prefix of the stream close to the target mixture, and needs no RNG.
def build_blend(weights, total):
"""Return arrays (source_of_p, draw_of_p) for global positions 0..total-1."""
w = np.asarray(weights, dtype=np.float64) / sum(weights)
counts = np.zeros(len(w), dtype=np.int64)
src = np.empty(total, dtype=np.int16)
draw = np.empty(total, dtype=np.int64)
for p_ in range(total): # vectorise or write in C for 1e9+ positions
d = int(np.argmax(w * (p_ + 1) - counts))
src[p_], draw[p_] = d, counts[d]
counts[d] += 1
return src, draw
def resolve(p_, src, draw, sources, seed):
"""Global position -> (source id, sample index inside that source)."""
d, j = int(src[p_]), int(draw[p_])
n = sources[d].n_samples
epoch, offset = divmod(j, n) # j-th draw from source d
return d, int(epoch_perm(n, seed, d, epoch)[offset])Look at the divmod in resolve. It makes repetition explicit. When a small, high-quality source is upweighted past its size, the draw counter runs past n, the epoch increments and that source is reshuffled. You can now compute before training how many times each source repeats, which is the number you need when you set repetition budgets. Publish that table with the mixture, and treat a source that repeats more than a few times as a decision someone signed off on. The mixture design article in the SLM series covers how to choose those budgets.
Partitioning across data-parallel ranks
The trainer consumes a global batch of B sequences per optimizer step. B is the product of data-parallel size, per-GPU micro-batch and gradient accumulation steps. With tensor or pipeline parallelism, only the data-parallel rank decides which samples a GPU gets. Every GPU in the same tensor-parallel group must receive identical inputs, so key the slice on the data-parallel rank, not the global rank.
def positions_for(consumed, dp_rank, dp_size, micro, accum):
"""Global positions this DP rank reads for the next optimizer step."""
B = dp_size * micro * accum
per_rank = B // dp_size
start = consumed + dp_rank * per_rank
flat = list(range(start, start + per_rank))
return [flat[i * micro:(i + 1) * micro] for i in range(accum)] # one list per micro-step
# Inside the training loop
for micro_positions in positions_for(state.consumed, dp_rank, dp_size, micro, accum):
batch = np.stack([sources[d].window(k)
for d, k in (resolve(q, src, draw, sources, SEED) for q in micro_positions)])
loss = train_micro_step(batch)
optimizer.step()
state.consumed += dp_size * micro * accumThe single number consumed is the pipeline's entire runtime state. Ranks read disjoint contiguous slices whose union is exactly positions consumed to consumed + B - 1. Wrap positions_for in a sampler so DataLoader workers can prefetch from it.
Exact and elastic resumption
At checkpoint time, save consumed, the seed, the mixture weights and a hash of the index files next to the model and optimizer state. On restart, rebuild or memory-map the same index, check the hash, and continue. There is no iterator state to serialise and no fast-forward loop that reads and discards millions of samples.
Two details catch people. First, consumed has to count samples the optimizer actually applied, not samples the loader prefetched. Increment it after optimizer.step() and save it in the same checkpoint as the weights. If the run dies between a data checkpoint and a model checkpoint, you replay or skip a step's worth of data. Second, count samples, not steps, so a mid-run batch-size change cannot shift positions.
Libraries solve the same problem. torchdata's StatefulDataLoader adds state_dict() and load_state_dict() to the DataLoader; by default it fast-forwards past batches already yielded, and it handles worker state but not state across ranks. MosaicML's StreamingDataset checkpoints the epoch, sample in epoch, shuffle seed and num_canonical_nodes. It keeps the global order deterministic across a change in node count, provided num_canonical_nodes stays the same between runs. Both solve the same problem as the index above. Make sure you can say which number is your consumed.
Throughput: latency, not bandwidth
Worked numbers make the bottleneck obvious. Take 512 GPUs, a global batch of 1,024 sequences of 4,096 tokens (4,194,304 tokens per step) and a 6-second step. That is about 700,000 tokens per second for the whole cluster, or 2.8 MB/s of uint32 tokens. A 2-trillion-token run is about 476,837 steps. Raw bandwidth is not the problem. Latency is: each sample is a random 16 KB read, and when shards live in object storage, a random read costs tens of milliseconds.
So the levers are locality and prefetch: stage shards to local NVMe, keep shards in the hundreds of megabytes, and prefetch a few steps ahead, which can be exact because positions are known in advance. Pin host buffers so the host-to-device copy overlaps compute, as described in the GPU data pipeline article. Tokenize offline. Tokenizing text inside the training loop is the most common reason an LLM loader that looked fine in a benchmark starves at scale.
Worked example: replaying a loss spike
Here is the payoff. The loss jumps from 2.31 to 3.9 at step 41,200 and recovers over 300 steps. Because batches are pure functions, you can rebuild that exact batch offline without touching the cluster:
B = 1024
step = 41_200
consumed = step * B # valid only if B never changed before this step; otherwise read it from the log
hits = {}
for q in range(consumed, consumed + B):
d, k = resolve(q, src, draw, sources, SEED)
for doc in sources[d].docs_in(k):
hits.setdefault((d, doc), 0)
hits[(d, doc)] += 1
top = sorted(hits.items(), key=lambda kv: -kv[1])[:10]
for (d, doc), n in top:
print(SOURCE_NAMES[d], doc, n, preview(sources[d], doc, 120))In this example the output showed one 9 MB document from the web source spanning 540 of the 1,024 windows. It was a log dump of one repeated line that passed the length filter because the filter capped characters per line, not repetition. Why did so much of it land in one step? The replay also showed that the run's loader shuffled shard order and then read each shard's windows sequentially, instead of using the per-source sample permutation shown above. Adjacent windows of one long document therefore stayed adjacent in the stream. The fix has two parts. Offline, add a repetition filter (unique-line ratio) and re-tokenize that source under a new index version. Online, permute at sample level, so even a giant document's windows spread across thousands of steps. Without a replayable index, both findings would have been guesses.
Log the step, consumed, batch size and index hash every step, so each future spike is a ten-minute query.
Failure modes
- Rank keyed on global rank. Tensor-parallel peers get different inputs. Activations go wrong silently and loss is worse than expected but not obviously broken. Key on the data-parallel rank.
- Seeding from process RNG.
np.random.seed(rank)plus worker forking gives different orders on restart. Seed every permutation from (seed, source, epoch). - Counting steps instead of samples. A batch-size change mid-run silently skips or repeats data.
- Index rebuilt with different inputs. A re-tokenized shard changes
n_samples, which shifts every permutation. Hash the index inputs and refuse to resume on mismatch. - Unplanned repetition. An upweighted small source loops ten times and the model memorises it. Compute repeats per source from the blend before launch.
- Shard-local shuffle. Neighbouring windows from one long document land in one batch and correlate the gradient, as in the worked example.
Trade-offs
| Choice | Gains | Costs |
|---|---|---|
| Precomputed global index (this article) | Exact resume, elastic resize, replayable batches | Index build time and storage; rebuild on any data change |
| Streaming with a shuffle buffer | No index, starts instantly, easy for new data | Order depends on timing; resume is approximate |
| Packed windows across documents | No padding waste, simple addressing | Cross-document attention unless masked |
| One document per sequence with padding | Clean boundaries | Wasted compute on pad tokens |
| Deterministic error-minimising blend | Every prefix matches the mixture | Mixture change mid-run needs a new blend from the current position |
What to do next
- Write your pipeline's contract (coverage, no rank overlap, determinism, exact resume) in the training README and name the one integer that is its runtime state.
- Check that sample selection keys on the data-parallel rank and that tensor-parallel peers receive identical batches. Assert it in a smoke test by hashing inputs per rank.
- Replace any step-based resume with a samples-consumed counter saved in the same checkpoint as the weights.
- Before launch, compute and publish the per-source repeat counts implied by your mixture and run length.
- Hash the index inputs, store the hash in every checkpoint and refuse to resume on a mismatch.
- Log step, consumed samples, batch size and index hash every step, and build the replay script now, before the first loss spike.
- Measure the loader in isolation at target batch size, then read GPU curation pipelines and checkpointing in depth to connect the offline and recovery sides.