使用Keras Embedding预测时出现索引不在列表中的错误
问题与解答
问题背景
我训练了一个基于Keras的模型,采用GloVe作为预训练嵌入词典,训练时通过Tokenizer构建文本序列并生成Embedding矩阵,核心代码如下:
common_embed = Embedding( name="synopsis_embedd", input_dim =len(t.word_index)+1, output_dim=len(embeddings_index['no']), weights=[embedding_matrix], input_length=len(X_train['asset_text_seq_pad'].tolist()[0]), trainable=True ) lstm_1 = common_embed(input_1) common_lstm = LSTM(64, input_shape=(100,2)) ...
训练阶段Tokenizer构建流程:
t = Tokenizer() t.fit_on_texts(all_text) text_seq= pad_sequences(t.texts_to_sequences(data['example_texts'].astype(str).values))
Embedding矩阵计算代码:
embeddings_index = {} for line in new_byte_string.decode('utf-8').split('\n'): if line: values = line.split() word = values[0] coefs = np.asarray(values[1:], dtype='float32') embeddings_index[word] = coefs embedding_vector = None not_present_list = [] vocab_size = len(t.word_index) + 1 print('Loaded %s word vectors.' % len(embeddings_index)) embedding_matrix = np.zeros((vocab_size, len(embeddings_index['no']))) for word, i in t.word_index.items(): if word in embeddings_index.keys(): embedding_vector = embeddings_index.get(word) else: not_present_list.append(word) if embedding_vector is not None: embedding_matrix[i] = embedding_vector else: embedding_matrix[i] = np.zeros(300)
使用新数据集预测时,重新执行了所有预处理步骤,出现如下错误:
Node: 'model/synopsis_embedd/embedding_lookup'
indices[38666,63] = 136482 is not in [0, 129872)
[[{{node model/synopsis_embedd/embedding_lookup}}]] [Op:__inference_predict_function_12452]
疑问
- 重新执行所有预处理步骤是否错误?
- 是否必须复用训练时的Tokenizer?
- 该错误的原因是什么?
解答
1. 重新执行所有预处理步骤是错误的
预测阶段不能从头重新执行完整预处理流程,尤其是不能重新训练Tokenizer。这样做会让预处理逻辑和训练阶段完全脱节,直接引发模型输入不兼容的问题。
2. 必须复用训练时的Tokenizer
Tokenizer的核心价值是训练阶段构建的词汇-索引映射表(word_index),模型的Embedding层是基于这个映射表生成的矩阵来工作的。如果预测时重新训练Tokenizer,新数据的词汇会生成全新的索引值,和模型训练时依赖的索引体系完全不匹配,必然导致后续嵌入查找失败。正确的做法是:训练时保存Tokenizer(比如用pickle序列化),预测时加载已保存的Tokenizer,仅调用texts_to_sequences和pad_sequences完成预处理。
3. 错误原因解析
报错信息明确指出:
- 模型训练时,Embedding层的
input_dim参数设置为129872(即len(t.word_index)+1),这意味着模型只接受0到129871之间的索引值; - 预测时重新训练Tokenizer,新生成的索引中出现了136482这个值,超出了模型Embedding层预设的索引范围,导致嵌入查找操作失败。
内容的提问来源于stack exchange,提问作者fsulser
相关产品推荐
相关产品推荐

