💻 策略一:特征提取
import torch
import torch.nn as nn
import torchvision.models as models
model = models.resnet50(pretrained=True)
for param in model.parameters():
param.requires_grad = False
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 10)
optimizer = torch.optim.Adam(model.fc.parameters(), lr=0.001)
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f'可训练参数: {trainable_params}')
🔧 策略二:微调训练
model = models.resnet50(pretrained=True)
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 10)
params = [
{'params': model.conv1.parameters(), 'lr': 1e-5},
{'params': model.layer1.parameters(), 'lr': 1e-4},
{'params': model.layer2.parameters(), 'lr': 1e-4},
{'params': model.layer3.parameters(), 'lr': 1e-3},
{'params': model.layer4.parameters(), 'lr': 1e-3},
{'params': model.fc.parameters(), 'lr': 1e-2},
]
optimizer = torch.optim.Adam(params)
optimizer = torch.optim.Adam([
{'params': model.fc.parameters(), 'lr': 0.01},
{'params': [p for n, p in model.named_parameters() if 'fc' not in n], 'lr': 0.001}
])
🌐 完整迁移学习示例
import torch
import torch.nn as nn
import torchvision.models as models
import torchvision.transforms as transforms
from torchvision.datasets import ImageFolder
from torch.utils.data import DataLoader
transform = {
'train': transforms.Compose([
transforms.RandomResizedCrop(224),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
]),
'val': transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
}
train_dataset = ImageFolder('data/train', transform=transform['train'])
val_dataset = ImageFolder('data/val', transform=transform['val'])
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=32)
model = models.resnet50(pretrained=True)
num_classes = len(train_dataset.classes)
model.fc = nn.Linear(model.fc.in_features, num_classes)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = model.to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)