【发布时间】:2021-06-20 20:25:09
【问题描述】:
我需要在我的设置中为 10K 点计算以下系列总和。所以我使用 numba 来加速我的计算
@jit(nopython=True)
def kernel(x, y, t, m, n):
return (1 + (-1)**(m+1))*(1 - np.cos(0.5*n*pi))*np.sin(0.5*m*pi*x)*np.sin(0.5*n*pi*y)*np.exp(-(pi**2)*(m**2 + n**2)*t/36)/(m*n)
@jit(nopython=True)
def Series_Sum(x, y, t, m, n):
res = 0
for i in np.linspace(1, m, m):
for j in np.linspace(1, n, n):
res += kernel(x, y, t, i, j)
# print(res)
return 200*res/(pi**2)
x = np.linspace(0, 2, 101)
y = np.linspace(0, 2, 101)
X, Y = np.meshgrid(x, y)
Z = np.concatenate((X.flatten()[:, None], Y.flatten()[:, None]), axis=1)
m, n = 100, 100
exact =[Series_Sum(i[0], i[1], 0, m, n) for i in Z]
但是,结果都是'nan'。
例如
Series_Sum(0.3,1.5,1,100,2) # returns nan
如果我执行以下操作
res = 0
for i in np.linspace(1, m, m):
for j in np.linspace(1, n, n):
res += kernel(x, y, t, i, j)
结果很好。
另外,如果我删除 '@jit' decrator,结果是合理的,但是计算结果需要几个小时。
有没有更好的方法来解决这个问题?
【问题讨论】:
标签: python python-3.x numpy numba