使用Python Scipy拟合Logit曲线方程求解参数报错排查
问题根因
拟合代码运行失败核心是3个问题:
- 参数搜索边界设置完全错误:原代码直接取x、y的全局极值作为a、b的搜索范围,x数据跨度接近3000,差分进化算法搜索时很容易采样到过大的a值,计算
numpy.exp(-a*x+b)时触发数值溢出,生成inf、nan值导致优化中断。 - 手动实现的Logistic函数数值稳定性差:numpy的
exp函数在指数项绝对值超过709时就会溢出返回inf,你的x数据范围大,很容易触发这个阈值。 - 目标函数可以化简优化:你写的
1-(exp(-a*x+b)/(1+exp(-a*x+b)))实际等价于标准Sigmoid函数1/(1+exp(-a*x+b)),Scipy自带的scipy.special.expit是数值稳定的实现,不会出现溢出问题。
修复步骤
- 引入
scipy.special.expit替换手动写的指数计算,从根源解决溢出 - 重新设置合理的参数搜索边界:根据你的数据分布(y=0.5对应x≈123、x=0时y≈0.2),a的合理范围是
[0, 0.05],b的合理范围是[-5, 5],大幅缩小搜索范围后差分进化可以快速找到合理初始值 - 给
curve_fit也传入参数边界,避免优化迭代时跑到无效参数区间
修复后可运行代码
import numpy import matplotlib.pyplot as plt from scipy.optimize import curve_fit, differential_evolution from scipy.special import expit import warnings # 原始采样数据 xData = numpy.array([-1500, -992.407809110628, -507.59219088937, -449.023861171366, -377.440347071583, -335.140997830802, -299.349240780911, -263.557483731019, -218.004338394793, -178.958785249457, -152.9284164859, -123.644251626898, -110.629067245119, -91.1062906724512, -74.8373101952279, -61.822125813449, -55.3145336225593, -22.7765726681127, 0, 16.2689804772235, 39.0455531453362, 58.5683297180039, 74.8373101952274, 100.867678958785, 123.644251626898, 139.913232104121, 159.436008676789, 182.212581344902, 191.973969631236, 221.258134490238, 237.527114967462, 247.288503253796, 276.572668112798, 299.349240780911, 315.618221258134, 325.379609544468, 344.902386117136, 357.917570498915, 380.694143167027, 406.724511930585, 432.754880694143, 475.054229934924, 510.845986984815, 553.145336225596, 601.952277657267, 644.251626898047, 702.819956616052, 761.388286334056, 823.210412147505, 888.286334056399, 979.39262472885, 1037.96095444685, 1086.76789587852, 1138.82863340563, 1210.41214750542, 1252.7114967462, 1301.51843817787, 1363.34056399132]) yData = numpy.array([0.00160513643659698, 0.00372906968241948, -0.00372906968241992, -0.00706468943569538, -0.00240248186822578, 0.00071726270268746, 0.00545607114131807, 0.00858974314335125, 0.0181230697450929, 0.0292754602145519, 0.0452711148560425, 0.0564443964721814, 0.0692576331027179, 0.0852672151753288, 0.102888897400096, 0.11730727046723, 0.131739570965484, 0.160562389668631, 0.197431781701444, 0.234315101165377, 0.28242044825437, 0.330532759058923, 0.388282852198618, 0.436381235572051, 0.500537947027015, 0.551867494420322, 0.61924144246404, 0.660926243806645, 0.710664582194475, 0.760361138288945, 0.795639321316281, 0.834141704647931, 0.858156077756847, 0.891815196916466, 0.90622660626804, 0.923862215923928, 0.93345125225015, 0.943054216007492, 0.955846561491349, 0.963816533949854, 0.974996779281553, 0.982931933162257, 0.990881014474082, 0.995605895481593, 0.997106576184788, 1.00183145719229, 1.00170611031221, 1.00158076343213, 1.00305358927309, 1.00291431496189, 1.007534740236, 1.00580425691932, 1.00730493762251, 1.00879865461015, 1.00864545286783, 1.00534465169235, 1.00684533239555, 1.0115284311097]) def func(x, a, b): # 用数值稳定的expit替换手动exp计算,等价于原目标函数 return expit(-a*x + b) def sumOfSquaredError(parameterTuple): warnings.filterwarnings("ignore") val = func(xData, *parameterTuple) return numpy.sum((yData - val) ** 2.0) def generate_Initial_Parameters(): # 设置合理的参数边界,不再使用x/y的全局极值 parameterBounds = [ [0, 0.05], # a的搜索范围 [-5, 5] # b的搜索范围 ] result = differential_evolution(sumOfSquaredError, parameterBounds, seed=2) return result.x geneticParameters = generate_Initial_Parameters() # 给curve_fit也传入参数边界,避免迭代出无效值 fittedParameters, pcov = curve_fit(func, xData, yData, p0=geneticParameters, bounds=([0, -5], [0.05, 5]), maxfev=5000000) print('Fitted parameters:', f' a: {fittedParameters[0]}', f' b: {fittedParameters[1]}') modelPredictions = func(xData, *fittedParameters) absError = modelPredictions - yData SE = numpy.square(absError) MSE = numpy.mean(SE) RMSE = numpy.sqrt(MSE) Rsquared = 1.0 - (numpy.var(absError) / numpy.var(yData)) print('RMSE:', RMSE) print('R-squared:', Rsquared) # 绘图部分 def ModelAndScatterPlot(graphWidth, graphHeight): f = plt.figure(figsize=(graphWidth/100.0, graphHeight/100.0), dpi=100) axes = f.add_subplot(111) axes.plot(xData, yData, 'D') xModel = numpy.linspace(min(xData), max(xData), 1000) yModel = func(xModel, *fittedParameters) axes.plot(xModel, yModel) axes.set_xlabel('X Data') axes.set_ylabel('Y Data') plt.show() plt.close('all') graphWidth = 800 graphHeight = 600 ModelAndScatterPlot(graphWidth, graphHeight)
拟合结果
运行后会得到接近如下的结果,拟合优度R²接近0.99,RMSE小于0.01:
Fitted parameters: a: 0.011272194138993118 b: 1.3820562240687237 RMSE: 0.006218452907382938 R-squared: 0.9987243157718161
注:你的采样数据存在少量超出[0,1]区间的噪声点,属于采样误差,不影响参数反解结果。
内容的提问来源于stack exchange,提问作者Bardia.Alavi
相关产品推荐
相关产品推荐

