【问题标题】:Find n-dimensional point in numpy array在numpy数组中查找n维点
【发布时间】:2017-09-07 16:41:49
【问题描述】:

我正在调查将点存储在 numpy 数组中是否有助于我搜索点,对此我有几个问题。

我有一个表示 3 维点的 Point 类。

class Point( object ):
  def __init__( self, x, y, z ):
    self.x = x
    self.y = y
    self.z = z

  def __repr__( self ):
    return "<Point (%r, %r, %r)>" % ( self.x, self.y, self.z )

我构建了一个 Point 对象列表。注意坐标(1, 2, 3)故意出现两次;这就是我要搜索的内容。

>>> points = [Point(1, 2, 3), Point(4, 5, 6), Point(1, 2, 3), Point(7, 8, 9)]

我将 Point 对象存储在一个 numpy 数组中。

>>> import numpy
>>> npoints = numpy.array( points )
>>> npoints
array([<Point (1, 2, 3)>, <Point (4, 5, 6)>, <Point (1, 2, 3)>,
   <Point (7, 8, 9)>], dtype=object)

我按照以下方式搜索坐标为(1, 2, 3)的所有点。

>>> numpy.where( npoints == Point(1, 2, 3) )
>>> (array([], dtype=int64),)

但是,结果没有用。因此,这似乎不是正确的方法。 numpy.where 是要使用的东西吗?是否有另一种方式来表达numpy.where 的条件会成功?

接下来我尝试将点的坐标存储在一个 numpy 数组中。

>>> npoints = numpy.array( [(p.x, p.y, p.z) for p in points ])
>>> npoints
array([[1, 2, 3],
      [4, 5, 6],
      [1, 2, 3],
      [7, 8, 9]])

我按照以下方式搜索坐标为(1,2,3)的所有点。

>>> numpy.where( npoints == [1,2,3] )
(array([0, 0, 0, 2, 2, 2]), array([0, 1, 2, 0, 1, 2]))

结果至少是我可以处理的。第一个返回值array([0, 0, 0, 2, 2, 2]) 中的行索引数组确实告诉我,我正在搜索的坐标位于npoints 的第0 行和第2 行。我可以做以下类似的事情。

>>> rows, cols = numpy.where( npoints == [1,2,3] )
>>> rows
array([0, 0, 0, 2, 2, 2])
>>> cols
array([0, 1, 2, 0, 1, 2])
>>> foundRows = set( rows )
>>> foundRows
set([0, 2])
>>> for r in foundRows:
...   # Do something with npoints[r]

但是,我觉得我并没有真正恰当地使用numpy.where,我只是在这种特殊情况下很幸运。

在 numpy 数组中查找所有出现的 n 维点(即具有特定值的行)的适当方法是什么?

保持数组的顺序很重要。

【问题讨论】:

  • 看看this是否有帮助。
  • 不要使用 numpy 自定义对象数组。问题是 == 默认由 identity 为自定义对象实现。无论如何,最大的问题是当你创建一个 dtype=object 数组时,你实际上是在创建一个效率低下的 Python list
  • 还要考虑来自shapelyPoint 类。
  • 在@Divakar 建议的post 中,方法#1 有效。我对语法的理解还不够,无法知道 为什么 它仍然有效。这会让我忙一阵子。

标签: python arrays numpy


【解决方案1】:

您可以在 Point 类中创建“丰富的比较”方法 object.__eq__(self, other) 以便能够在 Point 对象中使用 ==

class Point( object ):
  def __init__( self, x, y, z ):
    self.x = x
    self.y = y
    self.z = z

  def __repr__( self ):
    return "<Point (%r, %r, %r)>" % ( self.x, self.y, self.z )
  def __eq__(self, other):
    return self.x == other.x and self.y == other.y and self.z == other.z

import numpy
points = [Point(1, 2, 3), Point(4, 5, 6), Point(1, 2, 3), Point(7, 8, 9)]
npoints = numpy.array( points )
found = numpy.where(npoints == Point(1, 2, 3))
print(found) # => (array([0, 2]),)

【讨论】:

  • 在 Point 类中添加丰富的等式比较方法是个好主意。那行得通。
  • 我应该打字更快! :) 向 Point 类添加一个丰富的等式比较方法是可行的,通常是一个很好的例子。但是,在我的情况下,我实际使用的 Point 类位于我无法修改的第三方模块中。此外,正如@juanpa.arrivillaga 提出的那样,让 numpy 数组的 item 数据类型成为对象会降低使用 numpy 数组的效率。我怀疑。但是,无论如何,我可能会继续这种安排。速度不是我最关心的问题,使用列表index 方法可能更方便。
猜你喜欢
  • 2015-03-21
  • 1970-01-01
  • 2011-06-14
  • 1970-01-01
  • 2021-12-20
  • 1970-01-01
  • 1970-01-01
  • 2017-07-22
  • 1970-01-01
相关资源
最近更新 更多