如何用更简洁的方式在Numpy中筛选多列并排除指定列?
更简洁的NumPy二维数组列排除方法
你当前用列表推导式的方法可行,但有几种更简洁的实现方式:
方法1:使用np.delete()(最直接)
这是最贴合需求的写法,直接指定要删除的列索引和操作轴:
import numpy as np x = np.arange(20).reshape(2,10) result = np.delete(x, [3, 6, 9], axis=1) print(result)
输出:
[[ 0 1 2 4 5 7 8] [10 11 12 14 15 17 18]]
axis=1表示对列操作,第二个参数传入要排除的列索引列表即可。
方法2:布尔掩码索引
适合需要复用筛选规则或更复杂逻辑的场景:
import numpy as np x = np.arange(20).reshape(2,10) # 初始化掩码,默认保留所有列 mask = np.ones(x.shape[1], dtype=bool) # 标记要排除的列 mask[[3, 6, 9]] = False result = x[:, mask] print(result)
方法3:集合差集生成保留列索引
当需要排除的列较多时,用集合差集自动生成要保留的列索引:
import numpy as np x = np.arange(20).reshape(2,10) # 计算所有列索引与排除索引的差集 keep_cols = np.setdiff1d(np.arange(x.shape[1]), [3, 6, 9]) result = x[:, keep_cols] print(result)
注:np.setdiff1d返回的索引是有序的,和原列顺序一致(因原列索引为连续递增)。
内容的提问来源于stack exchange,提问作者nightstand
相关产品推荐
相关产品推荐

