如何结合解释方差与StopTrainingOnRewardThreshold实现RL模型早停?
问题:结合解释方差与奖励阈值实现RL训练早停
我正在用TensorFlow和StableBaselines3构建外汇交易强化学习机器人,希望同时基于解释方差阈值和奖励阈值实现训练早停,但自定义回调无法正常工作,也无法在训练过程中正确获取模型的解释方差。
训练代码
"""Train Model""" ################################################################ if user_action == 2: env_maker = lambda: gym.make('forex-v0', df=df, frame_bound=(15, 250), window_size=5) env = DummyVecEnv([env_maker]) model = A2C('MlpPolicy', env, verbose=1) callback_on_best = CustomCallback(explained_variance=0.7, reward_threshold=300, verbose=1) eval_callback = EvalCallback(env, callback_on_new_best=callback_on_best, verbose=1) model.learn(total_timesteps=1000000, callback=callback_on_best) model.save("A2C_trading_Ai") while True: user_action_after_train = int(input("\n===============================================\n" "Training Complete! Evaluate model now?\n" "\t1 = Yes\n" "\t2 = No\n\n" "Response = \t\t")) if user_action_after_train != 1 and user_action != 2: print("Invalid Input!\n") elif user_action_after_train == 1: user_action = 3 print("===============================================\n") break else: break ################################################################
失效的自定义回调代码
"""Custom Callback for Model Training""" ################################################################ class CustomCallback(BaseCallback): def __init__(self, explained_variance: float, reward_threshold: float, verbose: int = 0): super().__init__(verbose=verbose) self.explained_variance = explained_variance self.reward_threshold = reward_threshold def _on_step(self) -> bool: assert self.parent is not None, \ "``StopTrainingOnMinimumReward`` callback must be used " "with an ``EvalCallback``" # Convert np.bool_ to bool, otherwise callback() is False won't work continue_training = bool(self.parent.explained_variance < self.explained_variance and self.parent.best_mean_reward < self.reward_threshold) if self.verbose >= 1 and not continue_training: print( f"Stopping training because the mean explained variance {self.parent.explained_variance:.2f} " f"and the mean reward {self.parent.best_mean_reward:.2f}" f" are above the thresholds {self.explained_variance} and {self.best_mean_reward}" ) return continue_training ################################################################
解决方案
关键问题分析
- 解释方差获取错误:
EvalCallback(即self.parent)没有explained_variance属性,该指标属于模型本身,需从self.model获取。 - 早停逻辑反转:原代码中
continue_training的条件是两个指标都低于阈值时继续,实际需求是当两个指标都达标时停止训练,逻辑需要反转。 - 回调使用错误:训练时直接传入自定义回调,而非结合
EvalCallback传递,导致无法正确获取评估后的奖励值。
修正后的自定义回调
from stable_baselines3.common.callbacks import BaseCallback class CustomEarlyStopCallback(BaseCallback): def __init__(self, explained_variance_threshold: float, reward_threshold: float, verbose: int = 0): super().__init__(verbose=verbose) self.ev_threshold = explained_variance_threshold self.reward_threshold = reward_threshold def _on_step(self) -> bool: # 从模型获取当前解释方差 current_ev = self.model.explained_variance(self.model.get_env()) # 获取评估后的最佳平均奖励(依赖EvalCallback) if hasattr(self.parent, 'best_mean_reward'): current_reward = self.parent.best_mean_reward else: # 若未使用EvalCallback,可改用训练环境的实时奖励统计(示例) current_reward = self.model.get_env().get_attr('reward')[0] # 早停条件:解释方差达标 且 奖励达标时停止训练 stop_training = (current_ev >= self.ev_threshold) and (current_reward >= self.reward_threshold) continue_training = not stop_training if self.verbose >= 1 and stop_training: print( f"训练停止:解释方差 {current_ev:.2f} ≥ 阈值 {self.ev_threshold}," f"平均奖励 {current_reward:.2f} ≥ 阈值 {self.reward_threshold}" ) return continue_training
修正后的训练代码
"""Train Model""" ################################################################ if user_action == 2: env_maker = lambda: gym.make('forex-v0', df=df, frame_bound=(15, 250), window_size=5) env = DummyVecEnv([env_maker]) model = A2C('MlpPolicy', env, verbose=1) # 初始化自定义早停回调 early_stop_callback = CustomEarlyStopCallback( explained_variance_threshold=0.7, reward_threshold=300, verbose=1 ) # 结合EvalCallback,每10000步评估一次模型 eval_callback = EvalCallback( env, callback_on_new_best=early_stop_callback, verbose=1, eval_freq=10000 ) # 将EvalCallback传入训练流程 model.learn(total_timesteps=1000000, callback=eval_callback) model.save("A2C_trading_Ai") # 用户交互逻辑保持不变 while True: user_action_after_train = int(input("\n===============================================\n" "Training Complete! Evaluate model now?\n" "\t1 = Yes\n" "\t2 = No\n\n" "Response = \t\t")) if user_action_after_train not in [1, 2]: print("Invalid Input!\n") elif user_action_after_train == 1: user_action = 3 print("===============================================\n") break else: break ################################################################
内容的提问来源于stack exchange,提问作者Tian van Wyk
相关产品推荐
相关产品推荐

