ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

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

Flax NNX 入门指南:用 Python 引用语义构建 JAX 神经网络 Flax NNX 入门指南用 Python 引用语义构建 JAX 神经网络【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flaxFlax 是为 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, rngsrngs) self.bn nnx.BatchNorm(dmid, rngsrngs) self.dropout nnx.Dropout(0.2) self.linear_out nnx.Linear(dmid, dout, rngsrngs) def __call__(self, x, rngs): x nnx.relu(self.dropout(self.bn(self.linear(x)), rngsrngs)) 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) 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(...)构造即完成参数分配rngsnnx.Rngs(0)以一个根 PRNG key 驱动所有随机初始化nnx.Linear、nnx.BatchNorm、nnx.Dropout等层直接作为属性挂在模型上flax/nnx/nn/init.py 导出了这些内置层。Optimizer 持有模型引用nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param)接收模型引用而非参数副本wrtnnx.Param指定相对哪些变量求梯度——只对nnx.Param类型的可训练权重更新。源码见 flax/nnx/training/optimizer.py。自动状态传播nnx.jit是jax.jit的有状态版本源码见 flax/nnx/transforms/compilation.py允许函数输入输出为 NNX 对象nnx.value_and_grad是jax.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 githttps://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. 可检查性InspectionLinen Module 是惰性的setup()中的子模块在构造时不可访问只能在运行时获得导致检查与调试困难。NNX Module 是普通 Python 对象构造后立即可访问class Block(nnx.Module): def __init__(self, rngs): self.linear nnx.Linear(5, 10, rngsrngs) block Block(nnx.Rngs(0)) block.linear # Linear( # kernelParam(valueArray(shape(5, 10), dtypefloat32)), # biasParam(valueArray(shape(10,), dtypefloat32)), # ...代价是没有 shape 推断输入输出形状都必须显式提供——换来的是更显式、更可预期的行为。配合nnx.display(model)基于 Treescope 的可视化见 flax/nnx/visualization.py可以一键查看模型全貌。2. 运行计算Running computationLinen 中所有顶层计算必须通过init/apply完成参数作为独立结构与 Module 分离造成apply 内外代码不对称。NNX 中参数就是属性方法可以直接调用__init__和__call__与普通方法地位完全平等# Linen 需要: # y model.apply({params: params}, x) # z model.apply({params: params}, x, methodencode) # NNX 直接调用: y model(x) z model.encode(x) y model.decoder(z)子模块在 NNX 中也可以被直接调用因为它们在构造时就已经初始化完毕。3. 状态处理State handlingLinen 中一旦引入 Dropout 或 BatchNorm就必须手工维护batch_stats等额外状态结构并配置apply(mutable...)。NNX 中状态保存在nnx.Module内部且可变直接调用即可class Block(nnx.Module): def __init__(self, rngs): self.linear nnx.Linear(5, 10, rngsrngs) self.bn nnx.BatchNorm(10, rngsrngs) self.dropout nnx.Dropout(0.1, rngsrngs) 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.bias4. 模型手术Model surgeryLinen 中替换子模块困难重重一是惰性初始化不保证能替换二是参数结构与 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, rank5, rngsrngs)若要批量替换可以用nnx.iter_graph由 flax/nnx/graphlib.py 导出遍历对象图把模型中所有nnx.Linear换成LoraLinear这一点在 nnx_basics.md 中也有完整示例。5. 变换TransformsLinen 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( kernelrandom.uniform(random.key(seed), (2, 3)), biasjnp.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_axes0, out_axes0)(seeds) y nnx.vmap(vector_dot, in_axes(0, 0), out_axes1)(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是其子类一般存放可训练权重。还有BatchStat、Cache、Intermediate等预定义子类导出于 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支持按命名空间如params、dropout独立取随机数fork方法可为nnx.scan/nnx.vmap的每一层/每一分支切分独立随机流。Functional APIsplit / merge / updatennx.split把 Module 拆成静态的GraphDef类似 JAX 的PyTreeDef与动态的Statejax.Array的 pytreennx.merge反向重建 Modulennx.update用State原地更新对象。这个三元组是 NNX 与纯 JAX 代码互操作的桥梁——跨 JAX 变换边界时用它显式传递状态从而避免共享引用被静默丢失。完整说明见 nnx_basics.md。学习路径与更多资源docs_nnx/index.rst末尾以卡片形式列出了官方推荐的学习路径全部可从本仓库对应文档继续深入Flax NNX Basicsnnx_basics.md —— 从零讲解 Module 系统、状态计算、嵌套 Module、模型手术、变换与 Functional API。MNIST 教程mnist_tutorial.md —— 端到端训练 CNN 手写数字分类器覆盖nnx.Optimizer、nnx.MultiMetric指标、nnx.view训练/评估视图切换以及用 Orbax 导出 SavedModel 部署。Guides 指南guides/index.rst 汇总了基础与进阶指南包括 pytree、transforms、view、filters_guide、randomness、checkpointing、data_loaders 与 jax_and_nnx_transforms。Linen 迁移到 NNXguides/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),仅供参考
返回列表