google-deepmind/dm-haiku

JAX-based neural network library

What it solves

JAX は強力な数値計算ライブラリですが、変換(jitgrad など)を利用するには関数が純粋である必要があります。Haiku は、開発者が慣れ親しんだオブジェクト指向プログラミングモデルでニューラルネットワークを定義できるようにし、これら「不純」なオブジェクト指向定義を JAX が処理できる純粋関数へ自動的に変換することでこの問題を解決します。

How it works

Haiku は、オブジェクト指向設計と関数の純粋性のギャップを埋めるための主なツールを二つ提供します:

  • hk.Module:ネットワーク層やコンポーネントを定義するための Python オブジェクト。これらのモジュールはパラメータやメソッドへの参照を保持し、標準的なニューラルネットワークライブラリのようなコードを書けます。
  • hk.transformhk.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 Compatibilityjax.pmap と完全に互換性があり、複数アクセラレータ上での分散トレーニングが可能です。