资讯动态

Flax NNX Training 模块实战指南:Metric 指标、Optimizer 优化器与 EMA 指数移动平均

发布时间:2026/9/17 8:05:30 来源:尧图企业网站定制
Flax NNX Training 模块实战指南Metric 指标、Optimizer 优化器与 EMA 指数移动平均【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本文围绕 docs_nnx/api_reference/flax.nnx/training/index.rst 展开。Flax 的nnx.training是面向 NNX 编程范式的训练工具包提供指标统计Metrics、优化器状态管理Optimizer与指数移动平均EMA三大能力。读完本文你将掌握flax.nnx.MultiMetric、flax.nnx.Optimizer、flax.nnx.EMA的完整用法、关键参数与底层调用链能够直接构建一个可训练、可评估、可平滑参数的 NNX 训练循环。一、flax.nnx.training模块总览在 NNX 生态中训练逻辑由三块核心组件构成它们全部收敛在flax.nnx.training包中并通过 flax/nnx/init.py 以flax.nnx命名空间对外暴露组件导出名称源码文件职责指标统计nnx.metrics、nnx.Metric、nnx.MultiMetricflax/nnx/training/metrics.py训练/验证过程中的 loss、accuracy 等指标的增量式统计优化器nnx.Optimizer、nnx.ModelAndOptimizerflax/nnx/training/optimizer.py将 Optax 梯度变换与模型参数、优化器状态绑定封装一步更新指数移动平均nnx.EMAflax/nnx/training/ema.py维护参数的指数滑动平均影子副本用于稳定训练、提升推理表现对应的 API 参考文档分别位于 metrics.rst、optimizer.rst 与 ema.rst。需要说明的是该模块目前仍标注为Experimental API见 index.rst 第 4 行的说明使用时应留意后续版本可能发生的 API 调整。三个组件的共同设计哲学是状态即Variable。指标累计值、优化器状态、EMA 影子参数都以flax.nnx.Variable的形式挂在对象上因此它们天然可被nnx.split/nnx.merge、nnx.jit、nnx.vmap等 NNX 变换处理也与nnx.state状态提取、筛选器Filter机制无缝协作。二、Metrics增量式指标统计2.1 基类Metric与MetricState所有指标都继承自 metrics.py 中的MetricPytree子类它定义了一套就地in-place三方法协议reset()将指标清零就地修改update(**kwargs)用新一批数据增量更新累计状态就地修改compute()基于累计状态计算并返回当前指标值不修改状态。源码通过raise NotImplementedError强制子类实现这三个方法。与普通浮点累加不同指标的内部状态是MetricState——一个继承自nnx.Variable的包装类metrics.py这意味着累计值本身也是 NNX 变量可以被筛选器过滤、被nnx.split切出、在nnx.jit下自动追踪。2.2Average加权平均与 mask 过滤Average用于统计批内 loss 等标量/数组的均值构造参数为argname: str values它指定update从哪个关键字参数取值。源码实现metrics.py的关键点内部维护self.totalfloat32 累计和与self.countint32 累计样本数二者都是MetricStateupdate(maskNone, **kwargs)会先校验self.argname in kwargs否则抛出TypeError标量输入total values、count 1数组输入total values.sum()、count values.size传入mask时按values * mask规则过滤且要求values必须是 jax 数组对标量传 mask 会抛出ValueErrormask 形状需可广播到 values 形状compute()返回total / count。注意未做任何更新时count为 0compute()返回nan——这是设计行为测试与文档示例中均以此判断尚未更新。Average构造时的argname参数非常灵活avg Average(test)后即可用avg.update(testnew_value)更新。2.3Accuracy二分类与多分类准确率Accuracy继承自Average复用其reset/compute但重写了updatemetrics.py多分类默认thresholdNone要求logits.ndim labels.ndim 1先对 logits 在最后一维做argmax再与整数标签逐元素比较最后以values((logits.argmax(axis-1)) labels)调用父类update二分类threshold传入 float如0.5要求logits.ndim labels.ndim判定规则为(logits threshold) (labels 0)维度不匹配会抛出带明确提示的ValueError如For multi-class classification, expected logits.ndimlabels.ndim1见 tests/nnx/metrics_test.py 的test_accuracy_dims标签会统一转换为int32非整型标签 dtype 会报错同样支持mask参数。import jax, jax.numpy as jnp from flax import nnx # 多分类logits 形状 (5, 2)labels 形状 (5,) logits jax.random.normal(jax.random.key(0), (5, 2)) labels jnp.array([0, 1, 1, 1, 0]) acc nnx.metrics.Accuracy() acc.update(logitslogits, labelslabels) acc.compute() # Array(0.6, dtypefloat32) # 二分类阈值 0.5logits 与 labels 同形状 logits3 jax.random.normal(jax.random.key(2), (5,)) labels3 jnp.array([0, 1, 0, 1, 1]) acc2 nnx.metrics.Accuracy(threshold0.5) acc2.update(logitslogits3, labelslabels3) acc2.compute() # Array(0.8, dtypefloat32)2.4Welford数值稳定的均值与方差Welford采用 Welford 在线算法计算数据流的均值与方差metrics.py内部维护count、mean、m2平方差累计量三个MetricState每次update按递推公式更新delta new_mean - self.meanself.mean delta * count / new_countself.m2 m2_batch delta * delta * count * original_count / new_countcompute()返回一个flax.struct.dataclass定义的Statistics对象metrics.py包含三个字段字段含义mean累计均值standard_deviation标准差sqrt(m2 / count)standard_error_of_mean均值标准误std / sqrt(count)相比朴素的两趟算法Welford 对数值偏差如均值在1e16量级的数据流有显著更好的稳定性这一点由 tests/nnx/metrics_test.py 中的test_welford_large 1e16偏移与test_welford_many5 万样本、按 3 倍标准误校验均值两个测试用例印证。2.5MultiMetric一次调用更新多个指标实际训练中常常需要同时跟踪 loss 与 accuracyMultiMetric正是为此设计metrics.py构造时以关键字参数接收多个指标MultiMetric(accuracynnx.metrics.Accuracy(), lossnnx.metrics.Average())通过setattr将各指标挂为同名属性并记录在self._metric_names元组中update(**updates)会把所有关键字参数广播给每个子指标的update。特别地mask参数有两种形态直接传jax.Array则作用于所有子指标传字典{metric_name: metric_mask}则按指标名分别施加各自的 maskmetrics.py。MultiMetric内部的掩码分发逻辑正是来自mask.get(metric_name, None)这段实现compute()返回一个字典键为构造时的指标名值为各指标的compute()结果reset()遍历_metric_names逐个重置。from flax import nnx import jax, jax.numpy as jnp metrics nnx.MultiMetric( accuracynnx.metrics.Accuracy(), lossnnx.metrics.Average() ) batch_loss jnp.array([1, 2, 3, 4]) logits jax.random.normal(jax.random.key(0), (5, 2)) labels jnp.array([0, 1, 1, 1, 0]) metrics.update(logitslogits, labelslabels, valuesbatch_loss) metrics.compute() # {accuracy: Array(0.6, dtypefloat32), loss: Array(2.5, dtypefloat32)}值得注意的是MultiMetric的_metric_names是普通 Python 属性非变量因此其graphdef在nnx.split之后保持可哈希——tests/nnx/metrics_test.py 专门用hash(graphdef)验证了这一点同时test_multimetric_split_merge_under_jit证明了MultiMetric可以安全地在jax.jit/nnx.jit下完成split→ 更新 →merge的闭环tests/nnx/metrics_test.py。2.6 自定义指标与 JIT 兼容性要自定义指标只需继承Metric并实现__init__/reset/update/compute。例如 tests/nnx/metrics_test.py 中的CustomAccuracy演示了如何继承Accuracy改写update签名。测试同时提示了一个易错点MultiMetric.update会把mask关键字广播给所有子指标因此自定义指标的update要么接受**kwargs要么显式接收mask参数。此外test_vmap_reset_preserves_shapetests/nnx/metrics_test.py验证了nnx.vmap下reset不会破坏MetricState的 vmap 形状说明指标对象可直接用于nnx.vmap等批量变换场景。三、Optimizer绑定模型与 Optax 的一步优化器3.1 构造与核心参数Optimizerflax/nnx/training/optimizer.py是单 Optax 优化器的通用训练状态。构造签名nnx.Optimizer(model, tx, *, wrtnnx.Param, graphNone)参数含义model一个 NNX Module如nnx.Linear或自定义模块txOptax 梯度变换optax.GradientTransformation如optax.adam(1e-3)wrtNNX 筛选器指定优化器状态跟踪哪些Variable必须与nnx.grad的wrt一致默认nnx.ParamgraphTrue使用 graph 模式支持共享引用等完整 NNX 特性False使用 tree 模式把 Module 当普通 JAX pytree避免 graph 协议开销None由当前nnx.set_graph_mode上下文决定构造时执行三件事将step初始化为OptState(jnp.array(0, dtypejnp.uint32))用tx.init(nnx.state(model, wrt))初始化 Optax 状态并经to_opt_state包装记录wrt。其中to_opt_stateoptimizer.py会把每个叶子包装为两类状态变量OptArray普通数组优化器状态OptVariable从原始Variable派生的优化器状态并透传元数据——特别是把变量元数据中的optimizer_sharding改写为out_sharding从而让优化器状态继承参数的分片标注见 tests/nnx/optimizer_test.py 的test_sharding_propagation经过nnx.Optimizer后opt_state[0][mu][kernel]的out_sharding与参数一致nnx.get_partition_spec也能读出PartitionSpec(a, b)。Optimizer暴露三个核心属性见类 docstringstep步数OptState、txOptax 梯度变换、opt_stateOptax 优化器状态。自 Flax 0.11.0 起Optimizer不再持有model属性——源码通过自定义__getattribute__在访问model时抛出AttributeError并提示改用nnx.ModelAndOptimizer保留旧行为optimizer.py。3.2update的完整调用链updates optimizer.update(model, grads)update(model, grads, /, **kwargs)内部依次执行optimizer.py用nnx.as_pure(nnx.state(model, self.wrt))提取纯参数数组同理提取梯度数组与优化器状态数组调用self.tx.update(grad_arrays, opt_state_arrays, param_arrays, **kwargs)生成参数更新量与新优化器状态**kwargs用于支持GradientTransformationExtraArgs如optax.scale_by_backtracking_linesearch所需的grad/value/value_fn用optax.apply_updates(param_arrays, updates)应用更新通过nnx.update(model, new_params)与nnx.update(self.opt_state, ...)将结果写回模型与优化器状态self.step[...] 1步数自增返回updates——与wrt筛选后模型参数同构的更新 PyTreetests/nnx/optimizer_test.py 的test_update_returns_updates专门验证了返回值非空且结构与参数一致。_check_grads_arg_passed装饰器optimizer.py强制grads必须显式传入漏传会抛出TypeError提示Flax 0.11.0 起 update 需要同时传入 (model, grads)。此外wrt参数同样被强制校验_check_wrt_arg_passed。完整的最小示例与类 docstring 一致import jax, jax.numpy as jnp from flax import nnx import optax class Model(nnx.Module): def __init__(self, rngs): self.linear1 nnx.Linear(2, 3, rngsrngs) self.linear2 nnx.Linear(3, 4, rngsrngs) def __call__(self, x): return self.linear2(self.linear1(x)) model Model(nnx.Rngs(0)) x jax.random.normal(jax.random.key(0), (1, 2)) y jnp.ones((1, 4)) optimizer nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param) loss_fn lambda model: ((model(x) - y) ** 2).mean() grads nnx.grad(loss_fn)(model) optimizer.update(model, grads) loss_fn(model) # 2.3359997 - 2.310461损失下降3.3wrt筛选器与 LoRA 场景wrt的筛选器类型决定了哪些变量被优化。测试用例tests/nnx/optimizer_test.py覆盖了nnx.Param、nnx.LoRAParam、以及(nnx.Param, nnx.LoRAParam)组合三种情形在nnx.LoRA包装的模型上通过nnx.DiffState与nnx.grad(graphTrue)计算梯度optimizer.update后仅wrt命中的变量发生变化其余变量保持不变assert_not_equal/assert_equal分别校验。3.4ModelAndOptimizer已弃用ModelAndOptimizeroptimizer.py在Optimizer基础上额外保存self.modelupdate(grads)自动使用内部模型。其 docstring 明确标注deprecated将在未来版本移除官方建议使用Optimizer显式传model。测试中它仍被用于验证 jit 兼容性tests/nnx/optimizer_test.py 的test_jit与test_jit_linesearch覆盖了lambda f: f、nnx.compat.jit、jax.jit三种装饰器组合。四、EMA指数移动平均4.1 构造与更新规则EMAflax/nnx/training/ema.py维护模型参数的指数滑动平均影子副本常用于稳定训练并提升推理评估效果——用平滑后的参数做 inference 往往比用训练末尾的瞬时参数表现更好。构造签名nnx.EMA(params, decay, *, only..., graphNone)参数含义params任意 NNX Module / 节点其参数将被跟踪decay移动平均衰减率only筛选器指定跟踪哪些变量EMA 只跟踪nnx.Variable叶子默认匹配全部graph同Optimizer控制 graph / tree 模式update(updates)对每个被跟踪变量执行标准指数平均更新ema.pyema decay * ema (1 - decay) * update4.2apply_to生成带平滑参数的模型视图apply_to(model)是 EMA 的核心用法ema.py基于传入模型的结构构建一个追踪参数被替换为 EMA 平滑值、非追踪状态only排除的变量保持原样的新模型实例。其底层分三步graphlib.split(model)切出图结构与状态statelib.merge_state(state, self.params)将 EMA 影子参数合并进模型状态graphlib.merge(graphdef, merged_state)重建模型。由于ema_model与ema.params共享变量ema.update之后ema_model自动反映最新平滑值可直接用于评估。4.3 完整训练 评估示例源码 docstring 给出了一个完整的训练用原始模型、评估用 EMA 模型的nnx.jit示例from flax import nnx import jax, jax.numpy as jnp import optax model nnx.Linear(2, 2, rngsnnx.Rngs(0)) optimizer nnx.Optimizer(model, optax.sgd(0.1), wrtnnx.Param) ema nnx.EMA(model, decay0.9) ema_model ema.apply_to(model) def loss_fn(model, x, y): return jnp.mean((model(x) - y) ** 2) nnx.jit def train_step(model, optimizer, ema, x, y): grads nnx.grad(loss_fn)(model, x, y) optimizer.update(model, grads) ema.update(model) nnx.jit def eval_step(model, x, y): return loss_fn(model, x, y) x, y jnp.ones((1, 2)), jnp.ones((1, 2)) train_step(model, optimizer, ema, x, y) loss eval_step(ema_model, x, y)其中_to_ema_paramema.py在构造时用jax.tree.map_with_path为每个叶子创建同类型、同元数据、jnp.copy出的影子Variable若遇到非Variable叶子则抛出TypeError提示EMA only supports Variable leaves... 请用only筛选器选择 Variable 叶子。update时同样只对被only筛选器命中的变量做平均graphlib.state(updates, self.filter)未被跟踪的变量保持原样。五、三大组件组合成完整训练循环将三部分拼装即可得到 NNX 标准训练骨架import optax from flax import nnx # 1. 组装 model MyModel(rngsnnx.Rngs(0)) optimizer nnx.Optimizer(model, optax.adam(1e-3), wrtnnx.Param) ema nnx.EMA(model, decay0.999) metrics nnx.MultiMetric(accuracynnx.metrics.Accuracy(), lossnnx.metrics.Average()) nnx.jit def train_step(model, optimizer, ema, metrics, batch): grads nnx.grad(loss_fn)(model, batch) optimizer.update(model, grads) # 参数与优化器状态同步更新 ema.update(model) # 影子参数滑动平均 metrics.update(**batch) # 指标增量累计 nnx.jit def eval_step(ema_model, metrics, batch): metrics.update(**batch) for epoch in range(num_epochs): for batch in train_loader: train_step(model, optimizer, ema, metrics, batch) train_metrics metrics.compute() # 字典{accuracy: ..., loss: ...} metrics.reset() # 进入验证前清零 ema_model ema.apply_to(model) # 评估时切换到平滑参数 for batch in eval_loader: eval_step(ema_model, metrics, batch)该流程要点Optimizer.update驱动模型参数前进EMA.update持续吸收平滑参数decay越大平滑越强MultiMetric.update在每个 batch 增量累计指标epoch 末尾用compute()读数、reset()清零。三者都是Pytree/ 变量化对象可整体放进nnx.jit变换tests/nnx/optimizer_test.py 已覆盖nnx.Optimizer在各 jit 变体下的行为。六、源码与测试索引如需深入验证或继续研究可在仓库中查阅以下文件API 参考 training/index.rst、metrics.rst、optimizer.rst、ema.rst核心实现 flax/nnx/training/metrics.py、flax/nnx/training/optimizer.py、flax/nnx/training/ema.py命名空间导出 flax/nnx/init.py测试用例 tests/nnx/metrics_test.py、tests/nnx/optimizer_test.py、tests/nnx/ema_test.py七、使用注意事项小结API 状态nnx.training为 Experimental API接口可能随版本演进如 0.11.0 起Optimizer.update必须显式传入(model, grads)wrt必填空指标返回nanAverage/Accuracy/MultiMetric在未更新时compute()返回nan这是判断是否已累计过数据的约定信号wrt与only必须与梯度计算一致Optimizer的wrt要与nnx.grad的wrt匹配EMA的only只能筛选Variable叶子mask 语义Average/Accuracy/MultiMetric均支持按values * mask规则过滤MultiMetric的 mask 可传数组广播到全部指标或字典按指标名分别设置分片传播Optimizer构造时会把参数的optimizer_sharding元数据传递给对应优化器状态保证 GSPMD 场景下参数与动量/二阶矩同分片见 tests/nnx/optimizer_test.py。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价