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_hub 和 huggingface_sb3 库。该过程需要仓库 ID (repo-id) 以及该仓库中模型 zip 文件的特定文件名。
加载模型的示例工作流:
- 安装依赖项:
pip install huggingface_hub huggingface_sb3。 - 使用
huggingface_sb3库中的load_from_hub来检索检查点 (checkpoint)。 - 将检查点加载到 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