[论文解读] Podracer architectures for scalable Reinforcement Learning
本文提出了Podracer架构——Anakin与Sebulba——旨在利用JAX在TPU Pods上高效扩展强化学习训练。通过利用大规模并行计算、大批次训练以及优化的TPU资源利用,作者在完整TPU Pod上实现了高达每秒4300万帧的训练速度,训练智能体耗时不足一分钟,相比先前系统显著降低了成本并提升了数据效率。
Supporting state-of-the-art AI research requires balancing rapid prototyping, ease of use, and quick iteration, with the ability to deploy experiments at a scale traditionally associated with production systems.Deep learning frameworks such as TensorFlow, PyTorch and JAX allow users to transparently make use of accelerators, such as TPUs and GPUs, to offload the more computationally intensive parts of training and inference in modern deep learning systems. Popular training pipelines that use these frameworks for deep learning typically focus on (un-)supervised learning. How to best train reinforcement learning (RL) agents at scale is still an active research area. In this report we argue that TPUs are particularly well suited for training RL agents in a scalable, efficient and reproducible way. Specifically we describe two architectures designed to make the best use of the resources available on a TPU Pod (a special configuration in a Google data center that features multiple TPU devices connected to each other by extremely low latency communication channels).
研究动机与目标
- 为应对可扩展深度强化学习(RL)训练日益增长的计算需求。
- 通过利用TPU Pods实现快速原型设计与高吞吐量训练,推动强化学习研究。
- 设计兼顾易用性与生产级可扩展性及可复现性的系统。
- 通过优化TPU资源利用与架构选择,提升数据效率并降低训练成本。
- 支持多样化的强化学习工作负载,包括无模型与基于搜索的智能体(如MuZero)
提出的方法
- 设计Anakin与Sebulba作为Podracer架构,用于在TPU Pods上训练在线与演员-评论家RL智能体。
- 使用JAX实现可组合的程序变换、自动微分以及TPU上的硬件加速。
- 通过增加演员与学习者批次大小、延长轨迹长度以及增加网络深度来扩展训练,以最大化TPU资源利用率。
- 通过解耦动作与学习批次大小,实现吞吐量的线性扩展,尤其有利于基于搜索的智能体。
- 实现纯JAX版本的MCTS用于MuZero风格智能体,避免使用自定义C++代码进行搜索。
- 在多达2048个TPU核心上跨多个核心复制并扩展实验,以实现高吞吐量与低实际运行时间。
实验结果
研究问题
- RQ1如何有效利用TPU Pods扩展深度强化学习训练,同时保持研究敏捷性?
- RQ2在TPU上大规模强化学习训练中,哪些架构选择能最大化吞吐量与数据效率?
- RQ3纯JAX实现能否在性能上匹配或超越混合C++/Python系统,用于复杂强化学习智能体(如MuZero)?
- RQ4增加批次大小与网络容量如何影响基于TPU的强化学习训练中的数据效率与成本?
- RQ5在演员-评论家RL框架中,可扩展性在多大程度上可与数据效率解耦?
主要发现
- Sebulba在2048核TPU Pod上实现了每秒4300万帧的性能,训练Pong智能体耗时不足一分钟。
- 在8核TPU上训练V-trace智能体时,将演员批次大小从32提升至128,实现每秒20万帧的吞吐量。
- 使用更大网络而非更大批次可提升数据效率,且不增加TPU小时数或成本。
- 在16核TPU上,MuZero智能体通过Sebulba训练,在9小时内达到2亿个Atari帧,成本约为40美元(使用抢占式实例)。
- MuZero的吞吐量随TPU核心数量线性扩展,证明了有效的水平扩展能力。
- 解耦动作与学习批次大小在保持数据效率的同时,通过复制实现更快训练。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。