机器翻译LSTM模型训练:Loss骤降但BLEU始终为0的问题排查
问题描述
训练机器翻译LSTM模型时,首批次训练后验证集BLEU分数直接降至0并全程保持该值,但Loss却在持续大幅下降。此前怀疑是LSTM细胞状态传递问题,修复后仍未解决。
模型代码
import torch import torch.nn as nn class SimpleRNNTranslator(nn.Module): def __init__(self, inp_voc, out_voc, emb_size=64, hid_size=128): """ My version of simple RNN model, I use LSTM instead of GRU as in the baseline """ super().__init__() self.inp_voc = inp_voc self.out_voc = out_voc self.emb_inp = nn.Embedding(len(inp_voc), emb_size) self.emb_out = nn.Embedding(len(out_voc), emb_size) self.encoder = nn.LSTM(emb_size, hid_size, batch_first=True) self.decoder = nn.LSTM(emb_size, hid_size, batch_first=True) self.decoder_start = nn.Linear(hid_size, hid_size) self.logits = nn.Linear(hid_size, len(out_voc)) def forward(self, inp, out): """ Apply model in training mode """ encoded_seq = self.encode(inp) decoded_seq, _ = self.decode(encoded_seq, out) return self.logits(decoded_seq) def encode(self, seq_in): """ Take input symbolic sequence, compute initial hidden state for decoder :param seq_in: matrix of input tokens [batch_size, seq_in_len] :return: initial hidden state for the decoder """ embeddings = self.emb_inp(seq_in) output, (_, __) = self.encoder(embeddings) # last state isn't the actually last because of the padding, the next 2 lines find out the true last state seq_lengths = (seq_in != self.inp_voc.eos_ix).sum(dim=-1) last_states = output[range(seq_lengths.shape[0]), seq_lengths] return self.decoder_start(last_states) def decode(self, hidden_state, seq_out, previous_state=None): """ Take output symbolic sequence, compute logits for every token in sequence :param hidden_state: matrix of initial_hidden_state [batch_size, hid_size] :param previous_state: matrix of previous state [batch_size, hid_size] :param seq_out: matrix of output tokens [batch_size, seq_out_len] :return: logits for every token in sequence [batch_size, seq_len, out_voc] """ if not torch.is_tensor(previous_state): previous_state = torch.randn(*hidden_state.shape).to(device) embeddings = self.emb_out(seq_out) outputs, (_, cn) = self.decoder(embeddings, (hidden_state[None, :, :], previous_state[None, :, :])) return outputs, cn def inference(self, inp_tokens, max_len): """ Take initial state and return ids for out words :param initial_state: initial_state for a decoder, produced by encoder with input tokens """ initial_state = self.encode(inp_tokens) states = [initial_state] outputs = [torch.full([initial_state.shape[0]], self.out_voc.bos_ix, dtype=torch.int, device=device)] cn = None for i in range(100): hidden_state, cn = self.decode(states[-1], outputs[-1][:, None], previous_state=cn) hidden_state, cn = hidden_state.squeeze(), cn.squeeze() outputs.append(self.logits(hidden_state).argmax(dim=-1)) states.append(hidden_state) return torch.stack(outputs, dim=-1), torch.cat(states) def translate_lines(self, lines, max_len=100): """ Take lines and return translation :param lines: list of lines in Russian """ inp_tokens = self.inp_voc.to_matrix(lines).to(device) out_ids, states = self.inference(inp_tokens, max_len=max_len) return self.out_voc.to_lines(out_ids.cpu().numpy()), states
BLEU计算代码
from nltk.translate.bleu_score import corpus_bleu def compute_bleu(model, inp_lines, out_lines, bpe_sep='@@ ', **flags): """ Estimates corpora-level BLEU score of model's translations given inp and reference out Note: if you're serious about reporting your results, use https://pypi.org/project/sacrebleu """ with torch.no_grad(): translations, _ = model.translate_lines(inp_lines, **flags) translations = [line.replace(bpe_sep, '') for line in translations] actual = [line.replace(bpe_sep, '') for line in out_lines] return corpus_bleu( [[ref.split()] for ref in actual], [trans.split() for trans in translations], smoothing_function=lambda precisions, **kw: [p + 1.0 / p.denominator for p in precisions] ) * 100
训练曲线表现
训练阶段验证集的BLEU分数在首批次后直接降至0并保持不变,同时Loss呈现持续大幅下降的趋势。
问题排查与修复方案
编码器最后状态索引错误
你的encode方法中,seq_lengths计算的是非EOS的token数量,但LSTM输出的索引从0开始,比如有效序列长度为5时,最后一个有效token的索引是4,而你用seq_lengths直接索引会取到第5位(padding位置),导致编码器输出的初始状态完全无效。
修复:将last_states的索引改为seq_lengths - 1:last_states = output[range(seq_lengths.shape[0]), seq_lengths - 1]解码器细胞状态初始化错误
当previous_state不存在时,你用随机值初始化细胞状态,这会导致解码器初始状态完全混乱,无法承接编码器的语义信息。同时编码器丢弃了LSTM的细胞状态c_n,应该传递正确的初始细胞状态。
修复:- 修改
encode方法,返回编码器的隐藏状态和细胞状态:def encode(self, seq_in): embeddings = self.emb_inp(seq_in) output, (h_n, c_n) = self.encoder(embeddings) seq_lengths = (seq_in != self.inp_voc.eos_ix).sum(dim=-1) - 1 # 修正索引 last_h = output[range(seq_lengths.shape[0]), seq_lengths] last_c = c_n.squeeze(0) # 单层LSTM时h_n/c_n形状是[1, batch, hid] return self.decoder_start(last_h), last_c - 修改
decode方法,使用编码器传递的细胞状态初始化,而非随机值:def decode(self, hidden_state, cell_state, seq_out, previous_state=None): if previous_state is None: previous_state = cell_state embeddings = self.emb_out(seq_out) outputs, (_, cn) = self.decoder(embeddings, (hidden_state[None, :, :], previous_state[None, :, :])) return outputs, cn - 同步修改
forward和inference方法的调用逻辑,传递细胞状态。
- 修改
推理阶段无终止条件
inference方法循环固定100次,没有检测EOS token来终止生成,可能导致翻译结果充斥大量无效token(如PAD),直接拉低BLEU分数。
修复:添加EOS检测,遇到EOS就停止生成:for i in range(max_len): hidden_state, cn = self.decode(states[-1], outputs[-1][:, None], previous_state=cn) hidden_state, cn = hidden_state.squeeze(), cn.squeeze() next_token = self.logits(hidden_state).argmax(dim=-1) outputs.append(next_token) states.append(hidden_state) # 检查是否所有样本都生成了EOS if (next_token == self.out_voc.eos_ix).all(): breakLoss计算未忽略PAD token
如果Loss计算时没有忽略PAD token,模型可能会优先学习预测PAD来降低Loss,导致翻译结果全是PAD,BLEU自然为0。
修复:使用ignore_index指定PAD的索引:criterion = nn.CrossEntropyLoss(ignore_index=self.out_voc.pad_ix)手动检查翻译结果
先输出几个模型的翻译结果,看是否是全PAD、全重复token或完全无关内容,这能快速定位问题根源。比如如果全是PAD,那大概率是Loss没忽略PAD或者编码器状态错误。
内容的提问来源于stack exchange,提问作者Haski

