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

使用xlm-roberta-large-longformer时CrossEntropyLoss报错如何解决?

问题解决:CrossEntropyLoss使用错误修复

错误根源分析

你的代码存在两个核心问题,直接导致了损失计算失败:

  1. 提前对logits做softmax:torch.nn.CrossEntropyLoss内部已经集成了LogSoftmax和NLLLoss的计算逻辑,提前手动softmax会破坏损失的数值稳定性和计算逻辑。
  2. 目标张量格式错误: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 08:00:28