Scipy curve_fit调用自定义函数报错:unhashable type: 'numpy.ndarray'求解
自定义函数线性组合拟合报错解决
问题重现
定义了以下拟合函数:
from scipy.optimize import curve_fit import numpy as np def fit(f1, f2, x_data, y_data): def fun(x, w1, w2): return w1*f1(x) + w2*f2(x) return curve_fit(fun, x_data, y_data)
使用np.sin/np.cos作为基函数时代码可正常运行:
x_data, y_data = np.array(range(0, 5)), np.array([0, 1, 2, 3, 2]) f1 = np.sin f2 = np.cos fit(f1, f2, x_data, y_data)
但改用字典映射的lambda函数时,抛出错误unhashable type: 'numpy.ndarray':
d1 = {0: 1, 1:3, 2:4, 3:5, 4:3} d2 = {0: 2, 1:1, 2:6, 3:2, 4:1} f1 = lambda x: d1.get(x) f2 = lambda x: d2.get(x) # 修正原代码笔误,原代码写成了d1.get(x) fit(f1, f2, x_data, y_data)
报错原因
curve_fit会将x_data作为完整的numpy数组传入fun中的x参数,而非逐个元素传入。此时lambda x: d1.get(x)中的x是numpy数组,而字典的键要求是可哈希类型(数组不可哈希),因此触发报错。
解决方法
方法1:适配数组输入的基函数
修改基函数,使其能处理数组类型的输入,推荐两种方式:
- 用
np.vectorize包装lambda函数,将逐元素操作向量化:
f1 = np.vectorize(lambda x: d1.get(x, 0)) # 加入默认值避免KeyError f2 = np.vectorize(lambda x: d2.get(x, 0))
- 手动编写处理数组的函数:
def f1(x): return np.array([d1.get(val, 0) for val in x]) def f2(x): return np.array([d2.get(val, 0) for val in x])
修改后调用fit函数即可正常运行。
方法2:改用线性回归直接求解
由于需求是自定义函数的线性组合拟合,本质属于线性回归问题,可直接构造设计矩阵,用最小二乘法求解,无需使用curve_fit:
import numpy as np x_data = np.array(range(0, 5)) y_data = np.array([0, 1, 2, 3, 2]) d1 = {0: 1, 1:3, 2:4, 3:5, 4:3} d2 = {0: 2, 1:1, 2:6, 3:2, 4:1} # 构造设计矩阵:每一行对应[f1(xi), f2(xi)] X = np.column_stack([ np.array([d1.get(x, 0) for x in x_data]), np.array([d2.get(x, 0) for x in x_data]) ]) # 最小二乘法求解权重w1, w2 w, _, _, _ = np.linalg.lstsq(X, y_data, rcond=None) print(f"权重w1: {w[0]}, w2: {w[1]}")
这种方式更高效,也避免了curve_fit的函数适配问题,适合线性组合类的拟合场景。
内容的提问来源于stack exchange,提问作者user3603644
相关产品推荐
相关产品推荐

