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
  • 项目
  • 项目