skip to content

What does DataLoader's collate_fn do in PyTorch, and when must you write your own?

level: middleimportance: must knowfreq 60%

answer

  1. list of samples in, batch out
  2. stacking is where shapes must agree
  3. ragged data needs a hook
  4. padding to the batch maximum
  5. default_collate and pad_sequence

basics

~20 s

collate_fn turns a list of individual samples into one batch. The default stacks equally-shaped tensors along a new dimension 0 and recurses through tuples and dicts. Write your own for variable-length data, custom padding, or ragged structures.

solid answer

~40 s

After the sampler picks indices and `__getitem__` returns one sample each, `DataLoader` hands the resulting Python list to `collate_fn`. The default, `torch.utils.data.default_collate`, converts NumPy arrays and Python scalars to tensors and then `torch.stack`s them along a new leading dimension, recursing element-wise through tuples, lists, dicts and namedtuples so a batch of `(image, label)` becomes `(stacked_images, stacked_labels)`. You need a custom one whenever the samples are not stack-compatible: variable-length sequences (`torch.stack` raises because the shapes differ), ragged detection targets with different box counts, or a batch that needs extra derived fields such as lengths or an attention mask. A custom collate typically calls `torch.nn.utils.rnn.pad_sequence` to pad to the batch maximum and returns the lengths alongside. When `num_workers > 0`, collation runs inside the worker process, so that padding work is parallelized too.

code

python · 23 lines
python
import torch
from torch.nn.utils.rnn import pad_sequence
from torch.utils.data import DataLoader, Dataset


class RaggedDataset(Dataset):
    def __len__(self):
        return 4

    def __getitem__(self, idx):
        return torch.ones(idx + 1, dtype=torch.long), idx % 2


def collate(batch):
    seqs, labels = zip(*batch)
    lengths = torch.tensor([s.size(0) for s in seqs])
    padded = pad_sequence(seqs, batch_first=True, padding_value=0)
    return padded, lengths, torch.tensor(labels)


loader = DataLoader(RaggedDataset(), batch_size=4, collate_fn=collate)
padded, lengths, labels = next(iter(loader))
print(padded.shape, lengths)  # torch.Size([4, 4]) tensor([1, 2, 3, 4])

go deeper

for a junior

Know that DataLoader stacks samples into a batch for you, and that the fix for a stack-size error on variable-length data is a custom collate_fn.

for a middle

Explain the default's type-directed recursion, why unequal shapes raise, and write a pad_sequence-based collate that also returns lengths or a mask.

for a senior

Reason about placement: collate runs in the worker, pads to the batch maximum, and pairs with a length-bucketing batch_sampler to cut padding waste in real text or audio pipelines.

for a principal

Own the batch contract itself — a named structure rather than a positional tuple, consistent masking conventions, and picklable collate objects — so models, trainers and evaluation code across teams can share loaders without silent shape assumptions.

## Where collation sits in the pipeline One `DataLoader` iteration does four things: a sampler yields a list of indices; `__getitem__` is called once per index; the resulting list of samples is passed to `collate_fn`; the returned batch is optionally pinned and handed to the training loop. `collate_fn` is therefore the only place that ever sees *several samples at once* before the model does, which makes it the natural home for padding, masking and any cross-sample bookkeeping. ## What the default does `torch.utils.data.default_collate` is type-directed and recursive: - A list of tensors of identical shape and dtype becomes one tensor with a new dimension 0 (`torch.stack`). - NumPy arrays are converted with `torch.as_tensor` first — dtype is preserved, so `np.float64` arrays give you a `float64` batch that a float32 model will reject. - Python `int` and `float` become 1-D `int64` / `float32` tensors. - Strings and bytes are left as a Python list. - Tuples, lists, dicts and namedtuples are walked element-wise: a batch of dicts becomes a dict of batched values, with the same keys. When shapes disagree, the stack fails with a `RuntimeError` complaining that stack expects each tensor to be of equal size. That error message, appearing on the first batch of a sequence model, is the single most common reason people go looking for `collate_fn`. ## When you write your own **Variable-length sequences.** Text, audio and time series arrive at different lengths. The idiomatic collate zips the batch apart, records `lengths`, calls `pad_sequence(seqs, batch_first=True, padding_value=pad_id)` and returns the padded tensor plus the lengths (or a boolean mask derived from them). Padding to the *batch* maximum rather than a global maximum is why this belongs in collate: batches of short sequences stay small. **Ragged targets.** Object detection samples have different numbers of boxes per image. The usual answer is not to pad at all — return the images stacked and the targets as a plain Python list of per-image dicts, which most detection models expect. **Derived fields.** Attention masks, segment ids, per-batch normalization statistics, or a `torch.nested` / packed representation are all cheaper to build once per batch than once per sample. **Filtering.** A collate that receives `None` for corrupt samples (returned by a defensive `__getitem__`) can drop them and collate the rest — though a batch can then be empty, which the training loop must tolerate. ## Practical notes *It runs in the worker.* With `num_workers > 0`, both `__getitem__` and `collate_fn` execute in the worker process; only the finished batch crosses the process boundary through shared memory. Heavy padding is therefore parallelized, and — usefully — moving work from `__getitem__` into `collate_fn` does not cost you parallelism. *It must be picklable on spawn platforms.* A `lambda` as `collate_fn` fails to pickle when the start method is spawn (macOS, Windows, and increasingly Linux on newer Python versions). Use a module-level function, or a small class with `__call__` if it needs configuration such as a pad id. *Return whatever your loop expects.* There is no required return type. Returning a dataclass or a dict of named tensors is usually more readable than a five-element tuple, and it survives refactoring better. *It is not the place for GPU work.* Worker processes should not touch CUDA; move to the device in the training loop after the batch arrives. ## The bucketing alternative Padding to the batch maximum still wastes compute when a batch mixes a 5-token and a 500-token sample. The complementary fix lives on the sampler side: a `batch_sampler` that groups indices of similar length so each batch pads to a tight maximum. Collate and bucketing solve the same waste from two ends, and mature text pipelines use both.

  • How does default_collate handle a sample that is a dict?
    It recurses per key and returns a dict with the same keys, each value collated by the same rules — so a batch of `{"input_ids": tensor, "label": int}` becomes `{"input_ids": stacked tensor, "label": int64 tensor}`. Tuples, lists and namedtuples are walked the same way; strings pass through as a plain Python list rather than becoming tensors.
  • Why pad inside collate_fn rather than inside __getitem__?
    Inside `__getitem__` you only see one sample, so you must pad to a global maximum length, which wastes memory and compute on every short batch. Collate sees the whole batch and can pad to that batch's maximum. It also runs in the worker process when `num_workers > 0`, so the padding cost is parallelized rather than serialized in the training loop.
  • Why does a lambda collate_fn sometimes fail?
    With the spawn start method, DataLoader pickles the collate function to send it to each worker, and lambdas are not picklable — you get a pickling error at loader startup. Use a module-level function, or a small class with `__call__` when the collate needs parameters like a pad token id. Under fork it happens to work, which is why the failure often appears only on macOS or Windows.

saying these in an interview costs you the question

  • Thinks collate_fn transforms one sample at a time
  • Claims the default pads variable-length sequences automatically
  • Pads to a global max length inside __getitem__
  • Says collation happens in the main process regardless of num_workers
  • Moves the batch to CUDA inside collate_fn

context