使用flat_map与zip读取TensorFlow Dataset时的顺序一致性及重复读取疑问
这是个很典型的tf.data使用误区,我来帮你梳理清楚问题根源和解决办法:
首先,问题出在你使用zip(dataset, dataset.skip(1))的方式上——这里的两个dataset其实是完全独立的数据集实例。它们会各自从头执行list_files、flat_map等整个数据流水线操作:
- 默认情况下
tf.data.Dataset.list_files会随机打乱文件顺序,这就导致两个数据集可能读取文件的顺序不一样; - 就算你关闭了文件打乱,TensorFlow的异步数据读取机制也可能让两个数据集的读取进度出现交错,最终出现一个数据集在读取file1、另一个已经读到file2的情况,自然就产生了跨文件的连续样本。
要稳定实现你想要的「同一文件内相邻记录组成连续样本」的效果,你需要在同一个数据序列内生成连续样本,而不是用两个独立数据集做zip。这里推荐用tf.data.Dataset.window操作,具体分两种场景:
场景1:只保留同一文件内的相邻样本(不包含跨文件的连续对)
这种场景下,我们需要在每个文件的内部生成连续样本,而不是在整个数据集上操作:
import tensorflow as tf COLUMNS = ['image', 'label'] FIELD_DEFAULTS = [['empty'], [0]] def _line_parser(line): fields = tf.decode_csv(line, FIELD_DEFAULTS) data = dict(zip(COLUMNS, fields)) label = data.pop('label') return data, label filenames = ['file1.txt', 'file2.txt'] # 1. 关闭文件列表的随机打乱,确保按指定顺序读取文件 files = tf.data.Dataset.list_files(filenames, shuffle=False) # 2. 对每个文件单独处理:读取内容 -> 解析 -> 生成文件内的连续样本 dataset = files.flat_map( lambda filename: tf.data.TextLineDataset(filename) .map(_line_parser) # window(size=2):每个窗口包含2个相邻元素;shift=1:每次滑动1个元素 # drop_remainder=True:丢弃每个文件最后一个无法组成对的元素 .window(size=2, shift=1, drop_remainder=True) # 将窗口中的元素转换成((data1, label1), (data2, label2))的格式 .flat_map(lambda window_ds: window_ds.batch(2)) ) iterator = dataset.make_initializable_iterator() next_element = iterator.get_next() init_op = iterator.initializer with tf.Session() as sess: sess.run(init_op) try: while True: print(sess.run(next_element)) except tf.errors.OutOfRangeError: print("数据集遍历完成")
场景2:允许跨文件的连续样本(保持整体文件顺序)
如果你希望保留「file1最后一条 + file2第一条」这样的跨文件连续对,只需要把window操作移到整个数据集层面即可:
import tensorflow as tf COLUMNS = ['image', 'label'] FIELD_DEFAULTS = [['empty'], [0]] def _line_parser(line): fields = tf.decode_csv(line, FIELD_DEFAULTS) data = dict(zip(COLUMNS, fields)) label = data.pop('label') return data, label filenames = ['file1.txt', 'file2.txt'] files = tf.data.Dataset.list_files(filenames, shuffle=False) # 先读取所有文件的内容并解析成连续序列 dataset = files.flat_map( lambda filename: tf.data.TextLineDataset(filename).map(_line_parser) ) # 在整个数据集序列上生成连续样本 dataset = dataset.window(size=2, shift=1, drop_remainder=True).flat_map( lambda window_ds: window_ds.batch(2) ) iterator = dataset.make_initializable_iterator() next_element = iterator.get_next() init_op = iterator.initializer with tf.Session() as sess: sess.run(init_op) try: while True: print(sess.run(next_element)) except tf.errors.OutOfRangeError: print("数据集遍历完成")
核心原理
window操作是在同一个数据集序列上进行的,所有连续样本都是从这个统一的序列中截取的,完全避免了两个独立数据集读取交错的问题。再配合list_files(shuffle=False)固定文件读取顺序,就能100%保证样本顺序的稳定性。
内容的提问来源于stack exchange,提问作者prouast
相关产品推荐
相关产品推荐

