Scipy optimize调用brute函数时报参数数量不正确TypeError
简单问题/故障反馈
问题复现代码
def my_func(x,args): a = x[0] b = x[1] return a*args[0] + a*args[1] + b*args[2] import scipy.optimize as optimize params = [1,2,3] guess = [3,4] bounds = [(1,5),(2,7)] result = optimize.minimize(my_func, x0=guess, args=(params), method='SLSQP', bounds = bounds) print(result) print('done') result = optimize.brute(my_func, args = tuple(params), ranges = bounds) print(result)
问题现象
调用optimize.minimize的第一个优化任务可正常运行,但第二个调用optimize.brute的优化任务出现输入参数数量异常,对应运行输出与报错信息如下:
fun: 9.0 jac: array([3., 3.]) message: 'Optimization terminated successfully' nfev: 6 nit: 2 njev: 2 status: 0 success: True x: array([1., 2.]) done Traceback (most recent call last): File "c:\Users\bob\Desktop\optimiser.py", line 15, in <module> result = optimize.brute(my_func, args = tuple(params), ranges = bounds) File "C:\Users\bob\Anaconda3\lib\site-packages\scipy\optimize\optimize.py", line 3328, in brute Jout = np.array(list(mapper(wrapped_func, grid))) File "C:\Users\bob\Anaconda3\lib\site-packages\scipy\optimize\optimize.py", line 3400, in __call__ return self.f(np.asarray(x).flatten(), *self.args) TypeError: my_func() takes 2 positional arguments but 4 were given
报错提示向仅接收2个位置参数的my_func传入了4个参数,与函数定义冲突。
报错原因
两个scipy优化接口对args参数的处理逻辑存在差异,叠加代码中元组的不规范写法,共同导致问题:
optimize.minimize会将传入的args整体作为第二个位置参数传给目标函数。代码中写的args=(params)没有加尾逗号,Python不会将其识别为元组,实际传入的就是列表[1,2,3]本身,刚好匹配my_func(x, args)的入参要求,因此第一个调用可以正常运行。optimize.brute会将传入的args元组做解包处理,逐个作为位置参数传给目标函数。代码中写的args = tuple(params)会把[1,2,3]转为(1,2,3),brute内部调用函数时实际执行的是my_func(x, 1, 2, 3),算上x一共传入4个位置参数,和函数仅接收2个位置参数的定义冲突,因此触发报错。
解决方案
选择任意一种改法即可修复问题,推荐第一种,适配所有scipy优化接口的通用传参规范:
- 修改目标函数入参,规范
args传值
把目标函数改为接收可变长度额外参数,同时修正两处args的传值,统一传入单元素元组包裹params:# 改函数入参为*args,兼容解包传参 def my_func(x, *args): a = x[0] b = x[1] return a*args[0] + a*args[1] + b*args[2] import scipy.optimize as optimize params = [1,2,3] guess = [3,4] bounds = [(1,5),(2,7)] # 给args加尾逗号,明确传单元素元组 result = optimize.minimize(my_func, x0=guess, args=(params,), method='SLSQP', bounds = bounds) print(result) print('done') # 同样传包裹params的单元素元组,不要把params拆成多个值传 result = optimize.brute(my_func, args=(params,), ranges = bounds) print(result) - 不修改目标函数,仅修正brute的传参
如果不想改动原有函数定义,直接把brute调用时的args = tuple(params)改成args=(params,)即可,此时brute解包args只会拿到params这一个对象,刚好匹配函数的入参数量要求。
修改后两个优化调用都能正常运行,最终得到的范围内最优解为[1., 2.],对应最小函数值为9。
内容的提问来源于stack exchange,提问作者P227
相关产品推荐
相关产品推荐

