【问题标题】:How to replace only the first n elements in a numpy array that are larger than a certain value?如何仅替换numpy数组中大于某个值的前n个元素?
【发布时间】:2016-05-03 04:29:27
【问题描述】:

我有一个这样的数组myA

array([ 7,  4,  5,  8,  3, 10])

如果我想用 0 替换所有大于值 val 的值,我可以这样做:

myA[myA > val] = 0

这给了我想要的输出(对于val = 5):

 array([0, 4, 5, 0, 3, 0])

但是,我的目标不是替换所有,而是仅替换此数组中大于值 val 的第一个 n 元素。

所以,如果n = 2 我想要的结果是这样的(10 是第三个元素,因此不应被替换):

array([ 0,  4,  5,  0,  3, 10])

一个简单的实现是:

import numpy as np

myA = np.array([7, 4, 5, 8, 3, 10])
n = 2
val = 5

# track the number of replacements
repl = 0

for ind, vali in enumerate(myA):

    if vali > val:

        myA[ind] = 0
        repl += 1

        if repl == n:
            break

这行得通,但也许有人可以用一种聪明的方式来掩盖!?

【问题讨论】:

    标签: python arrays performance numpy


    【解决方案1】:

    以下应该有效:

    myA[(myA > val).nonzero()[0][:2]] = 0
    

    因为nonzero 将返回布尔数组myA > val 不为零的索引,例如True.

    例如:

    In [1]: myA = array([ 7,  4,  5,  8,  3, 10])
    
    In [2]: myA[(myA > 5).nonzero()[0][:2]] = 0
    
    In [3]: myA
    Out[3]: array([ 0,  4,  5,  0,  3, 10])
    

    【讨论】:

    • 非常优雅,谢谢。我现在赞成它,以后可能会根据其他答案的质量接受它。
    【解决方案2】:

    最终的解决方案很简单:

    import numpy as np
    myA = np.array([7, 4, 5, 8, 3, 10])
    n = 2
    val = 5
    
    myA[np.where(myA > val)[0][:n]] = 0
    
    print(myA)
    

    输出:

    [ 0  4  5  0  3 10]
    

    【讨论】:

    • 当 n = 3 时,这似乎失败了。
    • 是的,应该是np.where(mask)[0][n:]
    • 是的,现在它也可以正常工作了,所以我也赞成它,以后可能会接受它,具体取决于其他答案的质量。
    • 非常好。你可以让它像 JuniorCompressor 的答案一样单行:myA[np.where(myA > val)[0][:n]] = 0
    • @mtrw 是的,我可以,但会引发 VisibleDeprecationWarning: boolean index did not match indexed array along dimension 0; dimension is 6 but corresponding boolean dimension is 2 myA[mask[np.where(mask)[0][n:]]] = 0 警告。
    【解决方案3】:

    这是另一种可能性(未经测试),可能不比nonzero好:

    def truncate_mask(m, stop):
      m = m.astype(bool, copy=False) #  if we allow non-bool m, the next line becomes nonsense
      return m & (np.cumsum(m) <= stop)
    
    myA[truncate_mask(myA > val, n)] = 0
    

    通过避免构建和使用显式索引,您最终可能会获得稍微更好的性能...但您必须对其进行测试才能找到答案。

    编辑 1:虽然我们讨论的是可能性,但您也可以尝试:

    def truncate_mask(m, stop):
       m = m.astype(bool, copy=True) #  note we need to copy m here to safely modify it
       m[np.searchsorted(np.cumsum(m), stop):] = 0
       return m
    

    编辑 2(次日):我刚刚对此进行了测试,似乎 cumsum 实际上比 nonzero 差,至少与我使用的 kinds of values 相比(所以上述两种方法都不值得使用)。出于好奇,我也用numba试了一下:

    import numba
    
    @numba.jit
    def set_first_n_gt_thresh(a, val, thresh, n):
        ii = 0
        while n>0 and ii < len(a):
            if a[ii] > thresh:
                a[ii] = val
                n -= 1
            ii += 1
    

    这只对数组进行一次迭代,或者更确切地说,它只对数组的必要部分进行一次迭代,甚至从不触及后面的部分。这为小型n 提供了非常出色的性能,但即使对于n&gt;=len(a) 的最坏情况,这种方法也更快。

    【讨论】:

    • 代码中有一个小“错误”:它应该是&lt;=stop 而不是stop,它似乎工作正常。谢谢你的建议,我也赞成。
    • 啊,是的(我在考虑切片符号)。我为此添加了第二个变体,它也可能包含错误。
    【解决方案4】:

    您可以使用与here 相同的解决方案,将您的np.array 转换为pd.Series

    s = pd.Series([ 7,  4,  5,  8,  3, 10])
    n = 2
    m = 5
    s[s[s>m].iloc[:n].index] = 0
    
    In [416]: s
    Out[416]:
    0     0
    1     4
    2     5
    3     0
    4     3
    5    10
    dtype: int64
    

    分步说明:

    In [426]: s > m
    Out[426]:
    0     True
    1    False
    2    False
    3     True
    4    False
    5     True
    dtype: bool
    
    In [428]: s[s>m].iloc[:n]
    Out[428]:
    0    7
    3    8
    dtype: int64
    
    In [429]: s[s>m].iloc[:n].index
    Out[429]: Int64Index([0, 3], dtype='int64')
    
    In [430]: s[s[s>m].iloc[:n].index]
    Out[430]:
    0    7
    3    8
    dtype: int64
    

    In[430] 中的输出看起来与In[428] 相同,但在 428 中它是副本,在 430 原始系列中。

    如果您需要np.array,您可以使用values 方法:

    In [418]: s.values
    Out[418]: array([ 0,  4,  5,  0,  3, 10], dtype=int64)
    

    【讨论】:

    • 太好了,效果很好!也感谢您的详细解释。我现在赞成它,以后可能会接受它,具体取决于其他答案的质量。
    猜你喜欢
    • 2016-05-03
    • 2013-11-09
    • 2017-11-13
    • 2017-08-04
    • 1970-01-01
    • 2018-01-30
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多