PyTorch Transfer Learning Tutorial med eksempler

⚡ Smart opsummering

Transfer Learning genbruger et netværk, der allerede er trænet på et stort datasæt, så en ny, relateret opgave kan løses med langt færre mærkede billeder og en brøkdel af den oprindelige træningstid.

  • 🔘 Kerneidé: Vægte lært på ImageNet koder allerede kanter, teksturer og former, som de fleste visionsopgaver genbruger.
  • ☑️ To strategier: Funktion extraction fryser rygraden, mens finjustering fortsætter med at træne nogle eller alle de oprindelige lag.
  • Udarbejdet eksempel: Omkring 700 Alien- og Predator-billeder er nok til at genoptræne det sidste lag af VGG19.
  • 🧪 PyTorlm-stykker: ImageFolder, transforms.Compose og DataLoader samler de batches, som netværket forbruger.
  • 🛠️ Frysende lag: Hvis requires_grad sættes til Falsk, stoppes gradienter, så kun den erstattede klassifikator lærer.
  • ⚠️ Rapporteret resultat: Femogtyve epoker slutter på under tre minutter på eksempeldatasættet.

PyTorch Transfer Learning-vejledning med bearbejdede eksempler

Hvad er Transfer Learning?

Overfør læring er en teknik til at bruge en trænet model til at løse en anden relateret opgave. Det er en Maskinelæring en forskningsmetode, der lagrer den viden, der er opnået under løsningen af ​​et bestemt problem, og bruger den samme viden til at løse et andet, men relateret problem. Dette forbedrer effektiviteten ved at genbruge information indsamlet fra den tidligere lærte opgave.

Det er populært at genbruge vægtene fra en anden netværksmodel, fordi træning af et netværk fra bunden kræver en meget stor mængde data. For at reducere træningstiden tager man et eksisterende netværk og dets vægte og ændrer det sidste lag for at løse sit eget problem. Fordelen er, at dette sidste lag kan trænes med et lille datasæt.

Før du skriver nogen PyTorch-kode, er det nyttigt at vide, hvilken familie af Transfer Learning dit problem tilhører, fordi det afgør, hvor meget mærket data du har brug for.

Typer af transferlæring

Forskningslitteraturen opdeler teknikken i tre familier. Hvilken der gælder for dig, afhænger af hvilken side af problemet der bærer betegnelser, ikke af den ramme du bruger.

Type Kildedomæne Target domæne Typisk brug
Induktiv Mærket Mærket, men en anden opgave Omdirigerer en ImageNet-rygrad til et to-klasses Alien vs. Predator-problem
Transduktiv Mærket Umærket, samme opgave, forskellig datafordeling Domænetilpasning, såsom at flytte en model fra studiefotografier til telefonfotografier
Uden opsyn Umærket Umærket Clusterning eller dimensionsreduktion, hvor det er upraktisk at mærke hver enkelt post

Eksemplet i denne vejledning er induktiv transferlæring. VGG19 ankommer med mærket ImageNet-viden, og den er derefter rettet mod et mærket to-klasse problem, den aldrig har set før.

Funktion Eks.traction vs. finjustering

Når et præ-trænet netværk er valgt, er der to måder at tilpasse det på. Forskellen er simpelthen, hvor mange lag du tillader at fortsætte med at lære.

Aspect Funktion extraction Finjustering
Lag der træner Kun den erstattede klassifikator Klassifikatoren plus nogle eller alle konvolutionelle blokke
requires_grad på rygraden False Sandt for de blokke, der opdateres
Nødvendige data Små, ofte et par hundrede billeder pr. klasse Større, normalt tusinder
Udgifter til uddannelse Laveste, kører på en CPU Jo højere, desto mere værd bliver en GPU
Typisk nøjagtighed Godt når kilde- og målbilleder ligner hinanden Normalt bedre, når de to domæner er forskellige

Trinene nedenfor bruger funktionen f.eks.traction: hver VGG19-parameter fryses, og kun det nye, endelige lineære lag lærer. Skift til finjustering er en lille ændring, nemlig at lade requires_grad være sat til True på de blokke, du vil opdatere, og at sænke læringshastigheden, så de lånte vægte ikke ødelægges.

Indlæser datasæt

Før du begynder at bruge Transfer Learning med PyTorch, skal du forstå det datasæt, du skal bruge. I denne Transfer Learning PyTorFor eksempel skal du klassificere en Alien og en Predator ud fra næsten 700 billeder. Til denne teknik behøver du ikke rigtig en stor mængde data at træne. Du kan downloade datasættet fra Kaggle: Alien vs. Predator.

Samlingen er bevidst lille, og et udpluk af de billeder, den indeholder, vises nedenfor.

Alien vs Predator-billeddatasæt fra Kaggle brugt i denne PyToreksempel på læring med ch-overførsel

Kilde: Alien vs. Predator Kaggle

Næste i denne PyTorch Transfer Learning-tutorial, lærer du, hvordan du anvender Transfer Learning med PyTorch trin for trin.

Hvordan bruger man Transfer Learning?

Her er en trin-for-trin proces til, hvordan man bruger Transfer Learning til Deep Learning med PyTorch:

Trin 1) Indlæs dataene

Det første trin er at indlæse dataene og anvende nogle transformationer på billederne, så de matcher netværkets krav.

Du skal indlæse dataene fra en mappe med torchvision.datasets. Modulet itererer over mappen for at opdele dataene i tog- og valideringssæt. Den transformationspipeline, der bruges her, beskærer billederne fra midten, konverterer dem til en tensor og normaliserer dem for Deep Learning.

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

Visualiser nu datasættet. Visualiseringstrinnet tager den næste batch af billeder og etiketter fra træningsdataindlæseren og viser dem med 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()

Når du kører dette kodestykke, tegner du et fire gange fire-gitter af træningsbilleder, der hver især har titlen på det klassenavn, som indlæseren returnerede.

Gruppe af seksten træningsbilleder tegnet i et fire gange fire Matplotlib-gitter med klassetitler

Trin 2) Definer model

I denne Deep Learning-proces vil du bruge VGG19 fra Torchvision-modulet.

Du skal bruge torchvision.models til at indlæse vgg19 med de forudtrænede vægte aktiveret. Derefter fryser du lagene, så de ikke kan trænes. Derefter ændrer du det sidste lag med et lineært lag, der passer til problemet, hvilket her betyder 2 klasser. CrossEntropyLoss bruges som tabsfunktion, og optimeringsfunktionen er SGD med en læringsrate på 0.001 og et momentum på 0.9, som vist i Py-diagrammet nedenfor.Torch Eksempel på transferlæring.

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

Versionsnotat: og prætrænet=Sand argumentet virker stadig, men er blevet erstattet siden torchvision 0.13 af vægte argument, så nyere installationer forventer torchvision.models.vgg19(vægte=VGG19_Vægte.STANDARD) og udskriv en advarsel om udfasning ellers. Begge formularer indlæser de samme ImageNet-vægte.

Outputmodellens struktur

Udskrivning af modellen returnerer den fulde VGG19-graf. Læs den sidste linje i klassifikatorblokken for at bekræfte, at ombytningen fungerede: den producerer nu 2 output i stedet for de 1,000 ImageNet-klasser.

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

Trin 3) Træn og test model

Vi vil bruge nogle af funktionerne fra dette PyTorch-vejledning at hjælpe os med at træne og evaluere vores model.

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)

Endelig i denne Transfer Learning i PyTorFor eksempel starter træningsprocessen med antallet af epoker sat til 25 og evaluerer netværket bagefter. Ved hvert træningstrin tager modellen input og forudsiger outputtet. Forudsigelsen sendes til kriteriet for at beregne tabet, backpropagation beregner gradienterne, og optimeringsværktøjet opdaterer vægtene med autograd.

I visualiseringsfunktionen testes det trænede netværk med en batch af billeder for at forudsige etiketterne, og resultatet tegnes med Matplotlib.

vgg_based = train_model(vgg_based, criterion, optimizer_ft, num_epochs=25)

visualize_model(vgg_based)

plt.show()

Trin 4) Resultater

Den rapporterede nøjagtighed for denne kørsel er 92 %. Loggen, der udskrives ved træningens afslutning, viser løbetabet for de sidste to epoker sammen med den samlede træningstid.

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

Modellens forudsigelser visualiseres derefter med Matplotlib, som vist nedenfor.

Valideringsbilleder mærket med den forudsagte klasse og den sande klasse efter træning

Almindelige fejl ved transferlæring og hvordan man retter dem

De fleste fejl i et transfer learning script er mekaniske snarere end matematiske. Det er disse, der forhindrer ovenstående kode i at køre, og hvad hver enkelt betyder.

  • Størrelsesuoverensstemmelse i det sidste lag: Det erstattede lineære lag skal acceptere de in_features, der rapporteres af den oprindelige klassifikator, og udsende præcise len(class_names) output. Udskrivning af modellen, som i trin 2, er den hurtigste måde at bekræfte begge tal på.
  • Rygraden var aldrig frossen: Hvis requires_grad forbliver på True, opdateres alle VGG19-parametre, og kørslen bliver langsommere på en CPU. Sæt den til False, før klassifikatoren erstattes.
  • Normaliseringsmismatch: Middelværdien og standardafvigelsen, der bruges i transformationerne. Normalisering skal være de samme værdier ved trænings- og inferenstidspunktet, ellers afviger forudsigelserne uden nogen synlig grund.
  • Forældet vægtargument: På torchvision 0.13 og nyere udløser pretrained=True en advarsel om udfasning, og weights enum foretrækkes.
  • antal_arbejdere på Windows og notesbøger: En DataLoader med num_workers=4 kræver, at indgangspunktet er beskyttet, så sæt num_workers=0 hvis Python rejser en spawn- eller pickling-fejl.
  • Evaluering i træningstilstand: Kald model.eval() før scoring, så Dropout og BatchNorm opfører sig deterministisk, og skift tilbage med model.train() bagefter.

Ofte Stillede Spørgsmål

Et par hundrede billeder pr. klasse er normalt nok, når kun det sidste lag trænes; dette eksempel bruger cirka 700 billeder i alt. Finjustering af dybere blokke kræver flere - ofte flere tusinde - fordi langt flere parametre opdateres.

ResNet-, VGG-, EfficientNet- og Vision Transformer-backbones leveres alle med torchvision-vægte. ResNet50 er den almindelige standard, fordi den balancerer nøjagtighed mod størrelse. VGG19, der bruges her, er en tungere konvolutionelt netværk men det er nemt at dissekere lag for lag.

Nej. At fryse backbone-systemet og træne et enkelt lag kører acceptabelt på en CPU, hvilket er grunden til, at ovenstående kørsel afsluttes på under tre minutter. Finjustering af et helt netværk eller træning på tusindvis af billeder er, hvor en GPU ikke længere er valgfri.

Automatiseret modelsøgning benchmarker adskillige præ-trænede netværk på en stikprøve af dine data og rangerer dem efter nøjagtighed, latenstid og størrelse. Det fjerner gætteri fra backbone-valg og passer godt sammen med hyperparametersøgningerne i scikit-lære.

Den udarbejder standardteksten godt, fordi transformationer, DataLoader-opsætning og træningsløkker følger velkendte mønstre. Tjek de dele, den ikke kan udlede: antallet af outputklasser, normaliseringsstatistikken, og om requires_grad faktisk var slået fra.

Negativ overførsel sker, når kildedomænet er for forskelligt fra målet, så de lånte vægte gør snarere end hjælper. Vær opmærksom på valideringstab, der stagnerer tidligt, eller som sidder over et lille netværk, der er trænet fra bunden.

Ja. Sprogmodeller er præ-trænede på store tekstkorpora og derefter tilpasset til klassificering eller besvarelse af spørgsmål, og den samme idé gælder også sekvensmodeller, lyd og tabelindlejringer. Kun rygraden ændres.

Langt færre end et netværk, der er trænet fra bunden. Der bruges 25 epoker her, men ti til tyve er ofte nok til funktionsf.eks.tracStop når valideringstabet holder op med at forbedres, i stedet for at køre en fast optælling.

Opsummer dette indlæg med: