如何快速从PyTorch/Numpy二维数组按索引批量删除多行?
快速删除Numpy数组和PyTorch张量指定行的方案
你的问题出在原代码的[i not in remove_ixs for i in range(...)]——Python列表的in操作是线性遍历,500万次遍历5万个元素的列表,时间复杂度达到O(N*k),自然慢到无法结束。下面是基于Numpy和PyTorch的向量化优化方案,时间复杂度大幅降低。
Numpy 实现方案
方法1:布尔掩码(最快)
利用Numpy的向量化赋值直接生成掩码,避免Python循环:
import numpy as np # 将remove_ixs转为Numpy数组,确保类型匹配 remove_ixs = np.asarray(remove_ixs, dtype=np.int64) # 创建全True的掩码数组 mask = np.ones(my_array.shape[0], dtype=bool) # 将需要删除的索引对应位置设为False mask[remove_ixs] = False # 索引得到新数组 new_array = my_array[mask]
方法2:生成保留索引(内存更友好)
通过集合差运算直接得到需要保留的行索引,适合超大规模数组:
import numpy as np # 生成全量索引与待删索引的差集 keep_ixs = np.setdiff1d( np.arange(my_array.shape[0]), remove_ixs, assume_unique=True # 若remove_ixs无重复,开启此参数可提速 ) new_array = my_array[keep_ixs]
PyTorch 实现方案
方法1:布尔掩码(GPU/CPU通用)
和Numpy逻辑一致,利用PyTorch的张量操作实现:
import torch # 将remove_ixs转为同设备的长整型张量 remove_ixs = torch.tensor(remove_ixs, dtype=torch.long, device=my_tensor.device) # 创建全True的掩码张量 mask = torch.ones(my_tensor.shape[0], dtype=torch.bool, device=my_tensor.device) mask[remove_ixs] = False new_tensor = my_tensor[mask]
方法2:生成保留索引
使用PyTorch内置的集合差函数:
import torch # 生成全量索引与待删索引的差集 keep_ixs = torch.setdiff1d( torch.arange(my_tensor.shape[0], device=my_tensor.device), remove_ixs ) new_tensor = my_tensor[keep_ixs]
注意事项
- 确保
remove_ixs中的索引均为合法值(0 ≤ 索引 < N),且无重复(若有重复,两种方法均可自动处理); - 若张量在GPU上运行,PyTorch方案的速度优势会更明显;
- 超大规模数组下,
setdiff1d的内存占用略低于布尔掩码方案。
内容的提问来源于stack exchange,提问作者sanjeev mk
相关产品推荐
相关产品推荐

