Numba-Scipy中erfc不支持复数,实现wofz函数遇报错求助
Numba-Scipy中复数参数调用erfc函数报错解决
问题背景
在Numba-Scipy环境下编写并行函数,因Numba-Scipy不支持wofz函数,计划用erfc复现功能,但erfc仅能处理实数,传入复数时触发如下TypingError:
TypingError: No implementation of function Function(<ufunc 'erfc'>) found for signature: erfc(complex128) There are 2 candidate implementations: - Of which 2 did not match due to: Overload of function 'erfc': File: numba_scipy/special/overloads.py: Line 9. With argument(s): '(complex128)': No match.
原函数代码:
from numba import njit, parallel import numpy as np @njit(parallel=True) def faster(first,Hz, DRM, vpec,velaxis, bh, h1f, nh): t1=np.zeros(2048) t4=np.zeros(2048) t5=np.zeros(2048) for i in range(2048): va=0 va1=0 va1_new=0 for j in range(i-100, i+100): if j<0: ne=2048+j veldif=(i-j)*Hz*DRM - (vpec[ls,ne]) if j>2047: ne=j-2048 veldif=((i-j))*Hz*DRM*1 - (vpec[ls,ne]) if 0<= j <=2047: ne=j veldif=velaxis[i] - (velaxis[j]+vpec[ls,j]) powe=((veldif))/bh[ls,ne] u=powe x=u a=(6.265*10**8)*(1215.67*10**-13)*((4*3.14)**-1)*((bh[ls,ne])**-1) z = (u + 1j*a) zp=-1j*z wz1=np.exp(-a**2 -u**2) wz2=np.real(erfc(1.0+1j*4.9)) wz3=wz1*wz2 va=va+((h1f[ls,ne]*nh[ls,ne])/bh[ls,ne])*wz3 va1=va1+((h1f[ls,ne]*nh[ls,ne])/bh[ls,ne])*np.exp(-u**2) t1[i]=i t4[i]=first*va1 t5[i]=first*va return(t1,t4, t5)
解决方案
1. 自定义Numba兼容的复数erfc函数
Numba-Scipy的erfc仅支持实数输入,需手动实现复数版本。基于复误差函数的数学性质,利用实数erf/erfc和基本三角函数、指数函数实现:
@njit def complex_erfc(z): x = np.real(z) y = np.imag(z) # 利用复误差函数的对称性,先处理x≥0的情况 if x >= 0: erf_x = np.erf(x) erfc_x = 1 - erf_x exp_term = np.exp(-y**2) cos_term = np.cos(2 * x * y) sin_term = np.sin(2 * x * y) real_part = erfc_x * exp_term * cos_term imag_part = -erf_x * exp_term * sin_term else: # x<0时,利用erfc(-z) = 2 - erfc(z)转换为x>0的计算 erfc_pos = complex_erfc(-z) real_part = 2 - np.real(erfc_pos) imag_part = -np.imag(erfc_pos) return real_part + 1j * imag_part
2. 修改主函数调用
将原函数中调用erfc的部分替换为自定义的complex_erfc,同时修复原代码中未定义的ls变量(需作为参数传入函数):
@njit(parallel=True) def faster(first,Hz, DRM, vpec,velaxis, bh, h1f, nh, ls): # 新增ls参数 t1=np.zeros(2048) t4=np.zeros(2048) t5=np.zeros(2048) for i in range(2048): va=0 va1=0 va1_new=0 for j in range(i-100, i+100): if j<0: ne=2048+j veldif=(i-j)*Hz*DRM - (vpec[ls,ne]) if j>2047: ne=j-2048 veldif=((i-j))*Hz*DRM*1 - (vpec[ls,ne]) if 0<= j <=2047: ne=j veldif=velaxis[i] - (velaxis[j]+vpec[ls,j]) powe=((veldif))/bh[ls,ne] u=powe x=u a=(6.265*10**8)*(1215.67*10**-13)*((4*3.14)**-1)*((bh[ls,ne])**-1) z = (u + 1j*a) zp=-1j*z wz1=np.exp(-a**2 -u**2) wz2=np.real(complex_erfc(1.0+1j*4.9)) # 替换为自定义函数 wz3=wz1*wz2 va=va+((h1f[ls,ne]*nh[ls,ne])/bh[ls,ne])*wz3 va1=va1+((h1f[ls,ne]*nh[ls,ne])/bh[ls,ne])*np.exp(-u**2) t1[i]=i t4[i]=first*va1 t5[i]=first*va return(t1,t4, t5)
3. 精度与性能优化提示
- 若对计算精度要求极高,可替换
complex_erfc的实现为更精确的级数展开或积分近似(需确保代码可被Numba编译)。 - 自定义函数经
njit编译后,性能与Numba-Scipy内置函数接近,适配并行环境。
内容的提问来源于stack exchange,提问作者s.a
相关产品推荐
相关产品推荐

