【问题标题】:multi core processing in for loop using numpy and fft使用 numpy 和 fft 在 for 循环中进行多核处理
【发布时间】:2018-09-28 22:24:47
【问题描述】:

我使用 numpy 和 fft 计算了向量。 我使用了 numpy 广播方法和 for 循环。 两种方法的速度是相似的。 如何使用多核和 numpy 和 fft 计算向量?

import numpy as np
from numpy.fft import fft, ifft

num_row, num_col = 6000, 13572

ss = np.ones((num_row, num_col), dtype=np.complex128)
sig = np.random.standard_normal(num_col) * 1j * np.random.standard_normal(num_col)

# for loop    
for idx in range(num_row):
    ss[idx, :] = ifft(fft(ss[idx, :]) * sig)

# broadcast
ss = ifft(fft(ss, axis=1) * sig, axis=1)

结果

loop : 10.798867464065552 sec
broadcast : 11.298897981643677 sec

【问题讨论】:

  • 您可以将轴指定为fft,而不是使用循环np.fft.fft(a, axis=1)
  • 您也可以使用broadcasting 进行数组乘法ss*sig
  • 这到底是在做什么?您正在计算一个向量的 FFT(这只是 [N, 0, 0, 0...],其中 N 是向量的长度),乘以一些随机数据,然后是 IFFT???我想这里有一个XY problem
  • 无论如何,默认的 Numpy FFT 不是多线程的:您要么必须使用 Enthought 的 Python 发行版(其中 Numpy 是针对 Intel MKL 构建的,它具有高度优化的 FFT),要么使用 PyFFTW(它依赖于 FFTW 的多线程)。或者您想尝试绕过 GIL 并使用 Python 的内置线程功能?
  • 另外你可能打算做stdn(ncol) + 1j * stdn(ncol)(注意+)。

标签: python loops numpy fft multicore


【解决方案1】:

fftifft可以使用axis参数,也可以广播:

ss = np.ones((num_row, num_col), dtype=np.complex128)
sig = np.random.standard_normal(num_col) * 1j * np.random.standard_normal(num_col)

ss = ifft(fft(ss, axis=1) * sig, axis=1)

【讨论】:

  • ifft(fft(ss, axis=1) * sig, axis=1) 比 for 循环慢。
【解决方案2】:

我比较了广播、线程池和循环。 在这种情况下,ThreadPool具有最佳性能。

# %% Import
# Standard library imports
import time
from multiprocessing.pool import ThreadPool

# Third party imports
from numpy import zeros, complex128, allclose
from numpy.fft import fft, ifft
from numpy.random import standard_normal


# %% Generate data
n_row, n_col = 6000, 13572

ss = standard_normal((n_row, n_col)) + 1j * standard_normal((n_row, n_col))
sig = standard_normal(n_col) + 1j * standard_normal(n_col)
ss_loop = zeros((n_row, n_col), dtype=complex128)
ss_thread = zeros((n_row, n_col), dtype=complex128)

# %% Loop processing
start_time = time.time()
for idx in range(n_row):
    ss_loop[idx, :] = ifft(fft(ss[idx, :]) * sig)
print(f'loop elapsed time : {time.time() - start_time}')

# %% Broadcast processing
start_time = time.time()
ss_broad = ifft(fft(ss, axis=1) * sig, axis=1)
print(f'broadcast elapsed time : {time.time() - start_time}')


# %% ThreadPool processing
def filtering(idx_thread):
    ss_thread[idx_thread, :] = ifft(fft(ss[idx_thread, :]) * sig)


start_time = time.time()
pool = ThreadPool()
pool.map(filtering, range(n_row))
print(f'ThreadPool elapsed time : {time.time() - start_time}')


# %% Verify result
if allclose(ss_thread, ss_broad, rtol=1.e-8):
    print('ThreadPool Correct')

if allclose(ss_loop, ss_broad, rtol=1.e-8):
    print('Loop Correct')

结果

loop elapsed time : 5.102990627288818
broadcast elapsed time : 4.520442008972168
ThreadPool elapsed time : 1.6988463401794434

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2023-03-29
    • 2022-01-12
    • 2022-01-16
    • 1970-01-01
    • 2014-02-26
    • 1970-01-01
    • 2016-08-01
    • 2022-12-04
    相关资源
    最近更新 更多