PyTorTutoriel sur l'apprentissage par transfert avec exemples
โก Rรฉsumรฉ intelligent
L'apprentissage par transfert rรฉutilise un rรฉseau dรฉjร entraรฎnรฉ sur un grand ensemble de donnรฉes afin qu'une nouvelle tรขche connexe puisse รชtre rรฉsolue avec beaucoup moins d'images รฉtiquetรฉes et une fraction du temps d'entraรฎnement initial.
Qu'est-ce que l'apprentissage par transfert ?
Transfert d'apprentissage est une technique qui consiste ร utiliser un modรจle entraรฎnรฉ pour rรฉsoudre une autre tรขche connexe. Machine Learning Cette mรฉthode de recherche consiste ร capitaliser sur les connaissances acquises lors de la rรฉsolution d'un problรจme particulier et ร les rรฉutiliser pour rรฉsoudre un autre problรจme, diffรฉrent mais connexe. Elle permet ainsi d'amรฉliorer l'efficacitรฉ en rรฉutilisant les informations recueillies lors de la tรขche prรฉcรฉdemment apprise.
Il est courant de rรฉutiliser les poids d'un autre modรจle de rรฉseau, car l'entraรฎnement d'un rรฉseau ร partir de zรฉro nรฉcessite une trรจs grande quantitรฉ de donnรฉes. Pour rรฉduire le temps d'entraรฎnement, on utilise un rรฉseau existant et ses poids, puis on modifie la derniรจre couche pour rรฉsoudre le problรจme spรฉcifique. L'avantage est que cette derniรจre couche peut รชtre entraรฎnรฉe avec un petit ensemble de donnรฉes.
Avant d'รฉcrire du PyTorPour connaรฎtre le code ch, il est utile de savoir ร quelle famille d'apprentissage par transfert appartient votre problรจme, car cela dรฉtermine la quantitรฉ de donnรฉes รฉtiquetรฉes dont vous avez besoin.
Types dโapprentissage par transfert
La littรฉrature scientifique divise cette technique en trois familles. Le choix de celle qui s'applique ร votre cas dรฉpend de la nature du problรจme, et non du cadre thรฉorique utilisรฉ.
| Type | Domaine source | Target domaine | Utilisation typique |
|---|---|---|---|
| Inductif | รtiquetรฉ | รtiquetรฉ, mais tรขche diffรฉrente | Rรฉorientation d'une infrastructure ImageNet vers un problรจme ร deux classes Alien vs Predator |
| Transductrice | รtiquetรฉ | Non รฉtiquetรฉ, mรชme tรขche, distribution de donnรฉes diffรฉrente | Adaptation de domaine, comme le passage d'un modรจle de photos de studio ร des photos prises avec un tรฉlรฉphone portable. |
| Non supervisรฉ | Non รฉtiquetรฉ | Non รฉtiquetรฉ | Clusterrรฉduction de dimensionnalitรฉ ou รฉtiquetage de chaque enregistrement, lorsque cela est impraticable. |
L'exemple prรฉsentรฉ dans ce tutoriel est un apprentissage par transfert inductif. VGG19 arrive avec des connaissances รฉtiquetรฉes d'ImageNet, et il est ensuite appliquรฉ ร un problรจme รฉtiquetรฉ ร deux classes qu'il n'a jamais rencontrรฉ auparavant.
Fonctionnalitรฉ Extraction vs rรฉglage fin
Une fois un rรฉseau prรฉ-entraรฎnรฉ choisi, il existe deux maniรจres de l'adapter. La diffรฉrence rรฉside simplement dans le nombre de couches que l'on autorise ร continuer d'apprendre.
| Aspect | Fonctionnalitรฉ extracproduction | Rรฉglage fin |
|---|---|---|
| Couches qui s'entraรฎnent | Seul le classificateur remplacรฉ | Le classificateur plus certains ou tous les blocs convolutionnels |
| requires_grad sur le backbone | Faux | C'est vrai pour les blocs mis ร jour |
| Donnรฉes nรฉcessaires | Petites, souvent quelques centaines d'images par classe | Plus importants, gรฉnรฉralement des milliers |
| Coรปt de la formation | Niveau le plus bas, fonctionne sur un processeur | Plus le prix est รฉlevรฉ, plus un GPU devient intรฉressant ร possรฉder |
| Prรฉcision typique | C'est bien lorsque les images source et cible se ressemblent. | Il est gรฉnรฉralement prรฉfรฉrable que les deux domaines diffรจrent. |
Les รฉtapes ci-dessous utilisent la fonctionnalitรฉ extracModification : tous les paramรจtres de VGG19 sont gelรฉs et seule la nouvelle couche linรฉaire finale apprend. Passer au rรฉglage fin consiste simplement ร laisser `requires_grad` ร `True` sur les blocs ร mettre ร jour et ร rรฉduire le taux d'apprentissage afin de ne pas altรฉrer les poids empruntรฉs.
Chargement de l'ensemble de donnรฉes
Avant de commencer ร utiliser l'apprentissage par transfert avec PyTorch, vous devez comprendre l'ensemble de donnรฉes que vous allez utiliser. Dans ce Transfer Learning PyTorPar exemple, vous devrez classifier un extraterrestre et un prรฉdateur parmi prรจs de 700 images. Cette technique ne nรฉcessite pas une grande quantitรฉ de donnรฉes d'entraรฎnement. Vous pouvez tรฉlรฉcharger l'ensemble de donnรฉes depuis [lien manquant]. Kaggle : Alien contre Predator.
La collection est volontairement restreinte, et un รฉchantillon des images qu'elle contient est prรฉsentรฉ ci-dessous.
Source: Vs extraterrestres prรฉdateur Kaggle
Prochain article sur PyTorDans ce tutoriel sur l'apprentissage par transfert, vous apprendrez ร appliquer l'apprentissage par transfert avec Python.Torch รฉtape par รฉtape.
Comment utiliser lโapprentissage par transfert ?
Voici une procรฉdure รฉtape par รฉtape pour utiliser l'apprentissage par transfert dans le cadre de l'apprentissage profond avec Python.Torch:
รtape 1) Charger les donnรฉes
La premiรจre รฉtape consiste ร charger les donnรฉes et ร appliquer certaines transformations aux images afin qu'elles correspondent aux exigences du rรฉseau.
Vous chargerez les donnรฉes depuis un dossier contenant torchvision.datasets. Le module parcourt ce dossier pour diviser les donnรฉes en ensembles d'entraรฎnement et de validation. Le pipeline de transformation utilisรฉ ici recadre les images en les centrant, les convertit en tenseur et les normalise. L'apprentissage en profondeur.
from __future__ import print_function, division import os import time import torch import torchvision from torchvision import datasets, models, transforms import torch.optim as optim import numpy as np import matplotlib.pyplot as plt data_dir = "alien_pred" input_shape = 224 mean = [0.5, 0.5, 0.5] std = [0.5, 0.5, 0.5] #data transformation data_transforms = { 'train': transforms.Compose([ transforms.CenterCrop(input_shape), transforms.ToTensor(), transforms.Normalize(mean, std) ]), 'validation': transforms.Compose([ transforms.CenterCrop(input_shape), transforms.ToTensor(), transforms.Normalize(mean, std) ]), } image_datasets = { x: datasets.ImageFolder( os.path.join(data_dir, x), transform=data_transforms[x] ) for x in ['train', 'validation'] } dataloaders = { x: torch.utils.data.DataLoader( image_datasets[x], batch_size=32, shuffle=True, num_workers=4 ) for x in ['train', 'validation'] } dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'validation']} print(dataset_sizes) class_names = image_datasets['train'].classes device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
Visualisez maintenant l'ensemble de donnรฉes. Cette รฉtape consiste ร utiliser le lot suivant d'images et d'รฉtiquettes provenant du chargeur de donnรฉes d'entraรฎnement et ร les afficher avec Matplotlib.
images, labels = next(iter(dataloaders['train'])) rows = 4 columns = 4 fig=plt.figure() for i in range(16): fig.add_subplot(rows, columns, i+1) plt.title(class_names[labels[i]]) img = images[i].numpy().transpose((1, 2, 0)) img = std * img + mean plt.imshow(img) plt.show()
L'exรฉcution de cet extrait de code dessine une grille de quatre par quatre images d'entraรฎnement, chacune portant le nom de la classe renvoyรฉe par le chargeur.
รtape 2) Dรฉfinir le modรจle
Dans ce processus d'apprentissage profond, vous utiliserez VGG19 du module torchvision.
Vous utiliserez torchvision.models pour charger le modรจle vgg19 avec les poids prรฉ-entraรฎnรฉs activรฉs. Ensuite, vous figerez les couches afin qu'elles ne soient plus entraรฎnables. Vous modifierez ensuite la derniรจre couche par une couche linรฉaire adaptรฉe au problรจme, ici deux classes. La fonction de perte utilisรฉe est CrossEntropyLoss, et l'optimiseur est SGD avec un taux d'apprentissage de 0.001 et un momentum de 0.9, comme illustrรฉ dans le code Python ci-dessous.TorExemple d'apprentissage par transfert.
## Load the model based on VGG19 vgg_based = torchvision.models.vgg19(pretrained=True) ## freeze the layers for param in vgg_based.parameters(): param.requires_grad = False # Modify the last layer number_features = vgg_based.classifier[6].in_features features = list(vgg_based.classifier.children())[:-1] # Remove last layer features.extend([torch.nn.Linear(number_features, len(class_names))]) vgg_based.classifier = torch.nn.Sequential(*features) vgg_based = vgg_based.to(device) print(vgg_based) criterion = torch.nn.CrossEntropyLoss() optimizer_ft = optim.SGD(vgg_based.parameters(), lr=0.001, momentum=0.9)
Note de version : le prรฉ-entraรฎnรฉ=Vrai Cet argument fonctionne toujours, mais il a รฉtรฉ remplacรฉ depuis torchvision 0.13 par poids argument, donc les installations plus rรฉcentes s'attendent ร torchvision.models.vgg19(weights=VGG19_Weights.DEFAULT) et afficher un avertissement de dรฉprรฉciation dans le cas contraire. Les deux formulaires chargent les mรชmes poids ImageNet.
La structure du modรจle de sortie
L'impression du modรจle renvoie le graphe VGG19 complet. Lisez la derniรจre ligne du bloc classificateur pour confirmer que la modification a fonctionnรฉ : il produit dรฉsormais 2 sorties au lieu des 1 000 classes ImageNet.
VGG( (features): Sequential( (0): Conv2d(3, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (1): ReLU(inplace) (2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (3): ReLU(inplace) (4): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False) (5): Conv2d(64, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (6): ReLU(inplace) (7): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (8): ReLU(inplace) (9): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False) (10): Conv2d(128, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (11): ReLU(inplace) (12): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (13): ReLU(inplace) (14): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (15): ReLU(inplace) (16): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (17): ReLU(inplace) (18): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False) (19): Conv2d(256, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (20): ReLU(inplace) (21): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (22): ReLU(inplace) (23): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (24): ReLU(inplace) (25): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (26): ReLU(inplace) (27): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False) (28): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (29): ReLU(inplace) (30): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (31): ReLU(inplace) (32): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (33): ReLU(inplace) (34): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1)) (35): ReLU(inplace) (36): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False) ) (classifier): Sequential( (0): Linear(in_features=25088, out_features=4096, bias=True) (1): ReLU(inplace) (2): Dropout(p=0.5) (3): Linear(in_features=4096, out_features=4096, bias=True) (4): ReLU(inplace) (5): Dropout(p=0.5) (6): Linear(in_features=4096, out_features=2, bias=True) ) )
รtape 3) Former et tester le modรจle
Nous utiliserons certaines des fonctions de ceci PyTorTutoriel ch pour nous aider ร former et ร รฉvaluer notre modรจle.
def train_model(model, criterion, optimizer, num_epochs=25): since = time.time() for epoch in range(num_epochs): print('Epoch {}/{}'.format(epoch, num_epochs - 1)) print('-' * 10) #set model to trainable # model.train() train_loss = 0 # Iterate over data. for i, data in enumerate(dataloaders['train']): inputs , labels = data inputs = inputs.to(device) labels = labels.to(device) optimizer.zero_grad() with torch.set_grad_enabled(True): outputs = model(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step() train_loss += loss.item() * inputs.size(0) print('{} Loss: {:.4f}'.format( 'train', train_loss / dataset_sizes['train'])) time_elapsed = time.time() - since print('Training complete in {:.0f}m {:.0f}s'.format( time_elapsed // 60, time_elapsed % 60)) return model def visualize_model(model, num_images=6): was_training = model.training model.eval() images_so_far = 0 fig = plt.figure() with torch.no_grad(): for i, (inputs, labels) in enumerate(dataloaders['validation']): inputs = inputs.to(device) labels = labels.to(device) outputs = model(inputs) _, preds = torch.max(outputs, 1) for j in range(inputs.size()[0]): images_so_far += 1 ax = plt.subplot(num_images//2, 2, images_so_far) ax.axis('off') ax.set_title('predicted: {} truth: {}'.format(class_names[preds[j]], class_names[labels[j]])) img = inputs.cpu().data[j].numpy().transpose((1, 2, 0)) img = std * img + mean ax.imshow(img) if images_so_far == num_images: model.train(mode=was_training) return model.train(mode=was_training)
Enfin, dans cet apprentissage par transfert en PyTorPar exemple, commencez l'entraรฎnement avec 25 รฉpoques, puis รฉvaluez le rรฉseau. ร chaque รฉtape, le modรจle reรงoit les donnรฉes d'entrรฉe et prรฉdit la sortie. La prรฉdiction est ensuite transmise au critรจre de calcul de la perte, la rรฉtropropagation calcule les gradients et l'optimiseur met ร jour les poids avec l'algorithme autograd.
Dans la fonction de visualisation, le rรฉseau entraรฎnรฉ est testรฉ avec un lot d'images pour prรฉdire les รฉtiquettes, et le rรฉsultat est dessinรฉ avec Matplotlib.
vgg_based = train_model(vgg_based, criterion, optimizer_ft, num_epochs=25) visualize_model(vgg_based) plt.show()
รtape 4) Rรฉsultats
La prรฉcision enregistrรฉe pour cet entraรฎnement est de 92 %. Le journal imprimรฉ ร la fin de l'entraรฎnement affiche la perte cumulรฉe pour les deux derniรจres รฉpoques ainsi que la durรฉe totale de l'entraรฎnement.
Epoch 23/24 ---------- train Loss: 0.0044 train Loss: 0.0078 train Loss: 0.0141 train Loss: 0.0221 train Loss: 0.0306 train Loss: 0.0336 train Loss: 0.0442 train Loss: 0.0482 train Loss: 0.0557 train Loss: 0.0643 train Loss: 0.0763 train Loss: 0.0779 train Loss: 0.0843 train Loss: 0.0910 train Loss: 0.0990 train Loss: 0.1063 train Loss: 0.1133 train Loss: 0.1220 train Loss: 0.1344 train Loss: 0.1382 train Loss: 0.1429 train Loss: 0.1500 Epoch 24/24 ---------- train Loss: 0.0076 train Loss: 0.0115 train Loss: 0.0185 train Loss: 0.0277 train Loss: 0.0345 train Loss: 0.0420 train Loss: 0.0450 train Loss: 0.0490 train Loss: 0.0644 train Loss: 0.0755 train Loss: 0.0813 train Loss: 0.0868 train Loss: 0.0916 train Loss: 0.0980 train Loss: 0.1008 train Loss: 0.1101 train Loss: 0.1176 train Loss: 0.1282 train Loss: 0.1323 train Loss: 0.1397 train Loss: 0.1436 train Loss: 0.1467 Training complete in 2m 47s
Les prรฉdictions du modรจle sont ensuite visualisรฉes avec Matplotlib, comme indiquรฉ ci-dessous.
Erreurs courantes en matiรจre d'apprentissage par transfert et comment les corriger
La plupart des erreurs dans un script d'apprentissage par transfert sont d'ordre mรฉcanique plutรดt que mathรฉmatique. Ce sont celles qui empรชchent l'exรฉcution du code ci-dessus, et voici leur signification.
- Inadรฉquation de taille dans la couche finale : La couche linรฉaire remplacรฉe doit accepter les caractรฉristiques d'entrรฉe fournies par le classificateur d'origine et produire exactement `len(class_names)` sorties. Afficher le modรจle, comme ร l'รฉtape 2, est la mรฉthode la plus rapide pour vรฉrifier ces deux valeurs.
- La colonne vertรฉbrale n'a jamais รฉtรฉ gelรฉe : Si l'option `requires_grad` reste ร `True`, chaque paramรจtre de VGG19 est mis ร jour et l'exรฉcution devient extrรชmement lente sur un processeur. Dรฉfinissez-la sur `False` avant le remplacement du classificateur.
- Incohรฉrence de normalisation : La moyenne et l'รฉcart type utilisรฉs dans transforms.Normalize doivent รชtre identiques lors de l'entraรฎnement et de l'infรฉrence, sinon les prรฉdictions dรฉrivent sans raison apparente.
- Argument de pondรฉration obsolรจte : Sur torchvision 0.13 et versions ultรฉrieures, pretrained=True gรฉnรจre un avertissement de dรฉprรฉciation et l'รฉnumรฉration des poids est prรฉfรฉrรฉe.
- num_workers sur Windows et des cahiers : Un DataLoader avec num_workers=4 nรฉcessite que son point d'entrรฉe soit protรฉgรฉ ; il faut donc dรฉfinir num_workers=0 si Python gรฉnรจre une erreur de gรฉnรฉration ou de sรฉrialisation.
- รvaluation en mode entraรฎnement : Appelez model.eval() avant le scoring afin que Dropout et BatchNorm se comportent de maniรจre dรฉterministe, puis revenez ร model.train() ensuite.



