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

如何过滤张量中全列非零元素的行并获取被删行索引?

解决方案:保留全非零行并获取删除行索引

针对你需要保留所有列元素均非零的行,同时获取被删除行索引的需求,以下是适配任意行列数的实现方案,分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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 12:35:17