迁移学习的核心思想,是把一个已经在大规模数据上学到的视觉表示,迁移到新的任务中继续使用。对于数据量有限的项目,如果从随机初始化开始训练,模型往往需要更多样本和更长时间才能学到边缘、纹理、形状等通用特征;而预训练模型已经完成了其中相当一部分工作。
但“加载预训练权重并训练几轮”并不等于完整的迁移学习方案。真正影响结果的,通常是权重与数据预处理是否匹配、冻结范围是否合理、分类头是否正确替换,以及验证过程能否识别过拟合。下面以一个多分类图像任务为例,构建一个可执行、可检查的训练流程。
一、问题背景
假设项目需要识别若干种物体,每个类别只有几百张甚至几十张图片。数据规模不足时,直接训练完整卷积网络会面临三个问题:
- 模型参数量相对于样本量过大,容易记住训练集而不是学习类别规律。
- 训练初期的梯度同时承担“学习通用视觉特征”和“适应具体任务”两项工作,优化成本较高。
- 如果类别分布不均衡,仅看总体准确率,可能掩盖少数类别几乎无法识别的问题。
迁移学习将模型拆成两部分理解:前面的特征提取器负责把图像转换成高层表示,最后的分类头负责把表示映射为目标类别。预训练阶段学到的特征提取器可以复用,而分类头通常必须按照新任务的类别数重新构造。
这里的“预训练”只表示模型参数来自另一个训练任务,并不保证它一定适合当前数据。若源数据和目标数据差异很大,预训练特征的可迁移性会下降,此时需要解冻更多层进行微调,或者重新评估模型结构。
二、两种训练策略
1. 冻结特征提取器
冻结卷积骨干网络,只训练新建的分类头。它的优点是需要更新的参数少、训练速度较快,也更适合样本很少的场景。缺点是骨干网络不能适应目标领域,目标图像与预训练数据差异较大时,表达能力可能不足。
冻结参数的关键操作是将 requires_grad 设为 False。优化器只接收需要更新的参数,否则虽然不会更新冻结参数,但会增加不必要的计算和代码歧义。
2. 部分或全部微调
先训练分类头,使输出层进入合理状态,再解冻骨干网络的后几层,以较小学习率继续训练。后部网络通常包含更任务相关的语义特征,先解冻这些层可以在适应新任务与保留通用能力之间取得平衡。
微调时常用分层学习率:分类头使用较大的学习率,骨干网络使用较小的学习率。学习率没有脱离数据和模型的固定答案,下面的配置只是起点,实际值应通过验证集和训练曲线调整。
三、准备数据集
推荐采用按类别分目录的结构,便于使用 ImageFolder:
dataset/
train/
cat/
dog/
bird/
val/
cat/
dog/
bird/
训练集可以使用随机裁剪、水平翻转等增强;验证集应尽量只做确定性的缩放和裁剪。训练和验证必须使用相同的归一化参数,但不能把随机增强带入验证,否则评估结果会产生额外波动。
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
image_size = 224
mean = [0.485, 0.456, 0.406]
std = [0.229, 0.224, 0.225]
train_transform = transforms.Compose([
transforms.RandomResizedCrop(image_size, scale=(0.7, 1.0)),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean, std),
])
val_transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(image_size),
transforms.ToTensor(),
transforms.Normalize(mean, std),
])
train_set = datasets.ImageFolder("dataset/train", transform=train_transform)
val_set = datasets.ImageFolder("dataset/val", transform=val_transform)
train_loader = DataLoader(train_set, batch_size=32, shuffle=True, num_workers=2)
val_loader = DataLoader(val_set, batch_size=32, shuffle=False, num_workers=2)
num_classes = len(train_set.classes)
print(train_set.classes)
如果使用的预训练模型规定了不同的输入尺寸或归一化参数,应以该模型的公开配置为准。不能只因为图片能够进入网络,就认为预处理已经正确。
四、替换分类头并训练
下面使用 torchvision 提供的 ResNet 作为示例。模型名称和权重接口可能随安装版本不同而变化,因此实际运行时应以当前环境中的官方文档和接口为准。代码中的 weights 是模型权重配置对象,不是密钥。
import copy
import torch
from torch import nn
from torchvision.models import resnet18, ResNet18_Weights
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = resnet18(weights=ResNet18_Weights.DEFAULT)
# 第一阶段:冻结骨干网络
for parameter in model.parameters():
parameter.requires_grad = False
in_features = model.fc.in_features
model.fc = nn.Linear(in_features, num_classes)
model = model.to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.AdamW(
model.fc.parameters(), lr=1e-3, weight_decay=1e-4
)
训练和验证函数应明确切换 train() 与 eval()。这一步对包含 BatchNorm 或 Dropout 的模型尤其重要:训练模式会更新部分统计量并启用随机失活,验证模式则用于稳定评估。
def run_epoch(model, loader, criterion, optimizer=None):
is_train = optimizer is not None
model.train(is_train)
total_loss = 0.0
total_correct = 0
total_count = 0
for images, labels in loader:
images = images.to(device)
labels = labels.to(device)
if is_train:
optimizer.zero_grad(set_to_none=True)
with torch.set_grad_enabled(is_train):
logits = model(images)
loss = criterion(logits, labels)
if is_train:
loss.backward()
optimizer.step()
total_loss += loss.item() * labels.size(0)
total_correct += (logits.argmax(dim=1) == labels).sum().item()
total_count += labels.size(0)
return total_loss / total_count, total_correct / total_count
best_state = None
best_val_acc = -1.0
for epoch in range(8):
train_loss, train_acc = run_epoch(
model, train_loader, criterion, optimizer
)
with torch.no_grad():
val_loss, val_acc = run_epoch(model, val_loader, criterion)
print(
f"epoch={epoch + 1} "
f"train_loss={train_loss:.4f} train_acc={train_acc:.4f} "
f"val_loss={val_loss:.4f} val_acc={val_acc:.4f}"
)
if val_acc > best_val_acc:
best_val_acc = val_acc
best_state = copy.deepcopy(model.state_dict())
if best_state is not None:
model.load_state_dict(best_state)
torch.save(
{
"model": model.state_dict(), "classes": train_set.classes},
"classifier.pt",
)
这里保存验证集表现最好的权重,而不是无条件保存最后一个 epoch。验证集本身并不是完全独立的最终测试集,因此项目正式评估时还应保留一份没有参与调参的测试数据。
五、从冻结到微调
如果训练分类头后,训练集和验证集都表现不佳,可能是模型容量、数据质量或预处理存在问题;如果训练集持续提升而验证集停滞甚至下降,则需要优先检查过拟合。确认数据和标签没有明显问题后,可以解冻骨干网络的后几层。
# 解冻 layer4,分类头继续保持可训练
for parameter in model.layer4.parameters():
parameter.requires_grad = True
optimizer = torch.optim.AdamW([
{
"params": model.layer4.parameters(), "lr": 1e-5},
{
"params": model.fc.parameters(), "lr": 1e-4},
], weight_decay=1e-4)
for epoch in range(5):
train_loss, train_acc = run_epoch(
model, train_loader, criterion, optimizer
)
with torch.no_grad():
val_loss, val_acc = run_epoch(model, val_loader, criterion)
print(epoch + 1, train_loss, train_acc, val_loss, val_acc)
微调时需要注意三点。第一,解冻后应重新确认优化器参数确实包含新解冻的层。第二,学习率通常应小于只训练分类头时的学习率。第三,样本很少时,BatchNorm 的统计量可能不稳定,是否冻结其行为要结合模型结构、批大小和数据分布判断,不能一概而论。
六、推理与类别映射
推理阶段必须复用验证阶段的确定性预处理,并按照训练集生成的类别顺序解释输出。只保存权重而不保存类别映射,部署后可能出现数值正确但类别名称错位的问题。
from PIL import Image
checkpoint = torch.load("classifier.pt", map_location=device)
classes = checkpoint["classes"]
model.load_state_dict(checkpoint["model"])
model.eval()
image = Image.open("sample.jpg").convert("RGB")
input_tensor = val_transform(image).unsqueeze(0).to(device)
with torch.no_grad():
probability = model(input_tensor).softmax(dim=1)[0]
index = probability.argmax().item()
print({
"label": classes[index], "score": float(probability[index])})
输出的概率是模型在当前类别集合和训练分布下的相对置信度,不应直接当成真实正确率。对于安全、医疗或其他高风险场景,还需要额外的校准、拒识策略和人工复核机制。
七、常见问题
预训练模型为什么仍然识别错误?
迁移学习复用的是参数,不是目标领域的全部知识。拍摄角度、光照、分辨率、背景和类别定义发生较大变化时,通用特征可能不足。应先检查标签质量和类别边界,再考虑更合适的数据增强、解冻范围或模型结构。
训练准确率很高,验证准确率很低怎么办?
这通常符合过拟合特征,但也可能是训练集与验证集存在分布差异或数据泄漏。应检查同一对象的近似图片是否被拆到不同集合,减少不必要的增强强度,增加有效样本,并使用早停或正则化。不要只根据一次验证结果调整模型。
类别数量不平衡时只改 batch size 可以吗?
通常不够。可以根据训练集类别频率设置交叉熵权重,或使用合适的采样策略;同时报告每类召回率、精确率和混淆矩阵。具体采用哪种方法,需要结合少数类样本量和业务评价目标决定。
什么时候不适合冻结大部分层?
当目标数据与预训练数据差异显著,或者新任务依赖预训练模型未覆盖的细粒度特征时,完全冻结可能限制效果。此时可以逐步解冻后部网络,但应控制学习率,并使用独立测试集验证收益。
总结
迁移学习不是单一的模型调用,而是一套从数据、权重、训练策略到评估和部署的完整流程。可靠的实践至少应做到:确认预训练权重与输入预处理匹配;先训练新分类头,再根据验证表现决定是否微调;保存最佳权重和类别映射;用独立数据评估泛化能力;通过每类指标和混淆矩阵定位问题。
对于小样本视觉任务,冻结特征提取器往往是合理起点,但不是最终答案。只有把训练曲线、数据分布和错误样本结合起来,才能判断问题究竟来自模型、数据还是训练策略。