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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 11:05:15