【PyTorch】模型保存和加载
·
1. 模型保存
1.1 保存模型的完整结构和参数
使用 torch.save() 可以保存模型的完整结构和参数。这种方法保存的是模型的 state_dict 和模型的类定义。
import torch
import torch.nn as nn
# 定义一个模型
class SimpleModel(nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.fc = nn.Linear(10, 1)
def forward(self, x):
return self.fc(x)
# 创建模型实例
model = SimpleModel()
# 保存模型
torch.save(model, 'model.pth')
1.2 仅保存模型的参数
通常情况下,我们只需要保存模型的参数,即 state_dict。这种方法更加灵活,因为只需要保存和加载参数,而不需要保存模型的类定义。
# 保存模型的 state_dict
torch.save(model.state_dict(), 'model_state_dict.pth')
2. 模型加载
2.1 加载完整模型
如果保存了完整模型,可以直接使用 torch.load() 加载模型。
# 加载完整模型
model = torch.load('model.pth')
2.2 加载模型的参数
如果只保存了模型的参数,需要先定义模型的结构,然后使用 load_state_dict() 加载参数。
# 定义模型结构
class SimpleModel(nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.fc = nn.Linear(10, 1)
def forward(self, x):
return self.fc(x)
# 创建模型实例
model = SimpleModel()
# 加载模型的 state_dict
model.load_state_dict(torch.load('model_state_dict.pth'))
3. 注意事项
3.1 加载模型时的设备
在加载模型时,需要确保模型和数据在相同的设备上。如果保存时使用了 GPU,加载时也需要使用 GPU;如果保存时使用了 CPU,加载时也需要使用 CPU。
# 加载模型时指定设备
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.load_state_dict(torch.load('model_state_dict.pth', map_location=device))
3.2 模型结构的一致性
在加载模型参数时,必须确保模型的结构与保存时的结构完全一致。如果结构不一致,会报错。
3.3 保存和加载时的路径
确保保存和加载时的路径正确。如果路径不正确,会导致文件找不到的错误。
4. 总结
- 使用
torch.save()保存模型的完整结构和参数。 - 使用
torch.save()保存模型的参数(state_dict)。 - 使用
torch.load()加载完整模型。 - 使用
torch.nn.Module.load_state_dict()加载模型的参数。 - 注意加载时的设备和模型结构的一致性。
声明:
本文内容仅用于个人学习记录,不用于任何商业用途。部分代码、技术观点或示例可能来源于网络或其他公开资源,如有侵权,请联系我删除。
参考资料:kimi老师
更多推荐



所有评论(0)