【问题标题】:How to apply a function to a 2D numpy array with multiprocessing如何将函数应用于具有多处理的二维 numpy 数组
【发布时间】:2015-04-24 21:07:33
【问题描述】:

假设我有以下功能:

def f(x,y):
    return x*y

如何使用多处理模块将该函数应用于 NxM 2D numpy 数组中的每个元素?使用串行迭代,代码可能如下所示:

import numpy as np
N = 10
M = 12
results = np.zeros(shape=(N,M))
for x in range(N):
    for y in range(M):
        results[x,y] = f(x,y)

【问题讨论】:

  • 我认为这是一个玩具模型,您需要做的事情更复杂,但 numpy 具有高效的功能来完成您在代码中编写的操作
  • 我目前正在浏览文档,但还没有找到如何将函数应用于数组中的每个元素。有什么指导吗?
  • 看下面这个简单案例的答案
  • @JulienSpronck 看起来您忘记发布您提到的答案
  • Julien 是对的 - 当您尝试优化 numpy 代码时,多处理是 最后一个 工具。 Python 中的多处理既慢又麻烦,使用 broadcasting 和 BLAS 优化的线性代数运算通常可以获得更大的效率提升。除此之外,还有Cythonnumbanumexpr。由于您没有向我们展示您实际尝试优化的代码,因此很难给出更具体的建议。

标签: python arrays numpy multiprocessing


【解决方案1】:

以下是使用multiprocesssing 并行化示例函数的方法。我还包含了一个几乎相同的纯 Python 函数,它使用非并行 for 循环,以及一个实现相同结果的 numpy one-liner:

import numpy as np
from multiprocessing import Pool


def f(x,y):
    return x * y

# this helper function is needed because map() can only be used for functions
# that take a single argument (see http://stackoverflow.com/q/5442910/1461210)
def splat_f(args):
    return f(*args)

# a pool of 8 worker processes
pool = Pool(8)

def parallel(M, N):
    results = pool.map(splat_f, ((i, j) for i in range(M) for j in range(N)))
    return np.array(results).reshape(M, N)

def nonparallel(M, N):
    out = np.zeros((M, N), np.int)
    for i in range(M):
        for j in range(N):
            out[i, j] = f(i, j)
    return out

def broadcast(M, N):
    return np.prod(np.ogrid[:M, :N])

现在让我们看看性能:

%timeit parallel(1000, 1000)
# 1 loops, best of 3: 1.67 s per loop

%timeit nonparallel(1000, 1000)
# 1 loops, best of 3: 395 ms per loop

%timeit broadcast(1000, 1000)
# 100 loops, best of 3: 2 ms per loop

非并行纯 Python 版本比并行版本高出大约 4 倍,使用 numpy 数组广播的版本绝对碾压其他两个版本。

问题在于,启动和停止 Python 子进程会带来相当多的开销,而且您的测试函数非常琐碎,以至于每个工作线程只花费其生命周期的一小部分来做有用的工作。仅当每个线程在被杀死之前都有大量工作要做时,多处理才有意义。例如,您可以为每个工作人员分配更大的输出数组块来计算(尝试将chunksize= 参数设置为pool.map()),但是对于这样一个微不足道的示例,我怀疑您会看到很大的改进。

我不知道你的实际代码是什么样子的——也许你的函数大到足以保证使用多处理。但是,我敢打赌,有很多更好的方法可以提高其性能。

【讨论】:

  • 我确信通过其他方式可以获得性能提升,但目前一次运行大约需要 10 分钟,并且只使用一个逻辑核心。它最大化了核心,但只使用了一个。谢谢你的回答:)
  • @TraxusIV 相信我,我知道想要看到我所有的核心都在全速工作的诱惑,但相信我 - 这是进行优化的错误方法!始终从分析您的代码开始(例如使用line_profiler)并确定瓶颈在哪里,然后专注于这些。尽可能使用 BLAS 和广播,并使用 Cython 或 numba 处理任何无法通过广播摆脱的讨厌的内部 for 循环。
  • 附带说明:您可以简化为 pool.starmap,而不是使用 splat_fpool.map,它在内部与 splat_f 做同样的事情。
  • @sophros 是的,前提是您使用的是 Python 3.3 或更高版本(请参阅上面链接的 stackoverflow.com/q/5442910/1461210
【解决方案2】:

不确定您的情况是否需要多处理。在上面的简单示例中,您可以这样做

X, Y = numpy.meshgrid(numpy.arange(10), numpy.arange(12))
result = X*Y

【讨论】:

  • 我的应用程序需要多处理。这个例子被大大简化了,但确实代表了基本问题
  • 我是这么想的……那我恐怕不知道
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2018-10-20
  • 1970-01-01
  • 1970-01-01
  • 2021-05-14
  • 2017-01-11
  • 1970-01-01
相关资源
最近更新 更多