PyTorch加载DeiT预训练权重做迁移学习的梯度报错解决咨询
使用DEIT模型迁移学习冻结参数报错的解决方案
问题原因分析
你遇到的RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn错误,本质是反向传播时计算图中没有可训练的参数——要么所有参数都被冻结,要么新添加的分类头参数未被正确纳入优化器。虽然你替换了新的分类头,但优化器参数范围设置不当,导致没有可更新的参数触发梯度计算。
正确的实现步骤
1. 精准冻结旧层,确保新分类头可训练
冻结所有预训练层后,新添加的heads层参数默认是requires_grad=True,但要确保优化器只针对新层参数更新,避免不必要的参数遍历:
import torch import torch.nn as nn from torch.hub import load device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 加载预训练DEIT模型 pretrained_vit = load('facebookresearch/deit:main', 'deit_tiny_patch16_224', pretrained=True).to(device) # 冻结所有预训练层参数 for param in pretrained_vit.parameters(): param.requires_grad = False # 替换分类头,适配自定义数据集类别数 num_classes = len(class_names) pretrained_vit.heads = nn.Linear(in_features=192, out_features=num_classes).to(device) # 仅优化新分类头的参数(关键) optimizer = torch.optim.Adam(params=pretrained_vit.heads.parameters(), lr=1e-3) loss_fn = nn.CrossEntropyLoss() # 启动训练前确保模型处于train模式 pretrained_vit.train() results = engine.train(model=pretrained_vit, ...)
2. 验证新层参数的可训练状态
可以添加以下代码确认新分类头的参数确实处于可训练状态:
for name, param in pretrained_vit.named_parameters(): if 'heads' in name: print(f"{name}: requires_grad={param.requires_grad}")
正常输出应为heads.weight: requires_grad=True和heads.bias: requires_grad=True。
3. 针对小数据集的优化建议
由于你的数据集规模极小,仅训练分类头可能效果有限,可以尝试微调少量顶层Transformer层,同时分层设置学习率(预训练层用小学习率,新层用大学习率),平衡泛化性和适配性:
# 先冻结所有层 for param in pretrained_vit.parameters(): param.requires_grad = False # 解冻最后3层Transformer层(DEIT-Tiny有12层,取最后10组参数对应顶层2-3层,可根据效果调整) for param in list(pretrained_vit.parameters())[-10:]: param.requires_grad = True # 替换分类头 pretrained_vit.heads = nn.Linear(192, num_classes).to(device) # 分层设置学习率 optimizer = torch.optim.Adam([ {'params': pretrained_vit.heads.parameters(), 'lr': 1e-3}, {'params': [p for p in pretrained_vit.parameters() if p.requires_grad and p not in pretrained_vit.heads.parameters()], 'lr': 1e-5} ])
预训练权重加载的正确性确认
你使用torch.hub.load('facebookresearch/deit:main', 'deit_tiny_patch16_224', pretrained=True)加载权重是正确的,该方式会自动下载并加载官方预训练权重,无需额外处理。
内容的提问来源于stack exchange,提问作者Jose Ramon
相关产品推荐
相关产品推荐

