拓冰建站拓冰建站
首页 / 资讯中心 / 正文

图解async_deep_reinforce整体架构:1个全局网络+8个并行训练线程是如何协作的

图解async_deep_reinforce整体架构1个全局网络8个并行训练线程是如何协作的【免费下载链接】async_deep_reinforceAsynchronous Methods for Deep Reinforcement Learning项目地址: https://gitcode.com/gh_mirrors/as/async_deep_reinforceasync_deep_reinforce是 Google DeepMind 经典论文《Asynchronous Methods for Deep Reinforcement Learning》中A3C异步优势演员-评论家算法的 TensorFlow 实现。整个系统的架构可以概括为一句话1 个全局网络Global Network 8 个并行训练线程Worker Threads8 个线程各自在独立的 Atari Pong 游戏环境中收集经验再把梯度异步回传给全局网络从而大幅提升强化学习训练速度。下面带你完整看懂这套架构是怎么协作的。 快速上手克隆仓库并启动 A3C 训练先克隆代码仓库git clone https://gitcode.com/gh_mirrors/as/async_deep_reinforce启动异步训练只需一条命令python a3c.py训练完成后可用python a3c_display.py加载检查点观看训练好的网络实际打 Pong。依赖环境需要 TensorFlow、numpy、cv2、matplotlib并额外编译多线程版本的游戏环境 ALE。️ 架构全景Master-Worker 式协作这个项目是一个典型的Master-Worker主从架构由三类角色组成角色职责所在文件全局网络Master唯一的权威参数接收所有线程回传的梯度a3c.py 创建训练线程Worker × 8各持有一份局部网络 独立的 Pong 游戏环境a3c_training_thread.py梯度应用器用 RMSProp 把 Worker 的梯度应用到全局网络rmsprop_applier.py关键配置都集中在constants.py中两个最重要的参数PARALLEL_SIZE 8并行训练线程数量对应 8 个独立的 Pong 游戏进程LOCAL_T_MAX 20每个线程每轮最多收集 20 个时间步的经验再回传梯度。在a3c.py中程序先用thread_index -1创建全局网络FF 或 LSTM 两种结构均可切换见game_ac_network.py随后循环创建 8 个A3CTrainingThread每个线程内部再创建一份结构相同、变量独立的局部网络并挂上一个独立种子的GameStatePong 环境见game_state.py。最后通过threading.Thread把 8 个线程同时拉起。 一个 Worker 线程的完整训练循环理解 A3C 的关键是看懂单个线程每轮process()做的四件事a3c_training_thread.py同步权重线程开始新一轮前执行sync操作把全局网络的最新参数完整拷贝到自己的局部网络——保证大家基于同一份大脑出发独立收集经验用局部网络在自己的 Pong 环境里最多玩 20 步LOCAL_T_MAX记录状态、动作、奖励和价值估计计算 TD 误差对这段轨迹反向遍历用R r γ·Rγ0.99累积回报计算优势td R - V同时加入 0.01 的熵正则鼓励探索损失函数定义在game_ac_network.py的prepare_loss中回传梯度把局部网络计算出的梯度交给共享的RMSPropApplier动量为 0、梯度范数裁剪为 40直接更新全局网络的参数学习率随训练步数线性退火。 精髓所在8 个线程异步执行上述循环不需要互相等待。某个线程先完成梯度更新全局网络就立即变得聪明一点其他线程在下一轮同步时就能拿到新参数。这就是论文标题中 Asynchronous异步二字的含义。⚡ 为什么 8 个线程能显著加速8 路并行意味着游戏画面采集、前向推理、梯度计算几乎同时发生。作者在 GTX980Ti Core i7 6700 上的实测数据LOCAL_T_MAX 20网络结构GPU 训练速度CPU 训练速度A3C-FF1722 steps/sec1077 steps/secA3C-LSTM864 steps/sec540 steps/sec 训练结果Pong 分数曲线异步并行训练的效果在分数曲线上体现得很直观。A3C-LSTM 在LOCAL_T_MAX 20时全局步数约 10M 后分数就稳步上升到 20Pong 中 20 意味着每局净赢 2 个球而LOCAL_T_MAX 5经验更新鲜但每轮只走 5 步时曲线前期波动极大、整体爬升更慢说明每轮收集步数是 A3C 中一个值得权衡的超参数 核心文件导航想深入源码建议按以下顺序阅读a3c.py— 主入口创建全局网络、启动 8 个训练线程、处理中断与检查点保存a3c_training_thread.py— Worker 线程核心经验采集、TD 误差计算、梯度回传game_ac_network.py— A3C 网络结构GameACFFNetwork全连接版 /GameACLSTMNetworkLSTM 版与损失函数rmsprop_applier.py— 自定义 RMSProp 梯度应用器game_state.py— ALE Pong 游戏环境封装84×84 灰度画面、4 帧堆叠a3c_display.py— 加载检查点后可视化游戏过程constants.py— 全部超参数并行数、学习率、折扣因子等一句话总结async_deep_reinforce 用1 个全局网络 8 个并行 Worker的异步架构让数据收集和梯度更新不再串行排队是理解 A3C 异步深度强化学习原理的一个小而完整的经典实现。【免费下载链接】async_deep_reinforceAsynchronous Methods for Deep Reinforcement Learning项目地址: https://gitcode.com/gh_mirrors/as/async_deep_reinforce创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门