PyTorch Transfer Learning

Transfer Learning is a technique that involves taking a model pretrained on a large-scale dataset and transferring it to a new task with a smaller amount of data for training.

It is one of the most widely used techniques in deep learning practice today — in most cases, transfer learning performs better, is faster, and requires less data than training from scratch.


1. Core Idea of Transfer Learning

Features learned by deep neural networks on ImageNet are universal:

  • Shallow layers: learn general low-level features (edges, textures, color gradients)
  • Middle layers: learn general mid-level features (shapes, parts, texture combinations)
  • Deep layers: learn task-specific high-level features (faces, wheels, text)

These low- and mid-level features are effective for the vast majority of vision tasks and do not need to be relearned.

When to Use Transfer Learning?

Data amountSimilarity to source taskRecommended strategy
Small (< 1000)HighOnly replace the final classification head, freeze the entire backbone.
Small (< 1000)LowFine-tune shallower layers, freeze deeper layers.
Medium (1000~10000)HighFine-tune the entire network with a small learning rate.
Medium (1000~10000)LowFine-tune deeper layers, freeze shallower layers.
Large (> 10000)AnyFine-tune everything, or consider training from scratch.

Comparison of the Three Core Strategies

The layer structure of a pretrained model (e.g., ResNet50) is as follows:

  • Conv Layer 1~3 (low-level features: edges/textures): usually frozen
  • Conv Layer 4~6 (mid-level features: shapes/parts): optionally frozen
  • Conv Layer 7~N (high-level features: semantic information): fine-tuned
  • Classifier Head: replaced and trained
┌─────────────────────────────────────────────────┐
│              预训练模型(如 ResNet50)             │
│  ┌──────────────────────────────────────────┐   │
│  │  Conv Layer 1~3(低级特征:边缘/纹理)      │  ← 通常冻结
│  ├──────────────────────────────────────────┤   │
│  │  Conv Layer 4~6(中级特征:形状/部件)      │  ← 可选冻结
│  ├──────────────────────────────────────────┤   │
│  │  Conv Layer 7~N(高级特征:语义信息)       │  ← 微调
│  ├──────────────────────────────────────────┤   │
│  │  Classifier Head(分类头)                 │  ← 替换 & 训练
│  └──────────────────────────────────────────┘   │
└─────────────────────────────────────────────────┘

2. Loading Pretrained Models

PyTorch provides many official pretrained models through torchvision.models, and loading them is very simple.

Example

import torch
import torchvision.models as models

# Load pretrained model (automatically downloads weights)
# PyTorch >= 0.13 recommends the new syntax: use the weights parameter
from torchvision.models import ResNet50_Weights

model = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)

# Old syntax (still works, but will trigger a deprecation warning)
model = models.resnet50(pretrained=True)

# Do not load pretrained weights (use only the network architecture)
model = models.resnet50(weights=None)

Viewing the Model Structure

Example

# Print the complete structure
print(model)

# Only view the last few layers (classification head)
print(model.fc)
# Linear(in_features=2048, out_features=1000, bias=True)

# Count the number of parameters
total_params    = sum(p.numel() for p in model.parameters())
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"Total parameters: {total_params:,}")
print(f"Trainable parameters: {trainable_params:,}")

Names of Classification Heads for Various Models

Different models have different attribute names for their classification heads; when doing transfer learning, you need to replace the corresponding layer:

ModelClassification head attribute
ResNet / RegNetmodel.fc
VGG / AlexNetmodel.classifier[-1]
DenseNetmodel.classifier
EfficientNetmodel.classifier[-1]
MobileNetV2/V3model.classifier[-1]
ViT (Vision Transformer)model.heads.head
ConvNeXtmodel.classifier[-1]
Inception V3model.fc
Swin Transformermodel.head

3. Three Transfer Learning Strategies

3.1 Strategy 1: Feature Extraction (Freeze All)

Freeze all parameters of the pretrained model and train only the newly replaced classification head.

Suitable for scenarios with very little data (a few hundred images) or where the task is highly similar to the source task.

Example

import torch
import torch.nn as nn
import torchvision.models as models
from torchvision.models import ResNet18_Weights

NUM_CLASSES = 5   # Number of classes for the target task

# Step 1: Load the pretrained model
model = models.resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)

# Step 2: Freeze all parameters
for param in model.parameters():
    param.requires_grad = False

# Step 3: Replace the classification head (these parameters have requires_grad=True by default)
in_features = model.fc.in_features    # 512
model.fc = nn.Linear(in_features, NUM_CLASSES)

# Verify: only the classification head is trainable
trainable = [(n, p.shape) for n, p in model.named_parameters() if p.requires_grad]
print(f"Number of trainable layers: {len(trainable)}")
for name, shape in trainable:
    print(f"  {name}: {shape}")
# Output:
#   fc.weight: torch.Size([5, 512])
#   fc.bias:   torch.Size([5])

# Step 4: Pass only the trainable parameters to the optimizer (more efficient)
optimizer = torch.optim.Adam(
    filter(lambda p: p.requires_grad, model.parameters()),
    lr=1e-3
)
# Or an equivalent, cleaner approach:
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)

Feature extraction is the simplest and most commonly used transfer learning strategy, particularly suitable for scenarios with small amounts of data.

3.2 Strategy 2: Fine-tuning

Unfreeze all or some of the pretrained layers and train the entire model with a smaller learning rate.

Suitable for scenarios with a medium amount of data, or where the task differs from the source task.

Example

import torch
import torch.nn as nn
import torchvision.models as models

NUM_CLASSES = 10
model = models.resnet50(weights='IMAGENET1K_V2')

# Method A: Full fine-tuning (unfreeze all layers)
# First freeze
for param in model.parameters():
    param.requires_grad = False

# Then unfreeze (equivalent to full fine-tuning; this style is commonly used in gradual unfreezing scenarios)
for param in model.parameters():
    param.requires_grad = True

# Replace the classification head
model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)

# Full fine-tuning: use a small learning rate for the backbone and a large learning rate for the head (see Strategy 3)
optimizer = torch.optim.SGD(model.parameters(), lr=1e-4, momentum=0.9)


# Method B: Unfreeze the last N layers (partial fine-tuning)
model = models.resnet50(weights='IMAGENET1K_V2')

# First freeze all layers
for param in model.parameters():
    param.requires_grad = False

# Only unfreeze layer4 and fc (the last Block and the classification head of ResNet)
for param in model.layer4.parameters():
    param.requires_grad = True

model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)  # fc is trainable by default

print("Trainable parameters:")
for name, param in model.named_parameters():
    if param.requires_grad:
        print(f"  {name}")

3.3 Strategy 3: Layer-wise Differential Learning Rates

Use a small learning rate for the backbone (to preserve pretrained knowledge) and a large learning rate for the classification head (to quickly adapt to the new task).

This is the most commonly used fine-tuning strategy in the industry and provides the best overall performance.

Example

import torch
import torch.nn as nn
import torchvision.models as models

NUM_CLASSES = 8
model = models.resnet50(weights='IMAGENET1K_V2')
model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)

# Option 1: Two learning rates (backbone vs head)
optimizer = torch.optim.Adam([
    {'params': model.fc.parameters(),  'lr': 1e-3},     # Classification head: large learning rate
    {'params': [p for n, p in model.named_parameters()  # Backbone: small learning rate
                if not n.startswith('fc')],
     'lr': 1e-5},
])


# Option 2: Layer-wise decaying learning rate (most fine-grained)
# Layers closer to the output have larger learning rates
layer_groups = [
    (model.layer1, 1e-5),    # Shallowest layer, smallest learning rate
    (model.layer2, 3e-5),
    (model.layer3, 1e-4),
    (model.layer4, 3e-4),    # Deepest backbone layer
    (model.fc,     1e-3),    # Classification head, largest learning rate
]

param_groups = [
    {'params': layer.parameters(), 'lr': lr}
    for layer, lr in layer_groups
]
optimizer = torch.optim.Adam(param_groups)


# Option 3: Gradual unfreezing
# In early training, only train the head, then gradually unfreeze more layers (recommended by fastai)
model = models.resnet50(weights='IMAGENET1K_V2')
for param in model.parameters():
    param.requires_grad = False
model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)

def unfreeze_layers(model, num_layers):
    """Unfreeze the last num_layers layer blocks of ResNet"""
    layers = [model.layer4, model.layer3, model.layer2, model.layer1]
    for i in range(min(num_layers, len(layers))):
        for param in layers[i].parameters():
            param.requires_grad = True

# Epoch 1-5: train only the classification head
# Epoch 6-10: unfreeze layer4
unfreeze_layers(model, num_layers=1)
# Epoch 11+: unfreeze more layers
unfreeze_layers(model, num_layers=3)

4. Common Pretrained Models

4.1 Image Classification Models

Example

import torchvision.models as models

# ResNet series (most classic, suitable for most tasks)
resnet18  = models.resnet18(weights='IMAGENET1K_V1')   # Lightweight, suitable for edge devices
resnet50  = models.resnet50(weights='IMAGENET1K_V2')   # Balanced choice
resnet101 = models.resnet101(weights='IMAGENET1K_V2')  # Stronger, slower

# EfficientNet series (extremely high accuracy-efficiency ratio)
effnet_b0 = models.efficientnet_b0(weights='IMAGENET1K_V1')  # Most lightweight
effnet_b4 = models.efficientnet_b4(weights='IMAGENET1K_V1')  # Balanced
effnet_b7 = models.efficientnet_b7(weights='IMAGENET1K_V1')  # Strongest

# Vision Transformer (top choice for large-scale data tasks)
vit_b16 = models.vit_b_16(weights='IMAGENET1K_V1')    # ViT-Base/16
vit_l16 = models.vit_l_16(weights='IMAGENET1K_V1')    # ViT-Large/16

# MobileNet (mobile/embedded deployment)
mobilenet_v3 = models.mobilenet_v3_small(weights='IMAGENET1K_V1')

# ConvNeXt (modernized CNN, performance close to ViT)
convnext_t = models.convnext_tiny(weights='IMAGENET1K_V1')
convnext_b = models.convnext_base(weights='IMAGENET1K_V1')

Mainstream model performance comparison (ImageNet Top-1 Acc):

ModelTop-1 AccParametersInference speedApplicable scenarios
ResNet-1869.8%11.7MExtremely fastResource-constrained, fast prototyping
ResNet-5080.9%25.6MFastGeneral-purpose first choice
EfficientNet-B483.4%19.3MMediumAccuracy-efficiency balance
ConvNeXt-Base84.1%88.6MMediumHigh-accuracy CNN
ViT-B/1681.1%86.6MMediumLarge-data scenarios
ViT-L/1685.1%307MSlowHighest accuracy

4.2 Object Detection Models

Example

import torchvision.models.detection as detection

# Faster R-CNN (classic two-stage detector)
faster_rcnn = detection.fasterrcnn_resnet50_fpn(weights='DEFAULT')

# SSD (single-stage detector, fast)
ssd = detection.ssd300_vgg16(weights='DEFAULT')

# RetinaNet
retinanet = detection.retinanet_resnet50_fpn(weights='DEFAULT')

# FCOS
fcos = detection.fcos_resnet50_fpn(weights='DEFAULT')

# Replace Faster R-CNN's classification head (adapt to new number of classes)
from torchvision.models.detection.faster_rcnn import FastRCNNPredictor

NUM_CLASSES = 5 + 1   # 5 object classes + 1 background class

faster_rcnn = detection.fasterrcnn_resnet50_fpn(weights='DEFAULT')
in_features = faster_rcnn.roi_heads.box_predictor.cls_score.in_features
faster_rcnn.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)

4.3 Text Models (HuggingFace)

Transfer learning for NLP tasks typically uses the HuggingFace transformers library:

Example

# pip install transformers
from transformers import (
    BertForSequenceClassification,
    RobertaForSequenceClassification,
    AutoModelForSequenceClassification,
    AutoTokenizer,
)

NUM_CLASSES = 3

# BERT (top choice for text classification)
model = BertForSequenceClassification.from_pretrained(
    'bert-base-chinese',     # Chinese BERT
    num_labels=NUM_CLASSES
)
tokenizer = AutoTokenizer.from_pretrained('bert-base-chinese')

# RoBERTa (improved BERT, stronger)
model = RobertaForSequenceClassification.from_pretrained(
    'roberta-base',
    num_labels=NUM_CLASSES
)

# Generic loading method (automatically identifies model type)
model = AutoModelForSequenceClassification.from_pretrained(
    'hfl/chinese-roberta-wwm-ext',    # Chinese RoBERTa
    num_labels=NUM_CLASSES
)

5. Data Preprocessing and Augmentation

When using ImageNet pretrained weights in transfer learning, you must use the same normalization parameters as the pretraining; otherwise, feature distributions won't match and performance will drop significantly.

Example

from torchvision import transforms

# ImageNet standard normalization parameters (common to all torchvision pretrained models)
IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD  = [0.229, 0.224, 0.225]

# Training set: with data augmentation
train_transforms = transforms.Compose([
    transforms.RandomResizedCrop(224),          # Random crop and resize
    transforms.RandomHorizontalFlip(p=0.5),    # Random horizontal flip
    transforms.ColorJitter(                    # Color jitter
        brightness=0.2, contrast=0.2,
        saturation=0.2, hue=0.1
    ),
    transforms.RandomRotation(degrees=15),     # Random rotation
    transforms.ToTensor(),
    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
])

# Validation/test set: no random augmentation
val_transforms = transforms.Compose([
    transforms.Resize(256),                    # First resize to 256
    transforms.CenterCrop(224),                # Then center-crop to 224
    transforms.ToTensor(),
    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
])

# Dataset loading (directory structure: root/class_name/img.jpg)
from torchvision.datasets import ImageFolder
from torch.utils.data import DataLoader

train_dataset = ImageFolder(root='data/train', transform=train_transforms)
val_dataset   = ImageFolder(root='data/val',   transform=val_transforms)

train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True,
                          num_workers=4, pin_memory=True)
val_loader   = DataLoader(val_dataset,   batch_size=32, shuffle=False,
                          num_workers=4, pin_memory=True)

print(f"Number of training samples: {len(train_dataset)}")
print(f"Class list: {train_dataset.classes}")
print(f"Class mapping: {train_dataset.class_to_idx}")

Using torchvision's Official Recommended Preprocessing

Example

# PyTorch >= 0.13: obtain standard preprocessing directly from the weights object, no need to manually specify parameters
from torchvision.models import ResNet50_Weights

weights   = ResNet50_Weights.IMAGENET1K_V2
model     = models.resnet50(weights=weights)
preprocess = weights.transforms()   # Automatically returns the corresponding preprocessing pipeline

# preprocess already includes Resize(232), CenterCrop(224), Normalize, etc.
# For training, simply add data augmentation on top of this

You must use the ImageNet normalization parameters [0.485, 0.456, 0.406] and [0.229, 0.224, 0.225]; otherwise, the pretrained features will not align correctly.


6. Complete Hands-on: Image Binary Classification

Using cat vs. dog classification as an example, this demonstrates the complete transfer learning pipeline from data preparation to training and evaluation:

Example

import os
import torch
import torch.nn as nn
import torch.optim as optim
from torch.optim.lr_scheduler import CosineAnnealingLR
from torchvision import models, transforms, datasets
from torch.utils.data import DataLoader, random_split
from torchvision.models import EfficientNet_B0_Weights

# Configuration
DEVICE      = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
NUM_CLASSES = 2
BATCH_SIZE  = 32
EPOCHS      = 20
BASE_LR     = 1e-3
DATA_DIR    = 'data/cats_and_dogs'

# Data preparation
train_tfm = transforms.Compose([
    transforms.RandomResizedCrop(224),
    transforms.RandomHorizontalFlip(),
    transforms.ColorJitter(brightness=0.3, contrast=0.3),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406],
                         [0.229, 0.224, 0.225]),
])
val_tfm = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406],
                         [0.229, 0.224, 0.225]),
])

full_dataset = datasets.ImageFolder(DATA_DIR, transform=train_tfm)

n_val   = int(len(full_dataset) * 0.2)
n_train = len(full_dataset) - n_val
train_set, val_set = random_split(full_dataset, [n_train, n_val])
val_set.dataset = datasets.ImageFolder(DATA_DIR, transform=val_tfm)  # Replace val's transform

train_loader = DataLoader(train_set, BATCH_SIZE, shuffle=True,
                          num_workers=4, pin_memory=True)
val_loader   = DataLoader(val_set,   BATCH_SIZE, shuffle=False,
                          num_workers=4, pin_memory=True)

# Build transfer model
weights = EfficientNet_B0_Weights.IMAGENET1K_V1
model   = models.efficientnet_b0(weights=weights)

# Freeze backbone
for param in model.parameters():
    param.requires_grad = False

# Replace classification head (EfficientNet-B0 head structure)
in_features = model.classifier[1].in_features  # 1280
model.classifier = nn.Sequential(
    nn.Dropout(p=0.2, inplace=True),
    nn.Linear(in_features, NUM_CLASSES),
)

model = model.to(DEVICE)

# Optimizer: layer-wise learning rate
optimizer = optim.Adam([
    {'params': model.classifier.parameters(), 'lr': BASE_LR},
    {'params': model.features.parameters(),   'lr': BASE_LR * 0.1},
])
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)

# Training and validation functions
def train_epoch(model, loader, optimizer, criterion):
    model.train()
    total_loss, correct = 0.0, 0
    for imgs, labels in loader:
        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)
        optimizer.zero_grad()
        outputs = model(imgs)
        loss    = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        total_loss += loss.item() * imgs.size(0)
        correct    += (outputs.argmax(1) == labels).sum().item()
    n = len(loader.dataset)
    return total_loss / n, correct / n

def eval_epoch(model, loader, criterion):
    model.eval()
    total_loss, correct = 0.0, 0
    with torch.no_grad():
        for imgs, labels in loader:
            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)
            outputs = model(imgs)
            loss    = criterion(outputs, labels)
            total_loss += loss.item() * imgs.size(0)
            correct    += (outputs.argmax(1) == labels).sum().item()
    n = len(loader.dataset)
    return total_loss / n, correct / n

# Staged training
print("=== Stage 1: Train only the classification head (Epoch 1-5) ===")
best_acc = 0.0

for epoch in range(1, EPOCHS + 1):
    # Unfreeze the backbone at epoch 5, entering full fine-tuning stage
    if epoch == 6:
        print("\n=== Stage 2: Unfreeze backbone and fully fine-tune (Epoch 6-20) === ")
        for param in model.features.parameters():
            param.requires_grad = True

    train_loss, train_acc = train_epoch(model, train_loader, optimizer, criterion)
    val_loss,   val_acc   = eval_epoch(model,  val_loader,   criterion)
    scheduler.step()

    print(f"Epoch {epoch:2d}/{EPOCHS} | "
          f"Train Loss: {train_loss:.4f}, Acc: {train_acc:.4f} | "
          f"Val Loss: {val_loss:.4f}, Acc: {val_acc:.4f} | "
          f"LR: {scheduler.get_last_lr()[0]:.2e}")

    if val_acc > best_acc:
        best_acc = val_acc
        torch.save(model.state_dict(), 'best_model.pth')
        print(f" ✓ Saved best model Acc={best_acc:.4f}")

print(f"\n"Training complete, best validation accuracy: {best_acc:.4f}")

Inference and Prediction

Example

from PIL import Image

# Load the best model
model.load_state_dict(torch.load('best_model.pth', map_location=DEVICE))
model.eval()

# Predict a single image
def predict(image_path, model, class_names):
    img = Image.open(image_path).convert('RGB')
    tensor = val_tfm(img).unsqueeze(0).to(DEVICE)  # (1, 3, 224, 224)

    with torch.inference_mode():
        logits = model(tensor)
        probs  = torch.softmax(logits, dim=1)[0]
        pred   = probs.argmax().item()

    for cls, prob in zip(class_names, probs.tolist()):
        print(f"  {cls}: {prob:.4f}")
    print(f"Prediction result: {class_names[pred]}")
    return class_names[pred]

class_names = train_set.dataset.classes  # ['cat', 'dog']
predict('test_cat.jpg', model, class_names)

7. Complete Hands-on: Text Classification (BERT)

Example

# pip install transformers datasets
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from transformers import (
    BertTokenizer,
    BertForSequenceClassification,
    AdamW,
    get_linear_schedule_with_warmup,
)

DEVICE      = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
MODEL_NAME  = 'bert-base-chinese'
NUM_CLASSES = 3       # Sentiment classification: positive/neutral/negative
MAX_LEN     = 128
BATCH_SIZE  = 16
EPOCHS      = 5
LR          = 2e-5    # Recommended learning rate range for BERT fine-tuning: 1e-5 ~ 5e-5

# Custom Dataset
class SentimentDataset(Dataset):
    def __init__(self, texts, labels, tokenizer, max_len):
        self.texts     = texts
        self.labels    = labels
        self.tokenizer = tokenizer
        self.max_len   = max_len

    def __len__(self):
        return len(self.texts)

    def __getitem__(self, idx):
        encoding = self.tokenizer(
            self.texts[idx],
            max_length=self.max_len,
            padding='max_length',
            truncation=True,
            return_tensors='pt',
        )
        return {
            'input_ids':      encoding['input_ids'].squeeze(0),
            'attention_mask': encoding['attention_mask'].squeeze(0),
            'label':          torch.tensor(self.labels[idx], dtype=torch.long),
        }

# Load BERT pretrained model
tokenizer = BertTokenizer.from_pretrained(MODEL_NAME)
model     = BertForSequenceClassification.from_pretrained(
    MODEL_NAME,
    num_labels=NUM_CLASSES,
    hidden_dropout_prob=0.1,
)
model = model.to(DEVICE)

# Freeze bottom N layers (optional)
# BERT-base has 12 Transformer layers; you can freeze the first few layers to save computation
FREEZE_LAYERS = 6   # Freeze the first 6 layers

for i, layer in enumerate(model.bert.encoder.layer):
    if i < FREEZE_LAYERS:
        for param in layer.parameters():
            param.requires_grad = False

print(f"Froze the first {FREEZE_LAYERS} layers, reducing parameter updates by about {FREEZE_LAYERS/12*100:.0f}%")

# Data loading
# Example data (replace with real dataset in practice)
train_texts  = ["This movie is amazing!", "The service is terrible, disappointing", "Just so-so, nothing special"]
train_labels = [2, 0, 1]  # 0: negative 1: neutral 2: positive

train_dataset = SentimentDataset(train_texts, train_labels, tokenizer, MAX_LEN)
train_loader  = DataLoader(train_dataset, BATCH_SIZE, shuffle=True)

# Optimizer and scheduler
# BERT standard: AdamW + linear warmup
no_decay = ['bias', 'LayerNorm.weight']
optimizer_grouped = [
    {'params': [p for n, p in model.named_parameters()
                if not any(nd in n for nd in no_decay)], 'weight_decay': 0.01},
    {'params': [p for n, p in model.named_parameters()
                if     any(nd in n for nd in no_decay)], 'weight_decay': 0.0},
]
optimizer = AdamW(optimizer_grouped, lr=LR)

total_steps   = len(train_loader) * EPOCHS
warmup_steps  = int(total_steps * 0.1)   # 10% of steps used for warmup
scheduler = get_linear_schedule_with_warmup(
    optimizer,
    num_warmup_steps=warmup_steps,
    num_training_steps=total_steps,
)

# Training loop
for epoch in range(1, EPOCHS + 1):
    model.train()
    total_loss, correct = 0.0, 0

    for batch in train_loader:
        input_ids      = batch['input_ids'].to(DEVICE)
        attention_mask = batch['attention_mask'].to(DEVICE)
        labels         = batch['label'].to(DEVICE)

        optimizer.zero_grad()
        outputs = model(input_ids=input_ids,
                        attention_mask=attention_mask,
                        labels=labels)

        loss   = outputs.loss           # BERT already computes CE loss internally
        logits = outputs.logits

        loss.backward()
        # BERT fine-tuning standard: gradient clipping to prevent gradient explosion
        nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        optimizer.step()
        scheduler.step()

        total_loss += loss.item()
        correct    += (logits.argmax(1) == labels).sum().item()

    avg_loss = total_loss / len(train_loader)
    acc      = correct / len(train_dataset)
    print(f"Epoch {epoch}/{EPOCHS} | Loss: {avg_loss:.4f} | Acc: {acc:.4f}")

# Inference
def predict_sentiment(text, model, tokenizer, id2label):
    model.eval()
    encoding = tokenizer(text, max_length=MAX_LEN, padding='max_length',
                         truncation=True, return_tensors='pt')
    with torch.inference_mode():
        outputs = model(
            input_ids      = encoding['input_ids'].to(DEVICE),
            attention_mask = encoding['attention_mask'].to(DEVICE),
        )
    pred = outputs.logits.argmax(1).item()
    return id2label[pred]

id2label = {0: 'Negative', 1: 'Neutral', 2: 'Positive'}
print(predict_sentiment("The quality of this product is really good!", model, tokenizer, id2label))

8. Model Architecture Modification Tips

General Method for Replacing the Classification Head

Example

import torch.nn as nn
import torchvision.models as models

def build_transfer_model(arch, num_classes, pretrained=True, dropout=0.5):
    """
Generic transfer model building function that automatically identifies and replaces the classification head
Supported: resnet, efficientnet, densenet, mobilenet, vit, convnext
    """

    weights = 'IMAGENET1K_V1' if pretrained else None
    model   = getattr(models, arch)(weights=weights)
    name    = arch.lower()

    if 'resnet' in name or 'resnext' in name or 'inception' in name:
        in_f = model.fc.in_features
        model.fc = nn.Sequential(
            nn.Dropout(dropout),
            nn.Linear(in_f, num_classes)
        )

    elif 'efficientnet' in name or 'mobilenet' in name or 'convnext' in name:
        in_f = model.classifier[-1].in_features
        model.classifier[-1] = nn.Sequential(
            nn.Dropout(dropout),
            nn.Linear(in_f, num_classes)
        )

    elif 'densenet' in name:
        in_f = model.classifier.in_features
        model.classifier = nn.Linear(in_f, num_classes)

    elif 'vit' in name or 'swin' in name:
        in_f = model.heads.head.in_features
        model.heads.head = nn.Linear(in_f, num_classes)

    else:
        raise ValueError(f"Unsupported architecture: {arch}")

    return model


# Usage example
model_resnet  = build_transfer_model('resnet50',        num_classes=10)
model_effnet  = build_transfer_model('efficientnet_b3', num_classes=10)
model_vit     = build_transfer_model('vit_b_16',        num_classes=10)

Adding Intermediate Feature Extraction Layers

Example

import torch
import torch.nn as nn
import torchvision.models as models

class TransferWithAttention(nn.Module):
    """Add a custom attention module and classification head after the pretrained backbone"""
    def __init__(self, num_classes, dropout=0.5):
        super().__init__()
        backbone = models.resnet50(weights='IMAGENET1K_V2')

        # Remove the original classification head, keep the feature extractor
        self.backbone  = nn.Sequential(*list(backbone.children())[:-1])
        self.feat_dim  = 2048   # ResNet50 feature dimension

        # Custom attention gating
        self.attention = nn.Sequential(
            nn.Linear(self.feat_dim, 512),
            nn.Tanh(),
            nn.Linear(512, 1),
            nn.Sigmoid()
        )

        # Classification head
        self.classifier = nn.Sequential(
            nn.Dropout(dropout),
            nn.Linear(self.feat_dim, 256),
            nn.ReLU(),
            nn.Linear(256, num_classes),
        )

    def forward(self, x):
        feat   = self.backbone(x).flatten(1)          # (N, 2048)
        weight = self.attention(feat)                  # (N, 1)
        feat   = feat * weight                         # Weighted features
        return self.classifier(feat)

9. Transfer Learning Best Practices

Learning Rate Selection

Example

# Backbone (pretrained layers) learning rate
# Usually set to 1/10 ~ 1/100 of the head learning rate
backbone_lr = 1e-5   # Conservative strategy, recommended when data is scarce
head_lr     = 1e-3   # New layers start from random initialization and need a larger learning rate

# Learning rate for Transformers such as BERT / ViT
# These large models are extremely sensitive to learning rate; going out of range can damage pretrained knowledge
bert_lr = 2e-5        # Recommended range: 1e-5 ~ 5e-5

Common Issues and Solutions

Example

# Problem 1: Insufficient GPU memory
# Solution 1: Reduce batch size
# Solution 2: Use gradient checkpointing (trade time for memory)
from torch.utils.checkpoint import checkpoint_sequential
model.features = lambda x: checkpoint_sequential(model.features, 4, x)

# Solution 3: Freeze more layers (reduce backpropagation computation)

# Problem 2: High training accuracy, low validation accuracy (overfitting)
# Solution 1: Strengthen data augmentation
# Solution 2: Increase Dropout
model.classifier = nn.Sequential(nn.Dropout(0.5), nn.Linear(in_f, num_classes))
# Solution 3: Use Label Smoothing
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
# Solution 4: Freeze more backbone layers to reduce trainable parameters

# Problem 3: Unstable training, loss oscillation
# Solution 1: Reduce learning rate (try halving it)
# Solution 2: Add warmup
# Solution 3: Gradient clipping
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# Solution 4: Change optimizer (Adam → AdamW)

# Problem 4: Normalization parameters do not match
# Error: Using custom normalization doesn't match the pretrained model
transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])
# Correct: Must use ImageNet's mean and standard deviation
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])

Overall Strategy Quick Reference

ProblemRecommended Approach
Data size < 500Freeze all layers, only train the classification head
Data size 500~2000Freeze the first 2/3, fine-tune the last 1/3
Data size 2000~10000Full fine-tuning, backbone at 1e-5, head at 1e-3
Data size > 10000Full fine-tuning, consider a larger learning rate or training from scratch
Slow training speedFreeze and train for a few epochs first, then unfreeze; use AMP mixed precision
Performance bottleneckSwitch to a larger/newer backbone; try ConvNeXt / ViT
Model needs to be deployed to productionPrioritize lightweight models such as EfficientNet-B0 / MobileNetV3
Chinese NLP tasksBERT-base-chinese or chinese-roberta-wwm-ext

Recommended Complete Training Strategy

Phase 1 (Epoch 1~N/4):

  • Freeze the backbone, only train the classification head
  • Use a larger learning rate (1e-3) to quickly converge the head parameters

Phase 2 (Epoch N/4~N):

  • Unfreeze the backbone, full fine-tuning
  • Backbone uses a small learning rate (1e-5), head keeps (1e-3)
  • Combine with cosine annealing or ReduceLROnPlateau

Saving strategy:

  • Only save the checkpoint with the best validation metrics
  • Also save model / optimizer / scheduler states
Other Extensions