使用basinhopping优化Statsmodels MNLogit模型时触发报错
问题:MNLogit使用basinhopping拟合触发TypeError报错
复现代码
import numpy as np import statsmodels.api as sm x = np.random.randint(0, 100, 1000) y = np.random.randint(0, 3, 1000) model = sm.MNLogit(y, sm.add_constant(x)) results = model.fit(method='basinhopping') print(results.summary())
报错信息
Traceback (most recent call last): File "/Users/wagnerpf134/opt/anaconda3/lib/python3.9/site-packages/spyder_kernels/py3compat.py", line 356, in compat_exec exec(code, globals, locals) File "/Users/wagnerpf134/Documents/untitled0.py", line 7, in <module> results = model.fit(method='basinhopping') File "/Users/wagnerpf134/opt/anaconda3/lib/python3.9/site-packages/statsmodels/discrete/discrete_model.py", line 654, in fit mnfit = base.LikelihoodModel.fit(self, start_params = start_params, File "/Users/wagnerpf134/opt/anaconda3/lib/python3.9/site-packages/statsmodels/base/model.py", line 563, in fit xopt, retvals, optim_settings = optimizer._fit(f, score, start_params, File "/Users/wagnerpf134/opt/anaconda3/lib/python3.9/site-packages/statsmodels/base/optimizer.py", line 241, in _fit xopt, retvals = func(objective, gradient, start_params, fargs, kwargs, File "/Users/wagnerpf134/opt/anaconda3/lib/python3.9/site-packages/statsmodels/base/optimizer.py", line 1040, in _fit_basinhopping retvals = optimize.basinhopping(f, start_params, File "/Users/wagnerpf134/opt/anaconda3/lib/python3.9/site-packages/scipy/optimize/_basinhopping.py", line 728, in basinhopping callback(bh.storage.minres.x, bh.storage.minres.fun, True) TypeError: <lambda>() takes 1 positional argument but 3 were given
问题背景
使用其他优化方法无报错,basinhopping用于二元Logit或OrderedModel(distr='logit')时正常,仅MNLogit触发该错误。
原因分析
Statsmodels的_fit_basinhopping函数为MNLogit生成的默认lambda回调仅接受1个参数,但scipy的basinhopping会传递3个参数(最优参数、函数值、是否为新最小值),参数数量不匹配导致报错。而二元Logit等模型的回调逻辑适配了这个参数数量,因此无问题。
解决方法
方法1:自定义兼容的回调函数(推荐)
调用fit时,通过optimizer_kwargs传递接受3个参数的回调函数,覆盖默认的lambda:
import numpy as np import statsmodels.api as sm def basinhopping_callback(x, f, accept): # 可根据需求添加逻辑,比如打印进度,若不需要则留空即可 pass x = np.random.randint(0, 100, 1000) y = np.random.randint(0, 3, 1000) model = sm.MNLogit(y, sm.add_constant(x)) results = model.fit(method='basinhopping', optimizer_kwargs={'callback': basinhopping_callback}) print(results.summary())
方法2:修改Statsmodels源码(临时应急方案)
找到statsmodels/base/optimizer.py中的_fit_basinhopping函数,将默认回调的lambda改为接受3个参数:
原代码:
callback = kwargs.pop('callback', lambda x: None)
修改为:
callback = kwargs.pop('callback', lambda x, f, accept: None)
注意:修改源码会影响全局环境,Statsmodels升级后会失效,仅作为临时修复。
内容的提问来源于stack exchange,提问作者wagnerpf134
相关产品推荐
相关产品推荐

