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

基于Stable-Baselines3与Gym实现调用env.step时训练智能体的方法问询

手动交互同时训练智能体的实现方案

核心思路

手动调用env.step()与环境交互时,收集每一步的经验数据(观测、动作、奖励、下一观测、结束标记),将数据送入智能体的经验缓冲区(针对off-policy算法)或临时存储(针对on-policy算法),之后定期触发智能体的训练更新逻辑,以此实现智能体和旧版本自身对弈的自玩功能。

具体实现步骤

1. 初始化环境与智能体

以回合制场景适配性较好的DQN(off-policy算法)为例:

import gym
from stable_baselines3 import DQN

# 初始化自定义回合制环境(比如国际象棋类环境)
env = gym.make("YourCustomChessEnv-v0")
# 初始化待训练的智能体
agent = DQN("MlpPolicy", env, verbose=1)
# 保存初始版本智能体作为初始对手
agent.save("initial_agent")
old_agent = DQN.load("initial_agent")

2. 手动交互+经验收集+训练循环

total_timesteps = 10000
update_interval = 64  # 每收集64步经验触发一次训练

obs = env.reset()
for step in range(total_timesteps):
    # 回合制场景:当前智能体与旧版本对手交替行动
    if env.current_player == "train_agent":
        action, _ = agent.predict(obs, deterministic=False)
    else:
        action, _ = old_agent.predict(obs, deterministic=False)
    
    # 手动执行环境交互
    next_obs, reward, done, info = env.step(action)
    
    # 将当前步经验存入智能体的回放缓冲区
    agent.replay_buffer.add(obs, next_obs, action, reward, done, info)
    
    # 达到更新间隔且缓冲区数据足够时,触发训练
    if step % update_interval == 0 and agent.replay_buffer.size() > agent.batch_size:
        agent.train(batch_size=agent.batch_size)
    
    # 更新观测值
    obs = next_obs
    
    # 回合结束处理
    if done:
        obs = env.reset()
        # 每1000步更新一次对手(用当前训练的智能体替换旧对手)
        if step % 1000 == 0:
            agent.save(f"current_agent_{step}")
            old_agent = DQN.load(f"current_agent_{step}")

# 保存最终训练完成的自玩智能体
agent.save("final_self_play_agent")

3. On-Policy算法(如PPO)的适配调整

PPO这类on-policy算法需要收集一批完整轨迹后再训练,需临时存储轨迹数据:

from stable_baselines3 import PPO

agent = PPO("MlpPolicy", env, verbose=1)
agent.save("initial_ppo_agent")
old_agent = PPO.load("initial_ppo_agent")

trajectories = []
batch_size = 2048  # PPO的批次数据量

obs = env.reset()
for step in range(total_timesteps):
    # 交替行动逻辑同前
    action, _ = agent.predict(obs) if env.current_player == "train_agent" else old_agent.predict(obs)
    
    next_obs, reward, done, info = env.step(action)
    # 存储当前轨迹片段
    trajectories.append((obs, action, reward, next_obs, done))
    
    obs = next_obs
    
    if done:
        obs = env.reset()
    
    # 收集够批次数据后触发训练
    if len(trajectories) >= batch_size:
        # 调用learn方法训练,reset_num_timesteps=False避免重置步数计数
        agent.learn(total_timesteps=batch_size, reset_num_timesteps=False)
        trajectories = []
        
        # 定期更新对手
        if step % 1000 == 0:
            agent.save(f"ppo_current_agent_{step}")
            old_agent = PPO.load(f"ppo_current_agent_{step}")

关键注意事项

  • 回合制玩家切换:自定义环境需添加current_player属性,明确当前行动的是训练智能体还是旧版本对手。
  • 经验缓冲区管理:Off-Policy算法的回放缓冲区有容量限制,初始化智能体时可通过buffer_size参数调整。
  • 训练频率控制:避免每步都训练,根据算法特性设置合理的更新间隔,平衡训练效率与资源消耗。
  • 对手更新策略:定期用当前智能体替换旧对手,可设置固定步数间隔或根据胜率动态调整,平衡训练稳定性和迭代速度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 10:20:43