TensorFlow中如何将窗口化数据集传入StringLookup层
报错核心原因
两个维度不匹配问题直接触发报错:
- 你定义的模型输入层shape是
()(单值标量输入,对应无序列结构的单样本特征),但窗口化后的数据集每个特征的shape是(window_size,),每个输入是长度为窗口值的时间序列,输入维度和模型定义完全不匹配 StringLookup层在output_mode设为one_hot/multi_hot这类非整数输出模式时,最高仅支持2维输入(输出秩最大为2),如果直接传入3维的(批次大小, 窗口长度)字符串张量做one-hot编码,会生成秩为3的输出,超出该层的支持范围。
你之前单独测试预处理器能跑通,是因为直接传入原始DataFrame时,每个特征是1维的(总样本数,)格式,输入到shape为()的输入层会自动按逐样本标量处理,但窗口化后输入维度升了一级,自然触发报错。
可直接运行的调整方案
方案1:修改预处理器结构适配序列输入(推荐)
把预处理逻辑全部整合进模型,后续部署不需要单独处理数据,最稳妥。核心修改点是把输入层shape改为窗口长度,用TimeDistributed包装逐时间步处理的预处理层,让每个时间步的特征单独走编码逻辑,不会触发维度超限:
window_size = 3 # 输入层shape改为(window_size,),匹配窗口化后的序列长度 inputs = { 'col1': tf.keras.Input(shape=(window_size,), name='col1', dtype=tf.string), 'col2': tf.keras.Input(shape=(window_size,), name='col2', dtype=tf.float32), 'col3': tf.keras.Input(shape=(window_size,), name='col3', dtype=tf.float32) } vocab = sorted(set(x['col1'])) # 用TimeDistributed包装StringLookup,逐时间步做one-hot编码,输出shape(批次, 窗口, 词表大小+1) lookup = tf.keras.layers.TimeDistributed( tf.keras.layers.StringLookup(vocabulary=vocab, output_mode='one_hot') )(inputs['col1']) # 数值特征逐时间步拼接,输出shape(批次, 窗口, 数值特征数) numeric = tf.stack( [tf.cast(inputs[col], dtype=tf.float32) for col in ['col2', 'col3']], axis=-1 ) # 拼接后总特征shape为(批次, 窗口, 13) processed_feat = tf.concat([lookup, numeric], axis=-1) preprocessor = tf.keras.Model(inputs, processed_feat)
后续模型主体部分如果要处理序列结构,可以按需选LSTM/GRU等序列层,或者展平后走全连接层,对齐你原来的输出逻辑的写法如下:
body = tf.keras.models.Sequential([ tf.keras.layers.Flatten(), # 把(3,13)的序列特征展平为一维向量 tf.keras.layers.Dense(8), tf.keras.layers.Dense(window_size) ]) model_output = body(preprocessor(inputs)) model = tf.keras.Model(inputs, model_output) model.compile(loss='mae', optimizer='adam') model.fit(dataset) # 此时可正常训练
方案2:简化数据集窗口化逻辑
你原来手写字典遍历的窗口转换写法可以简化,减少冗余代码,输出和原来完全一致:
window_size = 3 dataset = tf.data.Dataset.from_tensor_slices((dict(x), y)) # 直接对整体数据集做窗口批处理,不需要单独遍历字典键 dataset = dataset.window(window_size, shift=1, drop_remainder=True)\ .flat_map(lambda feat, label: tf.data.Dataset.zip((feat, label)).batch(window_size))\ .batch(3)
内容的提问来源于stack exchange,提问作者Mykola Zotko
相关产品推荐
相关产品推荐

