如何将两个MapDataset合并为一个以实现数据集扩容
TensorFlow 两个MapDataset合并方案
你可以直接使用tf.data.Dataset内置的concatenate()方法完成合并,该方法要求两个数据集的元素结构、形状、数据类型完全匹配,你当前场景已经预处理保证了图片形状一致,且两个数据集均返回(图片张量, 标签)的二元组结构,可直接调用。
基础合并代码
# 拼接两个数据集,返回的combined_ds依然是MapDataset类型 combined_ds = ds.concatenate(ds1)
常用后续处理
合并后可直接执行常规的数据集流水线操作,示例:
BATCH_SIZE = 32 # 合并后全局打乱 + 分批 + 预取优化 combined_ds = combined_ds.shuffle( buffer_size=len(ds) + len(ds1) ).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
可选进阶操作
如果需要调整增广数据和原始数据的采样比例,不需要拼接,改用sample_from_datasets方法即可,示例为两类数据按1:1采样:
combined_ds = tf.data.Dataset.sample_from_datasets( datasets=[ds, ds1], weights=[0.5, 0.5] )
注意事项
- 合并前无需额外调整数据集结构,只要确认两个数据集的返回张量形状、数据类型完全匹配即可正常执行拼接
- 如果拼接时报类型/形状不匹配错误,可分别调用
ds.element_spec和ds1.element_spec打印结构,对比差异后调整一致再合并
内容的提问来源于stack exchange,提问作者Aloma85
相关产品推荐
相关产品推荐

