PyTorch多标签分类使用损失函数报错多目标不支持该如何处理
错误原因
你当前使用的nn.NLLLoss、F.cross_entropy都是为单标签多分类任务设计的,要求输入的标签是一维的类别索引张量,每个样本只能对应一个类别,因此传入多标签格式的多维标签时会抛出multi-target not supported错误。
解决方案
PyTorch提供了原生支持多标签分类的损失函数,你可以按以下步骤修改:
1. 选择适配的损失函数
多标签分类任务默认每个类别独立、不互斥,常用以下两类内置损失:
nn.BCEWithLogitsLoss(推荐):输入为模型输出的未激活logits,内部自动完成Sigmoid激活和二分类交叉熵计算,数值稳定性更好nn.BCELoss:输入为经过Sigmoid激活后的概率值,需要自己在模型输出层加Sigmoid激活
2. 代码修改示例
首先调整模型输出层:把原来的Softmax激活替换为Sigmoid,或者直接输出未激活的logits。
损失函数修改如下:
# 推荐用法:用BCEWithLogitsLoss,无需手动加Sigmoid def loss(self, pred, y_true): # 不要重复调用torch.tensor()包装pred和y_true,会导致计算图断开无法反向传播 loss_fn = nn.BCEWithLogitsLoss() # 多标签需要标签为浮点型,因此加.float()转换 return loss_fn(pred, y_true.float())
如果你的模型输出层已经加了Sigmoid,用以下写法:
def loss(self, pred, y_true): loss_fn = nn.BCELoss() return loss_fn(pred, y_true.float())
3. 注意事项
- 多标签分类的标签y_true需要是
[batch_size, num_classes]维度的多热编码,对应类别的位置为1、其余为0,如果你的原始标签是类别索引列表,需要先转换为该格式 - 不要在多标签任务中使用CrossEntropyLoss,该损失内部会做Softmax计算,默认所有类别互斥,不符合多标签可同时为正的需求
内容的提问来源于stack exchange,提问作者Newbie
相关产品推荐
相关产品推荐

