PyTorch如何统计Tensor内各唯一元素的出现次数
PyTorch原生统计张量元素出现次数方案
PyTorch有多个原生API可以实现该需求,无需引入Python标准库的collections.Counter,避免了张量转Python可迭代对象的拷贝开销,还能直接在GPU上运行,更适配深度学习训练、混淆矩阵构建等场景,常用方案如下:
- 方案1:
torch.bincount()(最优,适合非负连续整数统计)
这个API专门用于统计非负整数张量中各数值的出现频次,返回结果的索引对应原张量的数值,值对应出现次数,刚好匹配你的需求:
如果是高维张量,先拉平再传入即可:import torch y = torch.randint(0, 3, (10,)) counts = torch.bincount(y) # counts[0] 为0的出现次数,counts[1]为1的出现次数,counts[2]为2的出现次数counts = torch.bincount(y.flatten()) - 方案2:
torch.unique()+返回计数(适合数值范围不确定的场景)
如果你不知道张量里的数值范围,可以用这个方法同时拿到去重后的数值和对应频次:unique_vals, counts = torch.unique(y, return_counts=True) - 方案3:布尔掩码求和(适合自定义统计指定数值)
如果只需要统计某几个特定数值的出现次数,直接用布尔掩码加求和即可:count_0 = torch.sum(y == 0).item() count_1 = torch.sum(y == 1).item() count_2 = torch.sum(y == 2).item()
collections.Counter虽然能实现功能,但需要将张量转成Python列表/迭代器,大张量或者GPU张量场景下性能损失较大,更推荐用上述原生实现。
内容的提问来源于stack exchange,提问作者sachinruk
相关产品推荐
相关产品推荐

