You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何保存和加载已训练完成的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.25 00:36:00