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

层级LSTM自编码器训练无进展问题求助

段落层级自编码器训练无进展问题排查

我正在复现一篇段落层级自编码器论文,模型逻辑为:将段落拆分为句子,用LSTM编码每个句子,再将句子编码输入另一LSTM编码整个段落;解码器镜像该流程,先把段落编码解码为句子编码,再解码为单词,目标还原原始段落。

预处理后每个段落是形状为(maxSentence, maxWordsPerSentence, VocabSize)的独热编码张量,但模型完全不训练,loss始终不变,怀疑是loss计算或模型结构问题,以下是编码器、解码器及训练函数代码,求排查问题:

编码器代码

class Encoder(nn.Module):
    def __init__(self, input_dim, emb_dim, enc_hid_dim, dec_hid_dim, dropout):
        super().__init__()
        
        #self.embedding = nn.Embedding(input_dim, emb_dim)
        self.rnn_sent = nn.GRU(input_dim, enc_hid_dim, bidirectional = True)
        self.rnn_par = nn.GRU(enc_hid_dim*2, dec_hid_dim, bidirectional = True)

        
    def forward(self, src):
        
        outputs, hidden = self.rnn_sent(src[:,0,0])
        total_out = outputs.unsqueeze(0).permute(1,0,2)
        for i in range(1,src.shape[1]):
          for j in range(src.shape[2]):
            outputs, hidden = self.rnn_sent(src[:,i,j],hidden)
          total_out = torch.cat((total_out,outputs.unsqueeze(0).permute(1,0,2)),dim=1)  

        outputs_par, hidden_par = self.rnn_par(total_out[:,0])
        
        for i in range(total_out.shape[1]):
            outputs_par, hidden_par = self.rnn_par(total_out[:,i],hidden_par)

        return outputs_par, hidden_par

解码器代码

class Decoder(nn.Module):
    def __init__(self, output_dim, emb_dim, enc_hid_dim, dec_hid_dim, dropout, attention):
        super().__init__()

        self.output_dim = output_dim
        self.attention = attention
        #self.embedding = nn.Embedding(output_dim, emb_dim)
        self.rnn_par = nn.GRU((enc_hid_dim * 2), dec_hid_dim*2)
        self.rnn_sen = nn.GRU(output_dim, dec_hid_dim*2)
        self.fc_out = nn.Linear(dec_hid_dim*2, output_dim)
        self.dropout = nn.Dropout(dropout)
        
    def forward(self, input, hidden, encoder_outputs):
        output, hidden = self.rnn_par(encoder_outputs)
        all_par = output.unsqueeze(0).permute(1,0,2)
        for i in range(1,max_par_len):
          output,hidden = self.rnn_par(output,hidden)
          all_par = torch.cat((all_par,output.unsqueeze(0).permute(1,0,2)),dim=1)


        for i in range(max_par_len):
          output_arg = self.fc_out(all_par[:,i])
          #output_argmax = F.one_hot(output_arg.argmax(dim = 1), self.output_dim).to(torch.float)
          output_argmax = torch.softmax(output_arg,dim=1)
          output_sen, hidden_sen = self.rnn_sen(output_argmax)
          all_par_sen = output_argmax.unsqueeze(0).permute(1,0,2)
          for j in range(max_sen_len - 1):
            output_sen,hidden_sen = self.rnn_sen(output_argmax,hidden_sen)
            output_arg = self.fc_out(output_sen)
            output_argmax = torch.softmax(output_arg,dim=1)

            all_par_sen = torch.cat((all_par_sen,output_argmax.unsqueeze(0).permute(1,0,2)),dim=1)
          if i == 0:
            all_doc = all_par_sen.unsqueeze(0).permute(1,0,2,3)
          else:
            all_doc = torch.cat((all_doc,all_par_sen.unsqueeze(0).permute(1,0,2,3)),dim=1)

          i+=1
        return all_doc  ,hidden_sen

训练函数代码

def train(model, iterator, optimizer, criterion, clip, epoch):

    model.train()

    epoch_loss = 0
    data = tqdm(iterator)
    for i, batch in enumerate(data):
        src = batch[0].to(device)#.to(torch.long)#.reshape(batch[0].shape[0],-1)
        trg = batch[0].to(device)#.to(torch.long)#.reshape(batch[0].shape[0],-1)
        target = torch.argmax(trg,dim=3).view(-1)
        print(target)
        optimizer.zero_grad()
        output = model(src, trg).view(-1,OUTPUT_DIM)
        loss = criterion(output, target)        
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
        optimizer.step()
        epoch_loss += loss.item()


N_EPOCHS = 20
CLIP = 1
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
criterion = nn.CrossEntropyLoss(ignore_index = vocabulary['<pad>'])
best_valid_loss = float('inf')

for epoch in range(N_EPOCHS):

    start_time = time.time()
    train_loader, valid_loader = data_loaders['train_loader'], data_loaders['test_loader']
    train_loss = train(model, train_loader, optimizer, criterion, CLIP,f'{epoch+1}/{N_EPOCHS}')
    #valid_loss = evaluate(model, valid_loader, criterion)

    end_time = time.time()

    epoch_mins, epoch_secs = epoch_time(start_time, end_time)

    print(f'Epoch: {epoch+1:02} | Time: {epoch_mins}m {epoch_secs}s')
    print(f'\tTrain Loss: {train_loss:.3f} | Train PPL: {math.exp(train_loss):7.3f}')

核心问题排查与修正方案

一、编码器逻辑错误

  1. 句子编码完全违背序列模型设计
    原代码逐单词喂入句子GRU,且重复复用hidden状态,完全无法捕捉句子的序列依赖。正确做法是将整个句子序列输入GRU:

    # 假设src形状为 (batch_size, max_sentences, max_words, vocab_size)
    batch_size = src.shape[0]
    sent_encodings = []
    for i in range(src.shape[1]):
        # 调整句子维度为GRU要求的(seq_len, batch_size, input_dim)
        sentence = src[:, i, :, :].permute(1, 0, 2)
        _, hidden = self.rnn_sent(sentence)
        # 拼接双向GRU的hidden状态作为句子编码
        sent_encoding = torch.cat([hidden[0], hidden[1]], dim=1)
        sent_encodings.append(sent_encoding)
    # 转换为段落GRU需要的输入形状
    sent_encodings = torch.stack(sent_encodings).permute(1, 0, 2)
    # 用完整的句子编码序列训练段落GRU
    _, hidden_par = self.rnn_par(sent_encodings.permute(1, 0, 2))
    par_hidden = torch.cat([hidden_par[0], hidden_par[1]], dim=1)
    return sent_encodings, par_hidden
    
  2. 段落编码逻辑错误
    原代码逐句喂入段落GRU,正确做法是将完整的句子编码序列作为输入,一次性喂入段落GRU。

二、解码器逻辑错误

  1. 硬编码全局变量导致维度不匹配
    代码中max_par_len和max_sen_len未定义,需从输入张量或模型初始化参数中动态获取,避免维度错误。

  2. 句子解码无自回归逻辑
    原代码反复复用同一个output_argmax作为输入,完全无法学习单词序列依赖。正确的自回归解码应用上一步的预测作为下一步输入:

    # 解码单个句子示例
    current_input = initial_input  # 可使用<sos>标记或句子编码初始输入
    sentence_outputs = []
    for j in range(max_sen_len):
        output_sen, hidden_sen = self.rnn_sen(current_input.unsqueeze(0), hidden_sen)
        pred = self.fc_out(output_sen.squeeze(0))
        sentence_outputs.append(pred)
        # 训练阶段建议用teacher forcing,即使用真实单词作为下一个输入
        current_input = trg[:, i, j, :]  # trg为真实段落张量
    sentence_outputs = torch.stack(sentence_outputs, dim=1)
    
  3. 段落解码未利用编码器的hidden状态
    解码器的段落GRU应使用编码器输出的段落hidden作为初始状态,而非直接喂入编码器输出,遵循自编码器的镜像结构。

  4. 冗余的Softmax导致Loss计算错误
    CrossEntropyLoss默认输入为未经过Softmax的logits,原代码中解码器的Softmax会导致梯度消失,需去掉该操作,直接将logits传入损失函数。

三、训练函数与Loss问题

  1. Loss维度匹配问题
    需确保模型输出的logits形状与target完全匹配:output应为(batch_size * max_sentences * max_words, vocab_size),target应为(batch_size * max_sentences * max_words)。

  2. 优化器与学习率选择不当
    SGD学习率0.01对序列模型收敛过慢,建议替换为Adam优化器,初始学习率设为0.001。

  3. 训练函数未返回平均Loss
    train函数末尾需添加return epoch_loss / len(iterator),否则主循环中train_loss为None,无法正确监控训练状态。

四、其他细节问题

  1. 独热编码效率极低
    独热编码输入维度等于词汇量,会导致模型参数爆炸,应启用注释掉的Embedding层,将单词索引映射为低维嵌入向量。

  2. Dropout未启用
    编码器和解码器中定义的Dropout未在forward函数中使用,需添加x = self.dropout(x)实现正则化。

  3. Attention机制未集成
    解码器初始化传入的attention参数未使用,若论文包含Attention机制,需完成实现并集成到解码流程中。


内容的提问来源于stack exchange,提问作者sagi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 20:27:33