TensorFlow:加载SavedModel后继续训练异常及分阶段训练咨询
一、为什么加载模型后loss回到初始状态?如何从断点继续训练?
你遇到的核心问题大概率是优化器的状态没有被正确保存或恢复,或者加载后不小心重置了变量。下面是具体的排查和解决步骤:
1. 确认保存时包含了所有必要状态
tf.saved_model.builder.SavedModelBuilder.add_meta_graph_and_variables默认会保存当前会话里的所有变量,包括模型权重、偏置,以及优化器的内部状态(比如Adam的动量项m/v、SGD的动量缓存)。但要确保调用builder.save()前,你已经完成了训练步骤,这些变量已经更新到训练后的状态——别在初始化变量后立刻保存,必须先执行训练循环。
2. 加载后绝对不要重新初始化变量
加载模型后,一定不要调用tf.global_variables_initializer().run()或任何初始化变量的操作,否则会把刚恢复的训练后变量重新重置为初始随机值,直接导致loss回到初始状态。你的加载代码里没做这件事是对的,但要确保后续代码也不会触发初始化。
3. 验证张量/操作的名字是否正确
你通过graph.get_tensor_by_name获取loss、logits等张量时,必须保证名字完全匹配。很多时候问题出在名字写错了——比如原模型中loss张量的名字可能是loss:0而非loss(TensorFlow中张量名字会带:0后缀,表示该操作输出的第一个张量)。
可以在加载后打印所有张量名字来验证:
print([tensor.name for tensor in graph.get_operations()])
找到对应loss张量的正确名字,替换args.load_loss_name即可。
4. 验证变量是否正确恢复
加载后,你可以随机选一个模型变量(比如卷积层的权重),对比保存前和加载后的值,确认是否一致:
# 保存前打印变量值 print(sess.run('conv1/kernel:0')) # 加载后打印对应变量值 print(sess.run(graph.get_tensor_by_name('conv1/kernel:0')))
如果值一致,说明变量恢复正确。此时计算loss时要喂入和保存前相同的输入数据,别直接调用loss.eval()(无输入的话可能得到默认初始计算值,会误导你)。
二、能否分两个会话分别用原始数据和增强数据训练?
完全可以!这种分阶段训练的场景非常常见,具体实现步骤如下:
会话1:用原始数据训练并保存模型
# 构建模型、定义loss和优化器 input_image = tf.placeholder(tf.float32, shape=[None, 224, 224, 3], name='net_input') logits = your_model(input_image) # 替换为你的模型定义 loss = tf.losses.softmax_cross_entropy(labels, logits, name='loss') # 替换为你的loss计算 opt = tf.train.AdamOptimizer().minimize(loss, name='optimizer') # 初始化变量并开始训练 with tf.Session() as sess: sess.run(tf.global_variables_initializer()) # 用原始数据训练N轮 for epoch in range(10): train_on_original_data(sess, input_image, loss, opt) # 替换为你的原始数据训练逻辑 # 保存SavedModel,包含模型参数和优化器状态 builder = tf.saved_model.builder.SavedModelBuilder('./saved_model') tensor_info_input = tf.saved_model.utils.build_tensor_info(input_image) tensor_info_logits = tf.saved_model.utils.build_tensor_info(logits) prediction_signature = tf.saved_model.signature_def_utils.build_signature_def( inputs={'net_input': tensor_info_input}, outputs={'logits': tensor_info_logits}, method_name=tf.saved_model.signature_constants.PREDICT_METHOD_NAME ) builder.add_meta_graph_and_variables( sess, [tf.saved_model.tag_constants.SERVING], signature_def_map={ tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: prediction_signature } ) builder.save()
会话2:加载模型并用增强数据继续训练
import tensorflow as tf with tf.Session() as sess: # 加载SavedModel,恢复所有变量(包括优化器状态) meta_graph_def = tf.saved_model.loader.load( sess, [tf.saved_model.tag_constants.SERVING], './saved_model' ) graph = tf.get_default_graph() # 获取模型的输入、loss、优化器(确保名字和会话1中一致) net_input = graph.get_tensor_by_name('net_input:0') loss = graph.get_tensor_by_name('loss:0') opt = graph.get_operation_by_name('optimizer') # 用增强数据继续训练 for epoch in range(10): train_on_augmented_data(sess, net_input, loss, opt) # 替换为你的增强数据训练逻辑 # 可选:再次保存训练后的模型 builder = tf.saved_model.builder.SavedModelBuilder('./saved_model_augmented') # 重复会话1的保存步骤即可
这里的关键是:加载后的优化器状态是会话1结束时的状态,所以继续训练时会接着之前的梯度更新逻辑推进,完全相当于从断点处续训,而非从头开始。
内容的提问来源于stack exchange,提问作者jaycode

