You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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()

额外优化提示

  1. 把原来的for _ in xrange(10)替换为while True,让迭代器自然遍历完所有数据,不用依赖固定循环次数。
  2. 循环结束后不要再调用sess.run(loss)或sess.run(acc_mean),因为此时迭代器已经耗尽,会直接抛出OutOfRangeError。如果需要最终的loss和acc,可以在最后一次迭代中把值保存下来复用。

这样修改后,训练代码就能和测试代码一样,正常加载全部260条数据了。

内容的提问来源于stack exchange,提问作者Yu Zhiyuan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 09:29:02