You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何结合解释方差与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
################################################################

解决方案

关键问题分析

  1. 解释方差获取错误:EvalCallback(即self.parent)没有explained_variance属性,该指标属于模型本身,需从self.model获取。
  2. 早停逻辑反转:原代码中continue_training的条件是两个指标都低于阈值时继续,实际需求是当两个指标都达标时停止训练,逻辑需要反转。
  3. 回调使用错误:训练时直接传入自定义回调,而非结合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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 05:55:32