如何用TensorFlow迭代器获取不同类别样本用于BC学习?求替代方案
问题描述
我需要使用TensorFlow迭代器获取2个分属不同类别的样本,以开展BC学习。目前我尝试了tf.while_loop方案,但认为该方案并不合适,请问是否存在其他可行的解决方法?
以下是基于含5个类别的随机数数据集的示例代码:
import tensorflow as tf import numpy as np dataset = np.array([(np.random.rand(), i/20) for i in range(100)]) dataset = tf.data.Dataset.from_tensor_slices(dataset) dataset = dataset.shuffle(100) iterator = dataset.make_one_shot_iterator() a = iterator.get_next() b = iterator.get_next() loop_vars = [a, b] def cond(a, b): l1 = tf.gather(a, 1) l2 = tf.gather(b, 1) return tf.equal(l1, l2) def body(a, b): a = iterator.get_next() b = iterator.get_next() return a, b loop = tf.while_loop(cond, body, loop_vars) with tf.Session() as sess: for i in range(10): values = sess.run([loop]) print(values)
可行解决方案
方法1:按类别分组后交叉配对(推荐)
这种方法从根源上避免了反复采样的低效问题,核心思路是先将数据集按标签分组,再从不同类别组中抽取样本配对:
import tensorflow as tf import numpy as np # 构建数据集,修正标签为整数类别(0-4) dataset = np.array([(np.random.rand(), i//20) for i in range(100)]) dataset = tf.data.Dataset.from_tensor_slices(dataset) # 按类别分组 def group_by_label(element): label = tf.cast(element[1], tf.int32) return label grouped_ds = dataset.group_by_window( key_func=group_by_label, reduce_func=lambda key, ds: ds.batch(100), # 保留每个类别的所有样本 window_size=100 ) # 获取所有类别组的样本数据 group_list = list(grouped_ds.as_numpy_iterator()) # 生成所有有效类别对(i≠j) category_pairs = [(i,j) for i in range(5) for j in range(5) if i!=j] # 构建跨类别配对数据集 def generate_pairs(): pair_datasets = [] for i,j in category_pairs: # 从类别i和j中随机配对样本 pair_ds = tf.data.Dataset.from_tensor_slices( (group_list[i][:,0], group_list[j][:,0]) ).shuffle(len(group_list[i])) pair_datasets.append(pair_ds) # 合并所有配对数据集 return tf.data.Dataset.from_tensor_slices(pair_datasets).interleave(lambda x: x) final_ds = generate_pairs().batch(1) # 迭代获取10对样本 for pair in final_ds.take(10): print(pair.numpy())
方法2:TF1.x兼容的生成器式采样
如果你的项目仍在使用TensorFlow 1.x,可以用生成器逻辑替代tf.while_loop,更灵活地处理迭代器状态:
import tensorflow as tf import numpy as np def get_pair_iterator(dataset): iterator = dataset.make_initializable_iterator() next_elem = iterator.get_next() def pair_generator(sess): while True: try: a = sess.run(next_elem) b = sess.run(next_elem) # 直到拿到不同类别的样本 while a[1] == b[1]: b = sess.run(next_elem) yield (a, b) except tf.errors.OutOfRangeError: # 数据集耗尽后重置迭代器 sess.run(iterator.initializer) return iterator, pair_generator # 构建数据集 dataset = np.array([(np.random.rand(), i//20) for i in range(100)]) dataset = tf.data.Dataset.from_tensor_slices(dataset).shuffle(100) iterator, pair_gen = get_pair_iterator(dataset) with tf.Session() as sess: sess.run(iterator.initializer) gen = pair_gen(sess) for _ in range(10): print(next(gen))
方法3:TF2.x Eager模式简化版
如果可以升级到TensorFlow 2.x,利用Eager执行的灵活性,代码会更加直观:
import tensorflow as tf import numpy as np tf.config.run_functions_eagerly(True) # 构建可重复的数据集 dataset = np.array([(np.random.rand(), i//20) for i in range(100)]) dataset = tf.data.Dataset.from_tensor_slices(dataset).shuffle(100).repeat() # 创建迭代器 iterator = iter(dataset) def get_diff_class_pair(): a = next(iterator) b = next(iterator) # 循环直到拿到不同类别的样本 while tf.equal(a[1], b[1]): b = next(iterator) return a, b # 获取10对样本 for _ in range(10): pair = get_diff_class_pair() print(f"样本1: {pair[0].numpy()}, 样本2: {pair[1].numpy()}")
为什么tf.while_loop不合适?
你的原始方案中,tf.while_loop会在计算图中固化循环逻辑,一旦数据集耗尽就会抛出异常,且每次循环调用get_next()容易导致迭代器状态混乱。上面的方法要么通过预分组避免无效采样,要么用生成器逻辑灵活处理迭代器状态,更适配BC学习的样本配对需求。
内容的提问来源于stack exchange,提问作者Marc Moreaux
相关产品推荐
相关产品推荐

