如何子类化scikit-optimize的EarlyStopper作为gp_minimize回调?
正确子类化scikit-optimize EarlyStopper类的方法
你两次实现的错误点分析:
第一种实现的问题
- 方法缩进错误:
_criterion被定义为全局函数,而非StoppingCriterion类的成员方法 else语句缩进层级不匹配,与if逻辑脱节- 未导入
numpy却使用np.sort - Python3中父类构造器规范写法应为
super().__init__()(非直接报错原因,但不符合编码规范)
第二种实现的问题
EarlyStopper是抽象基类,它的__init__方法仅接收self一个参数,不存在patience、relative_improvement这类入参,因此你调用super().__init__(patience=patience, relative_improvement=relative_improvement)会触发参数不匹配的TypeError。
正确的子类化思路与实现
EarlyStopper的核心设计是:子类只需实现_criterion方法,该方法接收result对象作为参数,返回True(触发停止)、False(不停止)或None(继续优化)。父类的__call__方法会自动调用_criterion完成停止判断。
示例1:基于最优值差异的停止器
import numpy as np from skopt.callbacks import EarlyStopper class DiffEarlyStopper(EarlyStopper): def __init__(self, delta=0.05, n_best=10): super().__init__() self.delta = delta self.n_best = n_best def _criterion(self, result): if len(result.func_vals) >= self.n_best: sorted_vals = np.sort(result.func_vals)[:self.n_best] best = sorted_vals[0] worst_top_n = sorted_vals[-1] relative_diff = abs((best - worst_top_n) / worst_top_n) return relative_diff < self.delta return None
示例2:基于patience的停止器
from skopt.callbacks import EarlyStopper class PatienceEarlyStopper(EarlyStopper): def __init__(self, patience=15, relative_improvement=0.01): super().__init__() self.patience = patience self.relative_improvement = relative_improvement self.best_loss = float("inf") self.counter = 0 def _criterion(self, result): current_loss = result.fun if current_loss < self.best_loss * (1 - self.relative_improvement): self.best_loss = current_loss self.counter = 0 return False else: self.counter += 1 return self.counter >= self.patience
使用方式
将自定义停止器传入gp_minimize的callback参数即可:
from skopt import gp_minimize from skopt.space import Real def objective(x): return x[0]**2 + x[1]**2 space = [Real(-5, 5, name='x'), Real(-5, 5, name='y')] stopper = PatienceEarlyStopper(patience=10, relative_improvement=0.001) result = gp_minimize(objective, space, n_calls=100, callback=[stopper])
内容的提问来源于stack exchange,提问作者Henri
相关产品推荐
相关产品推荐

