PyTorch模型的存储与加载方式包括:仅存储和载入模型参数建议;完整地存储和载入包括结构与参数的整个模型;支持存储并加载CheckPoint文件;此外支持以. pt、. pth或. pkl为扩展名分别在CPU和GPU上进行存储与加载操作。
发布时间
阅读量:
阅读量
当提到保存和加载模型时,有三个核心功能需要熟悉:
- torch.save 函数用于将Python pickle库提供的序列化对象存储至磁盘上。
- torch.load 方法通过Python pickle模块中的unpickle功能从文件中提取已序列化的对象。
- PyTorch神经网络模块load_state_dict 方法通过加载state_dict中的键值对来更新模型参数。
一、模型保存与调用方式一:只保存模型参数
1、模型保存
model = TheModelClass(*args, **kwargs)
# ------------- 模型训练: 开始 -------------
......
# ------------- 模型训练: 结束 -------------
PATH = r'.\saved_model\model_state_dict_step_01.pt'
torch.save(model.state_dict(), PATH)
在保存模型进行推理时,只需要保存训练过的模型的学习参数即可。
一个常见的PyTorch约定是使用.pt或.pth文件扩展名保存模型。
2、模型加载
# 重构模型结构(与保存的模型结构
全部评论 (0)
还没有任何评论哟~
