如何替换numpy数组每一行中最小的N个元素并保持原数组结构
Numpy按行替换最小N个元素为0的实现方案
原有实现的问题是使用np.isin匹配值判断是否替换,遇到重复值时会将所有匹配到的值全部替换,导致替换数量超出要求的N个。
以下是可直接使用的实现代码,完全适配任意尺寸的矩阵,不会受重复值影响:
import numpy as np N = 2 x = np.array([[0,50,20],[2,0,10],[1,1,0]]) # 获取每行元素从小到大排序后的索引,取前N个最小元素的索引 min_n_indices = np.argsort(x, axis=1)[:, :N] # 按行将对应索引位置的元素替换为0 np.put_along_axis(x, min_n_indices, 0, axis=1) print(x)
运行后输出结果完全符合预期:
[[ 0 50 0] [ 0 0 10] [ 0 1 0]]
方案特点
- 全程基于Numpy向量化操作,无Python循环,大尺寸矩阵下性能优异
- 基于索引定位要替换的元素,不受重复值干扰,严格保证每行仅替换N个最小元素
- 不会改变原数组的元素顺序和形状,完全匹配需求
如果你的Numpy版本低于1.15无法使用np.put_along_axis,可以用高级索引写法替代,效果完全一致:
rows = np.arange(x.shape[0])[:, np.newaxis] x[rows, min_n_indices] = 0
内容的提问来源于stack exchange,提问作者EconMate
相关产品推荐
相关产品推荐

