使用Numba @njit优化Python薛定谔模拟函数时的编译错误求助
解决Numba @njit装饰薛定谔方程求解函数的问题
我来帮你一步步排查并解决这些问题,让Numba的优化能正常工作:
1. 修复np.zeros_like的类型匹配错误
你遇到的TypingError核心原因是:Numba不支持将字符串形式的'complex'作为dtype参数传入np.zeros_like,必须传入实际的numpy类型对象(比如np.complex128)。
把这行代码:
dphi = np.zeros_like(z, dtype = complex)
修改为:
dphi = np.zeros_like(phi, dtype=np.complex128)
(改用输入参数phi的形状来创建数组,避免依赖全局变量z,同时明确指定Numba能识别的复数类型)
2. 移除全局变量依赖,用闭包传递参数
Numba的@njit对全局变量的支持非常有限,而且全局变量会降低代码的可维护性和编译效率。我们可以用functools.partial把需要的参数(比如h、网格数m、a、b)封装成闭包,适配solve_ivp要求的(t, phi)回调签名:
首先修改res函数,将依赖的参数作为显式输入:
@njit def res(t, phi, h, m_grid, a, b): dphi = np.zeros_like(phi, dtype=np.complex128) for i in range(m_grid): if i != 0 and i != m_grid-1: dphi[i] = 1/h**2*1j/Ll(a,b,t)**2*(phi[i+1]-2*phi[i]+phi[i-1]) + 1/2/h*i*dLl(b)/Ll(a,b,t)*(phi[i+1] - phi[i-1]) else: if i == 0: dphi[i] = 1/h**2*1j/Ll(a,b,t)**2*(phi[i+1]-2*phi[i]+ 0) + 1/2/h*i*dLl(b)/Ll(a,b,t)*(phi[i+1] - 0) if i == m_grid-1: dphi[i] = 1/h**2*1j/Ll(a,b,t)**2*(0 -2*phi[i]+phi[i-1]) + 1/2/h*i*dLl(b)/Ll(a,b,t)*(0 - phi[i-1]) return dphi
然后用functools.partial包装函数,传入固定参数:
from functools import partial # 包装res函数,传入h、网格数、a、b等固定参数 res_wrapped = partial(res, h=h, m_grid=m, a=a, b=b)
最后调用solve_ivp时使用这个包装后的函数:
sol = ivp(res_wrapped, (t0,tf), y0, t_eval=t)
3. 确保所有被调用的辅助函数都被Numba编译
你已经尝试给辅助函数加@njit,这是正确的,但要确保函数签名是Numba能处理的,同时避免全局变量依赖(比如把omega作为参数传入Ls和dLs):
@njit def Ll(a,b,t): return a + b*t @njit def dLl(b): return b @njit def Ls(a,b,t, omega): return a + b*np.sin(omega*t) @njit def dLs(b,t, omega): return omega*b*np.cos(omega*t)
4. 其他优化与细节修复
- 移除print语句:Numba的
njit模式下,print对复数数组等复杂类型的支持不佳,还会严重拖慢性能,建议删除print(i,phi[i])和print(dphi)这类调试代码。 - 初始条件转复数:因为你的微分方程返回的是复数数组,最好把初始条件
y0转成复数类型,避免类型不匹配:y0 = f(z).astype(np.complex128) - 避免变量重定义:你的代码里两次定义了
m(一开始的粒子质量m=1/2和后面的网格数m=50),这会导致混淆,建议重命名其中一个变量(比如把网格数改成m_grid)。
完整修改后的代码示例
import numpy as np from scipy.integrate import solve_ivp as ivp from numba import njit from functools import partial # 物理参数(重命名避免和网格数冲突) mass = 1/2 # Un. Geom. h_bar = 1 # Unidades geometrizadas infty = 5000 # Esto sirve como infinito # ============================================================================= # # Parámetros # ============================================================================= z0 = 0 zf = 1 omega = 1 t0 = 0 tf = 1 a = 1 b = 3 m_grid = 50 # 网格数,重命名避免和质量变量冲突 # ============================================================================= # # Funciones # ============================================================================= @njit def f(z): return 2**0.5*np.sin(2*np.pi*z) @njit def Ll(a,b,t): return a + b*t @njit def dLl(b): return b @njit def Ls(a,b,t, omega): return a + b*np.sin(omega*t) @njit def dLs(b,t, omega): return omega*b*np.cos(omega*t) @njit def res(t, phi, h, m_grid, a, b): dphi = np.zeros_like(phi, dtype=np.complex128) for i in range(m_grid): if i != 0 and i != m_grid-1: dphi[i] = 1/h**2*1j/Ll(a,b,t)**2*(phi[i+1]-2*phi[i]+phi[i-1]) + 1/2/h*i*dLl(b)/Ll(a,b,t)*(phi[i+1] - phi[i-1]) else: if i == 0: dphi[i] = 1/h**2*1j/Ll(a,b,t)**2*(phi[i+1]-2*phi[i]+ 0) + 1/2/h*i*dLl(b)/Ll(a,b,t)*(phi[i+1] - 0) if i == m_grid-1: dphi[i] = 1/h**2*1j/Ll(a,b,t)**2*(0 -2*phi[i]+phi[i-1]) + 1/2/h*i*dLl(b)/Ll(a,b,t)*(0 - phi[i-1]) return dphi # ============================================================================= # # Resolución numérica # Esto se hace con ivp # ============================================================================= h = 1/m_grid z = np.linspace(z0+h,zf-h,m_grid) z = np.transpose(z) print('tamaño z = ', len(z)) y0 = f(z).astype(np.complex128) t = np.linspace(t0,tf,m_grid) # 包装res函数,传入固定参数 res_wrapped = partial(res, h=h, m_grid=m_grid, a=a, b=b) sol = ivp(res_wrapped, (t0,tf), y0, t_eval=t) print(sol.y)
这些修改应该能解决你遇到的两个错误,同时让Numba正常优化代码,大幅提升运行速度。
内容的提问来源于stack exchange,提问作者Julian Antonio Bertolini Suare
相关产品推荐
相关产品推荐

