LSTM Encoder-Decoder训练陷入平台期无法学习,求问题排查
字符序列元音识别:Encoder-Decoder训练平台期问题排查
问题背景
任务为识别随机字符序列中的元音,已生成100K条TSV样本(示例:molteyhpr对应010011000),所有数据转为等长独热编码。训练配置:10轮epoch、Adam优化器(学习率0.0001)、批量大小30、分类交叉熵损失。训练陷入平台期,准确率始终在0.46左右波动,调参无效。
模型代码
self.latent_dim = 256 enc_input_layer = Input(name="enc_input", shape=(None, self.source.enc_vocab_len)) enc_lstm_layer = LSTM(self.latent_dim, name="enc_lstm", return_state=True) enc_outputs, state_h, state_c = enc_lstm_layer(enc_input_layer) # We discard 'enc_outputs' and only keep the states. enc_states = [state_h, state_c] # Set up the decoder, using 'enc_states' as initial state. dec_input_layer = Input(name="dec_input", shape=(None, self.source.dec_vocab_len)) # We set up our decoder to return full output sequences, # and to return internal states as well. We don't use the # return states in the training model, but we will use them in inference. dec_lstm_layer = LSTM(self.latent_dim, name="dec_lstm", return_sequences=True, return_state=True) dec_outputs, _, _ = dec_lstm_layer(dec_input_layer, initial_state=enc_states) dec_dense_layer = Dense(self.source.dec_vocab_len, name="dec_dense", activation='softmax') dec_outputs = dec_dense_layer(dec_outputs) # Define the model that will turn # 'encoder_input_data' & 'decoder_input_data' into 'decoder_target_data' model = Model([enc_input_layer, dec_input_layer], dec_outputs)
数据生成器代码
def _generator(self, enc_data, dec_data, is_training): enc_oh_input_batch = None dec_oh_input_batch = None dec_oh_output_batch = None enc_space_token = self.enc_vocab[self.TOKEN_EMPTY] dec_space_token = self.dec_vocab[self.TOKEN_EMPTY] current_idx = 0 samples_len = len(enc_data) while True: # Create zero batch arrays enc_oh_input_batch = np.zeros( (self.batch_size, self.enc_max_seq_len, self.enc_vocab_len), dtype='int8') dec_oh_input_batch = np.zeros( (self.batch_size, self.dec_max_seq_len, self.dec_vocab_len), dtype='int8') dec_oh_output_batch = np.zeros( (self.batch_size, self.dec_max_seq_len, self.dec_vocab_len), dtype='int8') # Compile batch for i in range(self.batch_size): # when we get to the end of samples - start over if i + current_idx >= samples_len: current_idx = 0 if is_training: self.epoch += 1 tokens_in = enc_data[i + current_idx] tokens_out = dec_data[i + current_idx] # vectorize encoder input for t, token in enumerate(tokens_in): enc_oh_input_batch[i, t, token] = 1 enc_oh_input_batch[i, t + 1:, enc_space_token] = 1 # vectorize decoder input and output for t, token in enumerate(tokens_out): dec_oh_input_batch[i, t, token] = 1 if t > 0: # self.dec_oh_output will be ahead by one timestep # and will not include the start character. dec_oh_output_batch[i, t - 1, token] = 1 dec_oh_input_batch[i, t + 1:, dec_space_token] = 1 current_idx += self.batch_size yield [[enc_oh_input_batch, dec_oh_input_batch], dec_oh_output_batch]
训练日志
Training model ... Epoch 1/10 33/33 [==============================] - 6s 90ms/step - loss: 0.1759 - accuracy: 0.4380 Epoch 2/10 33/33 [==============================] - 3s 91ms/step - loss: 0.1370 - accuracy: 0.4533 Epoch 3/10 33/33 [==============================] - 3s 90ms/step - loss: 0.1258 - accuracy: 0.4634 Epoch 4/10 33/33 [==============================] - 3s 93ms/step - loss: 0.1220 - accuracy: 0.4602 Epoch 5/10 33/33 [==============================] - 3s 95ms/step - loss: 0.1199 - accuracy: 0.4602 Epoch 6/10 33/33 [==============================] - 3s 92ms/step - loss: 0.1218 - accuracy: 0.4625 Epoch 7/10 33/33 [==============================] - 3s 94ms/step - loss: 0.1208 - accuracy: 0.4643 Epoch 8/10 33/33 [==============================] - 3s 93ms/step - loss: 0.1202 - accuracy: 0.4619 Epoch 9/10 33/33 [==============================] - 3s 95ms/step - loss: 0.1199 - accuracy: 0.4601 Epoch 10/10 33/33 [==============================] - 5s 149ms/step - loss: 0.1207 - accuracy: 0.4630 - val_loss: 0.1195 - val_accuracy: 0.4630
核心问题分析与解决方案
1. 架构选型错误:用Seq2Seq做序列标注
当前任务本质是序列标注任务(每个输入字符对应一个0/1标签),但误用了用于机器翻译的Encoder-Decoder架构:
- 编码器丢弃了全序列输出
enc_outputs,只保留最终状态,导致解码器无法获取每个输入字符的位置信息,只能依赖全局语义生成标签,完全无法对应到每个输入位置 - 解码器的自回归输入(dec_input)是冗余设计,引入不必要的计算和噪声
修改方案:改用序列标注专用架构,去掉Decoder部分:
self.latent_dim = 256 # 输入层:字符序列独热编码 enc_input_layer = Input(name="enc_input", shape=(None, self.source.enc_vocab_len)) # 双向LSTM保留全序列输出,捕捉每个位置的上下文 enc_lstm_layer = Bidirectional(LSTM(self.latent_dim, return_sequences=True), name="enc_bilstm") enc_outputs = enc_lstm_layer(enc_input_layer) # 每个时间步对应二分类(元音/非元音) dec_dense_layer = Dense(self.source.dec_vocab_len, activation='softmax', name="dec_dense") outputs = dec_dense_layer(enc_outputs) # 模型仅需输入字符序列,输出对应标签序列 model = Model(enc_input_layer, outputs)
2. 数据生成器逻辑错误:误用Teacher Forcing
生成器中对decoder输入和输出的偏移处理(t>0时,dec_oh_output_batch[i, t-1, token] =1)是机器翻译的Teacher Forcing逻辑,完全不适合序列标注任务:
- 序列标注要求输入序列与输出序列长度一致、位置一一对应,偏移处理导致标签错位
- 填充逻辑存在bug:
enc_oh_input_batch[i, t + 1:, enc_space_token] = 1中的t是最后一个有效字符的索引,循环结束后t值固定,可能导致填充位置错误
修改方案:简化生成器,去掉dec_input逻辑,直接生成输入-标签对:
def _generator(self, enc_data, dec_data, is_training): enc_space_token = self.enc_vocab[self.TOKEN_EMPTY] dec_space_token = self.dec_vocab[self.TOKEN_EMPTY] current_idx = 0 samples_len = len(enc_data) while True: # 初始化批量数组 enc_oh_input_batch = np.zeros( (self.batch_size, self.enc_max_seq_len, self.enc_vocab_len), dtype='int8') dec_oh_output_batch = np.zeros( (self.batch_size, self.dec_max_seq_len, self.dec_vocab_len), dtype='int8') for i in range(self.batch_size): # 循环采样 if i + current_idx >= samples_len: current_idx = 0 if is_training: self.epoch += 1 tokens_in = enc_data[i + current_idx] tokens_out = dec_data[i + current_idx] # 编码器输入独热编码 seq_len_in = len(tokens_in) for t, token in enumerate(tokens_in): enc_oh_input_batch[i, t, token] = 1 # 填充剩余位置 enc_oh_input_batch[i, seq_len_in:, enc_space_token] = 1 # 标签序列独热编码(位置一一对应) seq_len_out = len(tokens_out) for t, token in enumerate(tokens_out): dec_oh_output_batch[i, t, token] = 1 dec_oh_output_batch[i, seq_len_out:, dec_space_token] = 1 current_idx += self.batch_size # 仅返回输入和标签,无需dec_input yield enc_oh_input_batch, dec_oh_output_batch
3. 损失与指标不匹配
当前使用的categorical_crossentropy和默认accuracy不适合二分类序列标注:
- 默认
accuracy是样本级多分类准确率,而任务需要的是每个时间步的二分类准确率 - 0.46的准确率大概率是随机猜测的结果(样本中元音占比约46%)
修改方案:
- 若标签为独热编码,改用
binary_crossentropy损失,搭配binary_accuracy指标 - 若标签为整数(0/1),改用
sparse_categorical_crossentropy损失,无需独热编码
编译示例:
model.compile(optimizer=Adam(learning_rate=0.001), loss='binary_crossentropy', metrics=['binary_accuracy'])
4. 训练配置不合理
- 学习率0.0001过小,初始学习率可调整为0.001,配合
ReduceLROnPlateau回调实现学习率衰减 - 批量大小30可增大至64/128,加快收敛速度
- 10轮epoch太少,建议增加至50-100轮,搭配
EarlyStopping回调防止过拟合
内容的提问来源于stack exchange,提问作者Lex Podgorny
相关产品推荐
相关产品推荐

