How to Use Transfer Learning with Pre-trained Models in PyTorch: A Microsoft AI for Beginners Guide
Transfer learning with pre-trained models in PyTorch allows you to adapt ImageNet-trained networks to new classification tasks by freezing convolutional feature extractors and replacing the final classifier head, significantly reducing training time and data requirements.
The microsoft/AI-For-Beginners curriculum demonstrates this technique in the 08-TransferLearning lesson, located at lessons/4-ComputerVision/08-TransferLearning/. By leveraging existing weights from models like ResNet-18, beginners can achieve high accuracy on small datasets such as Cats vs. Dogs without training deep networks from scratch.
Loading Pre-trained Models from torchvision
The first step involves importing a neural network that has already learned general visual features from millions of images. PyTorch provides these architectures through torchvision.models.
According to the source code in lessons/4-ComputerVision/08-TransferLearning/pytorchcv.py, you can load a pre-trained model by setting pretrained=True (legacy) or weights='DEFAULT' (recommended in newer torchvision releases):
import torchvision
# Load ResNet-18 with ImageNet weights
model = torchvision.models.resnet18(weights='DEFAULT')
# Legacy syntax also supported:
# model = torchvision.models.resnet18(pretrained=True)
This initializes the network with weights trained on the 1000-class ImageNet dataset, capturing universal patterns like edges, textures, and shapes in the early convolutional layers.
Freezing the Feature Extractor
The convolutional base of a pre-trained model acts as a feature extractor that recognizes generic visual patterns. To preserve these learned representations and prevent overfitting on your smaller dataset, freeze these parameters by disabling gradient computation:
for param in model.parameters():
param.requires_grad = False
Setting requires_grad = False ensures that only the parameters in layers you explicitly modify will update during backpropagation. This drastically reduces the computational cost and memory requirements of training.
Replacing the Classifier Head
The final fully-connected layer (model.fc in ResNet architectures) is specific to the original task (1000 ImageNet classes). You must substitute this classifier head with a new layer matching your target dataset's number of classes:
import torch.nn as nn
num_target_classes = 2 # Example: Cats vs. Dogs
model.fc = nn.Linear(model.fc.in_features, num_target_classes)
This new layer initializes with random weights and is the only portion of the network that will learn from your specific data during the initial training phase.
Training with Frozen Parameters
When configuring the optimizer, pass only the unfrozen parameters—in this case, the new classifier head. The helper functions train, train_epoch, and validate in pytorchcv.py demonstrate the standard PyTorch loop:
from torch import optim, nn
import torch
# Only optimize the new classifier parameters
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.fc.parameters(), lr=1e-3)
# Training loop pattern (as implemented in the curriculum)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = model.to(device)
for epoch in range(5):
model.train()
for imgs, labels in train_loader:
imgs, labels = imgs.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(imgs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
Use the load_cats_dogs_dataset() function from pytorchcv.py to obtain the pre-configured train_loader and val_loader with appropriate preprocessing transforms.
Optional Fine-Tuning Strategy
After the new head converges, you can improve performance through fine-tuning. Unfreeze selected deeper convolutional layers and continue training with a reduced learning rate:
# Unfreeze layer4 and fc for fine-tuning
for param in model.layer4.parameters():
param.requires_grad = True
# Use a lower learning rate to avoid destroying pre-trained features
optimizer = optim.Adam([
{'params': model.layer4.parameters(), 'lr': 1e-5},
{'params': model.fc.parameters(), 'lr': 1e-4}
])
This gradual unfreezing approach allows the network to subtly adjust its feature extraction to better match your specific domain while retaining generalizable low-level patterns.
Summary
- Import pre-trained weights using
torchvision.modelswithweights='DEFAULT'to leverage ImageNet features. - Freeze base parameters by setting
requires_grad = Falseon all layers except the classifier head to prevent overfitting and speed up training. - Replace the final layer to match your class count, adapting the model to your specific task.
- Optimize selectively by passing only unfrozen parameters to the optimizer, ensuring efficient updates.
- Fine-tune optionally by unfreezing deeper layers later with reduced learning rates for marginal accuracy gains.
Frequently Asked Questions
What is the difference between feature extraction and fine-tuning in transfer learning?
Feature extraction freezes all pre-trained layers and trains only the new classifier head, which is faster and requires less data. Fine-tuning unfreezes some or all of the pre-trained layers to adjust their weights to the new dataset, which can improve accuracy but requires more computation and risks overfitting if the dataset is small.
Why should I freeze layers when using a pre-trained model in PyTorch?
Freezing layers prevents the optimizer from updating weights that already capture universal visual features (edges, colors, textures). This preserves the knowledge gained from training on large datasets like ImageNet and prevents overfitting on smaller target datasets, as implemented in the 08-TransferLearning lesson.
Which optimizer parameters should I update during transfer learning?
Initially, update only the parameters of the new classifier head (e.g., model.fc.parameters()). After the head trains sufficiently, you can optionally include deeper convolutional layers in the parameter groups with a lower learning rate for fine-tuning, as shown in the multi-parameter group optimizer configuration.
How do I access the Cats vs. Dogs dataset used in the Microsoft AI for Beginners curriculum?
Import the load_cats_dogs_dataset() function from lessons/4-ComputerVision/08-TransferLearning/pytorchcv.py. This helper automatically downloads, preprocesses, and returns the dataset split into training and validation DataLoaders with the standard ImageNet normalization transforms required by pre-trained models.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →