您正在寻找 torch.topk 函数,该函数可计算维度上的前 k 个值。
torch.topk 的第二个输出是“arg top k”:前值的 k 个索引。
这是在语义分割上下文中的使用方法:
假设您有形状为b-h-w (dtype=torch.int64) 的基本事实预测张量y。
您的模型预测形状 b-c-h-w 的每像素类 logits,c 是类的数量(包括“背景”)。这些 logits 是 之前 softmax 函数将它们转换为类概率的“原始”预测。
由于我们只查看顶部的k,因此预测是“原始”还是“概率”并不重要。
# compute the top k predicted classes, per pixel:
_, tk = torch.topk(logits, k, dim=1)
# you now have k predictions per pixel, and you want that one of them will match the true labels y:
correct_pixels = torch.eq(y[:, None, ...], tk).any(dim=1)
# take the mean of correct_pixels to get the overall average top-k accuracy:
top_k_acc = correct_pixels.mean()
请注意,此方法不考虑“忽略”像素。这可以通过对上述代码稍作修改来完成:
valid = y != ignore_index
top_k_acc = correct_pixels[valid].mean()