共享参数的多数据集全局拟合提速方案及替代工具咨询
快速拟合共享参数的多数据集三次多项式模型
你遇到的问题其实是典型的混合线性-非线性多数据集拟合场景:a0、a1、a2、a3是所有数据集共享的线性参数,每个数据集独有的偏移量c则是非线性参数。symfit虽然用起来方便,但它的通用符号计算框架在数据集数量增多时,会带来不小的额外开销,导致拟合速度变慢。下面给你推荐几种高效的替代方案:
方案1:用scipy.optimize手动实现分离变量拟合
这里的核心思路是分离线性与非线性参数:线性参数(a0-a3)可以在固定c的情况下,通过最小二乘的闭式解直接算出,不需要迭代求解。这样我们只需要优化非线性的c参数,计算量会大幅降低。具体步骤如下:
- 对每个数据集,给定当前的c_i,构造对应的设计矩阵(每行是
[1, (x_ij - c_i), (x_ij - c_i)², (x_ij - c_i)³]) - 把所有数据集的设计矩阵和y值拼接成全局矩阵,用线性最小二乘求解a0-a3
- 以总残差平方和为目标,用scipy的非线性优化器优化所有c_i
示例代码:
import numpy as np from scipy.optimize import minimize def compute_total_residual(cs, all_x, all_y): num_datasets = len(cs) global_X = [] global_y = [] for idx in range(num_datasets): x_data = all_x[idx] c = cs[idx] x_shifted = x_data - c # 构造当前数据集的设计矩阵 design_matrix = np.column_stack([ np.ones_like(x_shifted), x_shifted, x_shifted ** 2, x_shifted ** 3 ]) global_X.append(design_matrix) global_y.append(all_y[idx]) # 拼接成全局矩阵 global_X = np.vstack(global_X) global_y = np.concatenate(global_y) # 求解线性参数的闭式解 linear_params, _, _, _ = np.linalg.lstsq(global_X, global_y, rcond=None) # 计算总残差平方和 y_pred = global_X @ linear_params return np.sum((global_y - y_pred) ** 2) # 初始化参数:用每个数据集x的均值作为c的初始值,收敛更快 num = 100 all_x = [...] # 你的x数据集列表,每个元素是一个数据集的x数组 all_y = [...] # 对应的y数据集列表 initial_c_values = [np.mean(x_set) for x_set in all_x] # 执行优化,L-BFGS-B适合这种参数较多的优化场景 optim_result = minimize( compute_total_residual, initial_c_values, args=(all_x, all_y), method='L-BFGS-B' ) # 提取最优参数 optimized_cs = optim_result.x # 用最优c计算最终的a0-a3 final_X = [] final_y = [] for idx in range(num): x_shifted = all_x[idx] - optimized_cs[idx] final_X.append(np.column_stack([ np.ones_like(x_shifted), x_shifted, x_shifted**2, x_shifted**3 ])) final_X = np.vstack(final_X) final_y = np.concatenate(all_y) a0, a1, a2, a3 = np.linalg.lstsq(final_X, final_y, rcond=None)[0]
方案2:用lmfit库(更推荐)
lmfit是专门为曲线拟合设计的库,它对多参数、参数共享的场景支持非常好,而且内部的优化器效率很高,不需要你手动处理矩阵拼接这类细节。
示例代码:
import lmfit import numpy as np # 定义三次多项式模型 def cubic_fit(x, a0, a1, a2, a3, c): x_shifted = x - c return a0 + a1 * x_shifted + a2 * x_shifted**2 + a3 * x_shifted**3 # 创建参数对象:共享的a0-a3,每个数据集独立的c params = lmfit.Parameters() # 给共享参数设合理初始值 params.add('a0', value=0) params.add('a1', value=1) params.add('a2', value=0) params.add('a3', value=0) # 给每个数据集的c设初始值(用对应x的均值) num = 100 for i in range(num): params.add(f'c{i}', value=np.mean(all_x[i])) # 定义目标函数:返回所有数据集的残差数组 def objective_function(params, all_x, all_y): residuals = [] # 提取共享参数 a0 = params['a0'].value a1 = params['a1'].value a2 = params['a2'].value a3 = params['a3'].value # 遍历每个数据集计算残差 for idx in range(num): c_val = params[f'c{idx}'].value y_pred = cubic_fit(all_x[idx], a0, a1, a2, a3, c_val) residuals.extend(all_y[idx] - y_pred) return np.array(residuals) # 执行拟合 minimizer = lmfit.Minimizer(objective_function, params, fcn_args=(all_x, all_y)) fit_result = minimizer.minimize(method='L-BFGS-B') # 查看拟合结果 lmfit.report_fit(fit_result) # 提取最优参数 a0_opt = fit_result.params['a0'].value a1_opt = fit_result.params['a1'].value a2_opt = fit_result.params['a2'].value a3_opt = fit_result.params['a3'].value optimized_cs = [fit_result.params[f'c{i}'].value for i in range(num)]
为什么symfit会慢?
symfit的核心是符号计算,它会先为所有数据集生成完整的符号表达式,再将其转换为数值计算代码。当数据集数量增加时,符号表达式的复杂度会线性上升,带来大量的编译和解析开销。而上面的两种方案都是直接操作数值数据,避免了符号计算的额外成本,所以在处理大量数据集时速度会快很多。
内容的提问来源于stack exchange,提问作者Sab
相关产品推荐
相关产品推荐

