pytorch模型存储的2种实现方法
更新时间:2020年02月14日 12:00:56 作者:慢行厚积
今天小编就为大家分享一篇pytorch模型存储的2种实现方法,具有很好的参考价值,希望对大家有所帮助。一起跟随小编过来看看吧
1、保存整个网络结构信息和模型参数信息:
torch.save(model_object, './model.pth')
直接加载即可使用:
model = torch.load('./model.pth')
2、只保存网络的模型参数-推荐使用
torch.save(model_object.state_dict(), './params.pth')
加载则要先从本地网络模块导入网络,然后再加载参数:
from models import AgeModel model = AgeModel() model.load_state_dict(torch.load('./params.pth'))
以上这篇pytorch模型存储的2种实现方法就是小编分享给大家的全部内容了,希望能给大家一个参考,也希望大家多多支持脚本之家。
相关文章
浅谈Python中的可迭代对象、迭代器、For循环工作机制、生成器
这篇文章主要介绍了Python中的可迭代对象、迭代器、For循环工作机制、生成器,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有一定的参考学习价值,需要的朋友们下面随着小编来一起学习学习吧2019-03-03
最新评论