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

Keras Embedding层输出尺寸计算及索引越界问题求助

解决Keras Embedding层索引越界与维度计算问题

错误原因分析

你遇到的indices[2,22] = 55 is not in [0, 30)错误,核心是Embedding层的输入维度(input_dim)设置过小,导致输入数据中的词索引(55)超出了Embedding层允许的范围(0到29)。同时X_train尺寸为(25000,3),说明pad_sequences的maxlen参数未按预期设置为30,大概率是代码中出现了参数混淆或笔误。

正确的Embedding层维度计算与设置

1. 核心参数说明

Embedding层的关键参数与输出维度逻辑:

  • input_dim:必须大于等于数据中最大的词索引。IMDB数据集用num_words=top_words加载时,会保留最常用的top_words个词,词索引范围是0(填充符)到top_words-1,因此input_dim直接设为top_words即可覆盖所有合法索引。
  • output_dim:每个词的嵌入向量维度,由你自定义(示例中是32)。
  • input_length:输入序列的固定长度(即pad_sequences设置的max_words),显式设置可让模型明确输入形状。

Embedding层的输出维度为:(batch_size, input_length, output_dim),示例中就是(None, 30, 32)(None表示批量大小可变),刚好能对接后续LSTM层的输入要求。

2. 修正后的代码

from tensorflow.keras.datasets import imdb
from tensorflow.keras.preprocessing.sequence import pad_sequences
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Embedding, LSTM, Dense

top_words = 2000
(X_train, y_train), (X_test, y_test) = imdb.load_data(num_words=top_words)
max_words = 30
# 确保pad_sequences使用正确的maxlen参数
X_train = pad_sequences(X_train, maxlen=max_words)
X_test = pad_sequences(X_test, maxlen=max_words)

# 验证数据维度,应为(25000, 30)
print(X_train.shape)
print(X_test.shape)

model = Sequential()
# 明确设置input_dim、output_dim和input_length
model.add(Embedding(input_dim=top_words, output_dim=32, input_length=max_words))
model.add(LSTM(128, dropout=0.2, recurrent_dropout=0.2))
model.add(Dense(1, activation='sigmoid'))

model.compile(optimizer='adam',
              loss='binary_crossentropy',
              metrics=['accuracy'])

# 启动训练
model.fit(X_train, y_train, epochs=5, batch_size=32, validation_split=0.2)

3. 额外排查步骤

如果仍出现索引越界:

  • 检查数据中的最大词索引:运行print(max([max(seq) for seq in X_train])),该值必须小于top_words,否则说明num_words=top_words未正确过滤数据。
  • 确认代码中top_words未被后续代码意外覆盖为30(这是触发错误的常见原因)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 05:27:25