如何清理LSTM多变量输入序列tf.data.Dataset中的NaN值
解决方案
1. 先修正数据集创建的错误
你分开创建输入和目标数据集再shuffle的做法会导致输入与目标序列不匹配——即使seed相同,两个数据集的窗口生成逻辑不同(输入是长度5的窗口,目标是长度1的窗口),shuffle后zip会打乱原本的时间对应关系。
正确的做法是直接用timeseries_dataset_from_array的targets参数,一次性生成对应输入-目标对:
import tensorflow as tf import numpy as np # 假设data是你的numpy数组,前6列特征,最后1列目标 input_features = data[:, :-1] target_values = data[:, -1] # 生成输入序列(长度5)和对应目标 # 每个输入序列是[t, t+1, t+2, t+3, t+4]的特征,对应目标是t+4时刻的标签 dataset = tf.keras.utils.timeseries_dataset_from_array( data=input_features, targets=target_values, sequence_length=5, sequence_stride=1, shuffle=True, seed=1 )
2. 修复NaN过滤的报错
你的过滤函数用了Python原生的not逻辑判断,这在TensorFlow的Graph执行模式下不允许——TensorFlow的符号张量不能直接当作Python布尔值使用。必须用TensorFlow提供的符号化逻辑操作:
方式一:定义过滤函数
def filter_invalid_samples(input_seq, target): # 检查输入序列所有元素都是有限值(非NaN/inf) input_is_valid = tf.reduce_all(tf.math.is_finite(input_seq)) # 检查目标值是有限值 target_is_valid = tf.math.is_finite(target) # 返回逻辑与的结果(Tensor类型的布尔值) return tf.logical_and(input_is_valid, target_is_valid) # 应用过滤 filtered_dataset = dataset.filter(filter_invalid_samples)
方式二:用lambda简化
filtered_dataset = dataset.filter( lambda x, y: tf.logical_and( tf.reduce_all(tf.math.is_finite(x)), tf.math.is_finite(y) ) )
为什么原来的代码报错?
tf.reduce_any(tf.math.is_nan(i))返回的是一个Tensor对象,而Python的not只能处理原生布尔值。在Graph执行模式下(TensorFlow默认的执行模式),所有操作都是符号化的,必须用TensorFlow的API(如tf.logical_not、tf.logical_and)来处理张量逻辑。
内容的提问来源于stack exchange,提问作者Jonathan Roy
相关产品推荐
相关产品推荐

