对批处理tf.data.Dataset使用map()时自定义函数的输入参数是什么?
核心结论
你自定义的处理函数输入是完整批次,每次迭代时函数作用在形状为(batch_size, sequence_length, features)的全量批次数据上,不是单个样本。
维度显示异常的原因
你打印输入参数得到(None, None, features)是TensorFlow静态形状推断的特性导致的:切片操作不会保留静态形状信息,维度位置的None仅代表编译阶段无法确定该维度的固定值,不是运行时的真实维度。你可以在函数内通过tf.print(features.shape)打印动态形状,就能看到符合预期的(32, 24, 你的特征数量)的实际结构。
示例代码逻辑佐证
你示例中split_window函数里写的features[:, self.input_slice, :]第一维用:全选,本身就是针对带批次维度的数据的操作:如果输入是单个样本,索引写法应该是features[self.input_slice, :],无需第一维的批次索引,这也能验证函数的输入是完整批次。
补充说明
tf.keras.utils.timeseries_dataset_from_array传入batch_size参数后生成的BatchDataset,每一个元素都是对应大小的批次数据,tf.data.Dataset.map方法默认对数据集的每个元素(也就是每个批次)做变换,所以处理函数的输入始终是完整批次。
内容的提问来源于stack exchange,提问作者ptmva
相关产品推荐
相关产品推荐

