You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 09:15:03