Flax NNX 入门指南:用 Python 引用语义构建 JAX 神经网络
2026/9/17 23:02:56 网站建设 项目流程

Flax NNX 入门指南:用 Python 引用语义构建 JAX 神经网络

【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax

Flax 是为 JAX 打造的神经网络库,其核心新 API——Flax NNX——通过一等公民的 Python 引用语义(reference semantics),让研究人员和开发者能够用普通 Python 对象创建、检查、调试和分析神经网络模型。本文以本仓库 docs_nnx/index.rst 为主干,结合源码与配套教程,系统讲解 NNX 的核心概念、特性、基础用法与安装方式,读完即可上手用 NNX 编写可训练、可检查、可改造的 JAX 模型。

Flax 与 Flax NNX 是什么

Flax 为使用 JAX 进行神经网络开发的用户提供了一套灵活、端到端的用户体验,让你能够充分释放 JAX 的全部能力。在 Flax 体系内部,NNX 是位于核心位置的简化 API,专门用于降低神经网络在创建、检查、调试与分析上的复杂度。

NNX 最重要的设计决策是:为 JAX 引入一等公民的 Python 引用语义,用户可以用普通 Python 对象直接表达模型,模型被建模为 PyGraph(而非 pytree),从而天然支持引用共享与可变性。NNX 是此前 Flax Linen API 的演进产物——团队在多年实践中积累经验后,推出了一个更简单、更友好的 API。

需要特别说明的是:根据 docs_nnx/index.rst 中的官方说明,Flax Linen API 在可预见的未来不会被弃用,因为绝大多数 Flax 用户仍在使用它;但新用户被鼓励优先使用 Flax NNX。关于两者差异及设计动机,可参阅 Why Flax NNX;从 Linen 迁移到 NNX 可先学习 NNX Basics,再参考迁移指南。

Flax NNX 的四大核心特性

NNX 的设计目标可以概括为四个关键词(对应 docs_nnx/index.rst 的 Features 章节):

特性含义
Pythonic支持使用普通 Python 对象,提供直观、可预期的开发体验
Simple依托 Python 的对象模型,对用户而言简单直接,提升开发速度
Expressive通过 Filter 系统 对模型状态进行细粒度控制
Familiar通过 Functional API 轻松将 NNX 对象与普通 JAX 代码集成

这些特性都源自同一个根基:一切显式(explicit)。与 Flax Linen 或 Haiku 的 Module 体系不同,NNX Module 本身直接持有状态(如参数),PRNG 状态由用户显式传入,初始化时必须提供全部形状信息(不做 shape 推断)。从源码看,flax/nnx/module.py 中Module继承自Pytree并使用自定义ModuleMeta元类,子模块可以直接作为属性在__init__中赋值,__call__等前向方法不享受任何特殊对待——这正是"Pythonic"与"Simple"的底层体现。

基础用法:一段完整的 NNX 训练代码

docs_nnx/index.rst 的 Basic usage 章节给出了一个高度浓缩的 NNX 示例,它同时演示了急切初始化、自动状态传播、原地更新三大特性:

from flax import nnx import optax class Model(nnx.Module): def __init__(self, din, dmid, dout, rngs: nnx.Rngs): self.linear = nnx.Linear(din, dmid, rngs=rngs) self.bn = nnx.BatchNorm(dmid, rngs=rngs) self.dropout = nnx.Dropout(0.2) self.linear_out = nnx.Linear(dmid, dout, rngs=rngs) def __call__(self, x, rngs): x = nnx.relu(self.dropout(self.bn(self.linear(x)), rngs=rngs)) return self.linear_out(x) model = Model(2, 64, 3, rngs=nnx.Rngs(0)) # eager initialization optimizer = nnx.Optimizer(model, optax.adam(1e-3), wrt=nnx.Param) @nnx.jit # automatic state propagation def train_step(model, optimizer, x, y): loss_fn = lambda model: ((model(x) - y) ** 2).mean() loss, grads = nnx.value_and_grad(loss_fn)(model) optimizer.update(model, grads) # in-place updates return loss

逐行拆解这段代码,可以看到 NNX 与传统 JAX 编程范式的关键差异:

  • 急切初始化Model(...)构造即完成参数分配,rngs=nnx.Rngs(0)以一个根 PRNG key 驱动所有随机初始化;nnx.Linearnnx.BatchNormnnx.Dropout等层直接作为属性挂在模型上(flax/nnx/nn/init.py 导出了这些内置层)。
  • Optimizer 持有模型引用nnx.Optimizer(model, optax.adam(1e-3), wrt=nnx.Param)接收模型引用而非参数副本,wrt=nnx.Param指定"相对哪些变量求梯度"——只对nnx.Param类型的可训练权重更新。源码见 flax/nnx/training/optimizer.py。
  • 自动状态传播@nnx.jitjax.jit的有状态版本(源码见 flax/nnx/transforms/compilation.py),允许函数输入输出为 NNX 对象;nnx.value_and_gradjax.value_and_grad的有状态版本(源码见 flax/nnx/transforms/autodiff.py)。BatchNorm 的均值方差、Dropout 的随机性等状态更新会被自动从loss_fn一路传播回model引用,无需手动返回和回填状态。
  • 原地更新optimizer.update(model, grads)直接就地更新参数,训练循环因此非常简洁。

这段代码在 Why Flax NNX 中还有更完整的变体,其中__call__直接调用model(x),进一步省略了显式传rngs的环节。

安装

NNX 随 Flax 主包一起分发,安装方式与 Flax 完全一致:

pip install flax

也可以直接从仓库安装最新开发版:

pip install git+https://github.com/google/flax.git

安装后通过from flax import nnx即可使用全部 NNX 功能。仓库中 pyproject.toml 定义了包的构建配置;NNX 依赖 JAX 及可选依赖 optax(优化器)、orbax(检查点/导出)等,具体可在对应的 requirements.txt 与 docs_nnx/mnist_tutorial.md 中看到搭配使用方式。

深入:NNX 相比 Linen 改进在哪里

Why Flax NNX 从五个维度系统对比了两代 API,这里提炼其核心论点,帮助你判断迁移价值。

1. 可检查性(Inspection)

Linen Module 是惰性的:setup()中的子模块在构造时不可访问,只能在运行时获得,导致检查与调试困难。NNX Module 是普通 Python 对象,构造后立即可访问:

class Block(nnx.Module): def __init__(self, rngs): self.linear = nnx.Linear(5, 10, rngs=rngs) block = Block(nnx.Rngs(0)) block.linear # Linear( # kernel=Param(value=Array(shape=(5, 10), dtype=float32)), # bias=Param(value=Array(shape=(10,), dtype=float32)), # ...

代价是没有 shape 推断,输入输出形状都必须显式提供——换来的是更显式、更可预期的行为。配合nnx.display(model)(基于 Treescope 的可视化,见 flax/nnx/visualization.py)可以一键查看模型全貌。

2. 运行计算(Running computation)

Linen 中所有顶层计算必须通过init/apply完成,参数作为独立结构与 Module 分离,造成"apply 内外代码不对称"。NNX 中参数就是属性,方法可以直接调用,__init____call__与普通方法地位完全平等:

# Linen 需要: # y = model.apply({'params': params}, x) # z = model.apply({'params': params}, x, method='encode') # NNX 直接调用: y = model(x) z = model.encode(x) y = model.decoder(z)

子模块在 NNX 中也可以被直接调用,因为它们在构造时就已经初始化完毕。

3. 状态处理(State handling)

Linen 中一旦引入 Dropout 或 BatchNorm,就必须手工维护batch_stats等额外状态结构,并配置apply(mutable=...)。NNX 中状态保存在nnx.Module内部且可变,直接调用即可:

class Block(nnx.Module): def __init__(self, rngs): self.linear = nnx.Linear(5, 10, rngs=rngs) self.bn = nnx.BatchNorm(10, rngs=rngs) self.dropout = nnx.Dropout(0.1, rngs=rngs) def __call__(self, x): return nnx.relu(self.dropout(self.bn(self.linear(x)))) model = Block(nnx.Rngs(0)) y = model(x)

最大的收益是:添加新的有状态层时,训练代码无需任何改动。自定义有状态层也非常简单——下面的简化版 BatchNorm 每次调用都会更新均值与方差(使用nnx.Param存放可训练的 scale/bias,用nnx.BatchStat存放统计量):

class BatchNorm(nnx.Module): def __init__(self, features: int, mu: float = 0.95): self.scale = nnx.Param(jax.numpy.ones((features,))) self.bias = nnx.Param(jax.numpy.zeros((features,))) self.mean = nnx.BatchStat(jax.numpy.zeros((features,))) self.var = nnx.BatchStat(jax.numpy.ones((features,))) self.mu = mu # Static def __call__(self, x): mean = jax.numpy.mean(x, axis=-1) var = jax.numpy.var(x, axis=-1) self.mean.value = self.mu * self.mean + (1 - self.mu) * mean self.var.value = self.mu * self.var + (1 - self.mu) * var x = (x - mean) / jax.numpy.sqrt(var + 1e-5) return x * self.scale + self.bias

4. 模型手术(Model surgery)

Linen 中替换子模块困难重重:一是惰性初始化不保证能替换,二是参数结构与 Module 结构分离需要手动同步。NNX 中直接按 Python 语义替换子模块即可,参数与 Module 同构、永不失同步。典型场景是给已有模型插入 LoRA 层:

class LoraParam(nnx.Param): pass class LoraLinear(nnx.Module): def __init__(self, linear, rank, rngs): self.linear = linear self.A = LoraParam(random.normal(rngs(), (linear.in_features, rank))) self.B = LoraParam(random.normal(rngs(), (rank, linear.out_features))) def __call__(self, x): return self.linear(x) + x @ self.A @ self.B rngs = nnx.Rngs(0) model = Block(rngs) model.linear = LoraLinear(model.linear, rank=5, rngs=rngs)

若要批量替换,可以用nnx.iter_graph(由 flax/nnx/graphlib.py 导出)遍历对象图,把模型中所有nnx.Linear换成LoraLinear;这一点在 nnx_basics.md 中也有完整示例。

5. 变换(Transforms)

Linen transforms 的局限包括:暴露了 JAX 之外的额外 API、只接受"Module 作为第一参数"的特定函数签名、只能在apply内使用。NNX transforms 则与对应 JAX transformsAPI 同构,只是额外支持 NNX Module——Module 可以作为任意位置的参数甚至返回值,并且可以出现在包括训练循环在内的任何地方。

nnx.vmap为例,既可以变换"创建权重"的函数来制造权重堆叠,也可以变换"向量点积"函数来对批量输入逐条应用:

class Weights(nnx.Module): def __init__(self, kernel, bias): self.kernel, self.bias = nnx.Param(kernel), nnx.Param(bias) def create_weights(seed): return Weights( kernel=random.uniform(random.key(seed), (2, 3)), bias=jnp.zeros((3,)), ) def vector_dot(weights, x): assert weights.kernel.ndim == 2, 'Batch dimensions not allowed' assert x.ndim == 1, 'Batch dimensions not allowed' return x @ weights.kernel + weights.bias weights = nnx.vmap(create_weights, in_axes=0, out_axes=0)(seeds) y = nnx.vmap(vector_dot, in_axes=(0, 0), out_axes=1)(weights, x)

与 Linen 变换不同,in_axes等参数会真实影响nnx.Module状态如何被变换。更妙的是,由于nnx.Module方法本质上就是"以 Module 为第一参数的函数",NNX transforms 可以直接作为方法装饰器使用(源码见 flax/nnx/transforms/iteration.py 的vmap与 flax/nnx/transforms/iteration.py 的scan)。

支撑 NNX 的底层机制

NNX 的易用性建立在一套清晰的抽象之上,理解它们能帮助你更好地驾驭这套 API。相关术语均可查阅 NNX 术语表:

  • Variable / Paramnnx.Variable是存放在 Module 中的权重/参数/数据/数组;nnx.Param是其子类,一般存放可训练权重。还有BatchStatCacheIntermediate等预定义子类,导出于 flax/nnx/variablelib.py。
  • Filter 系统:一种从 Module 中抽取特定Variable的方式,通常通过nnx.split配合类型过滤器(如nnx.Param、自定义Count类型)实现,用于把状态切成互斥的多个State分组——这正是 Filter 指南 讲解的内容,也是"Expressive"特性的落地。
  • Rngs / PRNG 管理nnx.Rngs持有根 PRNG 状态并可派发新 key(实现见 flax/nnx/rnglib.py),支持按命名空间(如paramsdropout)独立取随机数;fork方法可为nnx.scan/nnx.vmap的每一层/每一分支切分独立随机流。
  • Functional API(split / merge / update)nnx.split把 Module 拆成静态的GraphDef(类似 JAX 的PyTreeDef)与动态的Statejax.Array的 pytree);nnx.merge反向重建 Module;nnx.updateState原地更新对象。这个三元组是 NNX 与纯 JAX 代码互操作的桥梁——跨 JAX 变换边界时用它显式传递状态,从而避免共享引用被静默丢失。完整说明见 nnx_basics.md。

学习路径与更多资源

docs_nnx/index.rst末尾以卡片形式列出了官方推荐的学习路径,全部可从本仓库对应文档继续深入:

  • Flax NNX Basics:nnx_basics.md —— 从零讲解 Module 系统、状态计算、嵌套 Module、模型手术、变换与 Functional API。
  • MNIST 教程:mnist_tutorial.md —— 端到端训练 CNN 手写数字分类器,覆盖nnx.Optimizernnx.MultiMetric指标、nnx.view训练/评估视图切换,以及用 Orbax 导出 SavedModel 部署。
  • Guides 指南:guides/index.rst 汇总了基础与进阶指南,包括 pytree、transforms、view、filters_guide、randomness、checkpointing、data_loaders 与 jax_and_nnx_transforms。
  • Linen 迁移到 NNX:guides/linen_to_nnx.rst 为存量 Linen 代码提供分步迁移指导;背景动机见 why.rst。
  • API 参考:api_reference/index.rst 覆盖flax.nnx全部公开 API,包括 nn 模块(层)、transforms、state、graph 与 training 等。
  • 术语表:nnx_glossary.rst 可随时查阅 Filter、GraphDef、Split and merge、Variable 等核心概念。

仓库 examples 目录还提供了大量可直接运行的实战案例(MNIST、ImageNet、LM1B、WMT 翻译、PPO 强化学习等),其中 nnx_toy_examples 下的 10 个脚本按难度递进地演示了 NNX 的函数式 API、lifted transforms、训练状态、数据并行、VAE、层间 scan、数组叶子、检查点、参数手术与 FSDP 优化器,是与本文搭配的最佳动手练习素材。

【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询