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

如何快速从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 10:43:21