Le cœur de ResNet repose sur le concept de connexion sauteuse (skip connection). Deux variantes existent pour un bloc résiduel :
- Avec projection 1×1 : un convolution 1×1 transforme X afin d'alginer les dimensions avant l'addition, soit Y = Y + conv1x1(X).
- Sans projection 1×1 : l'addition directe est effectuée, soit Y = Y + X, lorsque les dimensions d'entrée et de sortie correspondent déjà.
Implémentation du bloc résiduel
import torch
import torch.nn as nn
import torch.nn.functional as F
class ResidualBlock(nn.Module):
"""Bloc résiduel contenant deux convolutions suivies de BatchNorm."""
def __init__(self, in_channels, out_channels, use_projection=False, stride=1):
super().__init__()
# Projection optionnelle via convolution 1x1
if use_projection:
self.projection = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride)
else:
self.projection = None
self.conv_a = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, stride=stride)
self.conv_b = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
self.norm_a = nn.BatchNorm2d(out_channels)
self.norm_b = nn.BatchNorm2d(out_channels)
def forward(self, x):
y = F.relu(self.norm_a(self.conv_a(x)))
y = self.norm_b(self.conv_b(y))
if self.projection is not None:
x = self.projection(x)
return F.relu(y + x)
Construction de l'architecture complète
Le réseau est organisé en cinq étapes successives. La première étage effectue un traitement initial du tenseur d'entrée, tandis que les suivantes empilent des blocs résiduels en doublant le nombre de filtres à chaque transition.
def build_residual_blocks(in_ch, out_ch, count, is_first=False):
"""Génère une séquence de blocs résiduels."""
blocks = []
for idx in range(count):
if idx == 0 and not is_first:
# Premier bloc : réduction spatiale et doublement des canaux
blocks.append(ResidualBlock(in_ch, out_ch, use_projection=True, stride=2))
else:
blocks.append(ResidualBlock(out_ch, out_ch))
return blocks
# Étage initial : convolution 7x7 + pooling
stage1 = nn.Sequential(
nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
)
stage2 = nn.Sequential(*build_residual_blocks(64, 64, 2, is_first=True))
stage3 = nn.Sequential(*build_residual_blocks(64, 128, 2))
stage4 = nn.Sequential(*build_residual_blocks(128, 256, 2))
stage5 = nn.Sequential(*build_residual_blocks(256, 512, 2))
resnet_model = nn.Sequential(
stage1, stage2, stage3, stage4, stage5,
nn.AdaptiveAvgPool2d((1, 1)),
nn.Flatten(),
nn.Linear(512, 10)
)
# Vérification des dimensions de sortie
dummy = torch.randn(1, 1, 224, 224)
for module in resnet_model:
dummy = module(dummy)
print(module.__class__.__name__, tuple(dummy.shape))
Entraînement sur Fashion-MNIST
import torch
from d2l import torch as d2l
# Chargement des données redimensionnées à 96x96
train_loader, test_loader = d2l.load_data_fashion_mnist(batch_size=256, resize=96)
# Lancement de l'entraînement sur GPU
d2l.train_ch6(resnet_model, train_loader, test_loader, num_epochs=10, lr=1, device=d2l.try_gpu())