TensorFlow Dataset如何在每个批次内实现数据重复与拼接
解决方案
问题说明
初始数据集定义:dataset = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5, 6])
设置batch_size=2做常规批次划分时,默认输出为[[1,2], [3,4], [5,6]],预期输出为[[1,2,1,2], [3,4,3,4], [5,6,5,6]]。核心要求是将每个批次沿batch维度扩充为原来的2倍:例如实际场景中形状为(64, 300)的输入批次,处理后需得到形状为(128, 300)的输出批次。
实现代码
直接在数据集流水线中通过map对每个批次做拼接即可,逻辑简单且无额外性能损耗:
import tensorflow as tf # 初始化基础数据集 dataset = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5, 6]) batch_size = 2 expand_ratio = 2 # batch维度扩充倍数 # 构建数据处理流水线 dataset = ( dataset .batch(batch_size) # 沿batch维度(第0维)将批次自身拼接expand_ratio次 .map(lambda x: tf.concat([x for _ in range(expand_ratio)], axis=0)) ) # 验证输出 for batch in dataset: print(batch)
效果验证
运行上述代码的输出完全匹配预期:
tf.Tensor([1 2 1 2], shape=(4,), dtype=int32) tf.Tensor([3 4 3 4], shape=(4,), dtype=int32) tf.Tensor([5 6 5 6], shape=(4,), dtype=int32)
对于实际场景中形状为(64, 300)的输入批次,该处理会保留原300维的特征长度,仅将batch维度复制扩充,最终输出形状恰好为(128, 300),无需额外reshape操作。
注意:如果你的需求是逐样本重复(即单个batch
[1,2]输出为[1,1,2,2]),将map中的逻辑替换为tf.repeat(x, repeats=expand_ratio, axis=0)即可,可根据业务场景选择对应实现。
内容的提问来源于stack exchange,提问作者Murali
相关产品推荐
相关产品推荐

