You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.27 15:39:13