Writing a collate_fn
collate_fn is the function that glues single samples into one batch — write your own when samples have different lengths and the default stacker crashes.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A collate function takes a handful of single samples and packs them into one batch; you write your own when the samples do not fit the standard packing.
Think of a courier packing parcels. Books — all identical size — stack into a carton in one obvious way. Now try packing five kitchen items: a ladle, a pressure cooker, two cups and a rolling pin. The standard book-stacking routine fails outright.
A good packer improvises: bubble wrap around the small things until everything is effectively the same size, and a note on the box saying what is where.
That packer is a collate function — "collate" meaning collect-and-arrange. The default one only knows book-stacking. When your samples are ladles and cookers, you write the packer yourself.
Why it exists
The DataLoader glues samples into a batch by stacking them — which demands every sample have exactly the same shape.
Much real data refuses. Sentences differ in word count. Audio clips differ in seconds. One medical record holds three test results; another, thirty. The stacker meets its first eleven-word sentence in a batch of nine-word ones, and crashes.
The DataLoader's answer is a socket: pass your own function as collate_fn, and the loader hands it the raw list of samples and steps aside. Your function decides what a batch even is. The usual decision: pad the short ones — the bubble wrap — and record each sample's true length — the note on the box.
How it works
fetcher output: [ five words ], [ two words ], [ three words ]
|
v
your collate_fn: pad to the longest: note true lengths:
five words 5
two + ... 2
three + .. 3
|
v
one rectangular batch + a lengths tensor → your training loopA real example you have seen
Every chatbot and translation service batches many users' requests together, and no two requests are the same length. A packer function — under this name or another — sits in every one of those serving stacks.
Remember this
- The collate function turns a list of samples into one batch — that step, and nothing else.
- The default one stacks; equal shapes only.
- Different-sized samples need your own packer: pad, plus note the true lengths.
What to learn next
- num_workers, prefetching and the Windows spawn trap — where your collate function actually executes.
- Padding and packing variable-length sequences — the model-side half of this exact story.
- Samplers and weighted sampling — the index-choosing stage, including batch samplers for bucketing.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on CPU. Fixed data, exact output.
The crash, then the fix
import torch
from torch.utils.data import Dataset, DataLoader
from torch.nn.utils.rnn import pad_sequence
class Sentences(Dataset):
"""Five sentences of different lengths, already turned into ids."""
def __init__(self):
self.data = [torch.tensor([4, 2, 9, 7, 1]),
torch.tensor([5, 3]),
torch.tensor([8, 6, 2]),
torch.tensor([1, 9, 9, 3]),
torch.tensor([7])]
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx], len(self.data[idx])
try:
next(iter(DataLoader(Sentences(), batch_size=3)))
except RuntimeError as e:
print("default collate fails:", str(e)[:80])
def pad_collate(batch):
seqs, lengths = zip(*batch) # unzip the list of (tensor, length) pairs
padded = pad_sequence(seqs, batch_first=True, padding_value=0)
return padded, torch.tensor(lengths)
loader = DataLoader(Sentences(), batch_size=3, collate_fn=pad_collate)
for padded, lengths in loader:
print("batch", tuple(padded.shape), "lengths", lengths.tolist())
print(padded)default collate fails: stack expects each tensor to be equal size, but got [5] at entry 0 and [2] at en
batch (3, 5) lengths [5, 2, 3]
tensor([[4, 2, 9, 7, 1],
[5, 3, 0, 0, 0],
[8, 6, 2, 0, 0]])
batch (2, 4) lengths [4, 1]
tensor([[1, 9, 9, 3],
[7, 0, 0, 0]])Memorise that error string. stack expects each tensor to be equal size at loader time means, almost always: variable-sized samples met the default collate. Now you know the cure on sight.
The walkthrough
What your function receives: a plain Python list of whatever __getitem__ returns — here, five (tensor, int) pairs become a list of three, then a list of two. The DataLoader does no preprocessing first; the list is raw.
zip(*batch) is the idiom worth internalising: it turns a list of pairs into a pair of lists. Samples-of-fields become fields-of-samples in one expression.
pad_sequence does the bubble-wrapping — the same tool as in padding and packing, where the downstream model-side story lives. Note the padding is per batch: the first batch padded to 5, the second to 4. Batches need not match each other; every batch only needs internal consistency.
Returning (padded, lengths) — your return value is what the training loop unpacks. The collate function is the last word on batch structure; return a dict, a triple, whatever the loop expects.
Where it runs matters: with num_workers > 0, your collate function runs inside worker processes. Keep it pure — no prints you rely on, no mutation of outside state, and it must be a module-level function (picklable) for Windows workers.
Sort before you pad, when it counts
Padding waste is the gap between longest and shortest in a batch. A one-word and a fifty-word sentence in one batch means 49 wasted columns for the short one. The standard remedy is grouping similar lengths — sort the dataset by length, or use a batch sampler that buckets. For heavily skewed length distributions, this alone can cut compute noticeably.
Common mistakes
Padding in the Dataset instead. Padding each sample to a global maximum in __getitem__ works but wastes compute on every batch that had no long sample. Padding is a batch-level fact; it belongs in collate.
Forgetting the lengths. A padded batch without true lengths forces the model to treat filler as data. Return the lengths; the model side needs them for masking or packing.
A pad value that is a real token. Pad with an id reserved for padding — the padding_idx contract from embedding layers. Padding with a real word id plants fake words in every short sentence.
A lambda as collate_fn on Windows. collate_fn=lambda b: ... cannot be pickled into spawned workers; training dies at startup with a pickling error. Define it with def at module level.
Try it yourself
Extend the dataset so each sample is (ids, length, label) with a 0/1 label, and update pad_collate to return (padded, lengths, labels) with labels as one tensor. The zip(*batch) line barely changes — that is the sign you have the idiom right.
What to learn next
- num_workers, prefetching and the Windows spawn trap — where your collate function actually executes.
- Padding and packing variable-length sequences — the model-side half of this exact story.
- Samplers and weighted sampling — the index-choosing stage, including batch samplers for bucketing.
Researcher — Mathematics and papers.
The collate contract, stated fully
collate_fn: list[T] → B where T is the dataset's sample type and B unconstrained — the loader treats it as opaque. With auto-batching on (batch_size set), it receives batch_size samples (fewer on the last batch); with batch_size=None, it receives single samples, acting as a per-sample transform at the loader layer. default_collate's recursion (tensors stack; numerics tensorise; mappings and sequences recurse per-field) was dissected in the internals lesson; a custom function replaces the recursion wholesale. Since torch 2.x, torch.utils.data.default_collate is public API precisely so custom functions can delegate the standard fields and hand-handle the ragged ones — the practical pattern for mixed dict samples.
Ragged batching beyond padding
Padding is one point in a design space:
- Bucketed batching: constrain each batch to similar lengths (a
batch_samplerconcern, not collate proper). Reduces pad fraction toward zero at a shuffling-quality cost; standard in machine translation since the RNN era. - Token-budget batching: batch size defined in tokens, not samples — variable sample count per batch, near-constant compute and memory per step. This is how sequence models train at scale (fairseq's
--max-tokenslineage). - Concatenate-and-chunk: for pretraining, samples are concatenated with separators and cut to fixed length; ragged-ness is eliminated upstream and collate degenerates to stacking. Attention masking then owes nothing to batch geometry.
- NestedTensor:
torch.nestedrepresents raggedness natively; as of torch 2.5 coverage is still partial (prototype-tier for many ops), so padded-plus-mask remains the production default — verify current status before adopting.
The padding-efficiency arithmetic from packing applies verbatim: expected pad fraction is one minus mean-over-max length within a batch, and bucketing attacks exactly that ratio.
Determinism and parallel execution
Collate executes in workers under multiprocess loading; its inputs are already materialised samples, so its own determinism reduces to being a pure function. What is not guaranteed across num_workers settings is RNG consumption patterns — augmentation randomness lives in __getitem__ under per-worker seeds, so changing worker count re-partitions the RNG streams and changes exact augmentations even with identical global seeds. Bit-exact reproducibility contracts must therefore pin num_workers alongside seeds — a real reproducibility footnote for papers publishing with "seed fixed".
A second-order effect: since collate output tensors are allocated in workers and moved via shared memory, very large batches interact with /dev/shm limits in containers (Linux) — the notorious "bus error" in dockerised training. On Windows, tensors serialise through named pipes instead; the analogous cost is copy bandwidth, another reason batch assembly should stay lean.
Reading
torch.utils.datadocs, "Working with collate_fn" — normative behaviour with and without auto-batching.- Ott et al. (2019), fairseq: A Fast, Extensible Toolkit for Sequence Modeling — token-budget batching in practice.
What to learn next
- num_workers, prefetching and the Windows spawn trap — where your collate function actually executes.
- Padding and packing variable-length sequences — the model-side half of this exact story.
- Samplers and weighted sampling — the index-choosing stage, including batch samplers for bucketing.