如何获取tf.keras.models.save_model()保存模型的源代码?
tf.keras.models.save_model() 默认导出的是SavedModel格式,这个格式只会存储模型的计算图拓扑、权重参数、训练时的编译配置(损失函数、优化器状态这类),本身就不包含你编写模型的原始Python源代码,加载后看不到源码是正常情况,不是加载出错。
要基于这种已保存的模型做调整开发,按你手里有没有原始源码分情况处理即可。
方案1:保留有原始模型源码(最推荐,改起来最灵活)
- 直接在原始源码上修改结构、调整逻辑即可,不需要从头重新训练,之前训好的权重可以直接复用
- 操作步骤:
- 先完成新模型的代码编写,注意需要复用旧权重的层,层名、输入输出维度必须和旧模型对应层完全一致
- 调用权重加载接口把旧模型的权重读入新结构,仅新增、修改的部分需要后续微调
- 加载权重时加
by_name=True、skip_mismatch=True参数,可以自动跳过名字不匹配、结构不匹配的层,不会因为你改了部分结构直接报错
- 参考代码:
import tensorflow as tf from tensorflow.keras import layers # 这里是你修改后的新模型代码,比如把原来的10分类输出改成5分类 new_model = tf.keras.Sequential([ layers.Conv2D(32, 3, activation='relu', name='conv_1'), # 和旧模型对应层层名保持一致 layers.MaxPooling2D(2), layers.Flatten(), layers.Dense(64, activation='relu', name='dense_1'), # 和旧模型对应层层名保持一致 layers.Dense(5, activation='softmax', name='new_output') # 新修改的输出层 ]) # 加载旧模型权重,匹配的层自动载入,不匹配的新层跳过 new_model.load_weights( './your_old_saved_model_dir', by_name=True, skip_mismatch=True ) # 加载完成就可以正常做微调、推理,不需要从头训练
方案2:已经丢失原始源码,仅保留保存的模型文件夹
- 这种情况没法直接还原你当初写的原始源码(你写的注释、自定义层的非计算逻辑、业务相关代码根本没存在模型文件里,别找工具瞎试,找不回来的),但可以基于加载后的模型对象做二次开发
- 轻量修改的操作方式:
- 正常用
tf.keras.models.load_model()加载旧模型,把整个旧模型当成一个可复用的模块嵌入新结构 - 按需求冻结旧模型的全部/部分层,在旧模型的输入输出端拼接新的层结构,完成你要的逻辑修改
- 正常用
- 参考代码:
old_model = tf.keras.models.load_model('./your_old_saved_model_dir') old_model.trainable = False # 先冻结旧模型权重,不想微调就不用改 # 搭建新模型 inputs = tf.keras.Input(shape=old_model.input_shape[1:]) x = old_model(inputs, training=False) # 把加载的旧模型当成单独一层调用 # 加你需要的新结构,比如加正则层、换输出头 x = layers.Dropout(0.2)(x) new_output = layers.Dense(5, activation='softmax')(x) new_model = tf.keras.Model(inputs, new_output) # 编译后就可以正常训练、推理
- 如果需要大改模型内部结构,可以先调用
old_model.summary()查看每一层的维度、参数、连接关系,手动反推写出结构代码,再用方案1的方式加载权重,就是工作量稍大。
后续存模型建议多做一步:除了存SavedModel格式用于部署,单独用
model.save_weights()存一份权重文件,同时备份好模型结构的源码文件,这是最稳妥的方式,不会出现想改代码找不到源文件的问题。
内容的提问来源于stack exchange,提问作者EduardoriosChicago
相关产品推荐
相关产品推荐

