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

如何检查numpy数组是否在Python序列中?

  •  0
  • jfaccioni  · 技术社区  · 6 年前

    我想检查给定数组是否在常规Python序列(列表、元组等)中。例如,考虑下面的代码:

    import numpy as np
    
    xs = np.array([1, 2, 3])
    ys = np.array([4, 5, 6])
    
    myseq = (xs, 1, True, ys, 'hello')
    

    我希望通过简单的会员资格检查 in 将起作用,例如:

    >>> xs in myseq
    True
    

    myseq ,例如:

    >>> ys in myseq
    ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
    

    那么我该如何进行这项检查呢?

    如果可能的话,我想做这件事,而不必投 迈赛克 进入numpy数组或任何其他类型的数据结构。

    0 回复  |  直到 6 年前
        1
  •  1
  •   dawg    6 年前

    any 通过适当的测试:

    import numpy as np
    
    xs = np.array([1, 2, 3])
    ys = np.array([4, 5, 6])
    zs = np.array([7, 8, 9])
    
    myseq = (xs, 1, True, ys, 'hello')
    
    def arr_in_seq(arr, seq):
        tp=type(arr)
        return any(isinstance(e, tp) and np.array_equiv(e, arr) for e in seq)
    

    测试:

    for x in (xs,ys,zs):
        print(arr_in_seq(x,myseq))
    True
    True
    False
    
        2
  •  1
  •   mapf    6 年前

    这可能不是最美好或最快速的解决方案,但我认为它是有效的:

    import numpy as np
    
    
    def array_in_tuple(array, tpl):
        i = 0
        while i < len(tpl):
            if isinstance(tpl[i], np.ndarray) and np.array_equal(array, tpl[i]):
                return True
            i += 1
        return False
    
    
    xs = np.array([1, 2, 3])
    ys = np.array([4, 5, 6])
    
    myseq = (xs, 1, True, ys, 'hello')
    
    
    print(array_in_tuple(xs, myseq), array_in_tuple(ys, myseq), array_in_tuple(np.array([7, 8, 9]), myseq))