PyTorch中如何获取掩码里第一个0的索引?
获取张量中第一个0的最优实现方式
假设是NumPy张量
情况1:张量是有序的(如示例中前全1后全0)
这种情况用二分查找效率最高,时间复杂度O(logn):
import numpy as np arr = np.array((1, 1, 1, 1, 1, 1, 1, 1, 0, 0)) # 将数组转为布尔数组,找第一个True的位置 first_zero_idx = np.searchsorted(arr == 0, True) print(first_zero_idx) # 输出8
情况2:张量是无序的
用np.flatnonzero直接取第一个匹配项,避免生成完整布尔数组后冗余遍历:
import numpy as np arr = np.array((1, 1, 1, 1, 1, 1, 1, 1, 0, 0)) first_zero_idx = np.flatnonzero(arr == 0)[0] print(first_zero_idx) # 输出8
若要处理“无0”的边界情况,可补充判断:
zero_indices = np.flatnonzero(arr == 0) first_zero_idx = zero_indices[0] if len(zero_indices) > 0 else -1
假设是PyTorch张量
情况1:张量是有序的
同样用二分查找优化性能:
import torch t = torch.tensor((1, 1, 1, 1, 1, 1, 1, 1, 0, 0)) # 转为浮点型布尔张量后,用searchsorted找第一个1的位置 first_zero_idx = torch.searchsorted((t == 0).float(), torch.tensor([1.0])).item() print(first_zero_idx) # 输出8
情况2:张量是无序的
用nonzero结合索引获取第一个结果,兼顾效率与可读性:
import torch t = torch.tensor((1, 1, 1, 1, 1, 1, 1, 1, 0, 0)) zero_indices = (t == 0).nonzero(as_tuple=True)[0] first_zero_idx = zero_indices[0].item() if len(zero_indices) > 0 else -1 print(first_zero_idx) # 输出8
核心思路总结
- 若张量有序,优先用二分查找(
searchsorted),大数据量下性能优势显著; - 若张量无序,用专门的非零索引函数直接取第一个匹配项,比手动遍历更简洁高效;
- 所有实现建议加入边界判断,避免无0时抛出索引错误。
内容的提问来源于stack exchange,提问作者Foobar
相关产品推荐
相关产品推荐

