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_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 上的 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