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

Scipy optimize.curve_fit无法正常执行线性拟合问题求助

解决线性拟合问题的可靠方案

问题根源

curve_fit 本质是面向非线性模型的最小二乘优化工具,对于 y = m*x + c 这类线性模型,它依赖初始参数猜测完成迭代。当自变量 x 取值跨度大(如本例中10到70)时,m*x 和 c 的量级差异悬殊,容易导致优化收敛到局部最优,无法得到预期结果。

方案1:解析法加权线性最小二乘(推荐)

对于线性模型,可直接通过矩阵运算求解最优参数,无需初始猜测,且能直接计算协方差矩阵,完全适配批量处理需求。

实现代码

import numpy as np
import matplotlib.pyplot as plt

# 原始数据
x = np.array([10.55977951, 11.45686089, 12.35423473, 13.30408861, 14.3528772,  15.45035217,
     16.64817782, 17.99661252, 19.49184541, 21.18832373, 23.08427703, 25.27645953,
     27.91555605, 31.14343828, 35.16744841, 40.31943045, 47.22644384, 57.23413423, 73.51606363])
y = np.array([ 4.50873148e-04, 2.29554869e-04, 4.62602769e-04, 1.57525181e-04,
     1.16334543e-04, 9.42105291e-05, 2.86606379e-04, 2.40194287e-04,
     4.74193673e-04, 6.83848270e-04, 6.73286482e-04, 2.03506304e-04,
     3.58126867e-04, 1.88155439e-04, 5.14133854e-04, 2.39990293e-04,
     -3.60391884e-05, -1.17866329e-04, 2.68954649e-05])
y_err = np.array([8.23397676e-05, 7.54222285e-05, 7.05355053e-05, 6.30493368e-05,
         5.73241555e-05, 5.56298862e-05, 5.00181328e-05, 4.76554758e-05,
         4.45081313e-05, 4.23716792e-05, 4.10516842e-05, 3.87066834e-05,
         3.67639901e-05, 3.51162489e-05, 3.39993704e-05, 3.29275562e-05,
         3.24743967e-05, 3.15296789e-05, 3.09144126e-05])

# 构造加权矩阵和设计矩阵
weights = 1 / y_err**2
X = np.vstack([x, np.ones_like(x)]).T
W = np.diag(weights)

# 计算最优参数
XTWX = X.T @ W @ X
XTWy = X.T @ W @ y
popt = np.linalg.inv(XTWX) @ XTWy
m, c = popt

# 计算协方差矩阵(若不信任y_err,可乘以残差方差)
cov = np.linalg.inv(XTWX)
# 可选:考虑残差方差的协方差计算
# residuals = y - (m*x + c)
# residual_variance = np.sum(weights * residuals**2) / (len(x) - 2)
# cov = residual_variance * np.linalg.inv(XTWX)

print(f"m = {m:.6e}, c = {c:.6e}")
print("协方差矩阵:")
print(cov)

# 绘图展示
line1 = m*x + c
plt.plot(x, line1, color="red", label="拟合线")
plt.plot(x, [0]*len(x), color="black")
plt.errorbar(x, y, yerr=y_err, label="数据", fmt="s", markersize=5, color="red")
plt.xscale("log")
plt.xlabel("dummy x ")
plt.ylabel("dummy y")
plt.legend()
plt.tight_layout()
plt.show()

方案2:优化curve_fit的拟合效果

若坚持使用curve_fit,可通过归一化自变量消除参数量级差异,避免收敛问题,无需手动指定p0。

实现代码

from scipy.optimize import curve_fit
import numpy as np
import matplotlib.pyplot as plt

def line(x, m, c):
    return x*m + c

# 原始数据(同上,省略重复定义)
x = np.array([...])
y = np.array([...])
y_err = np.array([...])

# 归一化x(除以均值)
x_mean = x.mean()
x_norm = x / x_mean

# 拟合归一化后的模型
popt_norm, cov_norm = curve_fit(line, x_norm, y, sigma=y_err)
m_norm, c = popt_norm
# 转换回原始尺度的m
m = m_norm / x_mean

# 转换协方差矩阵到原始尺度
cov = np.copy(cov_norm)
cov[0,0] /= x_mean**2
cov[0,1] /= x_mean
cov[1,0] /= x_mean

print(f"m = {m:.6e}, c = {c:.6e}")
print("协方差矩阵:")
print(cov)

# 绘图展示
line1 = m*x + c
plt.plot(x, line1, color="red", label="拟合线")
plt.plot(x, [0]*len(x), color="black")
plt.errorbar(x, y, yerr=y_err, label="数据", fmt="s", markersize=5, color="red")
plt.xscale("log")
plt.xlabel("dummy x ")
plt.ylabel("dummy y")
plt.legend()
plt.tight_layout()
plt.show()

方案3:使用scikit-learn的LinearRegression(带权重)

scikit-learn的LinearRegression支持样本权重,需手动计算协方差矩阵:

实现代码

from sklearn.linear_model import LinearRegression
import numpy as np
import matplotlib.pyplot as plt

# 原始数据
x = np.array([...]).reshape(-1,1)
y = np.array([...])
y_err = np.array([...])

# 权重为1/方差
weights = 1 / y_err**2

# 拟合模型
model = LinearRegression()
model.fit(x, y, sample_weight=weights)
m = model.coef_[0]
c = model.intercept_

# 手动计算协方差矩阵
X = np.hstack([x, np.ones_like(x)])
W = np.diag(weights)
XTWX = X.T @ W @ X
cov = np.linalg.inv(XTWX)
# 可选:乘以残差方差
residuals = y - model.predict(x)
residual_variance = np.sum(weights * residuals**2) / (len(x) - 2)
cov = residual_variance * cov

print(f"m = {m:.6e}, c = {c:.6e}")
print("协方差矩阵:")
print(cov)

# 绘图展示
line1 = m*x.flatten() + c
plt.plot(x, line1, color="red", label="拟合线")
plt.plot(x, [0]*len(x), color="black")
plt.errorbar(x.flatten(), y, yerr=y_err, label="数据", fmt="s", markersize=5, color="red")
plt.xscale("log")
plt.xlabel("dummy x ")
plt.ylabel("dummy y")
plt.legend()
plt.tight_layout()
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 16:24:59