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

如何基于稀疏张量计算CrossEntropyLoss?

如何基于稀疏张量计算CrossEntropyLoss?

你遇到的问题其实是PyTorch原生的CrossEntropyLoss目前确实不支持直接输入稀疏张量——不管是作为logits的输入还是目标标签输入都不行。不过结合你提到的场景(大batch但目标数据非常稀疏),我们可以通过只针对有效标签位置计算损失的方式来绕开这个限制,同时避免把稀疏张量转成密集张量导致内存爆炸。

核心思路

既然你的目标标签只有少数位置是有效的,我们完全不需要处理整个batch的所有空间位置,只需要:

  1. 从稀疏目标张量中提取出那些有有效标签的位置索引和对应的类别值
  2. 从logits张量(不管是密集还是稀疏)中提取这些位置对应的所有类别的logits值,形成一个小批量的logits集合
  3. 用原生的CrossEntropyLoss计算这个小批量的损失即可

代码实现(针对稀疏目标标签场景)

下面是贴合你场景的代码示例,假设你的目标标签只有少量有效位置,其余为无效值(比如用-1标记):

import torch
from torch.nn import CrossEntropyLoss

# 模拟你的大batch场景:batch=3000,假设h/w很大(比如1000),密集存储会爆内存
batch_size = 3000
num_classes = 20
h, w = 1000, 1000

# 生成logits(可以是密集张量,也可以是稀疏张量)
logits = torch.randn(batch_size, num_classes, h, w)

# 生成稀疏目标:大部分位置是无效标签(-1),仅随机选1000个位置有有效标签
dense_target = torch.full((batch_size, h, w), -1, dtype=torch.long)
valid_pos_flat = torch.randint(0, batch_size*h*w, (1000,))
dense_target.view(-1)[valid_pos_flat] = torch.randint(0, num_classes, (1000,))
# 转成稀疏张量,只保留有效标签的位置
mask = dense_target != -1
sparse_target = dense_target[mask].to_sparse_coo(indices=torch.nonzero(mask, as_tuple=False).T)

# 从稀疏目标中提取有效位置和标签
pos_indices = sparse_target.indices()  # 形状(3, N):分别是batch索引、高度索引、宽度索引
labels = sparse_target.values()        # 形状(N,):对应位置的类别标签

# 提取这些有效位置对应的所有类别的logits,得到形状(N, num_classes)的小批量logits
selected_logits = logits[pos_indices[0], :, pos_indices[1], pos_indices[2]]

# 计算损失
loss_fn = CrossEntropyLoss()
# 如果你的类别是从1开始的(比如你代码里的randint(1,20)),记得设置ignore_index忽略0或者调整标签到0开始
# loss_fn = CrossEntropyLoss(ignore_index=0)
loss = loss_fn(selected_logits, labels)

print(f"计算得到的损失值:{loss.item()}")

如果logits也是稀疏张量的处理方式

如果你的logits本身也是稀疏存储的(比如大量类别的logit值为0),可以先初始化一个全0的selected_logits张量(因为稀疏张量的默认值是0),然后把稀疏logits中对应有效位置的值填充进去,再计算损失。不过如果logits的稀疏度不是极高,直接把logits转成密集张量再提取有效位置反而更简单——毕竟小批量的logits(比如N=1000,num_classes=20,总共20000个元素)内存占用几乎可以忽略。

注意事项

  • 确保目标标签的类别索引合法:CrossEntropyLoss要求类别索引从0到num_classes-1,如果你的标签是从1开始的,记得要么调整标签到0起始,要么设置ignore_index参数忽略无效的0值。
  • 稀疏张量的索引格式要正确:目标稀疏张量的索引应该对应(batch, height, width)三维,values是对应的类别值。

备注:内容来源于stack exchange,提问作者MaKaNu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 13:28:02