PyTorch Transfer Learning Handledning med exempel
โก Smart sammanfattning
Transfer Learning รฅteranvรคnder ett nรคtverk som redan รคr trรคnat pรฅ en stor datamรคngd sรฅ att en ny, relaterad uppgift kan lรถsas med betydligt fรคrre mรคrkta bilder och en brรฅkdel av den ursprungliga trรคningstiden.
Vad รคr Transfer Learning?
รverfรถr lรคrande รคr en teknik fรถr att anvรคnda en trรคnad modell fรถr att lรถsa en annan relaterad uppgift. Det รคr en Maskininlรคrning en forskningsmetod som lagrar den kunskap som erhรฅllits vid lรถsning av ett visst problem och anvรคnder samma kunskap fรถr att lรถsa ett annat, men relaterat problem. Detta fรถrbรคttrar effektiviteten genom att รฅteranvรคnda informationen som samlats in frรฅn den tidigare inlรคrda uppgiften.
Det รคr populรคrt att รฅteranvรคnda vikterna frรฅn en annan nรคtverksmodell eftersom det krรคvs en mycket stor mรคngd data fรถr att trรคna ett nรคtverk frรฅn grunden. Fรถr att minska trรคningstiden tar man ett befintligt nรคtverk och dess vikter och modifierar det sista lagret fรถr att lรถsa sitt eget problem. Fรถrdelen รคr att detta sista lager kan trรคnas med en liten datamรคngd.
Innan du skriver nรฅgon PyTorch-kod, รคr det bra att veta vilken familj av Transfer Learning ditt problem tillhรถr, eftersom det avgรถr hur mycket mรคrkt data du behรถver.
Typer av รถverfรถringslรคrande
Forskningslitteraturen delar upp tekniken i tre familjer. Vilken som gรคller fรถr dig beror pรฅ vilken sida av problemet som bรคr etiketter, inte pรฅ vilket ramverk du anvรคnder.
| Typ | Kรคlldomรคn | Target domรคn | Typisk anvรคndning |
|---|---|---|---|
| Induktiv | Mรคrkt | Mรคrkt, men en annan uppgift | Att omdirigera en ImageNet-ryggrad mot ett tvรฅklassigt Alien vs. Predator-problem |
| Transduktiv | Mรคrkt | Omรคrkt, samma uppgift, annan datadistribution | Domรคnanpassning, som att flytta en modell frรฅn studiofotografier till telefonfotografier |
| Oรถvervakad | Omรคrkt | Omรคrkt | Clusterning eller dimensionsreduktion dรคr det รคr opraktiskt att mรคrka varje post |
Exemplet som byggs i den hรคr handledningen รคr induktiv รถverfรถringsinlรคrning. VGG19 anlรคnder med mรคrkt ImageNet-kunskap, och den riktar sig sedan mot ett mรคrkt tvรฅklassproblem som den aldrig sett fรถrut.
Funktion Extraction kontra finjustering
Efter att ett fรถrtrรคnat nรคtverk har valts finns det tvรฅ sรคtt att anpassa det. Skillnaden รคr helt enkelt hur mรฅnga lager du tillรฅter att fortsรคtta lรคra sig.
| Aspect | Funktion extraction | Finjustering |
|---|---|---|
| Lager som trรคnar | Endast den ersatta klassificeraren | Klassificeraren plus nรฅgra eller alla faltningsblock |
| requires_grad pรฅ ryggraden | Falsk | Sant fรถr blocken som uppdateras |
| Behรถvliga uppgifter | Smรฅ, ofta nรฅgra hundra bilder per klass | Stรถrre, vanligtvis tusentals |
| Utbildningskostnad | Lรคgst, kรถrs pรฅ en CPU | Ju hรถgre, desto mer vรคrt blir det att ha en GPU |
| Typisk noggrannhet | Bra nรคr kรคll- och mรฅlbilder ser likadana ut | Vanligtvis bรคttre nรคr de tvรฅ domรคnerna skiljer sig รฅt |
Stegen nedan anvรคnder funktionen extraction: varje VGG19-parameter fryses och endast det nya slutliga linjรคra lagret lรคr sig. Att byta till finjustering รคr en liten fรถrรคndring, nรคmligen att lรคmna requires_grad satt till True pรฅ de block du vill uppdatera och sรคnka inlรคrningshastigheten sรฅ att de lรฅnade vikterna inte fรถrstรถrs.
Laddar datauppsรคttning
Innan du bรถrjar anvรคnda Transfer Learning med PyTorch, du behรถver fรถrstรฅ den datauppsรคttning du ska anvรคnda. I denna Transfer Learning PyTorTill exempel kommer du att klassificera en utomjording och ett rovdjur frรฅn nรคstan 700 bilder. Fรถr den hรคr tekniken behรถver du egentligen inte en stor mรคngd data fรถr att trรคna. Du kan ladda ner datasetet frรฅn Kaggle: Alien vs. Predator.
Samlingen รคr avsiktligt liten, och ett urval av bilderna den innehรฅller visas nedan.
Kรคlla: Frรคmmande vs. Predator Kaggle
Nรคsta i denna PyTorch Transfer Learning-handledningen, lรคr du dig hur du tillรคmpar Transfer Learning med PyTorch steg fรถr steg.
Hur anvรคnder man Transfer Learning?
Hรคr รคr en steg-fรถr-steg-process fรถr hur man anvรคnder Transfer Learning fรถr djupinlรคrning med PyTorch:
Steg 1) Ladda data
Det fรถrsta steget รคr att ladda data och tillรคmpa nรฅgra transformationer pรฅ bilderna sรฅ att de matchar nรคtverkets krav.
Du kommer att ladda data frรฅn en mapp med torchvision.datasets. Modulen itererar รถver mappen fรถr att dela upp data i tรฅg- och valideringsuppsรคttningar. Transformationspipelinen som anvรคnds hรคr beskรคr bilderna frรฅn mitten, konverterar dem till en tensor och normaliserar dem fรถr 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")
Visualisera nu datamรคngden. Visualiseringssteget tar nรคsta omgรฅng bilder och etiketter frรฅn trรคningsdataladdaren och visar 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()
Genom att kรถra det dรคr snippet ritar man ett fyra gรฅnger fyra rutnรคt med trรคningsbilder, dรคr var och en har det klassnamn som laddaren returnerade.
Steg 2) Definiera modell
I den hรคr djupinlรคrningsprocessen kommer du att anvรคnda VGG19 frรฅn torchvision-modulen.
Du kommer att anvรคnda torchvision.models fรถr att ladda vgg19 med de fรถrtrรคnade vikterna aktiverade. Dรคrefter fryser du lagren sรฅ att de inte รคr trรคningsbara. Du modifierar sedan det sista lagret med ett linjรคrt lager som passar problemet, vilket hรคr betyder 2 klasser. CrossEntropyLoss anvรคnds som fรถrlustfunktion, och optimeraren รคr SGD med en inlรคrningshastighet pรฅ 0.001 och ett momentum pรฅ 0.9, som visas i Py nedan.Torch Exempel pรฅ รถverfรถring av lรคrande.
## 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)
Versionsnotering: d fรถrtrรคnad=Sant argumentet fungerar fortfarande men har ersatts sedan torchvision 0.13 av vikter argument, sรฅ nyare installationer fรถrvรคntar sig torchvision.models.vgg19(vikter=VGG19_Vikter.STANDARD) och skriv ut en varning om utfasning annars. Bรฅda formulรคren laddar samma ImageNet-vikter.
Utgรฅngsmodellens struktur
Att skriva ut modellen returnerar hela VGG19-grafen. Lรคs den sista raden i klassificeringsblocket fรถr att bekrรคfta att vรคxlingen fungerade: den producerar nu 2 utdata istรคllet fรถr de 1 000 ImageNet-klasserna.
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) ) )
Steg 3) Trรคna och testa modell
Vi kommer att anvรคnda nรฅgra av funktionerna frรฅn detta PyTorch-handledning fรถr att hjรคlpa oss att trรคna och utvรคrdera vรฅr modell.
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)
Slutligen i denna รถverfรถringslรคrande i PyTorTill exempel, starta trรคningsprocessen med antalet epoker satt till 25 och utvรคrdera nรคtverket efterรฅt. Vid varje trรคningssteg tar modellen indata och fรถrutspรฅr utdata. Fรถrutsรคgelsen skickas till kriteriet fรถr att berรคkna fรถrlusten, backpropagation berรคknar gradienterna och optimeraren uppdaterar vikterna med autograd.
I visualiseringsfunktionen testas det trรคnade nรคtverket med en sats bilder fรถr att fรถrutsรคga etiketterna, och resultatet ritas med Matplotlib.
vgg_based = train_model(vgg_based, criterion, optimizer_ft, num_epochs=25) visualize_model(vgg_based) plt.show()
Steg 4) Resultat
Noggrannheten som rapporterats fรถr denna kรถrning รคr 92 %. Loggen som skrivs ut i slutet av trรคningen visar lรถpfรถrlusten fรถr de tvรฅ senaste epokerna tillsammans med den totala trรคningstiden.
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 fรถrutsรคgelser visualiseras sedan med Matplotlib, som visas nedan.
Vanliga fel vid รถverfรถringsinlรคrning och hur man รฅtgรคrdar dem
De flesta fel i ett รถverfรถringsinlรคrningsskript รคr mekaniska snarare รคn matematiska. Det รคr dessa som hindrar koden ovan frรฅn att kรถras, och vad var och en betyder.
- Storleksavvikelse i det sista lagret: Det ersatta linjรคra lagret mรฅste acceptera de in_features som rapporteras av den ursprungliga klassificeraren och ge exakta len(class_names)-utdata. Att skriva ut modellen, som i steg 2, รคr det snabbaste sรคttet att bekrรคfta bรฅda siffrorna.
- Ryggraden har aldrig varit frusen: Om requires_grad lรคmnas pรฅ True uppdateras varje VGG19-parameter och kรถrningen saktar ner till en krypning pรฅ en CPU. Sรคtt den pรฅ False innan klassificeraren ersรคtts.
- Normaliseringsfel: Medelvรคrdet och standardavvikelsen som anvรคnds i transformationerna. Normalisering mรฅste vara samma vรคrden vid trรคnings- och inferenstid, annars avviker fรถrutsรคgelserna utan nรฅgon synlig anledning.
- Fรถrรฅldrat viktningsargument: Pรฅ torchvision 0.13 och senare genererar pretrained=True en varning om utfasning och weights enum รคr att fรถredra.
- antal_arbetare pรฅ Windows och anteckningsbรถcker: En DataLoader med num_workers=4 behรถver att startpunkten skyddas, sรฅ sรคtt num_workers=0 om Python genererar ett spawn- eller pickling-fel.
- Utvรคrdering i trรคningslรคge: anropa model.eval() innan poรคngsรคttning sรฅ att Dropout och BatchNorm beter sig deterministiskt, och vรคxla tillbaka med model.train() efterรฅt.



