5分类任务多ground truth标签场景下,适用损失函数咨询
适合该分类场景的损失函数方案
你的场景属于单样本对应多个可接受的正确类别:只要模型预测的argmax结果落在ground truth给出的类别集合中,就算预测正确。以下是几种适配的损失函数方案,按贴合度排序:
1. 目标类别最大得分的负对数损失(最贴合需求)
这种损失直接对齐你想要的判定逻辑:引导模型让ground truth类别中的最高得分尽可能高。当模型的argmax是目标类别之一时,该类别的得分就是目标类别中的最大值,损失会显著降低。
实现逻辑:
- 从预测向量中提取所有ground truth类别对应的得分
- 取这些得分的最大值
- 计算该最大值的负对数作为损失
示例代码(PyTorch):
import torch def max_target_score_loss(pred, ground_truth): # pred: 形状(1,5)的softmax得分张量 # ground_truth: 形状(3,)的类别索引张量(注意转换为模型输出对应的索引,比如原类别3对应索引2) # 提取目标类别的得分 target_scores = pred.gather(dim=1, index=ground_truth.unsqueeze(0)) # 取目标类别得分的最大值 max_score = target_scores.max() # 返回负对数损失 return -torch.log(max_score)
用你给出的示例计算:predicted_vector = tensor([0.0669, 0.1336, 0.3400, 0.3392, 0.1203])ground_truth转换为索引是tensor([2,1,4]),对应得分是[0.3400, 0.1336, 0.1203],最大值是0.34,损失为-log(0.34)≈1.07。如果模型后续把类别2的得分提升到0.5(成为argmax),损失会降到-log(0.5)≈0.69,完全符合“奖励正确预测”的需求。
2. 目标类别平均得分的负对数损失
这种损失会引导模型提升所有ground truth类别的平均得分,适合需要模型同时关注多个目标类别的场景(比如这些类别存在关联)。
实现逻辑:
- 提取目标类别的得分
- 计算得分的平均值
- 取负对数作为损失
示例代码(PyTorch):
import torch def mean_target_score_loss(pred, ground_truth): target_scores = pred.gather(dim=1, index=ground_truth.unsqueeze(0)) mean_score = target_scores.mean() return -torch.log(mean_score)
3. 软标签交叉熵损失
将ground truth转换为“软标签”向量:所有目标类别位置分配相等的概率(比如3个目标类别就每个设为1/3,其余为0),然后用常规交叉熵损失。这种方式会让模型学习向目标类别集合分配概率,间接实现“只要一个类别得分最高就算对”的效果。
实现逻辑:
- 创建和预测向量同形状的软标签张量,目标类别位置设为
1/目标类别数量,其余为0 - 计算预测向量和软标签的交叉熵损失
示例代码(PyTorch):
import torch.nn.functional as F def soft_label_ce_loss(pred, ground_truth): num_classes = pred.shape[1] num_targets = len(ground_truth) # 创建软标签 soft_target = torch.zeros_like(pred) soft_target[:, ground_truth] = 1.0 / num_targets # 计算交叉熵 return F.cross_entropy(pred, soft_target)
内容的提问来源于stack exchange,提问作者helloworld
相关产品推荐
相关产品推荐

