PyTorch加载预训练参数
发布时间
阅读量:
阅读量
1.在没有改变原网络结构的情况下
model = resnet50()
model = model.load_state_dict(torch.load('model.pth'))
torch.load('model.pth')负责将网络参数存储在一个有序字典中;随后将该有序字典传递给model.load_state_dict()。
2.在只改变原网络结构的某一层或某几层的情况下
面对不同的实际需求时,在深度学习模型的设计过程中往往需要相应地调整网络架构。例如ResNet50原本设计用于支持1,000个分类任务的应用场景,在本项目中由于任务仅需识别1个子类别的图像(即针对当前的任务仅需识别10个类别),则需要相应地调整其全连接层配置。其余层次则无需改动即可维持原有功能状态。此时我们可以利用预训练模型的权重以加快训练速度并提升模型性能
方法1:
# 将原始网络结构与修改后网络结构相同的键值对放到一个有序字典当中,不相同的键值对则被删除
pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model.state_dict()}
# 将这个有序字典传给load_state_dict
mod
全部评论 (0)
还没有任何评论哟~
