Skip to content
PyTorch

Transfer Learning

Fine-tune a pretrained torchvision model.

#transfer-learning#pretrained

Code

pytorch
import torch
import torch.nn as nn
from torchvision import models, transforms
from torch.utils.data import DataLoader

# Load pretrained ResNet
model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)

# Freeze backbone parameters
for param in model.parameters():
    param.requires_grad = False

# Replace the classifier head
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 10)

# Train only the head
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()

# Data augmentation for images
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225]),
])

# Later, unfreeze and fine-tune at a lower LR
# for param in model.parameters():
#     param.requires_grad = True
# optimizer = torch.optim.Adam(model.parameters(), lr=1e-5)