Datasets and DataLoaders

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.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

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 loop

A 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

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on CPU. Fixed data, exact output.

The crash, then the fix

pad_collate.py
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)
Output
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

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_sampler concern, 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-tokens lineage).
  • 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.nested represents 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.data docs, "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