TensorFlow Dataset迭代器重置时shuffle行为及reshuffle_each_iteration参数疑问
tf.Dataset.shuffle中reshuffle_each_iteration参数的困惑解析 我来帮你理清这个参数的工作机制,以及你代码里出现这个现象的原因:
一、reshuffle_each_iteration的预期行为
这个参数的作用并不是控制重新初始化迭代器时的打乱行为,而是针对同一个迭代器多次遍历数据集的场景(比如用repeat()生成多epoch的数据集)。具体来说:
- 当
reshuffle_each_iteration=True(默认值):同一个迭代器每完成一次数据集遍历(一个epoch),会重新打乱数据顺序,下一次遍历的顺序和上一次不同。 - 当
reshuffle_each_iteration=False:同一个迭代器的多次遍历会使用相同的打乱顺序,不会每次都重新打乱。
你代码里的问题在于:每次epoch都调用了it_train_x.initializer,这相当于重新创建了一个全新的迭代器状态——shuffle的缓冲区是在迭代器初始化时生成的,所以每次初始化都会重新执行打乱逻辑,自然顺序就变了,这和reshuffle_each_iteration=False没有关系,因为这个参数管的是同一个迭代器的多次遍历,不是重新初始化的情况。
二、如何让每次重置迭代器后顺序一致?
如果你希望每次初始化迭代器都得到相同的打乱顺序,需要给shuffle方法加上固定的seed参数,同时保持reshuffle_each_iteration=False:
import tensorflow as tf flist = ["trimg1", "trimg2", "trimg3", "trimg4"] filenames = tf.constant(flist) # 固定seed,同时设置reshuffle_each_iteration=False train_x_dataset = tf.data.Dataset.from_tensor_slices((filenames)).shuffle(buffer_size=10, reshuffle_each_iteration=False, seed=42) it_train_x = train_x_dataset.make_initializable_iterator() next_sample = it_train_x.get_next() with tf.Session() as sess: for epoch in range(3): sess.run(it_train_x.initializer) print("Starting epoch ", epoch) while True: try: s = sess.run(next_sample) print("Sample: ", s) except tf.errors.OutOfRangeError: break
这样每次初始化迭代器时,因为seed固定,打乱的逻辑会生成相同的顺序,每个epoch的输出就一致了。
三、适合reshuffle_each_iteration=False的场景
这个参数的正确使用场景是当你用repeat()来串联多epoch数据时,同一个迭代器遍历多次的情况。比如下面的代码,用repeat(3)生成3个epoch的数据集,同一个迭代器遍历三次,reshuffle_each_iteration=False会让每个epoch的顺序完全相同:
import tensorflow as tf flist = ["trimg1", "trimg2", "trimg3", "trimg4"] filenames = tf.constant(flist) train_x_dataset = tf.data.Dataset.from_tensor_slices((filenames)) # 固定seed + reshuffle_each_iteration=False + repeat(3) train_x_dataset = train_x_dataset.shuffle(buffer_size=10, reshuffle_each_iteration=False, seed=42).repeat(3) it_train_x = train_x_dataset.make_one_shot_iterator() next_sample = it_train_x.get_next() with tf.Session() as sess: epoch = 0 print(f"Starting epoch {epoch}") for i in range(12): if i %4 ==0 and i !=0: epoch +=1 print(f"Starting epoch {epoch}") s = sess.run(next_sample) print("Sample: ", s)
四、针对你“每个epoch样本数量不同”的场景
如果你的每个epoch样本数量不一样,没法用repeat(),除了固定seed的方式,还可以提前生成一次固定的打乱序列,之后每个epoch都基于这个序列创建数据集:
import tensorflow as tf flist = ["trimg1", "trimg2", "trimg3", "trimg4"] # 提前生成固定的打乱顺序(只执行一次) with tf.Session() as sess: filenames_tensor = tf.constant(flist) fixed_shuffled_filenames = sess.run(tf.random.shuffle(filenames_tensor, seed=42)) # 每个epoch都用这个固定序列构建数据集 for epoch in range(3): # 这里可以根据epoch调整样本数量,比如只取前N个 current_filenames = fixed_shuffled_filenames[:4] # 示例:保持全量,实际可按需调整 train_x_dataset = tf.data.Dataset.from_tensor_slices(current_filenames) it_train_x = train_x_dataset.make_initializable_iterator() next_sample = it_train_x.get_next() with tf.Session() as sess: sess.run(it_train_x.initializer) print(f"Starting epoch {epoch}") while True: try: s = sess.run(next_sample) print("Sample: ", s) except tf.errors.OutOfRangeError: break
这种方式更灵活,你可以在每个epoch自由调整样本数量,同时保持打乱顺序一致。
内容的提问来源于stack exchange,提问作者Łukasz Sromek

