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,並需要一個神經網路庫,滿足以下條件:

  • 保持 JAX 的全部彈性(無隱藏的圖編譯步驟)。
  • 允許您使用一般的 Python 語法撰寫模型。
  • 提供現成的層、訓練工具和範例程式碼。 那麼 Flax 是一個自然的選擇。

Citation: README 提供了學術引用用的 BibTeX 項目。

相關

  • 專案
  • 專案
  • Dispatch
  • 專案
  • 專案