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

如何提升Stable Baselines3中该强化学习场景的训练效果?

强化学习场景优化方案

场景与问题概述

观测空间为形状(1,10)的Box类型,观测值为0、1或2,其中0和2的出现概率各为2%,1的出现概率为96%。目标是让模型学会:当观测中存在2时选择其对应索引,无2则选择索引0。当前使用PPO算法训练,但存在训练耗时久、最终效果远未达最优的问题。

现有环境与训练代码如下:

环境实现代码

import numpy as np
import gym
from gym import spaces
from stable_baselines3 import PPO, DQN, A2C
from stable_baselines3.common.env_util import make_vec_env
from stable_baselines3.common.vec_env import VecFrameStack


action_length = 10

class TestBot(gym.Env):
    def __init__(self):
        super(TestBot, self).__init__()
        self.total_rewards = 0
        self.time = 0

        self.action_space = spaces.Discrete(action_length)
        self.observation_space = spaces.Box(low=0, high=2, shape=(1, action_length), dtype=np.float32)
    
    def generate_next_obs(self):
        p = [0.02, 0.02, 0.96]
        a = [0, 2, 1]
        self.observation = np.random.choice(a, size=(1, action_length), p=p)
        if 2 in self.observation[0][1:]:
            self.best_reward += 1

    def reset(self):
        if self.time != 0:
            print('Total rewards: ', self.total_rewards, 'Best possible rewards: ', self.best_reward)

        self.best_reward = 0
        self.time = 0
        self.generate_next_obs()
        self.total_rewards = 0
        self.last_observation = self.observation
        return self.observation

    def step(self, action):
        reward = 0
        if action != 0:
            last_value = self.last_observation[0][action]
            if last_value == 2:
                reward = 1
            else:
                reward = -1
        self.time += 1
        self.generate_next_obs()
        done = self.time == 4096
        info = {}
        self.last_observation = self.observation
        self.total_rewards += reward
        return self.observation, reward, done, info

训练代码

env = TestBot()
env = make_vec_env(lambda: env, n_envs=1)
model = PPO('MlpPolicy', env, verbose=0)

iters = 0
while True:
    iters += 1
    model.learn(total_timesteps=4096, reset_num_timesteps=True)

核心问题分析

  1. 奖励信号模糊:无2时选0无正向奖励,选错误索引的惩罚力度不足,模型难以区分有效与无效行为;同时忽略了索引0出现2时的正确奖励逻辑。
  2. 观测空间冗余:用Box存储离散值,模型需额外学习区分0/1/2,增加学习成本。
  3. 训练效率低下:单环境训练样本量不足,且reset_num_timesteps=True导致每次训练都从头开始,浪费前期学习成果。
  4. 正样本稀缺:2的出现概率仅2%,模型随机探索中难以获取足够的正确行为样本,学习速度慢。

优化方案

1. 重构奖励函数,强化行为反馈

调整奖励逻辑,让模型获得清晰的对错信号:

  • 观测含2时:选对应索引奖励+2,选错误索引奖励-2
  • 观测无2时:选索引0奖励+1,选其他索引奖励-1

同时修正最优奖励统计逻辑,将索引0的2纳入计算:

def generate_next_obs(self):
    p = [0.02, 0.02, 0.96]
    a = [0, 2, 1]
    self.observation = np.random.choice(a, size=(1, action_length), p=p)
    if 2 in self.observation[0]:
        self.best_reward += 1

def step(self, action):
    current_obs = self.last_observation[0]
    has_two = 2 in current_obs
    target_idx = np.where(current_obs == 2)[0][0] if has_two else 0

    if action == target_idx:
        reward = 2 if has_two else 1
    else:
        reward = -2 if has_two else -1

    self.time += 1
    self.generate_next_obs()
    done = self.time == 4096
    info = {}
    self.last_observation = self.observation
    self.total_rewards += reward
    return self.observation, reward, done, info

2. 优化观测空间,降低学习难度

将离散的观测值转为独热编码或MultiDiscrete类型,让模型直接识别每个位置的数值类型:

# 方案:使用MultiDiscrete空间
def __init__(self):
    super(TestBot, self).__init__()
    self.total_rewards = 0
    self.time = 0

    self.action_space = spaces.Discrete(action_length)
    self.observation_space = spaces.MultiDiscrete([3]*action_length)

def generate_next_obs(self):
    p = [0.02, 0.02, 0.96]
    a = [0, 2, 1]
    self.observation = np.random.choice(a, size=(1, action_length), p=p).astype(np.int32)
    if 2 in self.observation[0]:
        self.best_reward += 1

3. 调整训练参数,提升效率

  • 移除reset_num_timesteps=True,保留训练状态,避免重复学习
  • 增加并行环境数量,提升样本收集效率
  • 调整PPO超参数,让模型更充分利用样本
# 创建4个并行环境
env = make_vec_env(TestBot, n_envs=4)
# 初始化PPO并调整超参数
model = PPO(
    'MlpPolicy',
    env,
    verbose=1,
    learning_rate=3e-4,
    batch_size=256,
    n_epochs=10,
    gamma=0.99
)
# 连续训练10万步
model.learn(total_timesteps=100_000)

4. 主动增加正样本占比,加速探索

训练初期临时提高2的出现概率,让模型快速掌握正确行为逻辑,后期恢复原概率:

def __init__(self):
    super(TestBot, self).__init__()
    self.total_rewards = 0
    self.time = 0
    self.train_phase = True  # 标记训练初期

    self.action_space = spaces.Discrete(action_length)
    self.observation_space = spaces.MultiDiscrete([3]*action_length)

def generate_next_obs(self):
    # 训练初期提高2的概率
    p = [0.05, 0.1, 0.85] if self.train_phase else [0.02, 0.02, 0.96]
    a = [0, 2, 1]
    self.observation = np.random.choice(a, size=(1, action_length), p=p).astype(np.int32)
    if 2 in self.observation[0]:
        self.best_reward += 1

# 训练流程调整
model.learn(total_timesteps=30_000)
# 切换到原概率继续训练
for env_instance in env.envs:
    env_instance.train_phase = False
model.learn(total_timesteps=70_000)

5. 简化网络结构,适配简单任务

针对当前简单的分类任务,缩小网络规模,减少训练耗时:

from stable_baselines3.common.policies import ActorCriticPolicy

class SimplePolicy(ActorCriticPolicy):
    def __init__(self, *args, **kwargs):
        super().__init__(
            *args,
            **kwargs,
            net_arch=[dict(pi=[64], vf=[64])]  # 仅一层64单元全连接层
        )

# 使用自定义策略初始化PPO
model = PPO(SimplePolicy, env, verbose=1, learning_rate=3e-4)

6. 添加终止条件,避免无效训练

通过回调函数跟踪奖励与最优值的比值,达到目标后自动停止训练:

from stable_baselines3.common.callbacks import BaseCallback

class RewardCheckCallback(BaseCallback):
    def __init__(self, threshold=0.95, verbose=0):
        super().__init__(verbose)
        self.threshold = threshold
        self.total_best = 0
        self.total_reward = 0
        self.episodes = 0

    def _on_step(self) -> bool:
        if self.locals['dones'][0]:
            self.episodes += 1
            self.total_reward += self.locals['infos'][0]['episode']['r']
            # 从reset返回的info中获取最优奖励
            self.total_best += self.locals['infos'][0]['best_possible']
            avg_ratio = (self.total_reward / self.total_best) if self.total_best !=0 else 0
            if self.verbose >0:
                print(f"Episode {self.episodes}: Avg Reward Ratio {avg_ratio:.2f}")
            if avg_ratio >= self.threshold:
                print(f"Reached target ratio {self.threshold}, stopping training")
                return False
        return True

# 修改reset方法,返回最优奖励到info
def reset(self):
    if self.time != 0:
        print('Total rewards: ', self.total_rewards, 'Best possible rewards: ', self.best_reward)
    self.best_reward = 0
    self.time = 0
    self.generate_next_obs()
    self.total_rewards = 0
    self.last_observation = self.observation
    return self.observation, {"best_possible": self.best_reward}

# 使用回调训练
callback = RewardCheckCallback(threshold=0.95, verbose=1)
model.learn(total_timesteps=200_000, callback=callback)

内容的提问来源于stack exchange,提问作者Edhowler

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 11:25:30