google-deepmind/dm-haiku
JAX-based neural network library
What it solves
JAX 是一個功能強大的數值計算函式庫,但它要求函式必須是純函式才能使用其轉換(如 jit 與 grad)。Haiku 透過允許開發者使用熟悉的物件導向程式設計模型來定義神經網路,並自動將這些「不純」的物件導向定義轉換成 JAX 能處理的純函式,解決了這個問題。
How it works
Haiku 提供兩個主要工具,彌合物件導向設計與函式純度之間的鴻溝:
hk.Module:用於定義網路層與元件的 Python 物件。這些模組持有參數與方法的參照,讓使用者能寫出類似標準神經網路函式庫的程式碼。hk.transform:一個函式轉換器,將使用hk.Module的函式轉換為一對純函式:init(收集初始參數值)與apply(將這些參數注入函式以進行計算)。
對於需要內部可變狀態(如批次正規化)的模型,Haiku 提供 hk.transform_with_state,可分別管理參數與狀態。
Who it’s for
需要以物件導向 API 提升建構神經網路生產力,同時又想完整利用 JAX 的函式轉換與硬體加速的研究人員與開發者。
Highlights
- DeepMind Scale:已由 DeepMind 研究人員在大規模影像、語言與強化學習實驗中測試。
- Library, Not Framework:設計為輕量函式庫,專注於參數與狀態管理,並不強加自訂優化器或檢查點格式。
hk.next_rng_key():透過提供確定性的金鑰序列,簡化 JAX 中的隨機數產生。- JAX Compatibility:完全相容
jax.pmap,可在多個加速器上進行分散式訓練。