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

多分类任务可使用[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标签,只要满足两个要求就能正常传入计算:
    1. 标签张量必须是浮点类型(如torch.float32),不能使用整型
    2. 标签形状和模型输出形状完全一致,即当前场景下的[3,3],每一行对应一个样本的类别概率分布、行和为1

踩坑提示:如果在1.10以下的老版本直接传入[3,3]形状的标签,会直接触发形状不匹配报错;哪怕是新版本,如果传入的是整型的[3,3]张量,框架会默认按照类别索引的逻辑解析,直接出现计算错误。

两种正确写法的代码示例

  1. 类别索引写法(全版本兼容,推荐常规场景使用)
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)
  1. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 14:57:17