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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 03:10:55