🔄 Transfer Learning: Leveraging Pre-trained Models
📚 What You'll Learn
By the end of this lesson, you will be able to:
- Explain what transfer learning is and why it saves data, time, and compute
- Decide when transfer learning is appropriate for a problem
- Apply feature extraction by freezing a pre-trained model's convolutional base
- Fine-tune a pre-trained model by unfreezing and retraining its upper layers
- Choose among common pre-trained models (VGG, ResNet, Inception, MobileNet) for your task
- Build a custom image classifier on top of a pre-trained backbone
⏱️ Estimated Time: 45–60 minutes
🎯 Project: Build an image classifier from a pre-trained model — first as a feature extractor, then fine-tuned.
Introduction
Transfer learning is a powerful technique that leverages knowledge gained from pre-trained models to solve new, related problems. Instead of training a neural network from scratch, we can use models trained on large datasets like ImageNet and adapt them to our specific tasks. This dramatically reduces training time, computational resources, and data requirements while often achieving better performance.
Transfer Learning Overview
e.g., VGG16, ResNet] --> B{Transfer Learning Strategy} B --> C[Feature Extraction] C --> C1[Freeze Base Layers] C --> C2[Add New Classifier] C --> C3[Train Only New Layers] B --> D[Fine-tuning] D --> D1[Start with Feature Extraction] D --> D2[Unfreeze Top Layers] D --> D3[Train with Low Learning Rate] B --> E[Using as Initializer] E --> E1[Use Pre-trained Weights] E --> E2[Train Entire Network] E --> E3[Larger Learning Rate OK] C --> F[Quick Training
Less Data Needed] D --> G[Better Performance
More Customization] E --> H[Maximum Flexibility
Requires More Data] style A fill:#e1f5fe style F fill:#c8e6c9 style G fill:#fff9c4 style H fill:#ffccbc
When to Use Transfer Learning
Freeze Most Layers] C -->|Different Domain| E[Fine-tune More Layers
Or Train From Scratch] B -->|Large Dataset| F{Task Similarity} F -->|Similar| G[Fine-tuning
Unfreeze Many Layers] F -->|Different| H[Use as Initializer
Or Train From Scratch] D --> I[✓ Best Choice] G --> I style D fill:#4caf50,color:#fff style G fill:#4caf50,color:#fff style I fill:#2196f3,color:#fff
Setting Up Transfer Learning
import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers
from tensorflow.keras.applications import VGG16, ResNet50, MobileNetV2
import numpy as np
import matplotlib.pyplot as plt
# Check TensorFlow version
print(f"TensorFlow version: {tf.__version__}")
# Set random seeds for reproducibility
np.random.seed(42)
tf.random.set_seed(42)
Method 1: Feature Extraction
def create_feature_extraction_model(input_shape=(224, 224, 3), num_classes=10):
"""
Create a model using pre-trained VGG16 as feature extractor
"""
# Load pre-trained VGG16 without top layers
base_model = VGG16(
weights='imagenet', # Pre-trained on ImageNet
include_top=False, # Exclude final classification layer
input_shape=input_shape
)
# Freeze the base model layers (no training)
base_model.trainable = False
# Create new model
model = keras.Sequential([
base_model,
layers.GlobalAveragePooling2D(),
layers.Dense(128, activation='relu'),
layers.Dropout(0.5),
layers.Dense(num_classes, activation='softmax')
])
return model, base_model
# Create the model
model_fe, base_model_fe = create_feature_extraction_model(num_classes=10)
# Display model architecture
print("Feature Extraction Model Summary:")
print(f"Base model layers: {len(base_model_fe.layers)}")
print(f"Trainable parameters: {sum([tf.size(w).numpy() for w in model_fe.trainable_weights])}")
print(f"Non-trainable parameters: {sum([tf.size(w).numpy() for w in model_fe.non_trainable_weights])}")
# Compile model
model_fe.compile(
optimizer='adam',
loss='sparse_categorical_crossentropy',
metrics=['accuracy']
)
Method 2: Fine-tuning
def create_fine_tuning_model(input_shape=(224, 224, 3), num_classes=10):
"""
Create a model for fine-tuning with selective layer unfreezing
"""
# Load pre-trained ResNet50
base_model = ResNet50(
weights='imagenet',
include_top=False,
input_shape=input_shape
)
# First, freeze all layers
base_model.trainable = False
# Build the model
inputs = keras.Input(shape=input_shape)
# Pre-processing for ResNet50
x = keras.applications.resnet50.preprocess_input(inputs)
# Base model
x = base_model(x, training=False)
# Custom top layers
x = layers.GlobalAveragePooling2D()(x)
x = layers.Dense(256, activation='relu')(x)
x = layers.BatchNormalization()(x)
x = layers.Dropout(0.5)(x)
outputs = layers.Dense(num_classes, activation='softmax')(x)
model = keras.Model(inputs, outputs)
return model, base_model
# Create fine-tuning model
model_ft, base_model_ft = create_fine_tuning_model(num_classes=10)
print("\nFine-tuning Model Created")
print(f"Total layers: {len(model_ft.layers)}")
# Function to unfreeze top layers for fine-tuning
def unfreeze_top_layers(base_model, num_layers=20):
"""
Unfreeze the top 'num_layers' for fine-tuning
"""
# Freeze all layers first
base_model.trainable = True
# Freeze all but the top num_layers
for layer in base_model.layers[:-num_layers]:
layer.trainable = False
print(f"Unfroze top {num_layers} layers for fine-tuning")
# Count trainable layers
trainable_count = sum([1 for layer in base_model.layers if layer.trainable])
print(f"Trainable layers in base model: {trainable_count}/{len(base_model.layers)}")
# Example: Unfreeze top 20 layers for fine-tuning
# This would typically be done after initial training
# unfreeze_top_layers(base_model_ft, 20)
Transfer Learning Workflow
High Learning Rate] F --> G{Performance OK?} G -->|No| H[Unfreeze Top Layers] H --> I[Continue Training
Low Learning Rate] I --> G G -->|Yes| J[Save Model] J --> K[Deploy] style A fill:#f9f style K fill:#9f9
Practical Example: Custom Image Classifier
class TransferLearningClassifier:
"""
Complete transfer learning pipeline for image classification
"""
def __init__(self, num_classes, input_shape=(224, 224, 3)):
self.num_classes = num_classes
self.input_shape = input_shape
self.model = None
self.base_model = None
self.history = None
def build_model(self, base_model_name='MobileNetV2'):
"""
Build model with specified pre-trained base
"""
# Select base model
if base_model_name == 'MobileNetV2':
self.base_model = MobileNetV2(
weights='imagenet',
include_top=False,
input_shape=self.input_shape
)
preprocess_func = tf.keras.applications.mobilenet_v2.preprocess_input
elif base_model_name == 'VGG16':
self.base_model = VGG16(
weights='imagenet',
include_top=False,
input_shape=self.input_shape
)
preprocess_func = tf.keras.applications.vgg16.preprocess_input
else:
raise ValueError(f"Unknown base model: {base_model_name}")
# Freeze base model
self.base_model.trainable = False
# Build complete model
inputs = keras.Input(shape=self.input_shape)
x = preprocess_func(inputs)
x = self.base_model(x, training=False)
x = layers.GlobalAveragePooling2D()(x)
x = layers.Dense(128, activation='relu')(x)
x = layers.Dropout(0.5)(x)
outputs = layers.Dense(self.num_classes, activation='softmax')(x)
self.model = keras.Model(inputs, outputs)
print(f"Model built with {base_model_name} backbone")
print(f"Output classes: {self.num_classes}")
def compile_model(self, learning_rate=0.001):
"""
Compile the model with appropriate settings
"""
self.model.compile(
optimizer=keras.optimizers.Adam(learning_rate=learning_rate),
loss='sparse_categorical_crossentropy',
metrics=['accuracy']
)
print(f"Model compiled with learning rate: {learning_rate}")
def train_feature_extraction(self, train_data, val_data, epochs=10):
"""
Train only the top layers (feature extraction)
"""
print("\n--- Phase 1: Feature Extraction Training ---")
self.history_fe = self.model.fit(
train_data,
validation_data=val_data,
epochs=epochs,
callbacks=[
keras.callbacks.EarlyStopping(
monitor='val_loss',
patience=3,
restore_best_weights=True
)
]
)
return self.history_fe
def fine_tune(self, train_data, val_data, layers_to_unfreeze=20, epochs=10):
"""
Fine-tune the model by unfreezing top layers
"""
print("\n--- Phase 2: Fine-tuning ---")
# Unfreeze top layers of base model
self.base_model.trainable = True
for layer in self.base_model.layers[:-layers_to_unfreeze]:
layer.trainable = False
print(f"Unfroze top {layers_to_unfreeze} layers")
# Recompile with lower learning rate
self.compile_model(learning_rate=0.0001)
# Continue training
self.history_ft = self.model.fit(
train_data,
validation_data=val_data,
epochs=epochs,
initial_epoch=len(self.history_fe.history['loss']),
callbacks=[
keras.callbacks.EarlyStopping(
monitor='val_loss',
patience=3,
restore_best_weights=True
)
]
)
return self.history_ft
def plot_training_history(self):
"""
Plot training and validation metrics
"""
if not self.history_fe:
print("No training history available")
return
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
# Combine histories if fine-tuning was done
if hasattr(self, 'history_ft'):
acc = self.history_fe.history['accuracy'] + self.history_ft.history['accuracy']
val_acc = self.history_fe.history['val_accuracy'] + self.history_ft.history['val_accuracy']
loss = self.history_fe.history['loss'] + self.history_ft.history['loss']
val_loss = self.history_fe.history['val_loss'] + self.history_ft.history['val_loss']
# Mark fine-tuning start
ft_start = len(self.history_fe.history['accuracy'])
ax1.axvline(x=ft_start, color='red', linestyle='--', label='Fine-tuning start')
ax2.axvline(x=ft_start, color='red', linestyle='--', label='Fine-tuning start')
else:
acc = self.history_fe.history['accuracy']
val_acc = self.history_fe.history['val_accuracy']
loss = self.history_fe.history['loss']
val_loss = self.history_fe.history['val_loss']
# Plot accuracy
ax1.plot(acc, label='Training Accuracy')
ax1.plot(val_acc, label='Validation Accuracy')
ax1.set_xlabel('Epoch')
ax1.set_ylabel('Accuracy')
ax1.set_title('Model Accuracy')
ax1.legend()
ax1.grid(True, alpha=0.3)
# Plot loss
ax2.plot(loss, label='Training Loss')
ax2.plot(val_loss, label='Validation Loss')
ax2.set_xlabel('Epoch')
ax2.set_ylabel('Loss')
ax2.set_title('Model Loss')
ax2.legend()
ax2.grid(True, alpha=0.3)
plt.suptitle('Transfer Learning Training Progress')
plt.tight_layout()
plt.show()
# Example usage
print("\n" + "="*60)
print("TRANSFER LEARNING EXAMPLE")
print("="*60)
# Create classifier
classifier = TransferLearningClassifier(num_classes=5)
# Build model with MobileNetV2 (lightweight and efficient)
classifier.build_model('MobileNetV2')
# Compile
classifier.compile_model(learning_rate=0.001)
# Note: In practice, you would load your actual data here
print("\nModel ready for training!")
print("Steps:")
print("1. Load and preprocess your dataset")
print("2. Run train_feature_extraction() for initial training")
print("3. Run fine_tune() for performance improvement")
print("4. Use plot_training_history() to visualize results")
Comparing Transfer Learning Approaches
Common Pre-trained Models
# Available pre-trained models in Keras
models_info = {
'VGG16': {
'parameters': '138M',
'depth': 16,
'use_case': 'Good baseline, simple architecture',
'input_size': (224, 224, 3)
},
'ResNet50': {
'parameters': '25.6M',
'depth': 50,
'use_case': 'Good balance of performance and size',
'input_size': (224, 224, 3)
},
'MobileNetV2': {
'parameters': '3.5M',
'depth': 53,
'use_case': 'Lightweight, mobile deployment',
'input_size': (224, 224, 3)
},
'InceptionV3': {
'parameters': '23.8M',
'depth': 159,
'use_case': 'High accuracy, complex patterns',
'input_size': (299, 299, 3)
},
'EfficientNet': {
'parameters': '5.3M - 66M',
'depth': 'Variable',
'use_case': 'State-of-the-art, scalable',
'input_size': 'Variable'
}
}
print("Popular Pre-trained Models for Transfer Learning:")
print("-" * 60)
for model_name, info in models_info.items():
print(f"\n{model_name}:")
for key, value in info.items():
print(f" {key}: {value}")
Best Practices
🎯 Transfer Learning Guidelines
- Data Similarity: Choose pre-trained models from similar domains
- Start Simple: Begin with feature extraction before fine-tuning
- Learning Rates: Use lower learning rates for fine-tuning (1e-5 to 1e-4)
- Gradual Unfreezing: Unfreeze layers progressively, not all at once
- Monitor Overfitting: Watch validation metrics carefully
- Data Augmentation: Still important, especially with small datasets
- Batch Normalization: Be careful when fine-tuning models with BN layers
Practice Exercises
Exercise 1: Multi-Model Comparison
Compare different pre-trained models on your dataset:
- Test VGG16, ResNet50, and MobileNetV2
- Compare accuracy, training time, and model size
- Determine the best model for your use case
Exercise 2: Progressive Fine-tuning
Implement progressive unfreezing strategy:
- Start with all layers frozen
- Gradually unfreeze blocks of layers
- Track performance at each stage
- Find optimal number of trainable layers
Summary
✅ You've Learned
- What transfer learning is and why it's powerful
- Three main approaches: feature extraction, fine-tuning, and full training
- How to implement transfer learning with Keras/TensorFlow
- When to use each transfer learning strategy
- Common pre-trained models and their characteristics
- Best practices for successful transfer learning
📓 Learning Journal
Keep a learning journal — digital or physical. After this lesson, take a few minutes to write down:
- Key concepts you learned
- Techniques that clicked for you
- Questions or confusion points to revisit
- Ideas you want to try
- Your progress and feelings about learning this
✍️ This lesson's prompt: Transfer learning reuses knowledge a model gained from a completely different dataset. Where in your own learning have you transferred a skill from one area to speed up progress in another?
📝 Lesson Summary
🎓 Key Takeaways
- Transfer learning reuses features learned on large datasets (like ImageNet), so you need far less data and compute.
- Feature extraction freezes the pre-trained base and trains only a new head; fine-tuning also unfreezes and retrains upper layers.
- The more similar your task is to the original, the more of the pre-trained model you can reuse as-is.
- Choosing a base model trades off accuracy, size, and speed (for example, MobileNet for mobile, ResNet for accuracy).
🎉 What You've Accomplished
You can now stand on the shoulders of models trained on millions of images — building accurate classifiers quickly with feature extraction and fine-tuning.
❓ Common Questions at This Stage
What's the difference between feature extraction and fine-tuning?
Feature extraction freezes the pre-trained layers and trains only a new classifier head on top. Fine-tuning goes further, unfreezing some upper pre-trained layers and retraining them at a low learning rate to adapt them to your data.
When should I fine-tune instead of just extracting features?
Fine-tune when you have enough data and your task differs somewhat from the original. If your dataset is small or very similar to the source, feature extraction alone is safer and less prone to overfitting.
How do I pick a pre-trained model?
Match it to your constraints: MobileNet or EfficientNet for speed and small size, ResNet or Inception for higher accuracy. Start with a well-supported model and adjust based on accuracy and latency needs.
🔭 Looking Ahead
Transfer learning caps your deep-learning toolkit — from here you can combine these techniques on real projects and explore tuning, specialized domains, and deployment.
✅ Before the Next Lesson
- Train a classifier with a frozen base, then fine-tune the top layers and compare the accuracy.
- Swap in a different pre-trained backbone and compare size, speed, and accuracy.
- Write your Learning Journal entry for this lesson
🌟 Encouragement for the Journey
Transfer learning is how modern practitioners get strong results without massive datasets or budgets. You now know how to reuse the field's best models for your own goals — a genuine superpower. Well done.