How to Implement Transfer Learning with Pre-Trained CNN Models in PyTorch

Transfer learning allows you to repurpose a deep CNN like VGG-16 by freezing its pre-trained convolutional layers, substituting the final classifier for your specific number of classes, and training only the new layer before optionally unfreezing the entire network for fine-tuning.

The Microsoft AI-For-Beginners repository contains a complete implementation of transfer learning with pre-trained CNN models in lessons/4-ComputerVision/08-TransferLearning/TransferLearningPyTorch.ipynb. This guide walks through the exact workflow used to adapt ImageNet-trained architectures for custom datasets such as Cats vs. Dogs, leveraging the torchvision models library and standard PyTorch training patterns.

Loading a Pre-Trained VGG-16 Backbone

The first step involves instantiating a model with weights learned from the ImageNet dataset. In TransferLearningPyTorch.ipynb, the VGG-16 architecture serves as the demonstration backbone, though this pattern applies equally to ResNet, Inception, or EfficientNet.

Set pretrained=True to automatically download the trained weights from the PyTorch model zoo:

import torchvision
import torch

device = 'cuda' if torch.cuda.is_available() else 'cpu'

# Load VGG-16 with ImageNet weights

vgg = torchvision.models.vgg16(pretrained=True).to(device)

This initializes the full model including the features extractor (convolutional layers) and the original classifier (fully-connected layers) outputting 1000 classes.

Adapting the Model Architecture

Before training on your target dataset, you must modify the model to output the correct number of classes and prevent the pre-trained filters from being overwritten during initial training.

Replacing the Final Classifier

The original VGG-16 classifier maps the flattened feature vector (25,088 dimensions) to 1,000 ImageNet categories. For a binary classification task like Cats vs. Dogs, replace this with a single linear layer matching your class count:


# Replace classifier: 25088 input features -> 2 output classes

vgg.classifier = torch.nn.Linear(25088, 2).to(device)

As implemented in the notebook at line 138, this substitution discards the pre-trained classification head while preserving the convolutional feature extractor that detects generic visual patterns like edges and textures.

Freezing the Feature Extractor

To preserve the ImageNet-learned filters during the first training phase, disable gradient computation for all parameters in the features module:

for param in vgg.features.parameters():
    param.requires_grad = False

This freeze operation, shown at line 140 of TransferLearningPyTorch.ipynb, ensures that backpropagation updates only the weights of your new classifier, effectively using the pre-trained CNN as a fixed feature extractor.

Preparing the Dataset with ImageNet Normalization

Transfer learning requires that input images undergo the same preprocessing used during the original model's training. The AI-For-Beginners implementation uses torchvision.datasets.ImageFolder with the standardized ImageNet normalization values:

transform = torchvision.transforms.Compose([
    torchvision.transforms.Resize(256),
    torchvision.transforms.CenterCrop(224),
    torchvision.transforms.ToTensor(),
    torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],
                                     std=[0.229, 0.224, 0.225])
])

dataset = torchvision.datasets.ImageFolder('data/PetImages', transform=transform)

These specific mean and standard deviation values (line 133) ensure pixel intensities match the distribution the pre-trained model expects, preventing domain shift during inference.

Training the New Classifier

With the feature extractor frozen, configure the optimizer to update only the classifier parameters. This approach trains the new layer rapidly while keeping the expensive convolutional weights static:

from torch.utils.data import DataLoader, random_split

# Split dataset (e.g., 20,000 training samples)

train_set, val_set = random_split(dataset, [20000, len(dataset)-20000])
train_loader = DataLoader(train_set, batch_size=32, shuffle=True)
val_loader = DataLoader(val_set, batch_size=32)

# Optimizer affects only the classifier

optimizer = torch.optim.Adam(vgg.classifier.parameters(), lr=1e-3)
criterion = torch.nn.CrossEntropyLoss()

# Training loop (feature extraction phase)

for epoch in range(5):
    vgg.train()
    for imgs, labels in train_loader:
        imgs, labels = imgs.to(device), labels.to(device)
        optimizer.zero_grad()
        outputs = vgg(imgs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

This phase, starting at line 174 in the source notebook, typically converges quickly since only the small linear layer learns from scratch while the deep features remain optimal.

Fine-Tuning the Entire Network

After the classifier converges, you can optionally unfreeze the convolutional layers to allow the entire network to adapt to domain-specific nuances. Reduce the learning rate significantly to avoid catastrophic forgetting of the pre-trained weights:


# Unfreeze feature extractor

for param in vgg.features.parameters():
    param.requires_grad = True

# Lower learning rate for full-network fine-tuning

optimizer = torch.optim.Adam(vgg.parameters(), lr=1e-4)

# Continue training with all parameters updatable...

This fine-tuning stage, referenced at line 240 of TransferLearningPyTorch.ipynb, typically yields higher accuracy on complex datasets where low-level features might need subtle adjustment.

Saving the Adapted Model

Once validation accuracy satisfies your requirements, persist the complete model including both the fine-tuned features and the custom classifier:

torch.save(vgg, 'data/cats_dogs.pth')

This saves the entire model architecture and weights (line 184), enabling later loading without redefining the class replacement logic.

Summary

  • Load pre-trained weights using torchvision.models with pretrained=True to obtain a model with robust generic features.
  • Replace the classifier to match your dataset's class count, modifying the final linear layer while preserving the convolutional backbone.
  • Freeze features initially by setting requires_grad=False on convolutional parameters, training only the new classification head for efficiency.
  • Use ImageNet normalization (mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]) in your data pipeline to maintain input consistency.
  • Fine-tune optionally by unfreezing all layers and using a reduced learning rate (e.g., 1e-4) to adapt the entire network to your specific domain.
  • The complete reference implementation resides in lessons/4-ComputerVision/08-TransferLearning/TransferLearningPyTorch.ipynb, with additional practice available in the OxfordPets.ipynb lab file.

Frequently Asked Questions

Can I use a different pre-trained model instead of VGG-16?

Yes, the pattern works with any torchvision architecture. Replace torchvision.models.vgg16 with resnet50, inception_v3, or efficientnet_b0. Note that the input dimension to your new classifier layer will vary; for example, ResNet-50 features have 2,048 dimensions before the final fully-connected layer, requiring you to adjust the in_features parameter accordingly.

How do I know when to stop training the classifier and start fine-tuning?

Monitor validation loss or accuracy. Typically, you train the frozen feature extractor until the validation metric plateaus—usually 5-10 epochs for small datasets. Once the classifier converges, unfreeze the convolutional layers and continue with a lower learning rate. The Microsoft AI-For-Beginners notebook demonstrates this transition after the initial feature extraction phase achieves stable accuracy.

Why must I use the specific ImageNet normalization values?

Pre-trained models expect inputs with the same statistical distribution as their training data. The normalization parameters mean=[0.485,0.456,0.406] and std=[0.229,0.224,0.225] represent the average pixel values across the ImageNet dataset. Applying these transformations ensures that activations in the early convolutional layers fall within the range the network learned to process, preventing performance degradation on your custom dataset.

Where can I find additional practice datasets for this workflow?

The repository includes a dedicated lab exercise. Navigate to lessons/4-ComputerVision/08-TransferLearning/lab/OxfordPets.ipynb to apply these transfer learning steps to the Oxford-Pets dataset, which provides a different class structure and image distribution than the Cats vs. Dogs example, reinforcing the generalizability of the technique.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →