TD Architectures profondes

Author

Augustin Chevallier

1 Exercice 0 : diagnostiquer un entraînement

Reprendre le modèle de l’exercice 2 de la feuille précédente (CNN sur MNIST).

Tracer dans le temps:

  1. l’erreur sur le jeu d’entraînement;
  2. l’erreur sur un jeu de validation (ne pas oublier de desactiver le gradient);
  3. le taux de bonnes classifications (sur la validation) ;
  4. la norme du gradient, couche par couche.

Que diagnostique chacune de ces courbes ? En particulier : que signale un écart qui se creuse entre les courbes 1 et 2 ? Et une courbe 4 qui décroît fortement quand on va vers les premières couches ?

Indication pour la norme du gradient : après loss.backward(), chaque paramètre porte son gradient dans .grad.

# après loss.backward(), avant optimizer.step()
normes = [conv.weight.grad.norm().item() for conv in model.convs]

Remarque: Ces réseaux sont petits et entraînés peu de temps : l’écart entre deux graines aléatoires peut dépasser l’écart entre deux modèles. Relancer chaque configuration sur plusieurs graines (torch.manual_seed) et comparer les moyennes, pas un tirage unique.

2 Exercice 1 : effet des connexions résiduelles

Comparer des réseaux à 2, 5 et 10 blocs de convolution (avec 2 convolutions par blocs) — soit 4, 10 et 20 convolutions — avec et sans connexions résiduelles, d’abord sur MNIST puis sur CIFAR-10.

Utiliser les diagnostics de l’exercice 0 pour analyser ce qui se passe.

3 Exercice 2 : ajouter une normalisation

Le cours présente LayerNorm. Sur des images, nn.LayerNorm est peu pratique car il faut lui donner les dimensions spatiales. On utilise nn.GroupNorm, qui fait la même chose sans cette contrainte :

import torch
import torch.nn as nn

C = 32 # nombre de channels

# image -> c'est LayerNorm, mais sans avoir à préciser H et W
norm = nn.GroupNorm(1, C) # le 1 correspond à un nombre de groupes, ici on normalise sur tout les channels en même temps

# une normalisation se place entre la convolution et l'activation
conv = nn.Conv2d(3, C, 3, padding=1)
x = torch.randn(8, 3, 32, 32)
y = torch.relu(norm(conv(x)))
print(y.shape)

# variantes : nn.GroupNorm(8, C) normalise par groupes de canaux,
#             nn.BatchNorm2d(C)  normalise sur la dimension du batch

Reprendre l’exercice 1 en ajoutant cette normalisation. Que change-t-elle, sur les diagnostics comme sur la précision finale ?