TensorFlow如何保存带自定义层的不可序列化子类模型
问题解答
Checkpoint保存方式和可序列化模型的完整保存效果并不完全等效,二者有明确的适用场景差异:
核心差异
- 存储内容不同:Checkpoint仅存储模型、优化器等对象的状态数值(比如层的权重、Adam优化器的动量变量、训练步数等),不会保存任何模型计算逻辑和结构定义。恢复时必须依赖完全一致的Python源码定义的子类模型、优化器实例才能正常加载,缺失源码的场景下无法使用。而可序列化的Sequential/Functional模型保存时会同时导出计算图结构、层配置、权重参数,无需依赖原始模型定义代码即可直接加载、执行推理。
- 适用场景不同:
- Checkpoint专门适配训练阶段的断点续训需求,支持保存多份历史版本,训练中断后可以直接恢复到中断时的训练状态继续训练,不需要从头启动训练流程。
- 完整的序列化模型(如TensorFlow SavedModel格式)更适用于训练完成后的模型交付、部署、跨环境分享场景,支持在无Python源码的环境下直接运行推理,也可对接各种部署框架转换格式。
子类模型的额外保存方案
如果你使用自定义子类模型,训练完成后需要脱离源码部署,可以在训练结束后通过tf.saved_model.save()直接导出SavedModel格式,导出前可以指定input_signature固定模型输入形状,确保导出的模型能被正常追踪加载。
另外你最初的保存代码存在参数矛盾:model.save("model.h5", save_format='tf')同时指定了HDF5后缀和TF保存格式,属于错误写法,如需直接保存TF格式模型,修改为model.save("saved_model_dir", save_format='tf')即可。
内容的提问来源于stack exchange,提问作者OK 400
相关产品推荐
相关产品推荐

