By the end of this lesson, you will be able to adapt a pre-trained deep learning model for a new classification task using transfer learning, significantly reducing training time and data requirements.
What it is
Transfer learning is a machine learning method where a model developed for one task is reused as the starting point for a model on a second task. Instead of training a neural network from scratch (random initialization), you leverage weights learned from a large dataset (like ImageNet) that already understand general features such as edges, textures, and shapes. The mental model is "feature reuse": early layers capture generic patterns, while later layers specialize in specific classes. Related terms include fine-tuning (updating all or some weights) and feature extraction (freezing base weights and only training the final classifier).Why it matters
- Data Efficiency: Achieves high accuracy with small datasets by leveraging knowledge from massive pre-training sets.
- Computational Speed: Converges much faster than training from scratch because the model starts near an optimal solution.
- Performance Boost: Often outperforms models trained from scratch, especially when labeled data is scarce.
- Resource Conservation: Reduces energy consumption and hardware costs associated with long training epochs.
Syntax or steps
The standard workflow involves three main steps: 1. Load a pre-trained model (e.g., ResNet, VGG) without its top classification layer. 2. Freeze the base model's weights so they are not updated during training. 3. Add a new fully connected head tailored to your specific number of classes and train only this new part (or fine-tune the whole model with a low learning rate).Example
import torch
import torch.nn as nn
from torchvision import models
# 1. Load pre-trained ResNet18
model = models.resnet18(pretrained=True)
# 2. Freeze all parameters in the base model
for param in model.parameters():
param.requires_grad = False
# 3. Replace the final fully connected layer
# Original num_features depends on architecture; for resnet18 it is 512
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, 10) # Assuming 10 classes for new task
# 4. Define optimizer for ONLY the new layer
optimizer = torch.optim.SGD(model.fc.parameters(), lr=0.001, momentum=0.9)
print("Model adapted for transfer learning.")
Explanation: First, we load `resnet18` with weights trained on ImageNet. We then iterate through all parameters and set `requires_grad=False`, effectively freezing the feature extractor. Next, we replace the original output layer (`model.fc`) with a new linear layer matching our target class count (10). Finally, the optimizer is initialized specifically with `model.fc.parameters()`, ensuring gradients are computed and weights updated only for the new head, preserving the pre-trained features.
Common mistakes
- Forgetting to freeze weights: If you do not set `requires_grad=False`, the entire model updates, potentially destroying useful pre-trained features if your dataset is small.
- Mismatched input sizes: Pre-trained models expect specific image dimensions (e.g., 224x224). Failing to resize inputs correctly causes shape errors or poor performance.
- Using too high a learning rate: When fine-tuning (unfreezing layers), a high learning rate can destabilize the pre-trained weights. Use a much lower LR than typical fresh training.
- Ignoring normalization: Pre-trained models require inputs normalized with specific mean/std values (e.g., ImageNet stats). Skipping this step leads to inaccurate predictions.
When to use it
Compare transfer learning with training from scratch based on data availability and similarity.| Scenario | Recommended Approach | Reasoning |
|---|---|---|
| Small Dataset (< 10k images) | Transfer Learning | Scratch training overfits easily; pre-trained features provide regularization. |
| Large Dataset & Similar Task | Fine-Tuning | Leverage existing features but adapt deeper layers to specific nuances. |
| Very Different Domain | Feature Extraction | Early layers may still help, but avoid updating them if they conflict with new data distribution. |
| Huge Dataset & Unique Architecture | Train From Scratch | Pre-trained weights might introduce bias; full control allows custom optimization. |
Practice
Guided Exercise: Modify the example above to use `vgg16` instead of `resnet18`. Note that VGG uses a different structure for its classifier head. You must identify the correct attribute to replace (usually `classifier[6]`).Challenge: Implement a simple validation loop that checks if the frozen layers' weights remain unchanged after one epoch of training. Hint: Store a copy of a weight tensor before training and compare it after.
Quick check
Question: Why do we typically freeze the early layers of a pre-trained CNN when applying transfer learning?Answer: Early layers learn generic features (edges, colors) that are universal across many visual tasks. Freezing them preserves this valuable knowledge and prevents catastrophic forgetting, allowing the model to focus on learning task-specific patterns in the later layers.