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
    Zatim ćemo uporediti dve varijante:
    • 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)
    
    No description has been provided for this image
    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.
    U obe varijante koristimo isti pretreniran ResNet18. Razlika je samo u tome koje slojeve ostavljamo zamrznute, a koje dozvoljavamo da se dodatno treniraju.
    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()
    
    No description has been provided for this image

    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].
    
    No description has been provided for this image