
掌握 Flax NNX用 Pythonic 方式在 JAX 中构建与训练神经网络【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flaxNNXNeuralNetworks for JAX是 Flax 生态中面向 JAX 的神经网络库其核心目标是提供最佳开发体验让构建和实验神经网络像普通 Python 编程一样直观。本文以仓库中的 flax/nnx/README.md 为骨架结合 flax/nnx 源码与 examples/nnx_toy_examples 实战示例带你理解 NNX 的设计理念掌握其模块、变量、随机数与变换系统四大核心抽象并能够独立写出一个完整的模型定义与训练循环。NNX 是什么为 JAX 而生、为开发体验而设计的神经网络库NNX 是一个构建在 JAX 之上的神经网络库与 Flax 经典的 Linen API 相比它把「开发体验」放在了首要位置。README 明确概括了 NNX 的四大设计目标PythonicPython 风格模块就是标准 Python 类上手成本低开发体验贴近日常 Python 编程。Easy-to-use易于使用NNX 提供一组变换transforms负责状态管理用户可以把精力集中在模型结构与训练循环本身。Expressive富有表现力NNX 通过 lifted transforms 对模块状态进行细粒度控制支持定义复杂架构。Compatible兼容NNX 支持将模块状态功能化functionalizing需要时可以无缝使用原生 JAX 变换。「Pythonic」并非口号在源码中可以找到直接证据Module在 flax/nnx/module.py 中被定义为class Module(Pytree, metaclassModuleMeta)——它首先是一个标准的 Python 类继承自Pytree通过元类ModuleMeta在背后完成图节点graph node注册等机制但对用户而言写nnx.Module子类与写普通类没有区别。模块的属性子模块、参数、普通值会在类层面被自动追踪这是 NNX「声明即构建」体验的底层支撑。快速上手一个开箱即用的 NNX 模型README 给出了一段非常精炼的示例完整演示了 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, rngsrngs) self.bn nnx.BatchNorm(dmid, rngsrngs) self.dropout nnx.Dropout(0.2, rngsrngs) self.linear_out nnx.Linear(dmid, dout, rngsrngs) def __call__(self, x): x nnx.relu(self.dropout(self.bn(self.linear(x)))) return self.linear_out(x) model Model(2, 64, 3, rngsnnx.Rngs(0)) # eager initialization optimizer nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param) # reference sharing nnx.jit # automatic state management def train_step(model, optimizer, x, y): def loss_fn(model): y_pred model(x) # call methods directly return ((y_pred - y) ** 2).mean() loss, grads nnx.value_and_grad(loss_fn)(model) optimizer.update(model, grads) # inplace updates return loss模型定义标准 Python 类的组合Model继承nnx.Module在__init__中像搭积木一样把nnx.Linear、nnx.BatchNorm、nnx.Dropout直接赋给self属性。这些层均实现于 flax/nnx/nn 目录linear.py、normalization.py、stochastic.py并在 flax/nnx/init.py 中统一导出。注意__init__的写法每个子模块都需要传入rngs用于初始化参数权重、偏置、BatchNorm 的 scale/bias以及为 Dropout 提供随机流eager initialization急切初始化Model(2, 64, 3, rngsnnx.Rngs(0))在构造那一刻就完成了参数分配无需像 Linen 那样先init再apply__call__中直接self.linear(x)调用层方法、nnx.relu(...)使用激活函数BatchNorm与Dropout的训练/推理状态由模块内部管理。优化器通过引用共享跟踪参数optimizer nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param)nnx.Optimizer与 optax 梯度变换直接对接。其中wrtnnx.Param是一个过滤器filter它告诉优化器「只对Param类型的变量维护优化器状态」。在 flax/nnx/training/optimizer.py 中Optimizer.__init__会调用tx.init(nnx.state(model, wrt, graphgraph))来创建优化器状态并用OptState/OptArray/OptVariable这类Variable子类包装这些状态见 flax/nnx/training/optimizer.py使其同样受 NNX 图协议管理。README 注释中的 reference sharing 指的是优化器与模型之间通过引用共享参数——模型参数更新后优化器内部状态依然保持一致。训练步骤变换接管状态管理nnx.jit是 NNX 对 JAXjit的有状态封装实现在 flax/nnx/transforms/compilation.py。它自动处理自动状态管理被jit的函数内部直接调用model(x)即可BatchNorm的运行均值/方差、Dropout 的随机数等状态由变换框架自动「拆包—传入—回写」无需手工传递原样书写训练逻辑nnx.value_and_grad(loss_fn)(model)返回(loss, grads)其中grads是 NNX 状态对象inplace 更新optimizer.update(model, grads)就地更新模型参数函数返回loss供日志记录。这段示例浓缩了 NNX「状态自动管理」的完整闭环定义eager init→ 优化器reference sharing→ 变换automatic state management→ 更新inplace updates。核心抽象Module、Variable 与 Rngs深入源码可以看到NNX 的设计围绕三类基本构件展开。Module可嵌套、可遍历的图节点Module是 NNX 的顶层抽象。在 flax/nnx/module.py 中Module同时提供iter_children()、iter_modules()等遍历方法以及train()/eval()切换训练推理模式对 BatchNorm、Dropout 等层生效。模块树本质上是一张「图」共享引用同一对象被多处引用也能被正确表达。Variable带类型的可变状态NNX 中所有可训练/可更新的状态都是Variable的实例。flax/nnx/variablelib.py 定义了完整的变量类型体系flax/nnx/init.py 中导出Param可训练参数默认被nnx.grad求导BatchStatBatchNorm 等层的运行时统计量running mean/varianceCache推理缓存如自回归解码的 KV cacheIntermediate前向过程中间值Perturbation扰动项A通用数组变量基类。wrtnnx.Param、nnx.state(model, Count)这类用法中的过滤器正是按Variable类型对状态进行筛选。你还可以像 examples/nnx_toy_examples/02_lifted_transforms.py 那样自定义变量类型class Count(nnx.Variable): pass class MLP(nnx.Module): def __init__(self, din, dhidden, dout, *, rngs: nnx.Rngs): self.count Count(jnp.array(0)) # 自定义状态默认不参与求导 ... def __call__(self, x): self.count.value 1 # 就地修改 ...Variable还实现了丰富的运算符重载与value/raw_value访问器见 flax/nnx/variablelib.py让数组操作保持直觉化。Rngs按名字组织的随机数流NNX 用Rngs统一管理随机数。flax/nnx/rnglib.py 的实现中RngStream内部由RngKeyJAX PRNG key与RngCount计数两个变量构成且支持命名流Rngs(0)创建默认流Rngs(0, noise1)则额外创建名为noise的流VAE 示例 examples/nnx_toy_examples/05_vae.py 中self.rngs.noise()即如此使用。命名流让不同用途的随机性参数初始化、Dropout、采样噪声互不干扰且由于随机数也是Variable它们在变换与序列化时会被自动处理。变换系统自动状态管理与 lifted transformsNNX 的「变换」是其易用性的核心。除nnx.jit外flax/nnx/transforms 还提供自动微分nnx.grad、nnx.value_and_grad、nnx.vjp、nnx.jvp、nnx.custom_vjp、nnx.rematflax/nnx/transforms/autodiff.py编译与并行nnx.jit、nnx.pmap、nnx.shard_mapflax/nnx/transforms/compilation.py迭代结构nnx.scan、nnx.vmap、nnx.while_loop、nnx.fori_loopflax/nnx/transforms/iteration.py流程控制nnx.cond、nnx.switch、nnx.checkify等。这些变换内部依赖 flax/nnx/graphlib.py 提供的split/merge/state/update图协议进入变换前把模块图拆成「结构定义GraphDef 状态State」变换结束后把更新后的状态合并回模型。这就是 README 所述「lifted transforms 让你对模块状态拥有细粒度控制」的底层机制也是「Expressive」这一特性的实现基础。上图来自 docs_nnx/guides/images/stateful-transforms.png展示了 NNX 有状态变换的通用流程模型M经过 Partition分割进入jax transform函数变换y f(P)产出结果再经 Merge合并与状态更新得到新模型M_out。兼容 JAX把模块变成普通 PyTree「Compatible」体现在功能化functionalizing能力上。使用nnx.split/nnx.merge/nnx.state可以把模块拆成纯数据State从而直接使用原生jax.jit、jax.grad、jax.vmap等变换。examples/nnx_toy_examples/01_functional_api.py 是完整示例先用nnx.split(model, nnx.Param, Count)一次性拆出graphdef, params, counts然后训练步骤完全用原生 JAX 书写graphdef, params, counts nnx.split( MLP(din1, dhidden32, dout1, rngsnnx.Rngs(0)), nnx.Param, Count ) jax.jit def train_step(params, counts, batch): x, y batch def loss_fn(params): model nnx.merge(graphdef, params, counts) y_pred model(x) new_counts nnx.state(model, Count) loss jnp.mean((y - y_pred) ** 2) return loss, new_counts grad, counts jax.grad(loss_fn, has_auxTrue)(params) params jax.tree.map(lambda w, g: w - 0.1 * g, params, grad) return params, counts这套「拆—训—合」流程nnx.split→ 原生 JAX 变换 →nnx.merge意味着 NNX 模型可以在需要时退化为纯函数式编程与既有 JAX 生态无缝衔接。安装 NNXNNX 随 Flax 一起发布。要使用最新开发版可按 README 给出的方式直接从 Flax 官方仓库安装pip install githttps://github.com/google/flax.git安装后即可from flax import nnx。NNX 依赖 JAX 与 OptaxREADME 示例中optax.adam、nnx.Optimizer均直接使用 Optax 变换VAE 示例还使用了datasets库加载 MNIST见 examples/nnx_toy_examples/requirements.txt。从玩具示例到真实任务NNX 实战参考仓库 examples/nnx_toy_examples 提供了从入门到进阶的系列脚本覆盖 NNX 的主要 API 形态01_functional_api.py展示如何用功能化 APInnx.split 原生jax.jit/jax.grad训练一个简单模型对应 README 中的 Using the Functional API。它证明即便不借助任何 NNX 变换包装也能完成训练。02_lifted_transforms.pyREADME 中的 Basic Example。它用nnx.jit包装训练/测试步骤并通过nnx.cached_partial(train_step, model, optimizer)缓存模型与优化器避免每次调用重复传递同时演示了自定义Count变量在训练中被就地累加输出times called验证前向次数。03_train_state.py从文件名可以推断该示例聚焦如何使用TrainStateflax/nnx/helpers.py 中导出的训练状态容器组织训练循环。04_data_parallel_with_jit.py从文件名看它演示了借助nnx.jit进行数据并行的写法。05_vae.pyREADME 中的 Training a VAE。在二值化 MNIST 上训练 VAE亮点包括nnx.Rngs(0, noise1)命名随机流驱动重参数化采样自定义Loss变量在前向中sow播种KL 损失训练时用nnx.pop(model, Loss)取出并求和nnx.value_and_gradoptimizer.update完成训练。这是展示 NNX 变量类型体系实战价值的典型样例。06_scan_over_layers.pyREADME 中的 Scan over layers。该示例构造了一个含 Dropout 与共享 BatchNorm 层的多层 MLP初始化阶段用nnx.split_rngs(splitsn_layers)nnx.vmap(axis_sizen_layers)一次性创建n_layers份Block实例每个有独立的layer轴前向阶段用nnx.scan逐层迭代应用。README 特别指出它使用了功能化 API 配合jax.vmap与jax.lax.scan是理解 lifted transforms 如何实现的绝佳范例。此外README 还提及 LM1B 语言模型示例基于 1 Billion Word Benchmark 数据集训练本仓库中对应的实战代码位于 examples/lm1b含train.py、models.py、input_pipeline.py等见 examples/lm1b/README.md可作为将 NNX 应用到真实规模任务时的参考。从源码看 NNX 的设计与实现最后我们回到源码看 README 示例中optimizer.update(model, grads)背后发生了什么。在 flax/nnx/training/optimizer.py 中Optimizer.update的实现链路是param_arrays nnx.as_pure(nnx.state(model, self.wrt, graphself.graph))把模型中Param类型变量提取为纯数组同理提取梯度数组grad_arrays与优化器状态opt_state_arrays调用self.tx.update(grad_arrays, opt_state_arrays, param_arrays)得到更新量与新优化器状态new_params optax.apply_updates(param_arrays, updates)计算新参数nnx.update(model, new_params)与nnx.update(self.opt_state, ...)将结果回写到模型与优化器self.step[...] 1推进步数。可以看到NNX 的状态管理并不是黑魔法nnx.state提取、nnx.as_pure转为纯数组、nnx.update回写构成了一个清晰、可组合的管线任何一步都可以被用户在功能化 API 中手动复刻正如01_functional_api.py所做。这也是 NNX「既有高阶变换的便捷又有低阶原语的可控」的架构特点。小结NNX 以「Pythonic、易用、富有表现力、兼容 JAX」为设计目标通过标准 Python 类的Module、带类型的Variable、命名随机流Rngs与自动管理状态的变换系统把 JAX 中繁琐的状态搬运工作收敛到框架内部。无论你是想快速写出第一个模型README 示例用功能化 API 与既有 JAX 代码共存01_functional_api.py还是实现 VAE、scan over layers 这类复杂结构05_vae.py、06_scan_over_layers.pyNNX 都提供了对应的表达方式。想要继续深入仓库的 docs_nnx 目录是下一步的好去处nnx_basics.mddocs_nnx/nnx_basics.md系统讲解Module抽象docs_nnx/guides/transforms.md 深入变换系统docs_nnx/guides/filters_guide.md 详解过滤器语法docs_nnx/migrating/linen_to_nnx.rst 则帮助 Linen 用户平滑迁移。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考