You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.27 09:52:27