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

使用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是数值稳定的实现,不会出现溢出问题。
修复步骤
  1. 引入scipy.special.expit替换手动写的指数计算,从根源解决溢出
  2. 重新设置合理的参数搜索边界:根据你的数据分布(y=0.5对应x≈123、x=0时y≈0.2),a的合理范围是[0, 0.05],b的合理范围是[-5, 5],大幅缩小搜索范围后差分进化可以快速找到合理初始值
  3. 给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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 14:39:54