如何保存和加载已训练完成的CNN(卷积神经网络)模型
CNN模型保存与加载操作指南
不同主流深度学习框架的CNN模型保存加载逻辑可分为存完整模型、仅存权重参数两类,后者体积小、兼容性强,是官方推荐的通用用法,以下是两大常用框架的具体操作方式:
PyTorch 操作方法
推荐:仅保存/加载权重
仅存储模型训练得到的参数,不保存网络结构,加载时需要先初始化和训练时结构完全一致的模型实例。
- 保存代码:
torch.save(model.state_dict(), "trained_cnn_weights.pth")
- 加载代码:
# 先初始化和训练时结构完全相同的模型实例 model = YourCustomCNNClass() # 加载权重参数 model.load_state_dict(torch.load("trained_cnn_weights.pth")) # 切换到评估模式,保证Dropout、BatchNorm等层推理行为正确 model.eval()
可选:保存/加载完整模型
将网络结构和参数打包存储,仅适合同环境临时快速使用,兼容性较差,框架版本、模型定义路径变动都可能导致加载失败。
- 保存代码:
torch.save(model, "trained_cnn_full.pth") - 加载代码:
model = torch.load("trained_cnn_full.pth")
TensorFlow/Keras 操作方法
推荐:仅保存/加载权重
- 保存代码:
model.save_weights("trained_cnn_weights.h5") - 加载代码:
# 先初始化和训练时结构完全相同的模型实例 model = YourCustomCNNClass() # 加载权重参数 model.load_weights("trained_cnn_weights.h5")
可选:保存/加载完整模型
TF2+ 更推荐使用SavedModel格式存储完整模型,兼容性比H5格式更好。
- 保存为SavedModel格式:
model.save("trained_cnn_savedmodel") - 加载SavedModel格式模型:
model = tf.keras.models.load_model("trained_cnn_savedmodel")
注意:如果模型包含自定义层、自定义损失函数,加载时需要额外传入
custom_objects参数指定对应的自定义类/函数,否则会报错。
通用注意事项
- 加载权重前必须保证模型结构和训练时完全一致,包括层的顺序、输出维度、自定义层的实现逻辑,否则会出现权重形状不匹配的报错
- 推理前必须确认模型处于评估模式,避免Dropout、BatchNormalization等训练/推理行为不同的层输出异常结果
- 如果需要跨框架、跨设备部署模型,建议转换为ONNX通用格式,适配不同推理引擎
- 建议同时记录模型的训练超参数、验证集指标、训练数据版本等信息,避免不同迭代版本的模型混淆
内容的提问来源于stack exchange,提问作者sunita patel
相关产品推荐
相关产品推荐

