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.jitjax.gradjax.pmap)的研究人員和開發人員。

亮點

  • JAX 整合:與 JAX 的自動微分和 GPU/TPU 支援完全相容。
  • 類 Sonnet API:設計上與 Sonnet 2 API 近乎匹配,使從 TensorFlow/Sonnet 的遷移變得容易。
  • 簡化的 RNG 管理:提供 hk.next_rng_key(),以便在轉換後的函數中確定性地處理隨機數生成。
  • 可擴展性:由 DeepMind 研究人員在圖像、語言和強化學習任務中進行了大規模測試。
  • 是函式庫,而非框架:嚴格專注於參數和狀態管理,將優化器和檢查點保存留給其他函式庫。

相關

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