NAVIX:基于JAX扩展MiniGrid环境
机器学习
2024-07-30 v1 人工智能
摘要
随着深度强化学习(Deep RL)研究向解决大规模世界的方向发展,高效的环境模拟对于快速实验至关重要。然而,大多数现有环境在扩展到高吞吐量方面仍面临挑战,限制了有意义的进展。通常情况下,交互在CPU上计算,限制了训练速度和吞吐量,由于在跨多个机器分布任务时存在计算和通信开销。最终,深度强化学习训练受CPU限制,开发批处理、快速且可扩展的环境已成为推进的前沿。在最常用的强化学习(RL)环境中,MiniGrid是探索、课程学习、表征学习、多样性、元学习、信用分配和语言条件RL研究的基础,仍受上述限制的制约。本文提出NAVIX,将MiniGrid重新实现于JAX。NAVIX在批处理模式下实现了超过200 000倍的速度提升,支持在单块Nvidia A100 80 GB上并行运行最多2048个智能体。该做法将实验时间从一周缩短至15分钟,促进更快的设计迭代和更可扩展的RL模型开发。
引用
@article{arxiv.2407.19396,
title = {NAVIX: Scaling MiniGrid Environments with JAX},
author = {Eduardo Pignatelli and Jarek Liesen and Robert Tjarko Lange and Chris Lu and Pablo Samuel Castro and Laura Toni},
journal= {arXiv preprint arXiv:2407.19396},
year = {2024}
}