PyTorch带注意力的LSTM模型训练CrossEntropyLoss维度不匹配如何修复
错误原因
- CrossEntropyLoss参数顺序传反:PyTorch的
nn.CrossEntropyLoss要求第一个参数是模型输出的logits,第二个参数是真实标签,你代码里把两个参数写反了。 - 错误对预测结果取索引:你在损失函数里给
y_pred_argmax.float()加了[0],把形状为[32, 150]的预测结果取了第一个batch的内容,变成了形状为[150]的张量,和真实标签的batch维度不匹配。 - 模型输出不符合损失要求:
nn.CrossEntropyLoss输入要求是未经过softmax的原始logits,你在模型forward里提前做了softmax+argmax操作,不仅会导致梯度消失无法训练,输出的类别索引也完全不符合交叉熵损失的输入要求。 - 标签类型错误:交叉熵损失的真实标签要求是long类型的类别索引,你把它转成了float类型,也会导致计算异常。
- 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
相关产品推荐
相关产品推荐

