TensorFlow多GPU异步训练模型恢复后推理迭代器初始化失败求助
我来帮你分析下这个问题——你在单GPU下能正常恢复模型推理,但多GPU异步训练的模型就报迭代器未初始化的错误,核心问题大概率出在多GPU训练时用的MonitoredTrainingSession和恢复时用的普通tf.Session不兼容,或者迭代器的设备上下文、元图保存环节有遗漏。下面给你几个针对性的解决思路:
1. 确保正确获取加载的元图对象
你当前的恢复代码里,graph变量没有明确定义,单GPU时可能刚好默认图就是加载的模型图,但多GPU环境下默认图可能混杂了其他设备操作,导致找不到正确的迭代器操作。修改恢复代码,直接从saver获取加载的元图:
saver = tf.train.import_meta_graph('expr1.multi/train_logs/model.ckpt-44.meta') graph = saver.graph # 明确获取加载的模型图 sess = tf.Session(config=tf.ConfigProto(allow_soft_placement=True)) saver.restore(sess,'expr1.multi/train_logs/model.ckpt-44') # 后续获取张量和操作都基于这个graph logits = graph.get_tensor_by_name('strided_slice_1:0') logits_len = graph.get_tensor_by_name('strided_slice_2:0') targets = graph.get_tensor_by_name('evaluate/IteratorGetNext:2') targets_len = graph.get_tensor_by_name('evaluate/IteratorGetNext:3') init_op = graph.get_operation_by_name('evaluate/MakeIterator')
2. 匹配训练时的会话类型(改用MonitoredSession恢复)
你训练时用的是MonitoredTrainingSession,它会自动管理检查点、设备上下文和图的状态,普通tf.Session恢复时可能无法正确加载多GPU训练时的迭代器状态。改用MonitoredSession来恢复,兼容性更好:
from tensorflow.train import MonitoredSession, LatestCheckpoint checkpoint_dir = 'expr1.multi/train_logs' # 自动找到最新的检查点 latest_ckpt = LatestCheckpoint(checkpoint_dir) # 用MonitoredSession加载模型 sess = MonitoredSession(checkpoint_dir=checkpoint_dir) graph = tf.get_default_graph() # 获取所需张量和初始化操作 logits = graph.get_tensor_by_name('strided_slice_1:0') logits_len = graph.get_tensor_by_name('strided_slice_2:0') targets = graph.get_tensor_by_name('evaluate/IteratorGetNext:2') targets_len = graph.get_tensor_by_name('evaluate/IteratorGetNext:3') init_op = graph.get_operation_by_name('evaluate/MakeIterator') # 执行迭代器初始化(注意多GPU下可能需要指定设备) with tf.device("/cpu:0"): # 如果评估迭代器是在CPU上创建的,这里对应指定 sess.run(init_op) # 执行推理循环 while True: try: l, ll, t, tl = sess.run([logits, logits_len, targets, targets_len]) # 处理你的推理结果 except tf.errors.OutOfRangeError: break
3. 检查迭代器的设备上下文
多GPU训练时,评估迭代器可能被绑定到了某个特定GPU或者CPU上,恢复时如果设备上下文不匹配,执行init_op会导致实际没有初始化对应设备上的迭代器。你可以在训练时打印init_op的设备信息:
# 训练时添加:打印评估迭代器初始化操作的设备 init_op = graph.get_operation_by_name('evaluate/MakeIterator') print(f"Eval init op device: {init_op.device}")
然后在恢复时,用tf.device指定相同的设备来执行初始化,比如训练时显示设备是/cpu:0,就用上面代码里的with tf.device("/cpu:0")包裹sess.run(init_op)。
4. 训练时显式保存迭代器初始化操作到图集合
如果训练时评估迭代器的操作没有被正确包含到保存的元图中,可以在训练时把它添加到图的集合里,方便恢复时直接获取:
# 训练时,创建评估迭代器初始化操作后添加: tf.add_to_collection('eval_init_op', init_op)
恢复时直接从集合中获取:
# 恢复模型后 init_op = tf.get_collection('eval_init_op')[0] sess.run(init_op)
按照上面的步骤排查,应该能解决迭代器未初始化的问题。
内容的提问来源于stack exchange,提问作者Harry

