为何PyTorch内置损失函数要求targets为Long Tensor而非32位整数?
为什么
torch.nn.CrossEntropyLoss不支持32位整数类型的target? 核心原因主要有三点:
- 框架内索引操作的统一约定
PyTorch中所有涉及张量索引、取值的操作,原生默认要求索引张量为64位有符号整数(torch.long/int64)。分类场景下的target本质是类别索引,CrossEntropyLoss内部需要用该索引从输出的概率张量中取出对应类别的置信度做对数似然计算,统一使用int64类型可以避免内部频繁做类型转换,降低运算开销。 - 全场景兼容的溢出风险规避
虽然32位整数(torch.int/int32)最多可支持21亿类别的分类任务,足够覆盖绝大多数常规场景,但PyTorch作为通用深度学习框架,需要兼容超大类别分类、自定义索引映射、以及和其他算子(如nn.Embedding、NLLLoss)的链路串联,统一使用int64可以完全杜绝整数溢出导致的计算错误。 - 跨设备计算一致性保证
CrossEntropyLoss对应的底层CUDA核函数、CPU优化实现,都优先针对int64类型的索引做了性能适配。强制要求int64类型可以保证同一段代码在CPU、GPU、甚至其他加速设备上的计算结果完全一致,不会因为整数类型的差异出现精度偏差或者执行错误。
快速解决方案
如果你的target默认是int32类型,只需要调用.long()方法做一次类型转换即可,几乎不会产生额外的计算开销:
import torch loss_fn = torch.nn.CrossEntropyLoss() # 转换target为int64类型 targets = targets.long() loss = loss_fn(predictions, targets)
内容的提问来源于stack exchange,提问作者Joker
相关产品推荐
相关产品推荐

