访问numpy数组的相邻单元格

use*_*827 17 python numpy

如何以有效的方式访问和修改2D numpy阵列的周围8个单元?

我有一个像这样的2D numpy数组:

arr = np.random.rand(720, 1440)
Run Code Online (Sandbox Code Playgroud)

对于每个网格单元,我想减少中心单元的10%,周围的8个单元(角单元更少),但仅当周围单元值超过0.25时.我怀疑这样做的唯一方法是使用for循环但是想看看是否有更好/更快的解决方案.

- 编辑:对于基于循环的soln:

arr = np.random.rand(720, 1440)

for (x, y), value in np.ndenumerate(arr):
    # Find 10% of current cell
    reduce_by = value * 0.1

    # Reduce the nearby 8 cells by 'reduce_by' but only if the cell value exceeds 0.25
    # [0] [1] [2]
    # [3] [*] [5]
    # [6] [7] [8]
    # * refers to current cell

    # cell [0]
    arr[x-1][y+1] = arr[x-1][y+1] * reduce_by if arr[x-1][y+1] > 0.25 else arr[x-1][y+1]

    # cell [1]
    arr[x][y+1] = arr[x][y+1] * reduce_by if arr[x][y+1] > 0.25 else arr[x][y+1]

    # cell [2]
    arr[x+1][y+1] = arr[x+1][y+1] * reduce_by if arr[x+1][y+1] > 0.25 else arr[x+1][y+1]

    # cell [3]
    arr[x-1][y] = arr[x-1][y] * reduce_by if arr[x-1][y] > 0.25 else arr[x-1][y]

    # cell [4] or current cell
    # do nothing

    # cell [5]
    arr[x+1][y] = arr[x+1][y] * reduce_by if arr[x+1][y] > 0.25 else arr[x+1][y]

    # cell [6]
    arr[x-1][y-1] = arr[x-1][y-1] * reduce_by if arr[x-1][y-1] > 0.25 else arr[x-1][y-1]

    # cell [7]
    arr[x][y-1] = arr[x][y-1] * reduce_by if arr[x][y-1] > 0.25 else arr[x][y-1]

    # cell [8]
    arr[x+1][y-1] = arr[x+1][y-1] * reduce_by if arr[x+1][y-1] > 0.25 else arr[x+1][y-1]
Run Code Online (Sandbox Code Playgroud)

Wal*_*oss 2

这个答案假设您确实想完全按照您在问题中所写的操作。好吧,几乎完全一样,因为您的代码因索引越界而崩溃。解决这个问题最简单的方法是添加条件,例如,

\n\n
if x > 0 and y < y_max:\n    arr[x-1][y+1] = ...\n
Run Code Online (Sandbox Code Playgroud)\n\n

主操作无法使用 numpy 或 scipy 进行向量化的原因是,所有单元格都被一些已经\xe2\x80\x9creduced\xe2\x80\x9d 的邻居单元格\xe2\x80\x9creduced\xe2\x80\x9d 。Numpy 或 scipy 将在每个操作中使用邻居的未受影响的值。在我的另一个答案中,我展示了如何使用 numpy 执行此操作,如果允许您将操作分为 8 个步骤,每个步骤沿着一个特定邻居的方向,但每个步骤都使用该邻居的该步骤中不受影响的值。正如我所说,这里我认为你必须按顺序进行。

\n\n

在继续之前,让我交换您的代码中的x和。y您的阵列具有典型的屏幕尺寸,其中高度为 720,宽度为 1440。图像通常按行存储,默认情况下,ndarray 中最右边的索引是变化较快的索引,因此一切都有意义。诚然,这违反直觉,但正确的索引是arr[y, x].

\n\n

可以应用于代码的主要优化(在我的 Mac 上将执行时间从 ~9 秒减少到 ~3.9 秒)是在不需要时不要将单元格分配给自身,再加上就地乘法 和with[y, x]而不是[y][x]索引。像这样:

\n\n
y_size, x_size = arr.shape\ny_max, x_max = y_size - 1, x_size - 1\nfor (y, x), value in np.ndenumerate(arr):\n    reduce_by = value * 0.1\n    if y > 0 and x < x_max:\n        if arr[y - 1, x + 1] > 0.25: arr[y - 1, x + 1] *= reduce_by\n    if x < x_max:\n        if arr[y    , x + 1] > 0.25: arr[y    , x + 1] *= reduce_by\n    if y < y_max and x < x_max:\n        if arr[y + 1, x + 1] > 0.25: arr[y + 1, x + 1] *= reduce_by\n    if y > 0:\n        if arr[y - 1, x    ] > 0.25: arr[y - 1, x    ] *= reduce_by\n    if y < y_max:\n        if arr[y + 1, x    ] > 0.25: arr[y + 1, x    ] *= reduce_by\n    if y > 0 and x > 0:\n        if arr[y - 1, x - 1] > 0.25: arr[y - 1, x - 1] *= reduce_by\n    if x > 0:\n        if arr[y    , x - 1] > 0.25: arr[y    , x - 1] *= reduce_by\n    if y < y_max and x > 0:\n        if arr[y + 1, x - 1] > 0.25: arr[y + 1, x - 1] *= reduce_by\n
Run Code Online (Sandbox Code Playgroud)\n\n

另一项优化(在我的 Mac 上将执行时间进一步缩短至约 3.0 秒)是通过使用具有额外边界单元的数组来避免边界检查。我们不关心边界包含什么值,因为它永远不会被使用。这是代码:

\n\n
y_size, x_size = arr.shape\narr1 = np.empty((y_size + 2, x_size + 2))\narr1[1:-1, 1:-1] = arr\nfor y in range(1, y_size + 1):\n    for x in range(1, x_size + 1):\n        reduce_by = arr1[y, x] * 0.1\n        if arr1[y - 1, x + 1] > 0.25: arr1[y - 1, x + 1] *= reduce_by\n        if arr1[y    , x + 1] > 0.25: arr1[y    , x + 1] *= reduce_by\n        if arr1[y + 1, x + 1] > 0.25: arr1[y + 1, x + 1] *= reduce_by\n        if arr1[y - 1, x    ] > 0.25: arr1[y - 1, x    ] *= reduce_by\n        if arr1[y + 1, x    ] > 0.25: arr1[y + 1, x    ] *= reduce_by\n        if arr1[y - 1, x - 1] > 0.25: arr1[y - 1, x - 1] *= reduce_by\n        if arr1[y    , x - 1] > 0.25: arr1[y    , x - 1] *= reduce_by\n        if arr1[y + 1, x - 1] > 0.25: arr1[y + 1, x - 1] *= reduce_by\narr = arr1[1:-1, 1:-1]\n
Run Code Online (Sandbox Code Playgroud)\n\n

根据记录,如果可以使用 numpy 或 scipy 对操作进行矢量化,则该解决方案的速度至少会提高 35 倍(在我的 Mac 上测量)。

\n\n

注意:如果 numpy按顺序对数组切片进行操作,则以下结果将产生阶乘(即,正整数的乘积,最多一个数字) \xe2\x80\x93 但它不会:

\n\n
>>> import numpy as np\n>>> arr = np.arange(1, 11)\n>>> arr\narray([ 1,  2,  3,  4,  5,  6,  7,  8,  9, 10])\n>>> arr[1:] *= arr[:-1]\n>>> arr\narray([ 1,  2,  6, 12, 20, 30, 42, 56, 72, 90])\n
Run Code Online (Sandbox Code Playgroud)\n