CNN Backbones and Pretraining

Residual and skip connections

A skip connection adds a layer's input back onto its output, which is the single change that made networks hundreds of layers deep trainable at all.

On this page 7
  1. What was broken before
  2. How the fix works
  3. Why this rescues training
  4. Where you have already seen this
  5. The honest part
  6. Remember this
  7. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A skip connection takes what went into a layer and adds it back onto what came out.

You have played the whispering game at a party. One person whispers a sentence and it travels down a line of twenty people. What comes out the other end is nonsense. Every person adds a small error, and the errors pile up.

Now imagine each person also passes along the original slip of paper, unchanged, alongside their whisper. The message survives twenty people easily, because there is a clean path from one end to the other.

That slip of paper is a skip connection. It is the most important idea in this whole section.

What was broken before

Between 2012 and 2015, everyone believed deeper networks were better. Then people built very deep ones and found something odd.

A network with fifty-six layers scored worse than one with twenty layers. Not worse on new photos, which would mean it had memorised. Worse on the very photos it was trained on.

That is a strange failure. A deeper network can copy the shallower one and leave the extra layers doing nothing. So it should never be worse.

The problem was not capacity. The problem was that training could not find that solution. The signal that guides learning had to travel back through fifty-six layers. It faded to nothing on the way.

How the fix works

Instead of asking a layer to produce the answer, ask it to produce the change to the answer.

   plain layer:       input  ->  [ layer ]  ->  output

   residual layer:    input  ->  [ layer ]  ->  +  ->  output
                        \___________________/
                          the input, untouched

Now a layer that has nothing useful to add can output nothing, and the input passes through unchanged. Doing nothing is the easy option rather than the hard one.

The word residual means "what is left over". The layer learns only the leftover correction, not the whole answer.

Why this rescues training

Learning works by sending a correction signal backwards from the mistake to every layer. Through a plain stack, that signal gets multiplied by something at every step. Multiply by a small number forty times and the result is close to zero.

The skip connection gives the signal a second route with nothing in the way. It arrives at the early layers intact.

That is measured further down this page, and the difference is not small. It is a factor of many trillions.

Where you have already seen this

  • Every modern photo-tagging model in your phone's gallery is built from residual blocks.
  • The language model you chat with has skip connections in every one of its layers.
  • Speech recognition, image generation and protein folding models all use the same trick.

It is close to universal. If a network is deeper than about twenty layers, it almost certainly has skip connections.

The honest part

Adding the input back is not free. The two things being added must be the same shape.

When a layer shrinks the picture or changes the number of channels, the shortcut has to be reshaped too. Real networks solve this with a small extra layer on the shortcut. Getting it wrong is a common bug.

Remember this

  • A skip connection adds a layer's input back onto its output.
  • It lets a layer learn a small correction instead of the whole answer.
  • It gives the learning signal a clean path back to the early layers.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install "torch==2.5.1"

Run against PyTorch 2.5.1 on CPU, in about a second.

Watch the gradient die, then watch the skip save it

skip_connections.py
import torch
import torch.nn as nn

C, N = 8, 40                 # 8 channels, 40 stacked blocks


class Plain(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = nn.Conv2d(C, C, 3, padding=1)

    def forward(self, x):
        return torch.relu(self.conv(x))


class NaiveRes(Plain):
    def forward(self, x):
        return x + torch.relu(self.conv(x))       # the whole idea, in one +


class ResNetBlock(nn.Module):
    """conv-bn-relu-conv-bn, then add, then relu: the real BasicBlock."""
    def __init__(self, zero_init=False):
        super().__init__()
        self.conv1 = nn.Conv2d(C, C, 3, padding=1, bias=False)
        self.bn1 = nn.BatchNorm2d(C)
        self.conv2 = nn.Conv2d(C, C, 3, padding=1, bias=False)
        self.bn2 = nn.BatchNorm2d(C)
        if zero_init:
            nn.init.zeros_(self.bn2.weight)       # block starts as an exact identity

    def forward(self, x):
        out = torch.relu(self.bn1(self.conv1(x)))
        out = self.bn2(self.conv2(out))
        return torch.relu(out + x)


def probe(make_block, name):
    torch.manual_seed(0)
    net = nn.Sequential(*[make_block() for _ in range(N)])
    torch.manual_seed(1)
    x = torch.randn(16, C, 8, 8)
    acts, h = [], x
    for i, blk in enumerate(net):
        h = blk(h)
        if i in (0, 9, 39):
            acts.append(h.abs().mean().item())
    net.zero_grad()
    net(x).pow(2).mean().backward()
    first = next(p for p in net[0].parameters() if p.grad is not None and p.dim() == 4)
    print(f"{name:24s} {first.grad.norm().item():10.2e}   "
          f"{acts[0]:9.2e} {acts[1]:9.2e} {acts[2]:9.2e}")


print(f"{'block type':24s} {'grad at':>10s}   {'signal size after 1 / 10 / 40 blocks':>33s}")
print(f"{'':24s} {'layer 1':>10s}")
probe(Plain, "plain stack")
probe(NaiveRes, "naive x + f(x)")
probe(lambda: ResNetBlock(zero_init=False), "ResNet block")
probe(lambda: ResNetBlock(zero_init=True), "ResNet block, zero-init")

print("\nwhat 'zero-init' means, checked directly")
torch.manual_seed(0)
net = nn.Sequential(*[ResNetBlock(zero_init=True) for _ in range(N)])
torch.manual_seed(1)
x = torch.randn(16, C, 8, 8)
out = net(x)
print("  max |output - relu(input)| across 40 blocks:", (out - torch.relu(x)).abs().max().item())
out.pow(2).mean().backward()
print("  gradient on block 1 conv1 weight :", net[0].conv1.weight.grad.norm().item())
print("  gradient on block 1 batch-norm scale:", net[0].bn2.weight.grad.norm().item())
print("  gradient on block 40 batch-norm scale:", net[39].bn2.weight.grad.norm().item())
Output
block type                  grad at   signal size after 1 / 10 / 40 blocks
                            layer 1
plain stack                1.68e-20    2.14e-01  2.33e-02  5.07e-02
naive x + f(x)             1.25e+07    8.62e-01  3.91e+00  4.23e+03
ResNet block               3.67e+01    5.62e-01  2.06e+00  4.36e+00
ResNet block, zero-init    0.00e+00    4.01e-01  4.01e-01  4.01e-01

what 'zero-init' means, checked directly
  max |output - relu(input)| across 40 blocks: 0.0
  gradient on block 1 conv1 weight : 0.0
  gradient on block 1 batch-norm scale: 0.03991314396262169
  gradient on block 40 batch-norm scale: 0.04593270644545555

Reading the output line by line

The plain stack is dead. A gradient norm of 1.68e-20 at layer 1 means the first layer receives no usable learning signal. Forty multiplications by numbers below one did that. Deeper is not better when deeper means untrainable.

The naive skip fixes the gradient and creates a new problem. The gradient is now 1.25e+07, which is enormous, and the activations have grown to 4.23e+03 by block 40. Adding the input at every block makes the signal grow without limit. This is honest and worth seeing: the raw idea alone is not enough.

The real ResNet block is stable in both directions. Batch normalisation inside the block keeps the activation scale near 4.4 after forty blocks, and the gradient at layer 1 is a sane 36.7. Residual connections and normalisation were introduced together for this reason. Neither works well alone at this depth.

The zero-init row looks like a failure and is not. The forty-block network reproduces its own input exactly, to 0.0 difference. Every block starts as a perfect identity. The convolution weights get zero gradient at step one, but the batch-norm scale parameters do not, and they are what open each block up. This is the zero_init_residual trick from Goyal et al. (2017), and it means adding depth cannot hurt your starting point.

Common mistakes

Adding tensors of different shapes. When a block changes channel count or stride, out + x raises a size error. Real ResNets put a 1x1 convolution with matching stride on the shortcut. In torchvision it is called downsample, and it appears only in the first block of each stage.

Putting the activation in the wrong place. The original design does relu(f(x) + x). The pre-activation variant does x + f(relu(bn(x))) and keeps the shortcut completely clean, which trains better at extreme depth. Mixing the two is a common copying error.

Using concatenation and calling it a residual. DenseNet concatenates instead of adding. That is a different architecture with different memory behaviour. Addition keeps the channel count fixed; concatenation grows it.

Assuming skips remove the need for normalisation. The naive x + f(x) row shows what happens. Without normalisation the forward signal explodes even though the gradients survive.

Forgetting the shortcut in a custom block. return self.conv2(self.conv1(x)) inside a class named ResBlock is a bug that trains, converges slowly and is invisible until you read the code.

Try it yourself

Change N from 40 to 100 and rerun. The plain gradient underflows completely to 0.0, while the ResNet block barely moves. Then delete the two BatchNorm2d layers from ResNetBlock and watch the activation scale in the last column start to climb again.

What to learn next

Researcher — Mathematics and papers.

The formulation

He, Zhang, Ren and Sun (2015), Deep Residual Learning for Image Recognition, arxiv.org/abs/1512.03385, define a building block as

$$ y = \mathcal{F}(x, {W_i}) + x $$

where $x$ is the block input, $y$ the output, and $\mathcal{F}$ the residual mapping realised by the stacked layers with weights ${W_i}$. When shapes differ, a linear projection $W_s$ is applied to the shortcut:

$$ y = \mathcal{F}(x, {W_i}) + W_s x . $$

The paper's motivating observation is the degradation problem: a 56-layer plain network has higher training error than a 20-layer one on CIFAR-10. This is not overfitting and not vanishing gradients alone, since batch normalisation was already in use. It is an optimisation failure, and the hypothesis is that solvers find it hard to approximate the identity with a stack of non-linear layers, while driving $\mathcal{F}$ to zero is easy.

Their ImageNet result: a 152-layer ResNet, with an ensemble reaching 3.57% top-5 error, winning ILSVRC 2015, plus a 28% relative improvement on COCO detection.

Why the gradient survives

For a chain of residual units $x_{l+1} = x_l + \mathcal{F}(x_l, W_l)$, unrolling gives

$$ x_L = x_l + \sum_{i=l}^{L-1} \mathcal{F}(x_i, W_i) $$

and therefore

$$ \frac{\partial \mathcal{L}}{\partial x_l} = \frac{\partial \mathcal{L}}{\partial x_L}\left(1 + \frac{\partial}{\partial x_l}\sum_{i=l}^{L-1}\mathcal{F}(x_i, W_i)\right). $$

The leading $1$ is the whole argument. The gradient reaching layer $l$ contains an additive term that is not multiplied by any weight matrix, so it cannot vanish through depth. This derivation is from He et al. (2016), Identity Mappings in Deep Residual Networks, arxiv.org/abs/1603.05027, which also shows the clean derivation holds only when the shortcut is an exact identity. Gating, scaling or convolving the shortcut reintroduces a multiplicative factor and measurably hurts very deep models.

That paper introduces pre-activation: $x_{l+1} = x_l + \mathcal{F}(\mathrm{BN}(\mathrm{ReLU}(x_l)))$, which keeps the shortcut path free of any operation and enabled a trainable 1001-layer network on CIFAR.

Competing explanations

Veit, Wilber and Belongie (2016), Residual Networks Behave Like Ensembles of Relatively Shallow Networks, arxiv.org/abs/1605.06431, show that a network of $n$ residual blocks expands into $2^n$ paths of varying length, that deleting a single block barely changes accuracy, and that the gradient magnitude is dominated by short paths. Under this view depth buys ensemble diversity rather than a long computation.

Li et al. (2018), Visualizing the Loss Landscape of Neural Nets, show that skip connections remove the chaotic non-convexity that appears in deep plain networks, producing a landscape that gradient descent can traverse.

Balduzzi et al. (2017), The Shattered Gradients Problem, show that in plain networks the spatial correlation of gradients decays exponentially with depth, approaching white noise, while residual networks decay only as a square root.

These are complementary rather than competing. The empirical fact that all of them explain is the same: identity paths make depth optimisable.

Initialisation refinements

  • Zero-init residual (Goyal et al., 2017, Accurate, Large Minibatch SGD, arxiv.org/abs/1706.02677) sets the final batch-norm scale of each block to zero, so the network begins as the identity. Reported as worth roughly 0.2 points of ImageNet top-1 and better large-batch stability. It is exposed in torchvision as resnet50(zero_init_residual=True).
  • Fixup (Zhang, Dauphin and Ma, 2019) and SkipInit (De and Smith, 2020) show that careful scaling of the residual branch trains deep residual networks with no normalisation at all, which supports the view that normalisation's role here is scale control rather than anything deeper.
  • LayerScale (Touvron et al., 2021, CaiT) generalises the idea with a learnable per-channel multiplier on the residual branch, initialised near zero. It is standard in transformers and in ConvNeXt.

Where the shortcut is not an identity

The projection shortcut is used only where the shape changes: three times in a standard ResNet. He et al. (2019), Bag of Tricks for Image Classification, propose ResNet-D, which replaces the stride-2 1x1 projection with average pooling followed by a stride-1 1x1, on the grounds that a stride-2 1x1 convolution discards three quarters of its input outright. The change is close to free and is now common in timm variants.

What to learn next