You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.06 07:18:03