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と完全に互換性があり、複数アクセラレータ上での分散トレーニングが可能です。