使用pytorch-minimize/scipy时回调函数报错:需1参数却传2个
问题原因与解决方案
错误my_callback() takes 1 positional argument but 2 were given的核心原因是:当使用trust-constr优化方法时(pytorch-minimize的minimize_constr底层依赖该方法,scipy示例中显式指定了该方法),回调函数的参数签名要求与其他优化方法不同。
普通优化方法(如默认的L-BFGS-B)的回调函数仅需接收当前迭代的参数数组作为唯一参数,但trust-constr方法会向回调函数传入两个参数:
- 当前迭代的参数数组
x - 包含迭代状态信息的字典
state
解决方案:修改回调函数的参数签名
只需让回调函数接受第二个参数即可,即使你不需要使用这个参数也可以忽略它。
修正后的pytorch-minimize代码:
from torchmin import minimize_constr import torch import torch.nn as nn eps0 = torch.rand((2,3)) # 增加state参数,即使不用也保留 def my_callback(xi, state): print("hello world!") res = minimize_constr( lambda x : nn.L1Loss()(x.sum(), x.sum()), eps0, max_iter=100, callback = my_callback, disp=1 ) eps = res.x
修正后的scipy代码:
from scipy.optimize import minimize import numpy as np # 增加state参数 def my_callback(xi, state): print("hello world!") x0_np = np.random.rand(3) result = minimize( lambda x : x.sum(), x0_np, method='trust-constr', callback = my_callback )
补充说明
- 移除
callback参数后代码正常:因为没有触发回调函数的调用,自然不会出现参数不匹配的问题。 - scipy示例中移除
method参数后正常:此时scipy默认使用L-BFGS-B方法,该方法的回调函数仅需一个参数,与你最初定义的my_callback签名匹配。
内容的提问来源于stack exchange,提问作者Dudi Frid
相关产品推荐
相关产品推荐

