google-deepmind/mctx

Monte Carlo tree search in JAX

解決的問題

將搜尋演算法與深度神經網路結合時,通常需要使用 C++ 等快速編譯語言實作,這對研究人員來說可能難以使用與修改。Mctx 提供了一種高效率、JAX 原生的蒙特卡洛樹搜尋(MCTS)實作,可在 Python 中使用,彌補執行速度與研究者易用性之間的差距。

如何運作

Mctx 使用 JAX 實作 MCTS 演算法(包含 AlphaZero、MuZero 與 Gumbel MuZero)。它利用 JIT 編譯,並以平行方式處理批量輸入,以最大化硬體加速器的效率。使用者提供學習到的元件——例如根狀態的表示函數與環境動態的遞迴函數——該程式庫會使用這些元件來建構搜尋樹並提出動作。

適用對象

專為希望在不離開 Python 生態系的情況下,獲得編譯程式碼效能的 AI 研究人員設計,適用於研究基於搜尋的強化學習代理。

主要特色

  • JAX 原生:完全支援 JIT 編譯,實現顯著的計算加速。
  • 平行搜尋:以平行方式處理批量輸入,優化加速器使用效率。
  • 演算法支援:包含 AlphaZero、MuZero 與 Gumbel MuZero 的實作。
  • 高度可配置:設計為高度可配置,以利探索新的基於搜尋的構想。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案