skip to content

In PyTorch, how does broadcasting decide the output shape of a + b?

level: middleimportance: must knowfreq 72%

answer

  1. align from the right, not the left
  2. size 1 stretches, other mismatches raise
  3. stretching is a stride, not a copy
  4. (N,1) against (N,) becomes N by N

basics

~20 s

PyTorch aligns the two shapes from the trailing dimension backwards. Each pair must be equal, or one of them must be 1, or missing entirely; size-1 and missing dimensions are stretched to the other operand's size. Anything else raises.

solid answer

~50 s

Broadcasting lines the shapes up from the **right**. Walking right to left, two dimensions are compatible if they are equal, or one of them is 1, or one shape has run out of dimensions; a 1 (or a missing leading dimension) is stretched to the other size, and any other mismatch raises a RuntimeError about non-broadcastable shapes. Nothing is copied — the stretched dimension is given stride 0, so `torch.ones(3, 1) + torch.ones(1, 5000)` reads the same memory repeatedly rather than materialising it. The trap is that broadcasting is *too* permissive: `pred` of shape `(N, 1)` minus `target` of shape `(N,)` broadcasts to `(N, N)` instead of failing, which is why the loss functions warn when input and target sizes differ. Fix the shape explicitly with `squeeze(-1)`, `unsqueeze(1)` or `reshape`, and use `torch.broadcast_shapes` to check a pairing before trusting it.

code

python · 7 lines
python
import torch

a = torch.ones(3, 1)
b = torch.ones(3)
print((a - b).shape)                 # torch.Size([3, 3]) - not [3, 1]
print((a - b.unsqueeze(1)).shape)    # torch.Size([3, 1])
print(torch.broadcast_shapes((5, 1, 4), (3, 4)))  # torch.Size([5, 3, 4])

go deeper

for a junior

Be able to state the rule out loud: compare shapes from the right, a dimension of size 1 stretches, everything else must match. Know that a (3,1) and a (3,) tensor combine to (3,3).

for a middle

Explain that stretching is implemented as a stride-0 view rather than a copy, that the output is still a full dense tensor, and name the (N,1)-versus-(N,) loss bug as the classic silent failure.

for a senior

Show how you catch this in production code: shape asserts at component boundaries, torch.broadcast_shapes in library code, and treating the input/target size warning as a build-breaking error rather than noise.

for a principal

Own the convention. Decide whether the codebase keeps a trailing singleton dimension everywhere or squeezes it at one defined boundary, and enforce it, since most broadcasting incidents come from two components disagreeing about that one axis.

## What broadcasting is Broadcasting is the rule PyTorch uses to run an elementwise operation on two tensors whose shapes are not identical. Instead of demanding that you materialise both operands at the same size, PyTorch pretends the smaller one is repeated along the mismatched dimensions and computes the result at the combined shape. Almost every elementwise op (`+`, `-`, `*`, `/`, comparisons, `torch.where`, `torch.maximum`) and the batch dimensions of `torch.matmul` follow it. ## The rule, step by step Align the two shapes at their **trailing** (rightmost) dimension and walk leftwards. For each pair of dimensions: - if the sizes are equal, that dimension passes through unchanged; - if one of them is 1, it is stretched to the other size; - if one shape has run out of dimensions, it is treated as having size 1 there and stretched; - otherwise the operation raises `RuntimeError: The size of tensor a (...) must match the size of tensor b (...) at non-singleton dimension ...`. So `(5, 1, 4)` with `(3, 4)` gives `(5, 3, 4)`: the trailing 4s match, the 1 stretches to 3, and the missing leading dimension of the second operand becomes 5. But `(4, 3)` with `(4,)` fails, because alignment is from the right — the trailing 3 and 4 are both non-singleton. This asymmetry is the single most misremembered part of the rule: a `(4,)` tensor broadcasts against `(3, 4)`, never against `(4, 3)`. ## Why nothing is copied A PyTorch tensor is a shape plus a set of strides over a flat storage buffer. To "stretch" a size-1 dimension, PyTorch sets that dimension's stride to 0, so every index along it reads the same element. That is what `Tensor.expand` exposes directly: it returns a stride-0 **view**, costing no memory, whereas `Tensor.repeat` genuinely allocates and copies. The practical consequence is that broadcasting is cheap on the input side but not on the output side — the *result* is a full dense tensor. `torch.ones(10000, 1) * torch.ones(1, 10000)` allocates 100 million elements, which is a common accidental out-of-memory in a distance-matrix or attention-mask computation. ## The classic silent bug The damage broadcasting does is rarely an exception; it is a wrong number. Two canonical cases: - **Loss shapes.** Model output `(N, 1)` and labels `(N,)` broadcast to an `(N, N)` matrix of pairwise differences. The loss still reduces to a scalar, still decreases, and the model still trains — on nonsense. PyTorch emits a `UserWarning` that the target size differs from the input size and that this will likely lead to incorrect results due to broadcasting; treat that warning as an error. - **Normalisation.** Subtracting a per-feature mean of shape `(D,)` from data of shape `(N, D)` is correct; subtracting a per-sample mean of shape `(N,)` from the same data is not — you need `mean.unsqueeze(1)` to get `(N, 1)`. ## What does not broadcast `torch.mm` and `torch.bmm` are strict: they require exactly 2-D and exactly 3-D operands with matching inner dimensions, and raise rather than broadcast. `torch.matmul` (and the `@` operator) broadcasts the *batch* dimensions but still requires the matrix dimensions to line up. In-place ops broadcast only in the direction that keeps the destination's shape: `a.add_(b)` works if `b` broadcasts up to `a`'s shape, but raises if the broadcast result would be larger than `a`, because there is nowhere to write it. That asymmetry is occasionally useful as a cheap shape assertion. ## Checking and fixing shapes `torch.broadcast_shapes(s1, s2, ...)` returns the resulting shape (or raises) without allocating anything, which makes it a good assertion in library code. `torch.broadcast_tensors(a, b)` returns both operands expanded to the common shape as views, useful when debugging. To fix a mismatch, be explicit about intent rather than hoping: `unsqueeze(dim)` to insert a size-1 axis, `squeeze(dim)` with an explicit dim to drop one (bare `squeeze()` drops *every* size-1 axis and will silently do the wrong thing on a batch of size 1), `reshape` when you know the exact target, `None` indexing (`x[:, None]`) as shorthand for `unsqueeze(1)`. ## Habits that prevent the bug Write the expected shape in a comment or an assert at every boundary between components; prefer keeping a trailing singleton dimension consistently rather than sometimes squeezing it; and when a metric looks plausible but slightly off, print the shapes of both operands of every arithmetic line before suspecting the model.

  • Your loss decreases but the metric is nonsense, and PyTorch printed a warning about target size differing from input size. What happened?
    The model output was `(N, 1)` and the target `(N,)`, so the elementwise difference broadcast to an `(N, N)` matrix of every prediction against every label. The reduction still produced a scalar that decreases, so training looks healthy. Fix it by matching shapes explicitly — `output.squeeze(-1)` or `target.unsqueeze(1)` — and treat that warning as a hard failure in CI.
  • What is the memory difference between expand and repeat when you broadcast a row to a full matrix?
    `expand` returns a view whose stretched dimension has stride 0, so it allocates nothing and every row reads the same storage; it only works on size-1 dimensions and the result is read-only in practice, since writing through it would write the same element many times. `repeat` materialises a genuine copy at the full size. If you only need the value for an elementwise op, let broadcasting or `expand` do it.
  • Does torch.mm broadcast the way the + operator does?
    No. `torch.mm` requires two 2-D tensors with matching inner dimensions and raises otherwise; `torch.bmm` requires 3-D with matching batch size. Only `torch.matmul` and `@` broadcast, and only over the leading batch dimensions — the last two dimensions must still satisfy the matrix-multiply rule. Using `mm` is therefore a useful way to force a shape error instead of a silent broadcast.

saying these in an interview costs you the question

  • Says shapes align from the left, not the right
  • Thinks broadcasting copies the smaller tensor into memory
  • Believes mismatched shapes always raise an error
  • Uses bare squeeze() and breaks on batch size 1
  • Assumes torch.mm broadcasts like elementwise addition

context