代码之家  ›  专栏  ›  技术社区  ›  rhug123

将索引位置之前的所有值设置为np.NaN

  •  1
  • rhug123  · 技术社区  · 2 年前

    假设我们有一个数组,如下所示:

    a = np.array(np.arange(15)).reshape(3,-1)
    
    
    array([[ 0,  1,  2,  3,  4],
           [ 5,  6,  7,  8,  9],
           [10, 11, 12, 13, 14]])
    

    以及一个包含位置的列表,我们希望将这些位置之前的所有内容设置为 np.NaN

    l = [0,2,1]
    

    最终结果是:

    array([[ 1.,  2.,  3.,  4.,  5.],
           [nan, nan,  7.,  8.,  9.],
           [nan, 11., 12., 13., 14.]])
    

    有没有办法在numpy中做到这一点?我能想到的唯一不需要迭代的解决方案是创建一个带有索引位置的伪数组,但我想知道是否有更好的方法。非常感谢。

    当前解决方案:

    s = a.shape
    a = a.astype(float)
    a[np.where((np.array([[*np.arange(s[-1])]]*s[0]) < np.array(l)[:,None]))] = np.NaN
    
    1 回复  |  直到 2 年前
        1
  •  2
  •   Timeless    2 年前

    使用 broadcasting :

    a = a.astype("float")
    
    m = np.arange(a.shape[1]) < np.array(l)[:, None]
    
    a[m] = np.nan
    

    输出

    >>> a
    # array([[ 0.,  1.,  2.,  3.,  4.],
    #        [nan, nan,  7.,  8.,  9.],
    #        [nan, 11., 12., 13., 14.]])
    

    中间体:

    >>> np.arange(a.shape[1])
    # array([0, 1, 2, 3, 4])
    
    >>> np.array(l)[:, None]
    # array([[0],
    #        [2],
    #        [1]])
    
    >>> m
    # array([[False, False, False, False, False],
    #        [ True,  True, False, False, False],
    #        [ True, False, False, False, False]])
    

    另一种选择( 由@mozway建议 )避免铸造和浇注 nan where 它应该:

    out = np.where(
        np.arange(a.shape[1]) >= np.array(l)[:, None], a, np.nan
    )