PyTorch中如何索引获取像素分类任务各像素对应类别的概率值
正确提取对应类别像素概率的方法
问题原因
你原来直接用prob[y]的写法不符合PyTorch的张量索引规则:
- 你的概率张量
prob维度为[批量大小, 类别数, 高度, 宽度](对应你的示例是[2,3,5,5]) - 直接传入类别张量
y做索引时,PyTorch默认会优先索引第0维(批量维度),而你y里的取值是0-2,批量维度只有0、1两个取值,不仅索引逻辑不对,还可能触发越界报错。
正确实现方式
最简洁高效的方案是使用torch.gather方法,指定在类别维度(dim=1)上按标签索引取值:
import torch import torch.nn.functional as F y = torch.randint(0, 3, (2, 1, 5, 5)) # 标签维度 [B,1,H,W] logits = torch.randn(2, 3, 5, 5) prob = F.softmax(logits, dim=1) # 概率维度 [B,C,H,W] # 在类别维度(dim=1)上按y的索引取值,输出维度和y一致 [B,1,H,W] target_prob = torch.gather(prob, dim=1, index=y) # 如果不需要多余的单通道维度,可以压缩为 [B,H,W] target_prob = target_prob.squeeze(1)
如果你习惯用onehot编码的方式实现,也可以用如下写法,最终效果和gather一致:
# 把标签转成onehot格式,调整维度和prob对齐 y_onehot = F.one_hot(y.squeeze(1), num_classes=3).permute(0,3,1,2) # 对应位置相乘后对类别维度求和 target_prob = (prob * y_onehot).sum(dim=1)
内容的提问来源于stack exchange,提问作者sachinruk
相关产品推荐
相关产品推荐

