使用tf.data API训练时无法加载全部数据的问题求助
问题诊断与解决方案
我一眼就揪出了问题的根源——你的训练代码里多次调用sess.run()导致数据集迭代器被多次推进,每次调用都会消耗一个batch的数据,实际数据消耗速度是你预期的好几倍,这才出现了训练时只加载一半数据就耗尽的情况。
具体原因拆解
看你训练循环里的代码片段:
try: edges_value = sess.run(edges) # 第一次调用,消耗1个batch count = count+len(edges_value) _, step, summary = sess.run([train_all_op, global_step, merged]) # 第二次调用,再消耗1个batch writer.add_summary(summary, step) if step % 5 == 0: loss_value = sess.run(loss) # 第三次调用,消耗第3个batch print loss_value acc_mean_value = sess.run(acc_mean) # 第四次调用,消耗第4个batch print acc_mean_value saver.save(sess, params["save_dir"], step) except tf.errors.OutOfRangeError: print "end of data" break
每次循环里你调用了4次sess.run(),每一次调用都会触发iterator.get_next(),相当于一次循环就消耗了4个batch的数据。而你的测试代码每次循环只调用1次sess.run(),所以能完整遍历全部260条数据(10个batch)。
修复方案
把所有需要获取的张量打包到一次sess.run()调用里,确保每次循环只消耗1个batch的数据:
修改后的训练循环代码:
count = 0 with tf.Session(config=config) as sess: tf.global_variables_initializer().run() writer = tf.summary.FileWriter('./train', sess.graph) while True: try: # 把所有需要的张量一次性获取,避免多次推进迭代器 edges_value, _, step, summary, loss_value, acc_mean_value = sess.run( [edges, train_all_op, global_step, merged, loss, acc_mean] ) count += len(edges_value) writer.add_summary(summary, step) if step % 5 == 0: print loss_value print acc_mean_value saver.save(sess, params["save_dir"], step) except tf.errors.OutOfRangeError: print "end of data" break print count print "the final step is %d" % step # 此时迭代器已耗尽,无需再调用sess.run获取loss/acc saver.save(sess, params["save_dir"], step) writer.close()
额外优化提示
- 把原来的
for _ in xrange(10)替换为while True,让迭代器自然遍历完所有数据,不用依赖固定循环次数。 - 循环结束后不要再调用
sess.run(loss)或sess.run(acc_mean),因为此时迭代器已经耗尽,会直接抛出OutOfRangeError。如果需要最终的loss和acc,可以在最后一次迭代中把值保存下来复用。
这样修改后,训练代码就能和测试代码一样,正常加载全部260条数据了。
内容的提问来源于stack exchange,提问作者Yu Zhiyuan
相关产品推荐
相关产品推荐

