Python中如何过滤TensorFlow MapDataset内含NaN的窗口数据
问题描述
- 时间序列数据准备阶段需要将数据集划分为固定长度时间窗口,自定义实现了
WindowGenerator_with_nan窗口生成类,类定义如下:
class WindowGenerator_with_nan(): def __init__(self, input_width, label_width, shift, x_iter, train_df=cluster_concat_train_df, val_df=cluster_concat_val_df, test_df=cluster_concat_test_df, label_columns=None): # 存储原始数据 self.train_df = cluster_concat_train_df[x_iter] self.val_df = cluster_concat_val_df[x_iter] self.test_df = cluster_concat_test_df[x_iter] # 计算标签列索引 self.label_columns = label_columns if label_columns is not None: self.label_columns_indices = {name: i for i, name in enumerate(label_columns)} self.column_indices = {name: i for i, name in enumerate(train_df[x_iter].columns)} # 配置窗口参数 self.input_width = input_width self.label_width = label_width self.shift = shift self.total_window_size = input_width + shift self.input_slice = slice(0, input_width) self.input_indices = np.arange(self.total_window_size)[self.input_slice] self.label_start = self.total_window_size - self.label_width self.labels_slice = slice(self.label_start, None) self.label_indices = np.arange(self.total_window_size)[self.labels_slice] def __repr__(self): return '\n'.join([ f'Total window size: {self.total_window_size}', f'Input indices: {self.input_indices}', f'Label indices: {self.label_indices}', f'Label column name(s): {self.label_columns}'])
- 业务场景中迭代变量
i代表聚类编号,窗口生成器输出的MapDataset类型数据集中存在部分NaN值。 - 窗口实例列表生成代码如下:
wide_window_with_nan =[WindowGenerator(input_width=96, label_width=1, shift=1, label_columns = ['Labels'], x_iter = i) for i in range(len(df_without_impulate_before_RNN))]
- 执行
print(wide_window_with_nan[0].train)输出如下:
<MapDataset element_spec=(TensorSpec(shape=(None, 96, 112), dtype=tf.float32, name=None), TensorSpec(shape=(None, 1, 1), dtype=tf.float32, name=None))>
- 核心需求:从上述
MapDataset中移除所有包含NaN值的窗口,输出可直接输入不支持NaN值的预测模型的干净数据集;开发环境为Google Colab Pro,方案需要严格控制RAM占用,避免内存不足问题。
实现方案
全程基于TensorFlow Dataset原生流式API实现过滤,不需要把全量数据加载到内存,从根源上避免RAM溢出:
- 核心逻辑:使用
tf.data.Dataset.filter方法逐窗口校验NaN,保留无NaN的有效样本,整个过程保持数据集懒加载特性,内存占用始终维持在单窗口处理的水平。 - 第一步:给窗口生成类补充NaN校验和干净数据集获取方法,不需要修改原有类的初始化逻辑:
import tensorflow as tf # 给现有WindowGenerator_with_nan类补充两个方法 def _has_no_nan(self, features, labels): # 同时校验特征、标签张量的所有位置,任意位置存在NaN就过滤该窗口 return tf.logical_and( tf.reduce_all(~tf.math.is_nan(features)), tf.reduce_all(~tf.math.is_nan(labels)) ) def get_clean_split(self, split='train'): # 选择对应拆分的数据集 if split == 'train': raw_ds = self.train elif split == 'val': raw_ds = self.val else: raw_ds = self.test # 流式并行过滤,预取优化训练速度,不额外占用大量内存 return raw_ds.filter(self._has_no_nan).prefetch(tf.data.AUTOTUNE) # 把方法绑定到原有类上 WindowGenerator_with_nan._has_no_nan = _has_no_nan WindowGenerator_with_nan.get_clean_split = get_clean_split
- 第二步:训练/验证/测试时直接调用方法获取过滤后的数据集即可,不需要提前把所有聚类的数据集全部加载:
# 示例:获取第0个聚类的干净训练、验证、测试集 cluster_id = 0 clean_train = wide_window_with_nan[cluster_id].get_clean_split('train') clean_val = wide_window_with_nan[cluster_id].get_clean_split('val') clean_test = wide_window_with_nan[cluster_id].get_clean_split('test') # 校验过滤效果 for sample_feat, sample_label in clean_train.take(10): assert not tf.math.is_nan(sample_feat).numpy().any() assert not tf.math.is_nan(sample_label).numpy().any()
- 内存优化注意事项:
- 不要在过滤操作前调用
.cache(),否则会把带NaN的原始全量数据缓存到内存,浪费RAM空间 - batch操作放在过滤之后执行,避免单个窗口含NaN导致整个batch被丢弃,减少数据浪费
- 不需要提前把所有聚类的干净数据集全部生成,训练到哪个聚类再生成对应数据集即可,进一步降低内存峰值
- 不要在过滤操作前调用
内容的提问来源于stack exchange,提问作者Pradya Panyainkaew
相关产品推荐
相关产品推荐

