You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 23:32:32