Numba nopython模式类型错误求助:参数类型不匹配
问题解决:Numba TypingError 参数类型不匹配
问题根源
报错核心是funpp函数被错误定义为仅接受7个int32类型参数,但实际调用时传入的是complex128类型。原因是你给fun、funpp等函数的JIT签名完全错误——波矢参数(k_x、k_y、q_x、q_y等)应为浮点/复数类型,只有n、tau这类索引参数才是整数类型。
修复步骤
1. 修正所有函数的JIT签名
针对每个函数的参数用途,匹配正确的类型:
E/nf:前两个参数n、tau为整数(intc),后两个波矢参数为complex128I2/fun/funpp:前四个为波矢参数(complex128),后三个为整数参数(intc)
2. 处理积分回调的类型转换
jit_integrand_function中从CPointer(complex128)取出的所有参数都是复数类型,需要把整数参数(n_temp、m_temp、tau_temp)转换成intc类型,避免传递复数给要求整数的函数。
修复后的完整代码
import numpy as np import scipy.linalg as la from scipy import integrate, LowLevelCallable import matplotlib.pyplot as plt from numba import cfunc, jit from numba.types import intc, CPointer, complex128 # 全局参数 t0 = 2.61 t4 = 0.138 t3 = 0.283 t1 = 0.361 Delta = 0.015 d = 3.5 a = 2.46 D = 50e-3 mu = -0.02487 B = 0 theta = 0 omp = 1e-8 T = 0.9 # 修正签名:n、tau为intc,k_x、k_y为complex128 @jit(complex128(intc, intc, complex128, complex128), nopython=True) def E(n, tau, k_x, k_y): v = t0*a*np.sqrt(3)/2 v4 = t4*a*np.sqrt(3)/2 v3 = t3*a*np.sqrt(3)/2 p = tau*k_x + 1j*k_y p1 = tau*(k_x + B*d/2*np.sin(theta)) + 1j*(k_y - B*d/2*np.cos(theta)) p2 = tau*(k_x - B*d/2*np.sin(theta)) - 1j*(k_y - B*d/2*np.cos(theta)) H = np.array([[D/2-mu, v*np.conj(p1), -v4*np.conj(p), -v3*p], [v*p1, Delta + D/2 - mu, t1, -v4*np.conj(p)], [-v4*p, t1, Delta - D/2 - mu, v*np.conj(p2)], [-v3*np.conj(p), -v4*p, v*p2, -D/2 - mu]], dtype=np.complex128) eigenvals, eigenvecs = la.eig(H) # 确保索引为整数 n_idx = int(np.real(n)) return eigenvals[n_idx] # 修正签名 @jit(complex128(intc, intc, complex128, complex128), nopython=True) def nf(n, tau, k_x, k_y): return 1.0 / (1. + np.exp(-2. * (mu - E(n, tau, k_x, k_y)) / T)) # 修正签名:前四个为complex128,后三个为intc @jit(complex128(complex128, complex128, complex128, complex128, intc, intc, intc), nopython=True) def I2(k_x, k_y, q_x, q_y, n_temp, m_temp, tau_temp): numerator = nf(n_temp, tau_temp, q_x, q_y) - nf(m_temp, tau_temp, q_x + k_x, q_y + k_y) denominator = omp + E(n_temp, tau_temp, q_x, q_y) - E(m_temp, tau_temp, q_x + k_x, q_y + k_y) return numerator / denominator @jit(complex128(complex128, complex128, complex128, complex128, intc, intc, intc), nopython=True) def fun(k_x,k_y,q_x,q_y,n_temp,m_temp,tau_temp): return I2(k_x,k_y,q_x,q_y,n_temp,m_temp,tau_temp) @jit(complex128(complex128, complex128, complex128, complex128, intc, intc, intc), nopython=True) def funpp(k_x,k_y,q_x,q_y,n_temp,m_temp,tau_temp): return fun(k_x,k_y,q_x,q_y,n_temp,m_temp,tau_temp)/(2.*np.pi)**2. def jit_integrand_function(integrand_function): jitted_function = jit(integrand_function, nopython=True) @cfunc(complex128(intc, CPointer(complex128))) def wrapped(n, xx): # 提取参数并转换整数类型 k_x = xx[0] k_y = xx[1] q_r = xx[2] q_t = xx[3] n_temp = intc(np.real(xx[4])) m_temp = intc(np.real(xx[5])) tau_temp = intc(np.real(xx[6])) return jitted_function(k_x, k_y, q_r, q_t, n_temp, m_temp, tau_temp) return LowLevelCallable(wrapped.ctypes) @jit_integrand_function def funppr(k_x,k_y,q_r,q_t,n_temp,m_temp,tau_temp): q_x = q_r * np.cos(q_t) q_y = q_r * np.sin(q_t) return funpp(k_x,k_y,q_x,q_y,n_temp,m_temp,tau_temp)*q_r
额外注意事项
- 用
scipy.linalg.eig替代numpy.linalg.eig,Numba对SciPy线性代数函数的支持更稳定 - 显式指定矩阵
H的dtype=np.complex128,避免类型推断错误 - 所有JIT函数添加
nopython=True强制纯编译模式,提升性能同时规避类型问题
内容的提问来源于stack exchange,提问作者physshyp
相关产品推荐
相关产品推荐

