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

使用TensorFlow Keras做文本生成时LSTM维度不匹配报错如何解决

问题原因

LSTM层要求输入的维度是3维,格式为(batch_size, 序列长度, 词汇表大小),但你当前传入模型的输入只有2维(128, 100),本质原因是你定义了独热编码转换函数one_hot_samples,但没有将其应用到训练数据集上,数据集里存储的还是整数编码的字符序列,没有转换成模型需要的独热编码格式。

修复方法

找到代码中构造训练数据集ds的行:

ds = dataset.repeat().shuffle(1024).batch(BATCH_SIZE, drop_remainder=True)

修改为:

ds = dataset.map(one_hot_samples).repeat().shuffle(1024).batch(BATCH_SIZE, drop_remainder=True)

新增的map(one_hot_samples)会将数据集中所有的输入序列和目标字符都转换为独热编码格式,匹配LSTM层的输入要求。

可选优化方案

如果你的词汇表较大,独热编码会占用过多内存,可以考虑在模型第一层新增Embedding层替换独热编码,无需修改数据集构造逻辑,模型定义修改为:

model = Sequential([
    tf.keras.layers.Embedding(input_dim=n_unique_chars, output_dim=128, input_length=sequence_length),
    LSTM(256, return_sequences=True),
    Dropout(0.3),
    LSTM(256),
    Dense(n_unique_chars, activation="softmax"),
])

这种情况下也不需要再对数据集做独热编码转换,整数编码的序列可以直接输入模型训练。

内容的提问来源于stack exchange,提问作者Lee

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 10:36:07