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

PyTorch带注意力的LSTM模型训练CrossEntropyLoss维度不匹配如何修复

错误原因

  1. CrossEntropyLoss参数顺序传反:PyTorch的nn.CrossEntropyLoss要求第一个参数是模型输出的logits,第二个参数是真实标签,你代码里把两个参数写反了。
  2. 错误对预测结果取索引:你在损失函数里给y_pred_argmax.float()加了[0],把形状为[32, 150]的预测结果取了第一个batch的内容,变成了形状为[150]的张量,和真实标签的batch维度不匹配。
  3. 模型输出不符合损失要求:nn.CrossEntropyLoss输入要求是未经过softmax的原始logits,你在模型forward里提前做了softmax+argmax操作,不仅会导致梯度消失无法训练,输出的类别索引也完全不符合交叉熵损失的输入要求。
  4. 标签类型错误:交叉熵损失的真实标签要求是long类型的类别索引,你把它转成了float类型,也会导致计算异常。
  5. mask逻辑错误:你用预测结果做padding mask,正确应该用真实标签y_true来判断padding位置。

修复步骤

1. 修改模型forward代码

删除多余的softmax和argmax操作,直接返回线性层输出的原始logits:

def forward(self, enc_inputs, dec_inputs) : 
    enc_hidden = self.embedding(enc_inputs)
    dec_hidden = self.embedding(dec_inputs)
    enc_hidden , (enc_h_state,enc_c_state) = self.enc_lstm(enc_hidden)
    dec_hidden,(dec_h_state,dec_c_state) = self.dec_lstm(dec_hidden,(enc_h_state,enc_c_state))
    attn_score = torch.matmul(dec_hidden, torch.transpose(enc_hidden,2,1))
    attn_prob  = self.soft_prob(attn_score)
    attn_out = torch.matmul(attn_prob,enc_hidden)
    cat_hidden = torch.cat((attn_out, dec_hidden),-1)
    # 直接返回线性层输出的logits,形状为 [batch_size, seq_len, vocab_size] = [32,150, len(vocab)]
    y_pred = self.softmax_linear(cat_hidden)
    return y_pred

2. 修改损失函数代码

调整参数顺序、张量形状、mask逻辑和数据类型:

def lm_loss(y_true, y_pred):
    criterion = nn.CrossEntropyLoss(reduction="none")
    # 调整logits形状为 [batch_size*seq_len, vocab_size] 符合交叉熵输入要求
    y_pred = y_pred.view(-1, y_pred.size(-1))
    # 调整标签形状为 [batch_size*seq_len],转成long类型
    y_true = y_true.view(-1).long()
    # 注意参数顺序:logits在前,标签在后
    loss = criterion(y_pred, y_true)
    # 用真实标签判断padding位置
    mask = torch.not_equal(y_true, 0).type(torch.FloatTensor).to(device)
    loss *= mask
    loss = torch.sum(loss) / torch.maximum(torch.sum(mask), 1)
    return loss

3. 训练代码无需改动

你注释掉的y_pred = torch.argmax(y_pred,dim =-1)保持注释状态即可,训练阶段不需要对logits做argmax。

内容的提问来源于stack exchange,提问作者Yun-Ho Jang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 02:36:05