使用xlm-roberta-large-longformer时CrossEntropyLoss报错如何解决?
问题解决:CrossEntropyLoss使用错误修复
错误根源分析
你的代码存在两个核心问题,直接导致了损失计算失败:
- 提前对logits做softmax:
torch.nn.CrossEntropyLoss内部已经集成了LogSoftmax和NLLLoss的计算逻辑,提前手动softmax会破坏损失的数值稳定性和计算逻辑。 - 目标张量格式错误:
CrossEntropyLoss不接受one-hot编码的目标(形状[-1, num_labels]),它要求目标是类别索引(形状[-1],每个元素是0到num_labels-1的整数),这也是你转long类型后出现multi-target not supported错误的原因。
修复后的代码
根据你的b_labels格式,分两种情况处理:
情况1:b_labels是one-hot编码格式(如形状[batch_size, seq_len, num_labels])
import torch loss_func = torch.nn.CrossEntropyLoss() # 直接使用原始logits,调整形状为[-1, num_labels],无需softmax logits_flat = logits.view(-1, num_labels) # 将one-hot标签转为类别索引,调整形状为[-1]并转为long类型 target_flat = torch.argmax(b_labels, dim=-1).view(-1).long() # 计算损失 loss = loss_func(logits_flat, target_flat) train_loss_set.append(loss.item())
情况2:b_labels本身就是类别索引格式(如形状[batch_size, seq_len])
import torch loss_func = torch.nn.CrossEntropyLoss() logits_flat = logits.view(-1, num_labels) # 直接调整形状为[-1]并转为long类型 target_flat = b_labels.view(-1).long() loss = loss_func(logits_flat, target_flat) train_loss_set.append(loss.item())
关键注意点
- 永远不要给
CrossEntropyLoss传入经过softmax的输出,必须直接用模型输出的原始logits。 - 目标张量必须是单类别索引,不能是one-hot编码,否则会被模型判定为多目标输入,触发不支持的错误。
内容的提问来源于stack exchange,提问作者Niloufar Modir
相关产品推荐
相关产品推荐

