Stable-Baselines3 與 Hugging Face Hub 的整合
Hugging Face 已將 Stable-Baselines3 整合至 Hugging Face Hub 中,讓研究人員與愛好者能夠託管並載入 PyTorch Deep Reinforcement Learning (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 上的 Deep Reinforcement Learning 生態系統:
- 函式庫整合:整合
RL-baselines3-zoo與其他 Deep Reinforcement Learning 函式庫。 - 模型集合:上傳來自
rl-trained-agents集合的預訓練代理。 - 演算法實作:實作 Decision Transformers。
技術實作
此整合是透過 huggingface_hub 函式庫實現的,該函式庫提供了必要的 API 與小工具(widgets)以支援函式庫整合。Hugging Face 也為希望將其工具與 Hub 整合的其他函式庫維護者提供了指南。
Sources
相關
- Dispatch
- Dispatch
- Dispatch
- Dispatch
- Dispatch