如何编写回调函数在explained variance处于0.9-1范围时终止强化学习训练
问题需求
使用Stable Baselines3与gym_anytrading构建强化学习模型时,希望实现当explained variance处于0.9~1范围时自动终止训练,避免模型在指标不理想时过早停止。
改进后的回调实现
无需自定义Logger,直接利用SB3内置的日志机制即可实现需求,以下是完善后的代码:
from stable_baselines3.common.callbacks import BaseCallback class CustomStopCallback(BaseCallback): def __init__(self, min_explained_variance: float = 0.9, max_explained_variance: float = 1.0, starting_step: int = 10000): super().__init__() self.min_ev = min_explained_variance self.max_ev = max_explained_variance self.starting_step = starting_step def _on_step(self) -> bool: # 仅在训练步数超过起始阈值后检查 if self.num_timesteps < self.starting_step: return True # 从模型日志中获取最新的explained variance logger = self.model.logger if "explained_variance" in logger.name_to_value: explained_variance = logger.name_to_value["explained_variance"] # 检查是否满足终止条件 if self.min_ev <= explained_variance <= self.max_ev: print(f"训练终止:步数{self.num_timesteps},explained variance={explained_variance:.4f}") return False # 返回False会触发训练终止 # 未满足条件则继续训练 return True # 实例化回调 custom_stop_callback = CustomStopCallback(min_explained_variance=0.9, starting_step=10000) # 原环境与模型代码保持不变 env_maker = lambda: MyCustomEnv(df=df, frame_bound=(12,30660), window_size=12) env = DummyVecEnv([env_maker]) model = A2C('MlpPolicy', env, verbose=1, policy_kwargs=dict(net_arch=[dict(pi=[128, 256, 128], vf=[128, 256, 128])])) # 传入回调开始训练 model.learn(total_timesteps=1000000, callback=custom_stop_callback)
关键改进说明
- 移除自定义Logger:SB3的
self.model.logger.name_to_value会自动存储最新的日志指标,无需手动维护日志缓冲区,简化代码结构。 - 标准终止方式:在回调的
_on_step方法中返回False是SB3官方推荐的训练终止方式,比直接修改模型属性更可靠。 - 范围判断:明确实现0.9~1的区间检查,符合需求。
- 容错处理:先判断指标是否存在,避免因日志未记录该指标导致的
KeyError。 - 起始步数过滤:设置
starting_step确保训练初期(指标不稳定阶段)不会误触发终止。
内容的提问来源于stack exchange,提问作者ET4
相关产品推荐
相关产品推荐

