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

使用scipy.optimize.curve_fit时,非线性拟合函数的最优形式探究

非线性方程拟合形式对比:y = a + b * x ** c及其等价变换

在拟合形如y = a + b * x ** c的非线性方程时,函数形式会显著影响拟合效果,以下对比两种等价形式的表现:

形式1

  • 表达式:y = a + b * x ** c
  • 拟合时需要增大默认的最大函数评估次数(maxfev)才能收敛
  • 数据缩放会改变拟合结果

形式2(Box-Cox变换形式)

  • 表达式:y = d + e * (x ** f - 1) / f
  • 与形式1完全等价,参数满足转换关系:a = d - e/f、b = e/f、c = f
  • 使用默认maxfev即可完成拟合收敛
  • 拟合效果更优
  • 数据缩放对拟合结果无影响

注:当将数据缩放至标准差为1后,两种形式的拟合结果会完全一致。

测试数据与代码

from io import StringIO

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from scipy.optimize import curve_fit
from sklearn.metrics import mean_absolute_error
from sklearn.preprocessing import StandardScaler


def formulation_1(x, a, b, c):
    return a + b * x**c


def formulation2(x, d, e, f):
    return d + e * (x**f - 1) / f


TESTDATA = StringIO(
    '''x,y
410.0,73.06578085929756
25.0,205.29417389522575
72.0,110.48653325137172
51.0,168.52111516008628
15.0,119.75684720989004
72.0,164.46280991735537
73.0,145.53126751391522
36.0,161.41429319490925
40.0,219.91735537190084
190.0,89.18717804003897
21.0,203.3350571291969
47.0,170.12964176670877
21.0,197.1020714822368
18.0,96.43526170798899
53.0,117.55060034305319
50.0,189.89358650365415
43.0,179.03132807995385
863.0,69.63888057656149
71.0,131.42730764753813
40.0,205.2892561983471
65.0,131.3857292219426
401.0,133.81511047189076
50.0,115.65603442387814
58.0,151.99074870050802
50.0,165.8640803223824
21.0,210.87942861045792
236.0,124.21734182739671
53.0,180.11451429366744
12.0,320.77043917765184
36.0,244.3526170798898
25.0,202.41568198893515
21.0,184.03895128162597
29.0,165.64724945771087
25.0,218.1818181818182
72.0,161.8457300275482
130.0,107.38232466256511
84.0,177.52397865095088
38.0,57.524112172378224
50.0,168.132777815723
25.0,202.41568198893515
21.0,244.3978260657449
48.0,168.3167133392528
200.0,122.8403797554812
37.0,167.84185838731295
83.0,173.75445583988846
13.0,315.835929660122
11.0,314.47181327653976
32.0,203.68741889215215
200.0,123.96694214876034
39.0,110.59353869271226
39.0,190.81504521023686
40.0,235.53719008264466
37.0,181.71111484409758
25.0,215.55576804863057
40.0,235.53719008264466'''
)
df = pd.read_csv(TESTDATA)
sc = StandardScaler(with_mean=False)
df["x_scaled"] = sc.fit_transform(df[["x"]])

fig, ax = plt.subplots(1, 2, figsize=(10, 6), sharey=True)
for ax_i, scale in enumerate([False, True]):
    feature_name = "x"
    if scale:
        feature_name = "x_scaled"

    ax[ax_i].scatter(df[feature_name], df["y"], label="Data", alpha=0.5)

    for formulation_i, (func, linestyle) in enumerate(
        zip([formulation_1, formulation2], ["--", ":"])
    ):

        params, covariance = curve_fit(
            func,
            df[feature_name],
            df["y"],
            maxfev=10_000,
        )
        df["fit"] = func(df[feature_name], *params)
        mae = mean_absolute_error(df["y"], df["fit"])

        x = np.linspace(df[feature_name].min(), df[feature_name].max(), 100)
        fit = func(x, *params)

        ax[ax_i].plot(
            x,
            fit,
            linestyle=linestyle,
            label=f"Form {formulation_i + 1} (MAE: {mae:.2f})",
        )
    ax[ax_i].legend()
    ax[ax_i].set_title(f"Scaled: {scale}")
    ax[ax_i].set_xlabel(feature_name)
    ax[ax_i].set_ylabel("y")
plt.show()

内容的提问来源于stack exchange,提问作者3UqU57GnaX

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 15:18:09