为多时间序列创建无跨对象混合的TensorFlow Dataset
多独立时间序列滑窗切分方案(避免跨对象数据混合)
问题场景
现有按对象存储的多组独立时间序列数据,每个对象对应一段连续时序记录,示例数据构造代码如下:
import pandas as pd import numpy as np import tensorflow as tf df = pd.DataFrame({'Time': np.tile(np.arange(5), 2), 'Object': np.concatenate([[i] * 5 for i in [1, 2]]), 'Feature1': np.random.randint(10, size=10), 'Feature2': np.random.randint(10, size=10)})
示例数据结构:
| Time | Object | Feature1 | Feature2 |
|---|---|---|---|
| 0 | 1 | 3 | 3 |
| 1 | 1 | 9 | 2 |
| 2 | 1 | 6 | 6 |
| 3 | 1 | 4 | 0 |
| 4 | 1 | 7 | 7 |
| 0 | 2 | 4 | 8 |
| 1 | 2 | 3 | 7 |
| 2 | 2 | 1 | 1 |
| 3 | 2 | 7 | 5 |
| 4 | 2 | 1 | 7 |
实际场景共包含约2000个独立对象,需要对时序数据做滑窗切分,生成适合RNN/LSTM输入的定长窗口,核心约束为:
- 单个窗口内不能包含来自不同对象的时序数据
- 窗口数据需保留对象ID字段,方便后续接入Embedding层识别不同对象的序列特征
直接对全量数据集调用window方法会出现跨对象数据混合的问题,错误实现代码:
dataset = tf.data.Dataset.from_tensor_slices(df) for w in dataset.window(3, shift=1, drop_remainder=True): print(list(w.as_numpy_iterator()))
错误输出中存在跨对象混合的窗口:
[array([3, 1, 4, 0]), array([4, 1, 7, 7]), array([0, 2, 4, 8])] # 同时包含对象1和对象2的数据 [array([4, 1, 7, 7]), array([0, 2, 4, 8]), array([1, 2, 3, 7])] # 同时包含对象1和对象2的数据
期望输出为仅在单个对象的时序段内做滑窗,无跨对象混合,示例如下:
[array([0, 1, 3, 3]), array([1, 1, 9, 2]), array([2, 1, 6, 6])] [array([1, 1, 9, 2]), array([2, 1, 6, 6]), array([3, 1, 4, 0])] [array([2, 1, 6, 6]), array([3, 1, 4, 0]), array([4, 1, 7, 7])] [array([0, 2, 4, 8]), array([1, 2, 3, 7]), array([2, 2, 1, 1])] [array([1, 2, 3, 7]), array([2, 2, 1, 1]), array([3, 2, 7, 5])] [array([2, 2, 1, 1]), array([3, 2, 7, 5]), array([4, 2, 1, 7])]
实现方案
核心逻辑是先按对象拆分独立时序段,每个时序段单独做滑窗,最后合并所有窗口,从根源避免跨对象数据混合,切分后的窗口保留对象ID字段,可直接接入后续Embedding层。
WINDOW_SIZE = 3 SHIFT = 1 DROP_REMAINDER = True window_datasets = [] # 按对象ID分组遍历 for obj_id, group_df in df.groupby("Object"): # 单对象数据转为TF数据集 obj_ds = tf.data.Dataset.from_tensor_slices(group_df) # 仅在单对象时序范围内做滑窗 obj_window_ds = obj_ds.window(WINDOW_SIZE, shift=SHIFT, drop_remainder=DROP_REMAINDER) # 将嵌套窗口展平为定长张量批次 obj_window_ds = obj_window_ds.flat_map(lambda window: window.batch(WINDOW_SIZE)) window_datasets.append(obj_window_ds) # 合并所有对象的窗口得到最终数据集 final_dataset = window_datasets[0] for ds in window_datasets[1:]: final_dataset = final_dataset.concatenate(ds) # 验证输出 for batch in final_dataset: print(batch.numpy())
运行输出和预期结果完全一致,不存在跨对象数据混合的问题。
大规模场景适配
针对2000个对象的实际业务场景,可按需做以下优化:
- 若单对象时序长度较长、全量数据内存占用高,可提前按对象拆分存储为TFRecord格式,读取时直接对单对象序列做滑窗,无需全量加载pandas数据
- 合并完所有窗口后,可追加
shuffle、batch、prefetch等流水线优化操作,适配模型训练的输入效率要求
模型接入说明
切分后的窗口完整保留了Object字段,训练时可将该字段单独传入Embedding层,编码得到的对象嵌入向量和时序特征输出拼接后,再输入LSTM层即可完成不同对象时序特征的区分。
内容的提问来源于stack exchange,提问作者Mykola Zotko
相关产品推荐
相关产品推荐

