Skip to main content

🔄 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

flowchart TB A[Pre-trained Model
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

graph TD A[New Task] --> B{Dataset Size} B -->|Small Dataset| C{Task Similarity} C -->|Similar to Pre-trained| D[Feature Extraction
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

flowchart LR A[Start] --> B[Load Pre-trained Model] B --> C[Remove Top Layers] C --> D[Add Custom Layers] D --> E[Freeze Base Layers] E --> F[Train on New Data
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

graph TB subgraph "Feature Extraction" FE1[Pre-trained Base] --> FE2[Frozen] FE2 --> FE3[New Classifier] FE3 --> FE4[Train Only Top] end subgraph "Fine-tuning" FT1[Pre-trained Base] --> FT2[Initially Frozen] FT2 --> FT3[New Classifier] FT3 --> FT4[Train Top First] FT4 --> FT5[Unfreeze Some Layers] FT5 --> FT6[Train All with Low LR] end subgraph "Full Training" FL1[Pre-trained Weights] --> FL2[All Trainable] FL2 --> FL3[New Classifier] FL3 --> FL4[Train Everything] end style FE2 fill:#ffcdd2 style FT2 fill:#fff9c4 style FT5 fill:#c8e6c9 style FL2 fill:#bbdefb

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

🌟 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.