google/flax

Flax is a neural network library for JAX that is designed for flexibility.

Flax とは?

Flax は、JAX(高性能数値計算フレームワーク)の上に構築されたオープンソースのニューラルネットワークライブラリです。柔軟性があり、Pythonに優しいAPI(新しいFlax NNX API)を提供し、研究者が通常のPythonオブジェクトとしてモデルを記述できるようにします。参照共有や可変性もサポートしています。ライブラリには一般的なレイヤー(Linear、Conv、BatchNorm、Attention、LSTM/GRU、Dropoutなど)、複数ノードでのトレーニング、チェックポイント、メトリクス用のユーティリティ、MNISTやGemma言語モデルデモなどの教育用例が含まれています。

メンテナンスは誰が行っていますか?

Google DeepMind のエンジニアおよび研究者によって開発されており、JAXチームと密接に連携しています。公式Google製品ではありませんが、GitHubのディスカッションやプルリクエストを通じてコミュニティからの貢献を積極的に受け入れ、継続的にメンテナンスされています。

どうやって始めますか?

  1. JAX をインストールする(JAX CPU/GPU/TPUガイドに従ってください)。
  2. PyPI から Flax をインストールする:
    pip install flax
    
    オプションpip install "flax[all]" でMatplotlibなどの追加依存関係をインストールできます。
  3. nnx.Module を継承してモデルを記述し、提供されたレイヤーを使用して、標準的なJAXコードでトレーニングします。

例コード(READMEから)

class MLP(nnx.Module):
  def __init__(self, din, dmid, dout, *, rngs):
    self.linear1 = nnx.Linear(din, dmid, rngs=rngs)
    self.dropout = nnx.Dropout(rate=0.1, rngs=rngs)
    self.bn = nnx.BatchNorm(dmid, rngs=rngs)
    self.linear2 = nnx.Linear(dmid, dout, rngs=rngs)

  def __call__(self, x):
    x = nnx.gelu(self.dropout(self.bn(self.linear1(x))))
    return self.linear2(x)

詳しく学ぶには?

  • ドキュメントサイトhttps://flax.readthedocs.io/
  • チュートリアル:MNISTチュートリアル、Gemma LM推論例、および「Flax NNXの基礎」ガイド。
  • ディスカッション&サポート:GitHub Discussions、イシュー追跡、flax-dev@google.com メーリングアドレス。

Flaxを使うべきタイミングは?

JAXを使っている上で、以下の要件を満たすニューラルネットワークライブラリが必要な場合、Flaxは自然な選択です:

  • JAXの完全な柔軟性を維持(隠れたグラフコンパイルステップなし)。
  • 通常のPythonの文法でモデルを記述できる。
  • 事前に用意されたレイヤー、トレーニングユーティリティ、例コードが利用可能。

Citation: READMEには学術的な引用用のBibTeXエントリが提供されています。

関連

  • プロジェクト
  • プロジェクト
  • Dispatch
  • プロジェクト
  • プロジェクト