Pretraining i fine-tuning¶
Do sada smo se bavili modelima koje smo trenirali praktično od početka. Ipak, u savremenoj praksi vrlo često ne krećemo od nule. Umesto toga, polazimo od modela koji je već treniran na velikom skupu podataka, pa ga zatim prilagođavamo našem konkretnom zadatku.
Takav pristup naziva se pretraining i fine-tuning. Najpre se model trenira na velikom i opštem skupu podataka, kako bi naučio korisne reprezentacije. Nakon toga se isti model, ili njegov najveći deo, koristi kao polazna tačka za novi zadatak. Ovo je naročito važno kada nemamo ogroman sopstveni skup podataka, ali želimo dobar rezultat i razumnu brzinu učenja.
Šta znači da je model pretreniran?¶
Kažemo da je model pretreniran ako je već prošao proces učenja na velikom skupu podataka. Na primer, u obradi slika često se koriste modeli trenirani na skupu ImageNet, koji sadrži veoma veliki broj raznovrsnih slika. Takav model je već naučio da prepoznaje osnovne vizuelne obrasce: ivice, teksture, delove objekata, pa čak i složenije strukture.
Intuicija je sledeća: ako je model već naučio opšte vizuelne pravilnosti, onda mu za novi zadatak više nije potrebno da sve uči iz početka. Dovoljno je da postojeće znanje prilagodimo novoj situaciji.
Dve osnovne strategije¶
U praksi se veoma često javljaju dve osnovne strategije:
Feature extraction: najveći deo modela ostaje "zamrznut", a trenira se samo završni klasifikacioni sloj.
Fine-tuning: pored završnog sloja, otključavamo i deo dubljih slojeva modela kako bi se reprezentacije prilagodile novom zadatku.
Prvi pristup je jednostavniji i brži. Drugi pristup je fleksibilniji i često daje bolje rezultate, ali zahteva više računanja i više pažnje pri izboru hiperparametara.
Ako označimo sa $ \theta_{\mathrm{base}} $ parametre osnovnog dela mreže, a sa $ \theta_{\mathrm{head}} $ parametre završnog klasifikatora, onda važi sledeće.
Kod feature extraction pristupa fiksiramo parametre $ \theta_{\mathrm{base}} $ i optimizujemo samo $ \theta_{\mathrm{head}} $.
Kod fine-tuning pristupa optimizujemo i $ \theta_{\mathrm{head}} $ i deo parametara iz $ \theta_{\mathrm{base}} $.
Drugim rečima, razlika nije u arhitekturi, već u tome kojim parametrima dozvoljavamo da se menjaju tokom treniranja.
Kada koji pristup ima smisla?¶
Ako je novi skup podataka mali i sličan onome na kome je model originalno treniran, često je dovoljno koristiti feature extraction. Ako imamo više podataka ili je novi zadatak dovoljno drugačiji od početnog, onda je korisno uraditi fine-tuning, odnosno dozvoliti modelu da prilagodi i dublje slojeve.
Na primer, model treniran na opštim prirodnim slikama može biti dobra polazna tačka za različite zadatke iz computer vision-a. Međutim, ako pređemo na medicinske snimke ili satelitske slike, onda će često biti potrebno jače prilagođavanje modela.
Primer¶
Koristićemo pretrenirani ResNet18 model iz biblioteke
torchvision. Posmatraćemo jednostavan problem binarne klasifikacije na podskupu baze CIFAR-10. Radi preglednosti, zadržaćemo samo dve klase:
- airplane
- automobile
- feature extraction — treniramo samo završni sloj;
- fine-tuning — otključavamo i poslednji rezidualni blok modela.
In [ ]:
import numpy as np
import matplotlib.pyplot as plt
import torch
from torch import nn
from torch.utils.data import DataLoader, Subset
from torchvision import datasets, transforms, models
from torchvision.models import ResNet18_Weights
Učitavanje i priprema podataka¶
ResNet18 očekuje slike određene veličine i određenu normalizaciju. Zato koristimo transformacije usklađene sa pretreniranım težinama. Ovo je važan detalj: kada koristimo pretreniran model, ne menjamo samo arhitekturu, već često preuzimamo i način pretprocesiranja podataka.
In [ ]:
weights = ResNet18_Weights.DEFAULT
preprocess = weights.transforms()
In [ ]:
raw_transform = transforms.ToTensor()
train_raw = datasets.CIFAR10(root="data", train=True, download=True, transform=raw_transform)
test_raw = datasets.CIFAR10(root="data", train=False, download=True, transform=raw_transform)
train_full = datasets.CIFAR10(root="data", train=True, download=False, transform=preprocess)
test_full = datasets.CIFAR10(root="data", train=False, download=False, transform=preprocess)
selected_classes = [0, 1] # airplane, automobile
class_names = ["airplane", "automobile"]
In [ ]:
def collect_indices(dataset, allowed_labels, max_per_class):
counts = {label: 0 for label in allowed_labels}
indices = []
for idx, (_, label) in enumerate(dataset):
if label in allowed_labels and counts[label] < max_per_class:
indices.append(idx)
counts[label] += 1
if all(counts[label] >= max_per_class for label in allowed_labels):
break
return indices
train_indices = collect_indices(train_raw, selected_classes, max_per_class=1500)
test_indices = collect_indices(test_raw, selected_classes, max_per_class=300)
train_dataset_raw = Subset(train_raw, train_indices)
test_dataset_raw = Subset(test_raw, test_indices)
train_dataset = Subset(train_full, train_indices)
test_dataset = Subset(test_full, test_indices)
print("Broj trening instanci:", len(train_dataset))
print("Broj test instanci:", len(test_dataset))
Broj trening instanci: 3000 Broj test instanci: 600
In [ ]:
def show_examples(dataset, class_names, n=8):
plt.figure(figsize=(14, 3))
for i in range(n):
image, label = dataset[i]
plt.subplot(1, n, i + 1)
plt.imshow(np.transpose(image.numpy(), (1, 2, 0)))
plt.title(class_names[label])
plt.axis("off")
plt.tight_layout()
plt.show()
show_examples(train_dataset_raw, class_names)
In [ ]:
batch_size = 64
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
Kreiranje modela¶
Napravimo pomoćne funkcije za dve varijante modela:
- model za feature extraction,
- model za fine-tuning.
In [ ]:
def create_feature_extractor():
model = models.resnet18(weights=weights)
for param in model.parameters():
param.requires_grad = False
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 2)
return model
def create_finetuning_model():
model = models.resnet18(weights=weights)
for param in model.parameters():
param.requires_grad = False
for param in model.layer4.parameters():
param.requires_grad = True
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 2)
return model
In [ ]:
device = 'cpu'
feature_model = create_feature_extractor().to(device)
finetune_model = create_finetuning_model().to(device)
feature_trainable = sum(p.numel() for p in feature_model.parameters() if p.requires_grad)
finetune_trainable = sum(p.numel() for p in finetune_model.parameters() if p.requires_grad)
print("Broj trenirajućih parametara (feature extraction):", feature_trainable)
print("Broj trenirajućih parametara (fine-tuning):", finetune_trainable)
Broj trenirajućih parametara (feature extraction): 1026 Broj trenirajućih parametara (fine-tuning): 8394754
Pomoćne funkcije za treniranje i evaluaciju¶
In [ ]:
def run_epoch(model, loader, loss_fn, optimizer=None):
if optimizer is None:
model.eval()
else:
model.train()
total_loss = 0.0
correct = 0
total = 0
for X, y in loader:
X, y = X.to(device), y.to(device)
if optimizer is not None:
optimizer.zero_grad()
logits = model(X)
loss = loss_fn(logits, y)
if optimizer is not None:
loss.backward()
optimizer.step()
total_loss += loss.item() * X.size(0)
predictions = logits.argmax(dim=1)
correct += (predictions == y).sum().item()
total += y.size(0)
return total_loss / total, correct / total
def train_model(model, train_loader, test_loader, epochs=3, lr=1e-3):
optimizer = torch.optim.Adam([p for p in model.parameters() if p.requires_grad], lr=lr)
loss_fn = nn.CrossEntropyLoss()
history = {
"train_loss": [],
"train_acc": [],
"test_loss": [],
"test_acc": []}
for epoch in range(epochs):
train_loss, train_acc = run_epoch(model, train_loader, loss_fn, optimizer)
test_loss, test_acc = run_epoch(model, test_loader, loss_fn)
history["train_loss"].append(train_loss)
history["train_acc"].append(train_acc)
history["test_loss"].append(test_loss)
history["test_acc"].append(test_acc)
print(
f"Epohа {epoch + 1}/{epochs} | "
f"train loss = {train_loss:.4f}, train acc = {train_acc:.4f} | "
f"test loss = {test_loss:.4f}, test acc = {test_acc:.4f}")
return history
Varijanta 1: feature extraction¶
Konvolutivni deo mreže ostaje zamrznut, a trenira se samo završni linearni sloj koji odlučuje između dve klase.
In [ ]:
history_feature = train_model(feature_model, train_loader,test_loader, epochs=3, lr=1e-3)
# vreme treniranja: oko 6 minuta
# vreme treniranja raste linearno sa povecanjem broja epoha
Epohа 1/3 | train loss = 0.5038, train acc = 0.7630 | test loss = 0.3185, test acc = 0.8917 Epohа 2/3 | train loss = 0.2521, train acc = 0.9260 | test loss = 0.2277, test acc = 0.9267 Epohа 3/3 | train loss = 0.2047, train acc = 0.9377 | test loss = 0.1966, test acc = 0.9317
Varijanta 2: fine-tuning¶
Sada ćemo dozvoliti modelu da prilagodi i deo dubljih slojeva, tačnije poslednji rezidualni blok i završni klasifikacioni sloj. Na taj način model ne koristi samo unapred naučene osobine, već i menja deo svojih reprezentacija kako bi se bolje prilagodio novom zadatku.
In [ ]:
history_finetune = train_model(finetune_model, train_loader, test_loader, epochs=3, lr=1e-4)
# vreme treniranja: oko 8 minuta
# u ovom pristupu trening traje duže, ali dobijamo bolje rezultate, kao što smo i očekivali
Epohа 1/3 | train loss = 0.1477, train acc = 0.9433 | test loss = 0.0597, test acc = 0.9783 Epohа 2/3 | train loss = 0.0185, train acc = 0.9967 | test loss = 0.0503, test acc = 0.9867 Epohа 3/3 | train loss = 0.0055, train acc = 1.0000 | test loss = 0.0539, test acc = 0.9850
Grafik funkcije gubitaka i tačnosti¶
In [ ]:
epochs_range = np.arange(1, len(history_feature["train_loss"]) + 1)
plt.figure(figsize=(12, 4))
plt.subplot(1, 2, 1)
plt.title("Funkcija gubitaka")
plt.plot(epochs_range, history_feature["train_loss"], marker="o", label="Feature extraction - train")
plt.plot(epochs_range, history_feature["test_loss"], marker="o", label="Feature extraction - test")
plt.plot(epochs_range, history_finetune["train_loss"], marker="s", label="Fine-tuning - train")
plt.plot(epochs_range, history_finetune["test_loss"], marker="s", label="Fine-tuning - test")
plt.xlabel("Epoha")
plt.ylabel("Loss")
plt.legend()
plt.subplot(1, 2, 2)
plt.title("Tačnost")
plt.plot(epochs_range, history_feature["train_acc"], marker="o", label="Feature extraction - train")
plt.plot(epochs_range, history_feature["test_acc"], marker="o", label="Feature extraction - test")
plt.plot(epochs_range, history_finetune["train_acc"], marker="s", label="Fine-tuning - train")
plt.plot(epochs_range, history_finetune["test_acc"], marker="s", label="Fine-tuning - test")
plt.xlabel("Epoha")
plt.ylabel("Accuracy")
plt.legend()
plt.tight_layout()
plt.show()
Završno poređenje¶
Ako je novi zadatak dovoljno sličan originalnom problemu i ako nemamo mnogo podataka, feature extraction često daje veoma pristojan rezultat uz malo trenirajućih parametara. Sa druge strane, fine-tuning dozvoljava modelu da se dublje prilagodi novom zadatku, pa često daje bolje rezultate, ali je računarski skuplji i osetljiviji na izbor hiperparametara.
In [ ]:
print("Najbolja test tačnost (feature extraction):", max(history_feature["test_acc"]))
print("Najbolja test tačnost (fine-tuning):", max(history_finetune["test_acc"]))
Najbolja test tačnost (feature extraction): 0.9316666666666666 Najbolja test tačnost (fine-tuning): 0.9866666666666667
Važno je uočiti da u oba slučaja koristimo istu osnovnu arhitekturu i iste pretrenirane težine. Razlika je samo u tome koliko "slobode" dajemo modelu tokom dodatnog treniranja.
Pogledajmo nekoliko predikcija¶
In [ ]:
plt.figure(figsize=(14, 3))
for i in range(5):
img, label = test_dataset[i]
pred = finetune_model(img.unsqueeze(0).to(device)).argmax(1).item()
plt.subplot(1, 5, i + 1)
plt.imshow(img.permute(1, 2, 0))
plt.title(f"T:{class_names[label]}\nP:{class_names[pred]}")
plt.axis("off")
plt.tight_layout()
plt.show()
Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-1.4500387..2.2565577]. Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-2.1007793..2.622571]. Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-1.7069099..2.5179958]. Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-1.6555357..1.907974]. Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-2.117904..2.5179958].