如何从TensorFlow Dataset中提取配对的特征与标签?
解决TensorFlow Dataset迭代时特征与标签不匹配的问题
这个问题的核心原因很明确:每次调用.eval()(或者sess.run())都会触发一次TensorFlow图的执行,而每次执行iterator.get_next()都会让迭代器向前步进一个批次。你单独评估data[0]和data[1]时,相当于分别触发了两次迭代器步进,自然拿到的是不同批次的特征和标签。
最简单的解决方案:一次性评估整个元组
直接对data(也就是iterator.get_next()返回的元组)调用sess.run()或者.eval(),就能一次性获取同一批次对应的特征和标签配对。
修改后的代码如下:
import tensorflow as tf import numpy as np sess = tf.Session() def generator(): index = 0 while True: feature = np.ones([1,4]) * index label = feature[0:1,0:2] print('yielding:', feature, label) yield feature, label index +=1 dataset = tf.data.Dataset.from_generator( generator=generator, output_types=(tf.float64, tf.float64), output_shapes=(tf.TensorShape([1,4]),tf.TensorShape([1,2])), ) iterator = dataset.make_one_shot_iterator() data = iterator.get_next() # 一次性获取配对的特征与标签 feature_batch, label_batch = sess.run(data) print(feature_batch) print(label_batch) feature_batch, label_batch = sess.run(data) print(feature_batch) print(label_batch)
对应的输出会变成:
yielding: [[0. 0. 0. 0.]] [[0. 0.]] [[0. 0. 0. 0.]] [[0. 0.]] yielding: [[1. 1. 1. 1.]] [[1. 1.]] [[1. 1. 1. 1.]] [[1. 1.]]
原理说明
iterator.get_next()返回的是包含特征张量和标签张量的元组,当你对整个元组调用sess.run()时,TensorFlow会在同一次图执行中获取当前批次的特征和标签,迭代器只会步进一次,因此两者是完全匹配的。
如果需要在代码中多次使用这对配对,也可以先把它们存到变量里再进行后续操作,避免重复触发迭代器步进。
内容的提问来源于stack exchange,提问作者kaymes
相关产品推荐
相关产品推荐

