google-deepmind/dm-haiku
JAX-based neural network library
解決的問題
Haiku 是一個用於 JAX 的神經網路函式庫,它允許開發人員使用熟悉的物件導向程式設計模型(如 Sonnet 或 TensorFlow 中使用的模型),同時保持對 JAX 純函數轉換的完全存取。它簡化了模型參數和內部狀態的管理,而無需使用者為初始化編寫大量的樣板程式碼。
工作原理
Haiku 提供了兩種主要工具來彌合物件導向設計與函數純粹性之間的差距:
hk.Module:用於定義神經網路層和組件的 Python 物件。這些模組持有參數和方法的引用,但在定義過程中被視為函數上的「不純」物件。hk.transform:一種函數轉換,可將這些基於模組的不純函數轉換為一對純函數:init(收集初始參數值)和apply(將這些參數重新注入函數進行計算)。
對於需要可變狀態(如 Batch Normalization 的移動平均值)的模型,Haiku 提供了 hk.transform_with_state,它可以分別管理參數和狀態。
適用對象
想要利用物件導向 API 的生產力來構建神經網路,但又需要 JAX 的效能和轉換能力(如 jax.jit、jax.grad 和 jax.pmap)的研究人員和開發人員。
亮點
- JAX 整合:與 JAX 的自動微分和 GPU/TPU 支援完全相容。
- 類 Sonnet API:設計上與 Sonnet 2 API 近乎匹配,使從 TensorFlow/Sonnet 的遷移變得容易。
- 簡化的 RNG 管理:提供
hk.next_rng_key(),以便在轉換後的函數中確定性地處理隨機數生成。 - 可擴展性:由 DeepMind 研究人員在圖像、語言和強化學習任務中進行了大規模測試。
- 是函式庫,而非框架:嚴格專注於參數和狀態管理,將優化器和檢查點保存留給其他函式庫。
相關
- 專案
- 專案
- 專案
- 專案
- 專案