Stable-Baselines3 与 Hugging Face Hub 的集成

Hugging Face 已将 Stable-Baselines3 集成到 Hugging Face Hub 中,使研究人员和爱好者能够托管和加载 PyTorch 深度强化学习 (DRL) 模型。这种集成简化了在 Gym、Atari、MuJoco 和 Procgen 等环境中训练的预训练智能体 (agents) 的共享过程。

模型托管与分发

该集成允许用户通过 Hub 发现并共享保存的 Stable-Baselines3 模型。用户可以通过在 Hugging Face Hub 上过滤 stable-baselines3 来查找社区贡献的模型。

从 Hub 下载模型

要将保存的模型从 Hub 加载到 Stable-Baselines3 中,用户必须安装 huggingface_hubhuggingface_sb3 库。该过程需要仓库 ID (repo-id) 以及该仓库中模型 zip 文件的特定文件名。

加载模型的示例工作流:

  1. 安装依赖项:pip install huggingface_hub huggingface_sb3
  2. 使用 huggingface_sb3 库中的 load_from_hub 来检索检查点 (checkpoint)。
  3. 将检查点加载到 Stable-Baselines3 智能体中(例如,使用 PPO.load(checkpoint))。

将模型共享到 Hub

用户可以通过首先使用 huggingface-cli login 或针对 Jupyter/Colab 环境的 notebook_login() 进行身份验证,然后将训练好的智能体上传到 Hub。身份验证后,使用 huggingface_sb3 库中的 push_to_hub 函数将保存的模型 zip 文件上传到指定的仓库 ID。

未来路线图

Hugging Face 计划通过以下举措扩展 Hub 上的深度强化学习生态系统:

  • 库集成:集成 RL-baselines3-zoo 和其他深度强化学习库。
  • 模型集合:从 rl-trained-agents 集合中上传预训练智能体。
  • 算法实现:实现 Decision Transformers。

技术实现

这种集成是通过 huggingface_hub 库实现的,该库为库支持提供了必要的 API 和小部件 (widgets)。Hugging Face 为希望将其工具与 Hub 集成的其他库维护者提供了指南。

Sources

相关

  • Dispatch
  • Dispatch
  • Dispatch
  • Dispatch
  • Dispatch