google-deepmind/dm-haiku

JAX-based neural network library

What it solves

JAX 是一個功能強大的數值計算函式庫,但它要求函式必須是純函式才能使用其轉換(如 jitgrad)。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,可在多個加速器上進行分散式訓練。