如何按不同频率合并两个TensorFlow Dataset对象
实现方案
可以通过TensorFlow的tf.data原生API组合操作实现,全程无样本丢弃,最终得到标准Dataset对象以支持后续流水线操作。
具体代码实现
import tensorflow as tf ds1 = tf.data.Dataset.range(1, 10).repeat() ds10 = tf.data.Dataset.range(10, 100, 10).repeat() # 将ds1按每2个元素为一组打包 ds1_grouped = ds1.batch(2, drop_remainder=True) # 把分组后的ds1与ds10一一配对 merged_pairs = tf.data.Dataset.zip((ds1_grouped, ds10)) # 展开每组元素,按要求生成序列 ds = merged_pairs.flat_map(lambda ds1_pair, ds10_elem: tf.data.Dataset.from_tensor_slices([ds1_pair[0]+ds10_elem, ds1_pair[1]])) # 验证前7个元素 print(list(ds.take(7))) # 输出: [11, 2, 23, 4, 35, 6, 47]
代码逻辑说明
ds1.batch(2, drop_remainder=True):将ds1的元素按每2个一组打包,由于原数据集是无限重复的,drop_remainder=True不会导致样本丢失,只会确保每组都是完整的2个元素。Dataset.zip:将分组后的ds1与ds10进行配对,每一组对应(ds1的2个元素, ds10的1个元素)。flat_map:把每组元素按指定规则展开为单个元素序列——将ds1的第一个元素与ds10的元素相加,第二个元素直接保留,最终生成符合需求的连续序列。
该方案完全基于tf.data原生操作,效率无损耗,且最终的ds是标准Dataset对象,可直接用于后续的批处理、预处理等流水线操作。
内容的提问来源于stack exchange,提问作者user209974
相关产品推荐
相关产品推荐

