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 讨论和拉取请求接受社区贡献。
如何开始?
- 安装 JAX(请遵循 JAX CPU/GPU/TPU 指南)。
- 通过 PyPI 安装 Flax:
可选:pip install flaxpip install "flax[all]"以安装额外依赖项,如 Matplotlib。 - 通过继承
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
- 项目
- 项目