Implementing Transfer Learning with Pre-Trained PyTorch Models: A Complete Guide

You can implement transfer learning in PyTorch by loading a pre-trained model like VGG-16, replacing its final classifier layer to match your target classes, freezing the convolutional feature extractor, and training only the new head on your domain-specific data.

Implementing transfer learning with pre-trained PyTorch models allows you to leverage networks trained on millions of images (such as ImageNet) and adapt them to specialized tasks with minimal data and compute. The Microsoft AI-For-Beginners curriculum demonstrates this technique in lessons/4-ComputerVision/08-TransferLearning/TransferLearningPyTorch.ipynb using torchvision's VGG-16 architecture to achieve approximately 98% validation accuracy on a binary Cats vs. Dogs classification task.

Loading Pre-Trained VGG-16 from Torchvision

The foundation of transfer learning begins with importing a model pre-trained on ImageNet. In the AI-For-Beginners repository, the implementation loads VGG-16 with weights already optimized for visual feature extraction.

import torch
import torchvision

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

# Load pre-trained VGG-16

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

This single line downloads the model architecture along with learned parameters, enabling immediate use of sophisticated convolutional filters without training from scratch.

Preparing Data with ImageNet Normalization

Pre-trained models expect input tensors normalized to the same distribution used during their original training. The source code applies specific mean and standard deviation values derived from the ImageNet dataset.

import torchvision.transforms as T

# ImageNet normalization statistics

normalize = T.Normalize(
    mean=[0.485, 0.456, 0.406],
    std=[0.229, 0.224, 0.225]
)

transform = T.Compose([
    T.Resize(256),
    T.CenterCrop(224),
    T.ToTensor(),
    normalize,
])

Always use these specific normalization parameters when working with ImageNet pre-trained models to ensure feature compatibility and prevent distribution shift.

Modifying the Architecture for Transfer Learning

Transfer learning requires adapting the pre-trained network to output predictions for your specific task rather than the original 1000 ImageNet classes.

Replacing the Classifier Layer

The VGG-16 model contains a classifier attribute that houses the final fully connected layers. You replace the existing Linear layer with a new one matching your target class count.


# Replace classifier for binary classification (e.g., Cats vs. Dogs)

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

The input dimension 25088 corresponds to the flattened output of VGG-16's final pooling layer, while 2 represents the number of target classes in your custom dataset.

Freezing the Feature Extractor

To prevent overfitting and preserve low-level visual features (edges, textures, shapes), freeze the convolutional backbone by disabling gradient computation for all parameters in the features module.


# Freeze convolutional layers

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

This feature extraction approach drastically reduces training time and memory usage because gradients only flow through the newly initialized classifier layer.

Training and Fine-Tuning Strategies

The AI-For-Beginners curriculum implements a two-phase training strategy: initial classifier training followed by optional fine-tuning.

Phase 1: Training the New Classifier

With the feature extractor frozen, train only the classifier parameters using a standard optimizer configuration.

import torch.nn as nn
from torch.utils.data import DataLoader

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

def train_epoch(model, loader, optimizer, criterion):
    model.train()
    for images, labels in loader:
        images, labels = images.to(device), labels.to(device)
        
        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

This phase rapidly converges because only the final layer's weights require adjustment to map extracted features to your target classes.

Phase 2: Fine-Tuning the Full Network

After the classifier stabilizes, unfreeze the convolutional layers to allow subtle adjustments to pre-trained filters. Use a significantly lower learning rate to prevent catastrophic forgetting.


# Unfreeze feature extractor

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

# Reinitialize optimizer with lower learning rate and full parameter access

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

Fine-tuning adapts the network to domain-specific visual characteristics while maintaining the generalization benefits of ImageNet pre-training. According to the source code, this step is optional but can improve performance on datasets whose visual distribution differs from ImageNet.

Persisting Trained Models

Once training completes, persist the complete model state for future inference without retraining.


# Save the entire model

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

This creates a portable file containing both the architecture and learned weights, suitable for deployment or further experimentation as demonstrated in the notebook's final cell.

Summary

  • Load pre-trained architectures using torchvision.models to access sophisticated feature extractors trained on ImageNet.
  • Normalize input data using ImageNet statistics (mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) to maintain compatibility with pre-trained convolutional filters.
  • Replace the classifier by assigning a new torch.nn.Linear layer to vgg.classifier, using 25088 input features and adjusting output dimensions to match your class count.
  • Freeze the feature extractor by setting requires_grad=False on vgg.features.parameters() to implement efficient feature extraction with minimal compute.
  • Implement two-phase training: first train the classifier with frozen features, then optionally fine-tune the entire network with a reduced learning rate (e.g., 1e-4).
  • Persist results using torch.save() to create reusable model files for inference.

Frequently Asked Questions

How do I choose between freezing and fine-tuning in PyTorch transfer learning?

Freeze the feature extractor when your dataset is small or similar to ImageNet; this prevents overfitting and trains faster. Fine-tune the entire network when you have substantial domain-specific data or when target images differ significantly from natural photographs (e.g., medical imaging or satellite data).

Why must I use specific normalization values for pre-trained models?

Pre-trained models like VGG-16 learned their internal representations assuming input tensors normalized to ImageNet's mean and standard deviation. Using different normalization statistics shifts the input distribution away from what the convolutional filters expect, significantly reducing model accuracy.

What input dimension does the VGG-16 classifier expect?

The VGG-16 architecture expects inputs of size 25088, which equals 512 * 7 * 7 (the flattened output of the final max-pooling layer). When replacing vgg.classifier, your new torch.nn.Linear layer must accept 25088 features as its first argument to match the backbone's output.

Can I use this approach with other architectures besides VGG-16?

Yes. The same transfer learning pattern applies to ResNet, DenseNet, EfficientNet, and other torchvision models. Each architecture provides a classifier or fc attribute for modification, though the specific input dimension to the final layer varies by model (e.g., ResNet-50 uses 2048 features).

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 →