如何在scipy.optimize.minimize中实现损失函数早停
实现方案
你需要用到scipy.optimize.minimize的callback参数实现自定义提前终止逻辑,BFGS优化方法原生支持传入回调函数,每轮迭代结束后会自动执行该函数,你可以在回调函数中编写你的终止判断逻辑。
具体实现步骤如下:
- 定义一个存储容器,用来缓存最近4轮迭代的损失值
- 编写回调函数,每轮迭代结束后计算当前损失,按规则判断是否需要终止
- 满足终止条件时,直接抛出
StopIteration异常,scipy会自动捕获该异常并终止迭代(该特性在scipy 1.7.0及以上版本官方支持)
示例代码
import numpy as np from scipy.optimize import minimize # 存储最近4轮损失的容器 recent_losses = [] # 替换为你自己的自定义损失函数 def custom_loss(x): return np.sum(x**2) # 自定义回调函数 def callback(x): current_loss = custom_loss(x) # 不足4轮时只缓存损失不做判断 if len(recent_losses) < 4: recent_losses.append(current_loss) return # 计算变化率判断是否终止 avg_recent_loss = np.mean(recent_losses) change_rate = abs(current_loss - avg_recent_loss) / avg_recent_loss if change_rate >= 0.02: raise StopIteration(f"损失变化率达{change_rate:.2%},满足终止条件") # 更新最近4轮损失缓存 recent_losses.pop(0) recent_losses.append(current_loss) # 初始参数示例,替换为你自己的初始值 x0 = np.array([15.0, 25.0, 35.0]) # 调用minimize时传入callback参数 opt_result = minimize( fun=custom_loss, x0=x0, method="BFGS", options={"maxiter": 1000}, callback=callback ) print(opt_result)
注意事项
- 如果不想用全局变量存储损失缓存,可以把逻辑封装到类中,用类实例属性存储损失缓存,实例方法作为回调函数,代码整洁度更高
- 回调函数中计算损失的逻辑必须和传入minimize的损失函数完全一致,避免判断逻辑出现偏差
- 如果你使用的scipy版本低于1.7.0,可以通过全局标记位实现终止:设置一个初始为False的终止标记,满足条件时将标记改为True,每次损失函数执行前先判断标记,若为True直接返回任意值,迭代会自动结束
内容的提问来源于stack exchange,提问作者Oshin Patwa
相关产品推荐
相关产品推荐

