Architecture ResNet : implémentation et entraînement avec PyTorch

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())

Étiquettes: ResNet PyTorch DeepLearning ComputerVision NeuralNetwork

Publié le 11 août à 21h31