Tensor Notation, Einsum & Computational Graphs
(B, T, D), learn broadcasting and einsum once, and you'll read any paper or any PyTorch repo without squinting. Tensors are how GPUs talk; einsum is the dictionary.After this lesson, you will be able to:
- Read and write tensor shapes like (B, T, D) and (B, C, H, W) and explain what each axis means in real ML systems
- Apply broadcasting rules to combine tensors of different shapes without writing explicit loops, and spot the silent (N, 1) versus (N,) bug
- Translate any matrix or tensor operation into Einstein-summation (einsum) string syntax, and back again
- Explain a computational graph as a directed acyclic graph of operations, compute a gradient on one by hand, and describe why reverse-mode autodiff is the algorithm of choice for training neural networks
Before You Start
#Why Tensors
A tensor is just a list of lists of lists. That is genuinely all it is. The reason it gets a fancy name is that ML code constantly stacks more "list of" layers. A single image is a 2D list, a batch of images is a 3D list, a batch of color images is a 4D list, and a batch of color video frames is a 5D list. The math stays the same; only the bookkeeping gets harder.
(B, T, D) and returns (B, T, D)," "this loss reduces (B,) to a scalar," "this projection takes (B, T, D_in) to (B, T, D_out)." Many deep-learning bugs are shape mismatches. Get the shapes right and the math usually follows.#Common shapes you will see every day
(B, D)— a batch of B feature vectors of dimension D. Tabular data, classifier logits (a classifier's raw scores, one per class).(B, T, D). Transformer activations (the numbers a layer produces). B sequences, T tokens, D-dimensional embeddings. Vary T per batch and you get padding masks (flags that mark the filler positions).(B, C, H, W)— CNN inputs/activations in PyTorch (channels-first).B=32, C=3, H=224, W=224is the canonical ImageNet shape.(B, H, W, C)— same, channels-last (TensorFlow default, also faster on Apple Silicon).(B, H, T, D_h)— multi-head attention activations. Attention (covered in the Transformers track) runsHsmaller copies called heads side by side;D_h = D / His the per-head dimension.(L, B, D)— RNN/LSTM convention: time-major instead of batch-major.
(32, 3, 224, 224) could be a batch of color photos OR 32 separate 3-channel volumetric scans of size (224, 224). Convention and context decide.Try it! Open the Python REPL (bottom-right of the screen: click Quick Actions, then Python) and type these lines yourself.
A multi-head attention layer for a batch of B=8 sequences of length T=64 with hidden size D=512 and H=8 heads splits D into per-head dimension d_h = D/H = 64. What is the shape of the queries tensor Q inside the attention computation?
#Broadcasting: Math Without Loops
X = [[1, 2, 3], [4, 5, 6]] (shape (2, 3)) and b = [10, 20, 30] (shape (3,)). Writing X + b gives [[11, 22, 33], [14, 25, 36]]: the vector b was added to every row of X, with no loop in sight. That automatic stretching is broadcasting, and NumPy and PyTorch both do it.1 (or missing, which counts as 1). A "1" axis gets virtually expanded (no memory copy) to match the other tensor. In the example, (2, 3) against (3,) pairs 3 with 3, and the missing axis of b stretches to 2.Pick two shapes in the widget below and watch which axes stretch to match and which pairs get rejected.
#Worked example: per-position positional encoding
X of shape (B, T, D) and positional encodings P of shape (T, D), the same for every example in the batch. Adding them must "broadcast" P across the batch axis.#The (B, 1, D) and (1, T, D) trick
q of shape (B, D) and a memory M of shape (T, D) and you want pairwise distances ‖q - M_t‖ for every batch element and every memory slot. You need an output of shape (B, T, D) (per-coordinate differences) before reducing to (B, T).#The classic silent bug
pred of shape (32, 1) (one prediction per row, with a trailing axis of size 1) and target has shape (32,). Right-aligning pairs the 1 against the 32 and stretches both, so pred - target quietly becomes a (32, 32) table of every prediction minus every target. The cell below shows the bug on 5 numbers, then the fix.pred (5, 1) target (5,) pred - target (5, 5), then broken loss: 6.4 and fixed loss: 1.2. Check the fixed one by hand: the five squared errors are 0, 1, 0, 1, 4, which sum to 6 and average to 1.2. The broken version averaged 25 mismatched pairs and got a different number with no warning.You have logits L of shape (B, V) (batch B, vocabulary V: one raw score per word) and a temperature t of shape (B,), a different temperature per example. You write L / t. What happens if you do not reshape t first?
#From Index Notation to Einsum
Matrix multiplication in standard notation:
i and j appear on the left-hand side and survive. k appears twice on the right-hand side and not at all on the left. That is precisely the index we summed over. Einstein's idea: drop the Σ and let the bookkeeping of which indices appear where say what is summed.Σ symbols all day in general relativity equations was driving him crazy. In his original convention, an index repeated within a product is summed automatically.einsum use a slightly more explicit version: you write the letters for each input, an arrow, and the letters you want in the output. The string 'ik,kj->ij' is the same equation typed sideways. Read it left-to-right:'ik'— A has axes labeled i (rows) and k (cols).','— separator between operands.'kj'— B has axes labeled k (rows) and j (cols).'->'— separator before the output spec.'ij'— the output has axes i and j.
Three rules tell you what any einsum string does:
- The same letter in different inputs means those axes are lined up, so their sizes must be equal. Here
kties A's columns to B's rows. - A letter that appears after the arrow is kept: it becomes an axis of the output (a free or batch axis). Here
iandj. - A letter that appears in the inputs but NOT after the arrow is summed away (multiplied along that axis, then added up). Here
k.
'ik,kj->ij' means: line up A and B along k, multiply, sum over k, keep i and j. That is matrix multiplication. Notice that sharing a letter is not by itself a reason to sum: what decides is whether the letter appears after the arrow.Toggle through the presets below to see einsum operations execute one cell at a time. Watch the summed index disappear, and the kept indices fall into their slots in the output tensor. Internalise the visual once and you will never have to reason about einsum from scratch again.
#Worked Examples: Einsum Cookbook
Once you internalize the three rules, dozens of operations collapse into one syntax.
#Matrix multiplication
#Batched matrix multiplication
(B, M, K) tensor times a (B, K, N) tensor gives (B, M, N). The batch axis b rides along untouched.#Attention scores (the famous one)
Q and keys K both have shape (B, H, T, d_h): batch, heads, sequence positions, per-head dimension. The attention-score matrix is shape (B, H, T, T): how much each query position attends to each key position, separately for each head and each batch.'bhid,bhjd->bhij' with the three rules. The letters after the arrow are b, h, i, j, so all four are kept. The letter d is in both inputs but not after the arrow, so it is the only one summed. Note that b and h also appear in both inputs, yet they are not summed, because they are in the output. That is the whole difference between a batch axis and a summed axis.V of shape (B, H, T, d_h), is 'bhij,bhjd->bhid': the attention matrix (T, T) weights (T, d_h) to produce (T, d_h) per head per example, summing over the sequence index j.d collapse while batch b, head h, query position i, and key position j survive into the output. Every "scaled dot-product attention" diagram you have ever seen is this one contraction with a 1/√d_h scaling and a softmax tacked on the end.A tensor A has shape (32, 16, 8, 64) and B has shape (32, 16, 64, 8). What is the output shape of einsum('bhij,bhjk->bhik', A, B)?
#Trace (sum of diagonal)
#Outer product
The opposite of summing: take two vectors and produce a matrix.
#Sum along axis (one input only)
#Elementwise product (nothing is summed)
Does the order of indices on the right-hand side of an einsum string matter? Compare einsum('ik,kj->ij', A, B) vs einsum('ik,kj->ji', A, B).
A whole einsum cookbook in one runnable cell. Each pattern below shows you the einsum string, what it computes, and a shape sanity-check. After running, edit the strings: swap letters, drop indices, see what breaks.
#Computational Graphs
x = 2 and y = 3 and compute three things in turn: u = x * y, v = x + y, and L = u * v. The graph has the inputs x and y, then the nodes u, v and L. Going forward: u = 6, v = 5, L = 30. Note that x feeds two nodes (u and v) and so does y.L change if x changes a little? Put the local slope (how much each node's output moves per unit change of one of its inputs) on every arrow:L = u * vgivesdL/du = v = 5anddL/dv = u = 6.u = x * ygivesdu/dx = y = 3.v = x + ygivesdv/dx = 1.
x to L. Through u: 5 * 3 = 15. Through v: 6 * 1 = 6. Add the paths: dL/dx = 15 + 6 = 21. For y: through u, 5 * 2 = 10 (since du/dy = x = 2); through v, 6 * 1 = 6; total dL/dy = 16. Check with algebra: L = x²y + xy², so dL/dx = 2xy + y² = 12 + 9 = 21 and dL/dy = x² + 2xy = 4 + 12 = 16. The same numbers.The picture below is this kind of graph. Numbers flow forward along the arrows, and the local slopes sit on the edges. Move an input and watch which node values change, then multiply the slopes along a path to see where each gradient comes from.
What we just did by hand is the general rule: the gradient of the final result with respect to an input equals the sum, over every path from that input to the result, of the product of the local slopes along the path.
loss). The backward pass walks the same graph in reverse. Starting from loss with seed gradient dloss/dloss = 1, each node uses its own local slopes plus its already-computed children's gradients to compute the gradient for its parents. In the example: seed 1 at L, then u gets 5, v gets 6, then x collects 5*3 + 6*1 = 21.#Forward mode vs reverse mode (a teaser)
- Forward mode computes derivatives in the same direction as the forward pass: it pushes a nudge (a "tangent vector", meaning one chosen direction in which to wiggle the inputs) forward through the graph. Cost: about one pass per input variable. Cheap if you have few inputs and many outputs.
- Reverse mode (a.k.a. backpropagation) traverses the graph backward from a single output: it pulls a "cotangent vector" (the gradient seed,
1for a scalar loss) back through the graph. Cost: about one pass total, regardless of how many inputs there are.
What loss.backward() actually does
loss.backward() walks the computational graph in reverse, applying the chain rule node by node. This lesson's DAG is exactly the graph being walked. Every tensor with requires_grad=True has its .grad field populated with the corresponding partial derivative. The optimizer (SGD, Adam, etc.) then reads .grad and updates the parameter.nn.Embedding(vocab, dim, sparse=True) exists: the gradient of the loss w.r.t. an embedding row is sparse. Only the rows for tokens that appeared in the batch get nonzero gradients (the others were never read during the forward pass, so the chain rule yields zero). A naive dense gradient tensor would be (vocab_size, dim), which for a large vocabulary is a lot of zeros to allocate. The sparse version stores only the rows that changed, cutting memory and compute for large embedding tables.#Try It Yourself
Three short exercises: einsum by hand, a 20-line autograd engine, and a shape bug.
#Einsum by hand
Tests · Verify trace(A) = 5, outer(a,b)[1,1] = 40, batched matmul produces shape (4, 2, 5), attention scores shape is (2, 4, 6, 6), the Frobenius inner product equals 70, the row sums are [3 7], and the transposed product is [[19 43] [22 50]].
A @ B = [[19 22] [43 50]], trace(A) = 5, the outer product [[10 20] [20 40] [30 60]], then the shapes (4, 2, 5) and (2, 4, 6, 6). The solution adds <A, B> = 70, row sums of A = [3 7] and (A @ B)^T = [[19 43] [22 50]].#Build autograd: 20 lines
loss.backward(). Each Value holds a number, a grad slot, and the nodes it was made from. Each operation records a tiny function _back that pushes gradient to its parents using the local slopes from the worked example. Use += because a node like x can feed several others, and its gradient is the sum over paths. The test graph is the one you solved by hand: x = 2, y = 3, u = x * y, v = x + y, L = u * v.Tests · Verify u = 6, v = 5, L = 30, dL/dx = 21, dL/dy = 16, dL/du = 5, dL/dv = 6, and that the check against the hand calculation prints True.
u, v, L = 6.0 5.0 30.0, then dL/dx = 21.0 dL/dy = 16.0, then dL/du = 5.0 dL/dv = 6.0, then matches hand calculation: True. Those are exactly the numbers from the path sums above. Before you fill in the TODOs the gradients print as 0.0, which is a good sign the skeleton is wired correctly. Notice what reversed(order) guarantees: a node's _back runs only after every node that depends on it has already pushed its gradient in.#Debug a shape
The broken loss from the broadcasting section, now as a tiny debugging task. Run it, find the bug, and fix it with an assert.
Tests · Verify the shapes are (5, 1), (5,) and (5, 5), the broken loss is 6.4, the fixed loss is 1.2, and the by-hand check is 1.2.
shapes: (5, 1) (5,) (5, 5), broken loss: 6.4, fixed loss: 1.2 and by hand: 1.2. Reshaping target to (5, 1) instead (with target[:, None]) also fixes it and gives the same 1.2.Tensor Comprehensions: Framework-Agnostic High-Performance Machine Learning Abstractions
Nicolas Vasilache, Oleksandr Zinenko, Theodoros Theodoridis, Priya Goyal, Zachary DeVito, William S. Moses, Sven Verdoolaege, Andrew Adams, Albert Cohen (2018)
Pioneering paper introducing index-notation as a programming abstraction for ML.
#Key Takeaways
- Tensor = list of lists of lists. Rank counts how many axes it has; shape is the size at each axis. Many ML bugs are shape mismatches, so train your eye to read shapes.
- Broadcasting eliminates loops. Right-aligned shape compatibility lets you add a (D,) bias to a (B, D) batch, or compare a (B, 1, D) query to a (1, T, D) memory, without copying tensors. But a (N, 1) minus (N,) silently gives (N, N), so assert shapes before losses.
- Einsum, in two decisions. Letters after the arrow are kept; input letters missing from the output are summed. Matmul, batch matmul, trace, outer product and attention scores all follow.
- Attention scores. The string
'bhid,bhjd->bhij'keeps b, h, i, j and sums d: a dot product over the per-head dimension for every query-key pair. - Computational graphs make autodiff possible. Forward = topological evaluation of a DAG; backward = the gradient is the sum over paths of products of local slopes, computed in reverse. PyTorch builds it dynamically; JAX traces it once.
loss.backward()is just one DAG traversal.
#Quick Check
An einsum string is `'bhid,bhjd->bhij'`. Which axis is summed away?