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

机器翻译LSTM模型训练:Loss骤降但BLEU始终为0的问题排查

机器翻译LSTM模型训练异常:BLEU骤降为0但Loss持续下降

问题描述

训练机器翻译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,应该传递正确的初始细胞状态。
    修复:

    1. 修改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
      
    2. 修改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
      
    3. 同步修改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():
            break
    
  • Loss计算未忽略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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 07:20:33