时间序列数据准备时如何拼接不同试验的滑动窗口数据用于LSTM训练
错误原因
你遇到的报错由两个核心问题导致:
tf.data.Dataset.window()返回的是嵌套的WindowDataset结构,每个元素本身仍是子Dataset,无法直接拼接,也不符合LSTM的张量输入要求tf.concat是张量拼接API,不能直接作用于Dataset对象,Dataset拼接需要用自带的concatenate()方法
修正后完整代码
import numpy as np import tensorflow as tf participant1 = np.arange(0,40,1).reshape(10,4) # time * features participant2 = np.arange(40,84,1).reshape(11,4) # time * features window_size=2 # 处理参与者1的数据集 input1 = tf.data.Dataset.from_tensor_slices(participant1) input1 = input1.window(window_size, shift=1, drop_remainder=True) # 把嵌套的WindowDataset压平为单个窗口张量 input1 = input1.flat_map(lambda window: window.batch(window_size)) # 处理参与者2的数据集 input2 = tf.data.Dataset.from_tensor_slices(participant2) input2 = input2.window(window_size, shift=1, drop_remainder=True) input2 = input2.flat_map(lambda window: window.batch(window_size)) # 拼接两个数据集 dataset = input1.concatenate(input2) # 验证结果:可打印所有元素的形状和值 for window in dataset: print(window.shape, window.numpy())
关键修改说明
- 新增
flat_map操作:将window生成的每个子Dataset批量打包为形状为(窗口大小, 特征数)的二维张量,直接适配LSTM的输入格式要求 - 替换拼接API:用
Dataset.concatenate()方法完成两个同结构Dataset的拼接,拼接后总样本量为两个参与者的窗口数之和(参与者1有9个窗口,参与者2有10个窗口,拼接后共19个样本)
内容的提问来源于stack exchange,提问作者Zhengliang Xia
相关产品推荐
相关产品推荐

