首页
学习
活动
专区
圈层
工具
发布
社区首页 >问答首页 >用Parakeet优化Python函数

用Parakeet优化Python函数
EN

Stack Overflow用户
提问于 2013-11-23 20:36:55
回答 1查看 374关注 0票数 3

我需要对这个函数进行优化,因为我试图使我的OpenGL模拟运行更快。我想使用鹦鹉,但我不太明白我需要以何种方式修改下面的代码才能这样做。你能看到我该怎么做吗?

代码语言:javascript
复制
def distanceMatrix(self,x,y,z):
    " ""Computes distances between all particles and places the result in a matrix such that the ij th matrix entry corresponds to the distance between particle i and j"" "
    xtemp = tile(x,(self.N,1))
    dx = xtemp - xtemp.T
    ytemp = tile(y,(self.N,1))
    dy = ytemp - ytemp.T
    ztemp = tile(z,(self.N,1))
    dz = ztemp - ztemp.T

    # Particles 'feel' each other across the periodic boundaries
    if self.periodicX:
        dx[dx>self.L/2]=dx[dx > self.L/2]-self.L
        dx[dx<-self.L/2]=dx[dx < -self.L/2]+self.L
    if self.periodicY:
        dy[dy>self.L/2]=dy[dy>self.L/2]-self.L
        dy[dy<-self.L/2]=dy[dy<-self.L/2]+self.L
    if self.periodicZ:
        dz[dz>self.L/2]=dz[dz>self.L/2]-self.L
        dz[dz<-self.L/2]=dz[dz<-self.L/2]+self.L

    # Total Distances
    d = sqrt(dx**2+dy**2+dz**2)

    # Mark zero entries with negative 1 to avoid divergences
    d[d==0] = -1

    return d, dx, dy, dz

据我所知,Parakeet应该能够不用修改就使用上面的函数--它只使用Numpy和数学。但是,在从Parakeet jit包装器调用函数时,我总是会得到以下错误:

代码语言:javascript
复制
AssertionError: Unsupported function: <bound method Particles.distanceMatrix of <particles.Particles instance at 0x04CD8E90>>
EN

回答 1

Stack Overflow用户

回答已采纳

发布于 2013-11-24 19:34:47

Parakeet还很年轻,它的NumPy支持是不完整的,您的代码涉及到几个尚未工作的特性。

1)您正在包装一个方法,而Parakeet到目前为止只知道如何处理函数。常见的解决方法是创建一个@jit包装的助手函数,并让您的方法调用包含所有必需的成员数据。方法不起作用的原因是将有意义的类型分配给“self”是非常重要的。这不是不可能的,但足够棘手的是,方法不会进入鹦鹉,直到较低的悬挂水果被摘下。说到低挂水果..。

2)布尔索引。尚未实现,但将在下一个版本中实现。

3) np.tile:也不起作用,可能也会出现在下一个版本中。如果您想查看哪些内置程序和NumPy库函数可以工作,请查看Parakeet的映射模块。

我重写了你的代码,以便对鹦鹉更友好一点:

代码语言:javascript
复制
@jit 
def parakeet_dist(x, y, z, L, periodicX, periodicY, periodicZ):
  # perform all-pairs computations more explicitly 
  # instead of tile + broadcasting
  def periodic_diff(x1, x2, periodic):
    diff = x1 - x2 
    if periodic:
      if diff > (L / 2): diff -= L
      if diff < (-L/2): diff += L
    return diff
  dx = np.array([[periodic_diff(x1, x2, periodicX) for x1 in x] for x2 in x])
  dy = np.array([[periodic_diff(y1, y2, periodicY) for y1 in y] for y2 in y])
  dz = np.array([[periodic_diff(z1, z2, periodicZ) for z1 in z] for z2 in z])
  d= np.sqrt(dx**2 + dy**2 + dz**2)

  # since we can't yet use boolean indexing for masking out zero distances
  # have to fall back on explicit loops instead 
  for i in xrange(len(x)):
    for j in xrange(len(x)):
      if d[i,j] == 0: d[i,j] = -1 
  return d, dx, dy, dz 

在我的机器上,N= 2000的运行速度仅比NumPy快3倍( NumPy为0.39s,鹦鹉为0.14s )。如果重写数组遍历以更显式地使用循环,那么性能将比NumPy快6倍(Parakeet运行在~0.06s):

代码语言:javascript
复制
@jit 
def loopy_dist(x, y, z, L, periodicX, periodicY, periodicZ):
  N = len(x)
  dx = np.zeros((N,N))
  dy = np.zeros( (N,N) )
  dz = np.zeros( (N,N) )
  d = np.zeros( (N,N) )

  def periodic_diff(x1, x2, periodic):
    diff = x1 - x2 
    if periodic:
      if diff > (L / 2): diff -= L
      if diff < (-L/2): diff += L
    return diff

  for i in xrange(N):
    for j in xrange(N):
      dx[i,j] = periodic_diff(x[j], x[i], periodicX)
      dy[i,j] = periodic_diff(y[j], y[i], periodicY)
      dz[i,j] = periodic_diff(z[j], z[i], periodicZ)
      d[i,j] = dx[i,j] ** 2 + dy[i,j] ** 2 + dz[i,j] ** 2 
      if d[i,j] == 0: d[i,j] = -1
      else: d[i,j] = np.sqrt(d[i,j])
  return d, dx, dy, dz 

通过一些创造性的重写,您也可以在Numba中运行上面的代码,但它的速度比NumPy快1.5倍(0.25秒)。编译时间为:鹦鹉w/理解:1秒,鹦鹉w/循环:.5秒,Numba w/循环:0.9秒。

希望接下来的几个版本能够更实际地使用NumPy库函数,但就目前而言,理解或循环通常是可行的。

票数 4
EN
页面原文内容由Stack Overflow提供。腾讯云小微IT领域专用引擎提供翻译支持
原文链接:

https://stackoverflow.com/questions/20167572

复制
相关文章

相似问题

领券
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档