pytorch模型存储和恢复
发布时间
阅读量:
阅读量
关于模型或权重参数保存格式的后缀问题:
在PyTorch框架中,模型数据通常采用.t7、.pth或.pkl等格式进行存储,其中.t7格式继承自torch7的模型权重读取方式,而.pth则是Python环境下常用的文件存储形式。相比之下,在Keras框架中,模型参数则多以.h5格式进行保存。
来源:<>
两种实现方式如下:
(1)保存模型参数
#保存
torch.save( model.state_dict(), path)
#加载
the_model = CNN()
the_model.load_state_dict(torch.load(path))
在模型加载过程中,需要在代码中重新构建CNN的结构,之后才能将已保存的参数(w和b)导入模型中,用于后续的训练过程。
以下为采用上述方式加载已存储模型的具体代码实现,该代码可对任意一张图片进行识别(所使用的数据集为手写体数据集),其中CNN的结构参考了莫烦提供的架构。
from PIL import Image
import torch.nn as nn
import torch
import numpy as np
d
全部评论 (0)
还没有任何评论哟~
