使用scipy.optimize.fmin拟合正弦函数报错:setting an array element with a sequence
问题分析与修复方案
这个报错的核心原因是你的目标函数func_error返回的是数组,但scipy.optimize.fmin要求目标函数必须返回一个单个标量值(即误差的总和,一个数字)。下面一步步拆解问题并给出修复代码:
问题具体点
func_noise的噪声生成逻辑错误:当传入单个x时,np.random.randn(100)会生成长度为100的数组,但func_model(x, para_fact)是单个数值,两者相加后得到的是长度100的数组,而非对应x的带噪声的单个y值。func_error返回值类型错误:循环中每次累加的是长度100的数组,最终error_sum还是数组,不符合fmin对目标函数返回值的要求。- 循环效率低下:用for循环遍历x序列计算误差,不如用numpy的向量化操作高效。
修复后的完整代码
import numpy as np import scipy.optimize as opt import matplotlib.pyplot as plt def func_model(x, para): ''' Model: y = a*sin(2*k*pi*x+theta)''' a, k, theta = para return a*np.sin(2*k*np.pi*x + theta) def func_noise(x, para): # 生成和x长度一致的噪声,而非固定100个 noise = np.random.randn(len(x)) return func_model(x, para) + noise def func_error(para_guess): '''误差函数:返回误差平方和的标量''' x_seq = np.linspace(-2*np.pi, 0, 100) para_fact = [10, 0.34, np.pi/6] # 向量化计算,避免循环,直接得到所有误差的平方和 y_true_noise = func_noise(x_seq, para_fact) y_pred = func_model(x_seq, para_guess) error_sum = np.sum((y_true_noise - y_pred)**2) return error_sum # 初始猜测参数 para_guess_init = np.array([7, 0.2, 0]) # 使用fmin求解 solution = opt.fmin(func_error, para_guess_init) print("拟合得到的参数:", solution) # 可选:可视化拟合效果 x_test = np.linspace(-2*np.pi, 0, 200) y_true = func_model(x_test, [10, 0.34, np.pi/6]) y_fit = func_model(x_test, solution) y_noise = func_noise(x_test, [10, 0.34, np.pi/6]) plt.scatter(x_test, y_noise, label="带噪声数据", s=5) plt.plot(x_test, y_true, label="真实曲线", color="r") plt.plot(x_test, y_fit, label="拟合曲线", color="g", linestyle="--") plt.legend() plt.show()
关键修改说明
func_noise调整:用len(x)获取输入x的长度,生成对应长度的噪声数组,确保输出和输入x的维度匹配。func_error重构:去掉for循环,直接用numpy的向量化计算所有x的误差平方和,最终返回一个标量值,满足fmin的要求。- 可选优化:如果是做最小二乘拟合,推荐使用
scipy.optimize.leastsq,它专门针对最小平方误差的场景,效率更高。示例如下:
# 用leastsq的写法(误差函数返回残差数组而非平方和) def residuals(para_guess, x, y): return y - func_model(x, para_guess) x_seq = np.linspace(-2*np.pi, 0, 100) para_fact = [10, 0.34, np.pi/6] y_true_noise = func_noise(x_seq, para_fact) solution_leastsq, _ = opt.leastsq(residuals, para_guess_init, args=(x_seq, y_true_noise)) print("leastsq拟合参数:", solution_leastsq)
内容的提问来源于stack exchange,提问作者sun0727
相关产品推荐
相关产品推荐

