使用scipy.optimize.curve_fit结合numpy.piecewise为何抛出属性与运行时错误?
分段曲线拟合中curve_fit触发numpy内部错误问题
我正在开发一个曲线拟合的简化示例,该分段函数由两个常数函数和两个线性函数组成。调用函数生成y_space时运行正常,但使用scipy.optimize.curve_fit时触发了numpy内部错误。我对numpy内部机制了解不足,无法着手排查,也未见过其他人报告此类错误。(不使用np.piecewise编写函数时也遇到过类似但不同的错误,目前仅聚焦此问题)。
代码示例
# Imports from scipy.optimize import curve_fit import numpy as np import matplotlib.pyplot as plt plt.style.use("dark_background") # Piecewise function def my_piecewise( x: float, a: float, b: float, c: float, d: float, precision: float = 64, ): domain_length = 512 if precision == 32: domain_length = 1024 if precision not in (32, 64): raise ValueError("Precision must be either 64 or 32") y = np.piecewise( x, [ x <= 0, x < domain_length, (x >= domain_length) & (x < 2 * domain_length), x >= 2 * domain_length, ], [ lambda x: 0, lambda x: a + ((b - a) / (domain_length - 8)) * x, lambda x: c + ((d - c) / domain_length) * (x - domain_length), lambda x: d, ], ) return y # p0 parameters for curve_fit init_guess = [ 0.006, 0.0065, 0.006, 0.0065, ] # x and y data x_space = np.linspace(8, 4096, 512, endpoint=True) y_space = my_piecewise(x_space, *init_guess) # plotting with matplotlib works fine plt.plot(x_space, y_space) for i in [512, 1024, 2048]: plt.vlines(i, 0.004, 0.008, color="white", linewidth=0.75) plt.show() # curve_fit param, param_cov = curve_fit( f=my_piecewise, xdata=x_space, ydata=y_space, p0=init_guess, )
Matplotlib绘图结果

错误信息
Traceback (most recent call last): File "/path/to/conda/env/lib/python3.12/site-packages/numpy/_core/arrayprint.py", line 34, in <module> from . import numerictypes as _nt File "/path/to/conda/env/lib/python3.12/site-packages/numpy/_core/numerictypes.py", line 102, in <module> from ._type_aliases import ( File "/path/to/conda/env/lib/python3.12/site-packages/numpy/_core/_type_aliases.py", line 38, in <module> allTypes[_abstract_type_name] = getattr(ma, _abstract_type_name) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ AttributeError: module 'numpy.core.multiarray' has no attribute 'signedinteger' The above exception was the direct cause of the following exception: Traceback (most recent call last): File "/home/garrett/Documents/projects/benchmark-data-fitting/temp.py", line 63, in <module> param, param_cov = curve_fit( ^^^^^^^^^^ File "/path/to/conda/env/lib/python3.12/site-packages/scipy/optimize/_minpack_py.py", line 1007, in curve_fit res = leastsq(func, p0, Dfun=jac, full_output=1, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/path/to/conda/env/lib/python3.12/site-packages/scipy/optimize/_minpack_py.py", line 439, in leastsq retval = _minpack._lmdif(func, x0, args, full_output, ftol, xtol, ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/path/to/conda/env/lib/python3.12/site-packages/scipy/optimize/_minpack_py.py", line 519, in _memoized_func if np.all(_memoized_func.last_params == params): ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/path/to/conda/env/lib/python3.12/site-packages/numpy/core/_internal.py", line 855, in array_ufunc_errmsg_formatter args_string = ', '.join(['{!r}'.format(arg) for arg in inputs] + ^^^^^^^^^^^^^^^^^^ RuntimeError: Unable to configure default ndarray.__repr__
内容的提问来源于stack exchange,提问作者fngarrett
相关产品推荐
相关产品推荐

