基于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
相关产品推荐
相关产品推荐

