Lesson 8 / 25
Tensors and Shapes
Most bugs are shape bugs.
Multi-dimensional arrays with a convention
A tensor is a multi-dimensional array with a data type and a device (CPU or GPU). Deep learning code follows shape conventions: images as (batch, channels, height, width), sequences as (batch, length, features) when batch_first=True, and fully connected layers expect (batch, features). Matrix multiplication needs the inner dimensions to match; broadcasting stretches smaller tensors (like a bias vector) across a batch. Printing shapes at each step is the fastest way to understand and debug a model.
Tensors, modules, loops
Every PyTorch project uses the same building blocks: tensors with shapes, modules with parameters, and a training loop.
Shapes, flattening and broadcasting, run
I ran this on CPU with Python 3, PyTorch 2.14.1, numpy 2.5.3 and scikit-learn 1.9.1, with fixed seeds and one thread. A batch of 32 RGB 28x28 images has 75,264 numbers. Flattening gives 32 rows of 2,352 values; multiplying a (32, 784) batch by a (784, 10) weight matrix gives (32, 10), and adding a bias of length 10 broadcasts across the batch.
import torch
batch = torch.zeros(32, 3, 28, 28) # 32 images, 3 channels, 28x28 pixels
print("shape:", tuple(batch.shape), "| dtype:", batch.dtype, "| elements:", batch.numel())
flat = batch.view(32, -1)
print("flattened per image:", tuple(flat.shape))
a = torch.randn(32, 784); W = torch.randn(784, 10)
print("matrix multiply (32x784) @ (784x10) ->", tuple((a @ W).shape))
print("broadcast add of a bias (10,) ->", tuple((a @ W + torch.zeros(10)).shape))
Output:
shape: (32, 3, 28, 28) | dtype: torch.float32 | elements: 75264 flattened per image: (32, 2352) matrix multiply (32x784) @ (784x10) -> (32, 10) broadcast add of a bias (10,) -> (32, 10)
Assert shapes in forward
Add a few assert x.shape[1:] == (...) checks while developing; they turn silent shape mistakes into clear errors.
Quick check: What shape convention do PyTorch image layers expect?
- (channels, batch)
- (height, width) only
- (batch, channels, height, width)
- (pixels, labels)
Answer
(batch, channels, height, width) — NCHW is the default layout.