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 amount | Similarity to source task | Recommended strategy |
|---|---|---|
| Small (< 1000) | High | Only replace the final classification head, freeze the entire backbone. |
| Small (< 1000) | Low | Fine-tune shallower layers, freeze deeper layers. |
| Medium (1000~10000) | High | Fine-tune the entire network with a small learning rate. |
| Medium (1000~10000) | Low | Fine-tune deeper layers, freeze shallower layers. |
| Large (> 10000) | Any | Fine-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 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(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:
| Model | Classification head attribute |
|---|---|
| ResNet / RegNet | model.fc |
| VGG / AlexNet | model.classifier[-1] |
| DenseNet | model.classifier |
| EfficientNet | model.classifier[-1] |
| MobileNetV2/V3 | model.classifier[-1] |
| ViT (Vision Transformer) | model.heads.head |
| ConvNeXt | model.classifier[-1] |
| Inception V3 | model.fc |
| Swin Transformer | model.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.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.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.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
# 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):
| Model | Top-1 Acc | Parameters | Inference speed | Applicable scenarios |
|---|---|---|---|---|
| ResNet-18 | 69.8% | 11.7M | Extremely fast | Resource-constrained, fast prototyping |
| ResNet-50 | 80.9% | 25.6M | Fast | General-purpose first choice |
| EfficientNet-B4 | 83.4% | 19.3M | Medium | Accuracy-efficiency balance |
| ConvNeXt-Base | 84.1% | 88.6M | Medium | High-accuracy CNN |
| ViT-B/16 | 81.1% | 86.6M | Medium | Large-data scenarios |
| ViT-L/16 | 85.1% | 307M | Slow | Highest accuracy |
4.2 Object Detection Models
Example
# 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
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
# 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
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 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
# 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
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 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.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
# 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
# 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
| Problem | Recommended Approach |
|---|---|
| Data size < 500 | Freeze all layers, only train the classification head |
| Data size 500~2000 | Freeze the first 2/3, fine-tune the last 1/3 |
| Data size 2000~10000 | Full fine-tuning, backbone at 1e-5, head at 1e-3 |
| Data size > 10000 | Full fine-tuning, consider a larger learning rate or training from scratch |
| Slow training speed | Freeze and train for a few epochs first, then unfreeze; use AMP mixed precision |
| Performance bottleneck | Switch to a larger/newer backbone; try ConvNeXt / ViT |
| Model needs to be deployed to production | Prioritize lightweight models such as EfficientNet-B0 / MobileNetV3 |
| Chinese NLP tasks | BERT-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