gather, scatter and index_select
gather picks a different element from each row using a tensor of positions, scatter writes the same way — the pattern behind every cross-entropy loss.
- 7 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.
gather picks one element from each row, using a list that says which position to take from that particular row.
Think of serving a thali to a table of guests. Each guest points at a different dish — paneer for one, dal for the next. You walk down the line with their requests and serve each person exactly what they pointed at. One pass, every guest served differently.
Ordinary indexing cannot do this. It takes the same column from every row. gather takes a different position per row, guided by a list of requests.
Why it exists
The situation appears constantly in AI. A model scores every possible answer, for every example in a batch. Then you need each example's score for its own correct answer — a different column for every row.
With a loop: slow, and one line per example. With gather: one operation for the whole batch. Its mirror twin scatter does the reverse — it writes a value into a chosen position per row.
How it works
scores (one row per sample): requests: gather takes:
[ 2.0 0.5 1.0 ] position 0 -> 2.0
[ 0.1 3.0 0.2 ] position 1 -> 3.0Read gather's motto as: "for each row, take the element my request list names." Scatter's motto: "for each row, write into the slot my request list names."
A real example you have seen
An exam-marking program holds every student's chosen option — A, B, C or D — and the mark sheet holds marks for all options per question. Marking is one gather: for each student-question pair, pick the mark of the option they chose. Nobody marks by looping through students one by one.
Remember this
- Plain indexing takes the same position from every row.
- gather reads a different position per row; scatter writes one.
- The request list is itself a tensor, so the whole thing runs in one pass.
What to learn next
- einsum for tensor operations — a notation that replaces whole families of shape gymnastics.
- Loss functions — cross-entropy, the gather you will use most.
- Indexing and boolean masks — the simpler selection tools these generalise.
Developer — Code and libraries.
Setup
pip install torchOutputs verified with torch 2.5.1, CPU.
The three tools side by side
import torch
scores = torch.tensor([[2.0, 0.5, 1.0], # sample 0: score per class
[0.1, 3.0, 0.2]]) # sample 1: score per class
labels = torch.tensor([0, 1]) # the true class of each sample
# pick each sample's score for ITS OWN label, in one shot
picked = scores.gather(1, labels.unsqueeze(1))
print(picked)
# scatter_ writes instead of reads: build one-hot rows
one_hot = torch.zeros(2, 3)
one_hot.scatter_(1, labels.unsqueeze(1), 1.0)
print(one_hot)
# index_select takes whole rows (or columns) by position
print(scores.index_select(0, torch.tensor([1, 0])))tensor([[2.],
[3.]])
tensor([[1., 0., 0.],
[0., 1., 0.]])
tensor([[0.1000, 3.0000, 0.2000],
[2.0000, 0.5000, 1.0000]])The walkthrough
gather(1, index) — the first argument is dim, the direction the requests point along. Here dim=1: requests name columns, one per row. The rule for the shapes: index has the same number of dimensions as the input, and the output has exactly the index's shape. That is why labels needed unsqueeze(1) — (2,) became (2, 1): two rows, one request each, giving a (2, 1) result.
Sample 0's request was 0, so it got 2.0. Sample 1's request was 1, so it got 3.0. This is the beating heart of cross-entropy: pick each sample's score for the true class. See loss functions for where it lands.
scatter_(1, index, value) — same request format, opposite direction: write 1.0 into each named slot. Starting from zeros, that builds one-hot rows. The trailing underscore means in-place — it edits one_hot directly. (For real projects, torch.nn.functional.one_hot(labels, 3) says the same thing more readably.)
index_select(0, order) — whole rows by position, allowed to repeat and reorder. It is the tensor form of fancy indexing scores[[1, 0]], and the operation behind shuffling and mini-batch assembly.
Common mistakes
Float requests. Index tensors must be int64. A float index raises:
RuntimeError: gather(): Expected dtype int64 for index
The usual source is labels that passed through an average or a division somewhere. Cast with .long() — but ask why they became floats first.
Wrong dim. gather(0, ...) with the same index runs happily and reads down columns instead of across rows — wrong numbers, no error. Say the motto aloud: requests along dim 1 name columns; along dim 0 they name rows.
Forgetting the unsqueeze. Handing gather a flat (batch,) label tensor fails the "same number of dimensions" rule with a size-mismatch error. labels.unsqueeze(1) is part of the idiom; type them together.
Out-of-range requests. A request of 3 against 3 columns raises an index-out-of-bounds error on CPU. On GPU the same bug can surface later, as a vague device-side assert triggered. When you see that on CUDA, suspect your indices and re-run once on CPU for an honest message — the same trick as in device debugging.
Scatter with duplicate requests. Two requests naming the same slot means one silently wins. If you meant "add both", the operation is scatter_add_, which is also the one with defined behaviour under duplicates.
Try it yourself
Reverse the demo: given picked requests labels, use scatter_ to write each sample's picked score back into a zeros tensor at the label position — reconstructing a masked version of scores. Then break something on purpose: request position 5 and read the error you get.
What to learn next
- einsum for tensor operations — a notation that replaces whole families of shape gymnastics.
- Loss functions — cross-entropy, the gather you will use most.
- Indexing and boolean masks — the simpler selection tools these generalise.
Researcher — Mathematics and papers.
Exact semantics
For 2-D tensors (generalisation to k dimensions is index-wise):
out = input.gather(dim, index), dim=1: out[i][j] = input[i][ index[i][j] ]- dim=0: out[i][j] = input[ index[i][j] ][j]
self.scatter_(dim, index, src), dim=1: self[i][ index[i][j] ] = src[i][j]
Where index has the same rank as input, every index value lies in [0, size(dim)), and out.shape = index.shape. gather is a generalised read; scatter its adjoint write. index_select(d, idx) is the rank-preserving special case selecting whole slices along d.
Complexity: O(numel(index)) reads or writes, embarrassingly parallel, memory-bound. On GPU, locality of the index pattern determines coalescing — random gathers pay far more per element than contiguous ones.
Autograd: each is the other's backward
The differential of gather with respect to input is a scatter-add of the upstream gradient into the gathered positions; the differential of scatter (with respect to src) is a gather. This adjoint pairing is why both exist as primitives: implementing either's backward requires the other. Duplicate indices in the backward scatter-add accumulate — exactly the multivariate chain rule summing over paths.
Determinism
scatter_ with duplicate indices has undefined winner-selection; scatter_add_ and index_add_ are well-defined in value but implemented with atomic adds on CUDA — float atomics reorder, so results vary in the last bits between runs. Under torch.use_deterministic_algorithms(True), some of these ops switch to deterministic (slower) implementations or raise, per operation, as documented in the reproducibility notes. Relevant when you chase bit-identical reproducible runs.
Where the pattern shows up at scale
- Cross-entropy: log-softmax then gather of the target column — fused inside
F.cross_entropy. - Embedding lookup:
nn.Embeddingis index_select on the weight matrix; its backward is the sparse scatter-add that motivatessparse=Truefor huge vocabularies. - Beam search / top-k decoding: repeated gathers reorder candidate states by beam index each step.
- Graph neural networks: message passing is gather (neighbour features) followed by scatter-add (aggregate at nodes) — libraries like PyTorch Geometric (Fey and Lenssen, 2019, Fast Graph Representation Learning with PyTorch Geometric) are built on exactly these primitives.
What to learn next
- einsum for tensor operations — a notation that replaces whole families of shape gymnastics.
- Loss functions — cross-entropy, the gather you will use most.
- Indexing and boolean masks — the simpler selection tools these generalise.