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

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),后两个波矢参数为complex128
  • I2/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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 10:55:40