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老师

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐