如何过滤张量中全列非零元素的行并获取被删行索引?
解决方案:保留全非零行并获取删除行索引
针对你需要保留所有列元素均非零的行,同时获取被删除行索引的需求,以下是适配任意行列数的实现方案,分PyTorch和NumPy两种常用场景:
PyTorch 实现
核心逻辑是通过按行判断所有列元素是否非零生成布尔掩码,再用掩码过滤张量、提取删除索引:
import torch input_tensor = torch.tensor([ [-0.5535, 0.0000], [ 0.0000, 0.0000], [-1.1370, -0.2736], [-1.2300, 0.9185] ]) # 生成每行全非零的布尔掩码(dim=1表示按行判断所有列) mask = torch.all(input_tensor != 0, dim=1) # 过滤得到保留的行 filtered_tensor = input_tensor[mask] # 获取被删除的行索引(~mask取反掩码,torch.where返回索引张量) deleted_indices = torch.where(~mask)[0] print("过滤后的张量:") print(filtered_tensor) print("被删除的行索引:") print(deleted_indices)
输出结果:
过滤后的张量: tensor([[-0.5535, 0.0000], [-1.1370, -0.2736], [-1.2300, 0.9185]]) 被删除的行索引: tensor([1])
NumPy 实现
逻辑与PyTorch一致,仅API略有差异:
import numpy as np input_arr = np.array([ [-0.5535, 0.0000], [ 0.0000, 0.0000], [-1.1370, -0.2736], [-1.2300, 0.9185] ]) # 生成每行全非零的布尔掩码(axis=1表示按行判断所有列) mask = np.all(input_arr != 0, axis=1) # 过滤得到保留的行 filtered_arr = input_arr[mask] # 获取被删除的行索引 deleted_indices = np.where(~mask)[0] print("过滤后的数组:") print(filtered_arr) print("被删除的行索引:") print(deleted_indices)
注意事项
如果处理的是浮点型张量/数组,直接用!=0可能因浮点精度误差导致误判(比如极小的数值被当作非零),此时可以用阈值判断替代:
# PyTorch 浮点精度处理 mask = torch.all(torch.abs(input_tensor) > 1e-6, dim=1) # NumPy 浮点精度处理 mask = np.all(np.abs(input_arr) > 1e-6, axis=1)
上述方法自动适配任意行列数,无需针对不同维度修改逻辑。
内容的提问来源于stack exchange,提问作者user1340852
相关产品推荐
相关产品推荐

