PyTorch如何获取Tensor中不全为0的行对应的索引下标
PyTorch提取非全0行索引的最优实现方案
你当前使用的torch.unique(torch.nonzero(y>0,as_tuple=True)[0])虽然可以得到正确结果,但存在冗余开销:torch.nonzero会为所有非零元素单独输出对应的行索引,后续还要通过unique去重,当张量规模大、非零元素多的时候,性能损耗非常明显。
推荐最优实现
直接逐行判断是否存在非零元素,一步得到目标索引:
indices = torch.where(y.any(dim=1))[0]
你也可以用nonzero替代where,性能没有明显差异:
indices = torch.nonzero(y.any(dim=1), as_tuple=True)[0]
方案说明
y.any(dim=1)会沿着行维度(dim=1)判断每一行是否存在至少一个非零值,直接输出和行数等长的布尔张量,对应位置为True代表该行非全0。和求和再判断的方案相比,any只要找到该行第一个非零值就会停止校验,运算效率更高torch.where/torch.nonzero直接提取所有为True的位置的索引,没有重复值,不需要额外去重操作,整体运算效率远高于原有方案
兼容场景补充
如果你的张量存在负数,需要判断所有不全为0的行(不限定非零值为正),可以调整为:
indices = torch.where(y.abs().any(dim=1))[0]
效果验证
针对你给出的示例张量:
import torch y = torch.tensor([ [0., 0., 0., 0.], [0., 1., 1., 0.], [0., 0., 0., 0.], [0., 0., 1., 0.], [0., 0., 0., 0.], [0., 0., 1., 0.], [1., 0., 0., 1.], [0., 0., 0., 0.] ]) indices = torch.where(y.any(dim=1))[0] print(indices) # 输出:tensor([1, 3, 5, 6])
完全符合要求的输出格式。
内容的提问来源于stack exchange,提问作者lima0
相关产品推荐
相关产品推荐

