【问题标题】:Apply argsort per row in array skipping certain elements based on threshold - NumPy / Python对数组中的每行应用 argsort,根据阈值跳过某些元素 - NumPy / Python
【发布时间】:2020-05-23 06:56:49
【问题描述】:

我想应用排序操作,每行一行,只保持高于给定阈值的值。

为此,我看到我可以使用掩码数组来应用阈值。 但是,argsort 一直在考虑掩码值(低于阈值)并将其替换为 fill_value

但是,如果值已被替换为 NaN,我根本不想要任何结果。

a = np.array([[0.522235,0.128270,0.708973],
              [0.994557,0.844426,0.366608],
              [0.986669,0.143659,0.395891],
              [0.291339,0.421843,0.278869],
              [0.250303,0.861475,0.904534],
              [0.973436,0.360466,0.751913]])

threshold = 0.5
m_a = np.ma.masked_less_equal(a, threshold)
argsorted = m_a.argsort(-1)

这给了我:

array([[0, 2, 1],
       [1, 0, 2],
       [0, 1, 2],
       [0, 1, 2],
       [1, 2, 0],
       [2, 0, 1]])

但我想得到:

array([[0,   NaN,   1],
       [1,     0, NaN],
       [0,   NaN, NaN],
       [NaN, NaN, NaN],
       [NaN,   0,   1],
       [  1, NaN,   0]])

有什么想法可以得到这个结果吗?

感谢您的帮助! 最好的,

【问题讨论】:

  • 是的,你是对的,我马上改正。

标签: python numpy np.argsort


【解决方案1】:

我们可以再添加一个argsort,以便更轻松地获得我们想要的输出 -

sidx = argsorted.argsort(1)
mask = sidx >= (a.shape[1]-m_a.mask.sum(1,keepdims=True))
out = np.where(mask,np.nan,sidx)

我们也可以从头开始避免masked-arrays-

def thresholded_argsort(a, threshold):
    m = a<threshold
    ac = a.copy()
    ac[m] = ac.max()+1
    sidx = ac.argsort(1).argsort(1)
    mask = sidx>=(ac.shape[1]-m.sum(1,keepdims=True))
    return np.where(mask,np.nan,sidx)

示例运行 -

In [46]: a
Out[46]: 
array([[0.522235, 0.12827 , 0.708973],
       [0.994557, 0.844426, 0.366608],
       [0.986669, 0.143659, 0.395891],
       [0.291339, 0.421843, 0.278869],
       [0.250303, 0.861475, 0.904534],
       [0.973436, 0.360466, 0.751913]])

In [47]: thresholded_argsort(a, threshold=0.5)
Out[47]: 
array([[ 0., nan,  1.],
       [ 1.,  0., nan],
       [ 0., nan, nan],
       [nan, nan, nan],
       [nan,  0.,  1.],
       [ 1., nan,  0.]])

注意:我们可以避免使用 array-assignment 进行额外的 argsort 以提高使用 argsort_unique 的性能。因此,对于沿第二个轴的 2D 数组,它将是 -

def argsort_unique2D(idx):
    m,n = idx.shape
    idx_out = np.empty((m,n),dtype=int)
    np.put_along_axis(idx_out, idx, np.arange(n), axis=1)
    return idx_out

因此,argsorted.argsort(1) 可以替换为 argsort_unique2D(argsorted),而在之前发布的解决方案中,ac.argsort(1).argsort(1) 可以替换为 argsort_unique2D(ac.argsort(1))

【讨论】:

  • 嗨@Divakar,感谢您的帮助。我正在研究这个方向。总是有兴趣学习新技巧并获得更高性能的代码:)。你介意用argsort_unique 详细说明你的食谱吗?已经感谢您的帮助!
  • @pierre_j 已添加到帖子中。
  • 谢谢@Divakar。我开始使用您之前发布的argsort_unique 的帖子,并且确实遇到了索引问题。 argsort_unique2D 处理这个,非常感谢!
【解决方案2】:

如果我理解正确,您不想考虑 NaN 进行排序。在这种情况下,我不确定您预期结果背后的逻辑。你可以试试下面的代码。我相信这就是您正在寻找的:-

import numpy as np
a = np.array([[0.522235,0.128270,0.708973],
              [0.994557,0.844426,0.366608],
              [0.986669,0.143659,0.395891],
              [0.291339,0.421843,0.278869],
              [0.250303,0.861475,0.904534],
              [0.973436,0.360466,0.751913]])

threshold = 0.5
m_a = np.ma.masked_less_equal(a, threshold).filled(np.nan)
result = np.where(
        np.isnan(m_a),
        np.nan, m_a.argsort(-1)
    )
result

它应该给你以下结果:-

array([[ 0., nan,  1.],
       [ 1.,  0., nan],
       [ 0., nan, nan],
       [nan, nan, nan],
       [nan,  2.,  0.],
       [ 2., nan,  1.]])

希望这会有所帮助!

【讨论】:

  • 最后一行似乎不匹配。
  • 我不明白预期结果中倒数第三行背后的逻辑。
  • AFAIU,对于所有上述阈值元素,索引将依次为 0...n。
  • 不考虑 NaN 进行排序不应更改元素 IMO 的索引。
  • @TuhinSharma 感谢您的帮助。但我也不明白第三行背后的逻辑。我可以理解应该可以设计一个在过滤 NaN 之前保留索引的函数。但是这个逻辑也应该适用于前 3 行。因此,要么在排序后过滤 nan,在这种情况下,索引从 2 开始按降序排列(我们在你的示例中遇到的情况是在第 5 行和第 6 行),或者在排序之前进行过滤,并且索引从 0 开始按升序排列,并且这是示例的第一行。还是我错过了什么?
【解决方案3】:
a = np.array([[0.522235,0.128270,0.708973],
              [0.994557,0.844426,0.366608],
              [0.986669,0.143659,0.395891],
              [0.291339,0.421843,0.278869],
              [0.250303,0.861475,0.904534],
              [0.973436,0.360466,0.751913]])

threshold = .5


def tri(ligne):
    s = sorted(ligne, key=lambda x: x < threshold and float('inf') or x)
    nv_liste = [s.index(v) for v in ligne]
    for i in range(len(ligne)):
        if ligne[i] < threshold:
            nv_liste[i] = np.nan
    return nv_liste

np.apply_along_axis(tri, 1, a)

给你:

array([[ 0., nan,  1.],
       [ 1.,  0., nan],
       [ 0., nan, nan],
       [nan, nan, nan],
       [nan,  0.,  1.],
       [ 1., nan,  0.]])

【讨论】:

  • 嗨@David,感谢您的帮助。将您的建议与 Divakar 的建议进行比较,它大约慢了 80 倍。在 100k 阵列上,这开始是可见的 ;) 。不过还是谢谢!
猜你喜欢
  • 2018-12-20
  • 1970-01-01
  • 1970-01-01
  • 2014-08-09
  • 2019-05-17
  • 1970-01-01
  • 2017-07-10
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多