You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何按不同频率合并两个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]

代码逻辑说明

  1. ds1.batch(2, drop_remainder=True):将ds1的元素按每2个一组打包,由于原数据集是无限重复的,drop_remainder=True不会导致样本丢失,只会确保每组都是完整的2个元素。
  2. Dataset.zip:将分组后的ds1与ds10进行配对,每一组对应(ds1的2个元素, ds10的1个元素)。
  3. flat_map:把每组元素按指定规则展开为单个元素序列——将ds1的第一个元素与ds10的元素相加,第二个元素直接保留,最终生成符合需求的连续序列。

该方案完全基于tf.data原生操作,效率无损耗,且最终的ds是标准Dataset对象,可直接用于后续的批处理、预处理等流水线操作。


内容的提问来源于stack exchange,提问作者user209974

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.14 18:43:12