如何使用scipy的curve_fit处理多组带共享参数的二维x,y数据集
scipy.optimize.curve_fit 拟合多组二维x带共享参数数据集的解决方案
错误根因
原示例的适用前提是所有组共用同一个一维x数组,而你的场景中3组数据各对应独立的x序列,直接复用原逻辑会导致拟合函数输出的预测值长度和打平后的y数组长度不匹配,触发广播错误。
解决方案
方案1:打平x数组后传入,在函数内拆分分组
该方案符合curve_fit对数值型xdata的常规传参规则,适配性更强。
首先定义基础拟合函数与适配后的全局拟合函数:
import numpy as np from scipy.optimize import curve_fit # 自定义基础拟合函数,a为共享参数,b为组独有参数,可根据需求替换为其他函数形式 def f(x, a, b): return a * x + b # 全局拟合函数,负责分组计算后拼接结果 def g(x_flat, a, b_1, b_2, b_3): # 将打平的x还原为原始(3,4)结构 x_group = x_flat.reshape(3, 4) # 分别计算每组预测值后拼接为一维数组,和y.ravel()长度对齐 return np.concatenate([ f(x_group[0], a, b_1), f(x_group[1], a, b_2), f(x_group[2], a, b_3) ])
调用拟合逻辑:
# 你的示例输入 x = np.random.rand(3,4) y = np.random.rand(3,4) # 传入打平的x和打平的y完成拟合 params, _ = curve_fit(g, x.ravel(), y.ravel()) # 提取参数:第一个元素为共享参数a,后续为各组独有参数 a, b1, b2, b3 = params
方案2:直接传入完整二维x数组(适合组数量可变场景)
利用curve_fit支持xdata为任意对象的特性,无需打平x,代码扩展性更强:
def g_obj(x_full, a, *b_list): # 自动遍历所有组计算预测值 pred_list = [f(x_full[i], a, b_list[i]) for i in range(x_full.shape[0])] return np.concatenate(pred_list) # 调用时直接传入完整二维x,y仍需打平,可通过p0指定初始参数提升拟合稳定性 params, _ = curve_fit(g_obj, x, y.ravel(), p0=[1, 0, 0, 0])
两种方案的核心都是保证拟合函数输出的一维预测值长度为3*4=12,和y.ravel()的长度完全匹配,即可解决维度不匹配的报错问题。
内容的提问来源于stack exchange,提问作者omanuelcosta
相关产品推荐
相关产品推荐

