多分类任务可使用[3,3]形状独热标签搭配nn.CrossEntropyLoss训练吗
关于
nn.CrossEntropyLoss使用one-hot格式标签的说明 核心结论
默认场景下不支持直接传入你写的[3,3]形状one-hot矩阵作为标签,满足特定版本和格式要求时可以正常使用,具体规则如下:
- 所有PyTorch版本通用的稳妥写法,是使用一维类别索引作为标签
你提到的常规格式[1,2,3]其实存在错误:nn.CrossEntropyLoss的类别索引默认从0开始计数,对应3类花卉的场景,3个样本分别属于flowerA、flowerB、flowerC的话,正确的索引标签应该是形状为[3]的整型张量[0,1,2],和形状为[3,3]的模型输出(第0维是batch大小3,第1维是类别数3)完全匹配,直接传入损失即可正常计算。 - PyTorch 1.10及以上版本,支持传入和预测张量同形状的概率分布标签(包括one-hot硬标签)
你给出的3阶单位矩阵就是标准的one-hot标签,只要满足两个要求就能正常传入计算:- 标签张量必须是浮点类型(如
torch.float32),不能使用整型 - 标签形状和模型输出形状完全一致,即当前场景下的
[3,3],每一行对应一个样本的类别概率分布、行和为1
- 标签张量必须是浮点类型(如
踩坑提示:如果在1.10以下的老版本直接传入[3,3]形状的标签,会直接触发形状不匹配报错;哪怕是新版本,如果传入的是整型的[3,3]张量,框架会默认按照类别索引的逻辑解析,直接出现计算错误。
两种正确写法的代码示例
- 类别索引写法(全版本兼容,推荐常规场景使用)
import torch import torch.nn as nn # 模拟ResNet输出的3个样本、3类的预测结果 logits = torch.randn(3, 3, requires_grad=True) # 类别索引标签,dtype必须为long/int类型 labels = torch.tensor([0, 1, 2], dtype=torch.long) loss_fn = nn.CrossEntropyLoss() loss = loss_fn(logits, labels)
- one-hot标签写法(仅PyTorch≥1.10可用)
import torch import torch.nn as nn logits = torch.randn(3, 3, requires_grad=True) # 你构造的one-hot标签,必须用浮点类型 labels = torch.tensor( [[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=torch.float32 ) loss_fn = nn.CrossEntropyLoss() loss = loss_fn(logits, labels)
在标签是严格one-hot、无标签平滑的场景下,两种写法计算出的损失值完全一致,没有数值差异。
内容的提问来源于stack exchange,提问作者Lorrain
相关产品推荐
相关产品推荐

