【问题标题】:"bucketsort" with pythons multiprocessing带有python多处理的“桶排序”
【发布时间】:2014-08-25 22:19:40
【问题描述】:

我有一个均匀分布的数据系列。我希望利用分布对数据进行并行排序。对于 N 个 CPU,我基本上定义了 N 个存储桶并对存储桶进行并行排序。我的问题是,我没有加速。

怎么了?

from multiprocessing import Process, Queue
from numpy import array, linspace, arange, where, cumsum, zeros
from numpy.random import rand
from time import time


def my_sort(x,y):
    y.put(x.get().argsort())

def my_par_sort(X,np):
 p_list=[]
 Xq = Queue()
 Yq = Queue()
 bmin = linspace(X.min(),X.max(),np+1) #bucket lower bounds
 bmax = array(bmin); bmax[-1] = X.max()+1 #bucket upper bounds
 B = []
 Bsz = [0]
 for i in range(np):
  b = array([bmin[i] <= X, X < bmax[i+1]]).all(0)
  B.append(where(b)[0])
  Bsz.append(len(B[-1]))
  Xq.put(X[b])
  p = Process(target=my_sort, args=(Xq,Yq))
  p.start()
  p_list.append(p)

 Bsz = cumsum(Bsz).tolist()
 Y = zeros(len(X)) 
 for i in range(np):
   Y[arange(Bsz[i],Bsz[i+1])] = B[i][Yq.get()]
   p_list[i].join()

 return Y


if __name__ == '__main__':
 num_el = 1e7
 mydata = rand(num_el)
 np = 4 #multiprocessing.cpu_count()
 starttime = time()
 I = my_par_sort(mydata,np)
 print "Sorting %0.0e keys took %0.1fs using %0.0f processes" % (len(mydata),time()-starttime,np)
 starttime = time()
 I2 = mydata.argsort()
 print "in serial it takes %0.1fs" % (time()-starttime)
 print (I==I2).all()

【问题讨论】:

    标签: python sorting numpy parallel-processing multiprocessing


    【解决方案1】:

    看起来你的问题是当你把原始数组分成几部分时增加的开销。我拿走了你的代码,只是删除了multiprocessing的所有用法:

    def my_sort(x,y): 
        pass
        #y.put(x.get().argsort())
    
    def my_par_sort(X,np, starttime):
        p_list=[]
        Xq = Queue()
        Yq = Queue()
        bmin = linspace(X.min(),X.max(),np+1) #bucket lower bounds
        bmax = array(bmin); bmax[-1] = X.max()+1 #bucket upper bounds
        B = []
        Bsz = [0] 
        for i in range(np):
            b = array([bmin[i] <= X, X < bmax[i+1]]).all(0)
            B.append(where(b)[0])
            Bsz.append(len(B[-1]))
            Xq.put(X[b])
            p = Process(target=my_sort, args=(Xq,Yq, i)) 
            p.start()
            p_list.append(p)
        return
    
    if __name__ == '__main__':
        num_el = 1e7 
        mydata = rand(num_el)
        np = 4 #multiprocessing.cpu_count()
        starttime = time()
        I = my_par_sort(mydata,np, starttime)
        print "Sorting %0.0e keys took %0.1fs using %0.0f processes" % (len(mydata),time()-starttime,np)
        starttime = time()
        I2 = mydata.argsort()
        print "in serial it takes %0.1fs" % (time()-starttime)
        #print (I==I2).all()
    

    在完全没有排序发生的情况下,multiprocessing 代码与序列代码一样长:

    Sorting 1e+07 keys took 2.2s using 4 processes
    in serial it takes 2.2s
    

    您可能认为启动进程和在它们之间传递值的开销是开销的原因,但如果我删除所有使用 multiprocessing,包括 Xq.put(X[b]) 调用,它最终会稍微快一点:

    Sorting 1e+07 keys took 1.9s using 4 processes
    in serial it takes 2.2s
    

    因此,您似乎需要研究一种更有效的方法来将您的数组分解成碎片。

    【讨论】:

      【解决方案2】:

      在我看来有两个主要问题。

      1. 多个进程的开销以及它们之间的通信

        生成几个 Python 解释器会导致一些开销,但主要是在“工作”进程之间传递数据会降低性能。您通过Queue 传递的数据需要“腌制”和“取消腌制”,这对于较大的数据来说有点慢(您需要这样做两次)。

        如果您使用线程而不是进程,则无需使用Queues。在 CPython 中使用线程处理 CPU 繁重的任务通常被认为是低效的,因为通常你会遇到Global Interpreter Lock,但并非总是如此!幸运的是 Numpy 的排序功能似乎正在释放 GIL,因此使用线程是一个可行的选择!

      2. 数据集的划分和连接

        对数据进行分区和连接是这种“桶排序方法”不可避免的成本,但可以通过更有效地执行来减轻一些负担。特别是这两行代码

        b = array([bmin[i] <= X, X < bmax[i+1]]).all(0)
        
        Y[arange(Bsz[i],Bsz[i+1])] = ...
        

        可以改写为

        b = (bmin[i] <= X) & (X < bmax[i+1])
        
        Y[Bsz[i] : Bsz[i+1]] = ...
        

        进一步改进我还发现np.take 比“花式索引”更快,np.partition 也很有用。

      总结一下,我能做到的最快速度如下(但它仍然不能像您想要的那样随内核数量线性扩展..):

      from threading import Thread
      
      def par_argsort(X, nproc):
          N = len(X)
          k = range(0, N+1, N//nproc)
          I = X.argpartition(k[1:-1])
          P = X.take(I)
      
          def worker(i):
              s = slice(k[i], k[i+1])
              I[s].take(P[s].argsort(), out=I[s])
      
          t_list = []
          for i in range(nproc):
              t = Thread(target=worker, args=(i,))
              t.start()
              t_list.append(t)
      
          for t in t_list:
              t.join()
      
          return I
      

      【讨论】:

        猜你喜欢
        • 2014-11-21
        • 2018-11-17
        • 1970-01-01
        • 2014-08-04
        • 1970-01-01
        • 1970-01-01
        • 2018-04-23
        • 1970-01-01
        • 1970-01-01
        相关资源
        最近更新 更多