如何在torch.any中指定多维度?PyTorch语义分割张量判断问题
语义分割张量的False值判断方案
一、最简洁的实现方式
针对形状为(B×H×W)的布尔张量,判断每个样本是否包含False值,有以下几种简洁实现:
利用
torch.all直接判断(推荐)
因为"存在False"等价于"并非全为True",而torch.all支持传入元组类型的维度参数,一步就能得到结果:has_false = ~torch.all(input, dim=(1, 2))多次调用
torch.any
如果你的PyTorch版本不支持多维度参数,可分两次对H、W维度做归约:has_false = ~torch.any(torch.any(input, dim=1), dim=1)展平维度后调用
torch.any
先将每个样本的H×W维度展平为一维,再做归约判断:has_false = ~torch.any(input.flatten(start_dim=1), dim=1)
二、为什么早期torch.any不支持多维度参数?
在PyTorch 1.10之前的版本,torch.any确实仅支持单个int类型的dim参数,原因主要有三点:
- 历史实现遗留:早期PyTorch的算子设计以单维度归约为基础,多维度归约属于后续扩展功能。
- 底层优化难度:多维度归约需要适配不同设备(CPU/GPU)的底层计算逻辑,实现和优化成本更高,因此优先完善单维度核心功能。
- 替代方案存在:用户可通过多次单维度归约实现相同效果,所以多维度参数的优先级相对较低。
不过在PyTorch 1.10及以上版本,torch.any已经支持传入元组类型的dim参数,升级后直接使用~torch.any(input, dim=(1, 2))即可得到结果。
内容的提问来源于stack exchange,提问作者artyomuiii
相关产品推荐
相关产品推荐

