Flax 优化器 API 演进:基于 FLIP 1009 从 flax.optim 全面迁移到 Optax 梯度变换体系 📅 发布时间:2026/9/17 22:20:42 👁 浏览次数: Flax 优化器 API 演进基于 FLIP 1009 从 flax.optim 全面迁移到 Optax 梯度变换体系【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本指南以 Flax 官方提案 FLIP 1009《Optimizer API》 为骨架系统讲解 Flax 如何将内置的flax.optim优化器 API 替换为 DeepMind 的 Optax 库覆盖梯度变换gradient transformation组合、Optax 训练步、多优化器Multi Optimizer与TrainState四个核心主题并完整对照旧 API 的Optimizer/OptimizerDef模式。读完本文你将掌握用 Optax 组合任意优化器、编写简洁的 Linen 训练循环、通过optax.masked()实现参数分域优化以及基于TrainState组织可断点续训的训练状态的全部实战方案。一、提案背景为什么 Flax 要用 Optax 取代自带优化器FLIP 1009Start Date: 2021-02-08提出将当时的flax.optimAPI文档中称为 previous API替换为 Optax。动机有三点今天依然成立旧 API 模式复杂flax.optim采用「Optimizer数据类 OptimizerDef定义」的模式即由一个Optimizerdataclass 持有 target 参数的 pytree并由OptimizerDef定义如何更新优化器状态、超参数和目标参数。对于实现一个简单的优化器而言这种模式相对复杂而在典型的 Linen 训练步中尤其是涉及可变状态集合mutable时代码相当冗长。内置优化器覆盖不足flax.optim虽然包含若干优化器但列表远非详尽理想做法是使用来自独立 PyPI 包的 JAX 优化器避免 Flax 自身维护一套优化器实现。复用生态DeepMind 已有专门库 Optax实现了大量优化器并提供了从可复用的梯度变换组合出新优化器的框架。从当前仓库源码结构看flax/目录下已不存在flax/optim/模块说明该提案落地后旧 API 已被移除如今 Flax 官方示例全部基于 Optax 编写见 examples/mnist/train.py、examples/imagenet/train.py、examples/wmt/train.py 等这一演进还进一步延续到了新的 NNX 接口中。二、核心概念梯度变换Gradient TransformationOptax 虽然提供了预定义的优化器如optax.adam、带动量的optax.sgd但它本质上是一个梯度变换库。实例化优化器的惯用方式是组合若干梯度变换。例如要复刻旧 API 示例中的动量优化器无 Nesterov 动量可以写成import optax tx optax.chain( optax.trace(decay0.9, nesterovFalse), optax.scale_by_schedule(lambda step: -get_learning_rate(step)), )几点关键说明上述变换与 Optimizer 和 OptimizerDef 一节中定义的 Momentum 优化器等价旧 API 的beta参数对应optax.trace()的decay参数学习率则由第二条链式变换应用注意optax.scale_by_schedule需要取负号因为我们习惯把梯度更新加到参数上即param - lr * update。超参数如decay、nesterov只存在于返回GradientTransformation的高阶函数的内部作用域中。梯度变换目前定义为由init()和update()两个函数组成的NamedTuple原则上这一模式可扩展为同时存储超参数这一点可在 Optax 仓库中进一步讨论。get_learning_rate(step)可以返回依赖步数的学习率从而在定义变换时就完成学习率调度成为旧训练步中学习率函数的即插即用替代注意符号取反。除手工写调度函数外Optax 还提供inject_hyperparams()来对任意超参数做调度。这一点已在当前仓库中得到充分印证ImageNet 示例 examples/imagenet/train.py 使用optax.linear_schedule、optax.cosine_decay_schedule与optax.join_schedules拼出「线性预热 余弦退火」的两段式调度再整体作为optax.sgd(learning_ratelearning_rate_fn, ...)的学习率参数传入WMT 示例 examples/wmt/train.py 则以optax.adamw(learning_ratelearning_rate_fn, b10.9, b20.98, eps1e-9, weight_decay...)使用调度后的学习率。三、Optax 训练步把整个 variables 作为输入输出FLIP 给出的 Optax 训练步完整实现如下functools.partial(jax.jit, static_argnums(4, 5)) def train_step(opt_state, variables, inputs, labels, apply_fn, tx_update_fn): def loss_fn(params): logits, new_model_state apply_fn( {**variables, params: params}, inputs, mutable[batch_stats]) loss xent_loss(logits, labels) return loss, new_model_state variables, params variables.pop(params) (loss, new_model_state), grads jax.value_and_grad(loss_fn, has_auxTrue)( params) updates, new_opt_state tx_update_fn(grads, opt_state, params) new_params optax.apply_updates(params, updates) new_variables {**variables, **new_model_state, params: new_params} return new_opt_state, new_variables, loss opt_state tx.init(variables[params]) for batch in ds.as_numpy_iterator(): opt_state, variables, loss train_step( opt_state, variables, batch[image], batch[label], model.apply, tx.update) print(loss)要点解析tx.update()只变换梯度它返回(updates, new_opt_state)必须再调用optax.apply_updates(params, updates)把变换后的梯度应用到参数上。两步分离是 Optax 与旧Optimizer.apply_gradient()最大的心智差异。与旧 API 相比现在可以把整个variables含params作为train_step()的输入与输出无需像旧版那样把参数包进optimizer.target。步内仍然要把params从variables中拆分出来因为我们只对params求梯度而不是对全部variables。只要 Optax 变换在各自状态中暴露信息就可以记录内部优化器状态如学习率。例如optax.scale_by_schedule()当前只暴露opt_state.count但很容易扩展为同时暴露step_size对其他随时间变化的内部状态同理。对照当前仓库的 ImageNet 训练步 examples/imagenet/train.py可以看到同样的「loss_fn内state.apply_fn({params: params, batch_stats: state.batch_stats}, batch[image], mutable[batch_stats])→jax.value_and_grad(loss_fn)→state.apply_gradients(gradsgrads, ...)」结构已被固化为标准范式且额外加入了 L2 权重惩罚等自定义逻辑。四、多优化器从 ModelParamTraversal 到 optax.masked旧 API 通过flax.optim.MultiOptimizer对参数树的不同部分使用不同优化器biases_traversal flax.optim.ModelParamTraversal( lambda path, _: path.endswith(/bias)) not_biases_traversal flax.optim.ModelParamTraversal( lambda path, _: not path.endswith(/bias)) optimizer_def flax.optim.MultiOptimizer( (biases_traversal, flax.optim.GradientDescent(learning_rate0.1)), (not_biases_traversal, flax.optim.GradientDescent(learning_rate0.05)), )这里的思路是先用一个 traversal遍历器基于参数路径模块作用域与变量名的拼接选择参数再为每条独立遍历绑定不同优化器。Optax 提供了optax.masked()来指定只作用于梯度子集的变换配合flax.traverse_util.flatten_dict/unflatten_dict见 flax/traverse_util.py可以表达同样的语义def flattened_traversal(fn): def mask(data): flat traverse_util.flatten_dict(data) return traverse_util.unflatten_dict({k: fn(k, v) for k, v in flat.items()}) return mask tx optax.chain( optax.masked(optax.sgd(learning_rate0.1), maskflattened_traversal(lambda path, _: path[-1] bias)), optax.masked(optax.sgd(learning_rate0.05), maskflattened_traversal(lambda path, _: path[-1] ! bias)), )mask返回一个与data结构相同的布尔掩码 pytreeTrue的位置应用对应的梯度变换False的位置梯度原样传递等价于恒等变换。这种「按路径掩码 链式组合」的写法比旧的MultiOptimizer更通用也天然与optax.chain的组合模型兼容。从仓库现状看多优化器组合也大量出现在真实示例中例如 examples/sst2/train.py 用optax.chain(optax.sgd(learning_rate..., momentum...), optax.add_decayed_weights(weight_decay...))实现「SGD 权重衰减」的组合。五、TrainState统一封装优化器状态与参数更新在 Flax 中通常用一个可断点保存的TrainState对象在训练循环中来回传递。它简化了上面的 Optax 训练步减少参数个数并去掉static_argnums。FLIP 提出了flax.training.train_state.TrainState的设计当前仓库的实现位于 flax/training/train_state.py# 对应于当前 flax/training/train_state.py 中的实现 class TrainState(struct.PyTreeNode): step: int | jax.Array apply_fn: Callable struct.field(pytree_nodeFalse) params: core.FrozenDict[str, Any] struct.field(pytree_nodeTrue) tx: optax.GradientTransformation struct.field(pytree_nodeFalse) opt_state: optax.OptState struct.field(pytree_nodeTrue) def apply_gradients(self, *, grads, **kwargs): updates, new_opt_state self.tx.update( grads_with_opt, self.opt_state, params_with_opt ) new_params_with_opt optax.apply_updates(params_with_opt, updates) return self.replace( stepself.step 1, paramsnew_params, opt_statenew_opt_state, **kwargs, ) classmethod def create(cls, *, apply_fn, params, tx, **kwargs): opt_state tx.init(params) return cls( step0, apply_fnapply_fn, paramsparams, txtx, opt_stateopt_state, **kwargs, )要点apply_gradients()内部依次调用self.tx.update(grads, self.opt_state, self.params)与optax.apply_updates()一次性完成步数递增、参数更新与优化器状态更新并可通过**kwargs顺带replace其他字段。create()用tx.init(params)初始化优化器状态step从 0 开始。apply_fn与tx被标记为pytree_nodeFalse保证TrainState本身是一个可被jax.jit、jax.grad处理的 pytree同时不会被无谓地当作张量树遍历。用户可以直接继承该数据类添加字段例如可变模型状态batch_statsfrom flax.training import train_state class TrainState(train_state.TrainState): batch_stats: flax.core.FrozenDict[str, Any]这正是当前 ImageNet 示例 examples/imagenet/train.py 的做法——它在官方TrainState之上追加了batch_stats与dynamic_scale用于混合精度训练的动态缩放两个字段。5.1 带可变状态mutable的训练步有了扩展后的TrainState带batch_stats的训练步变为jax.jit def train_step(state, inputs, labels): def loss_fn(params): outputs, new_model_state state.apply_fn( {params: params, batch_stats: state.batch_stats}, inputs, mutable[batch_stats]) loss xent_loss(outputs, labels) return loss, new_model_state (loss, new_model_state), grads jax.value_and_grad( loss_fn, has_auxTrue)(state.params) new_state state.apply_gradients( gradsgrads, batch_statsnew_model_state[batch_stats], ) return new_state, loss state TrainState.create( apply_fnmodel.apply, paramsvariables[params], txtx, batch_statsvariables[batch_stats], ) for batch in ds.as_numpy_iterator(): state, loss train_step(state, batch[image], batch[label])注意state.apply_gradients(gradsgrads, batch_stats...)通过**kwargs把更新后的 BatchNorm 统计量一并写回这就是「TrainState 简化参数列表」的体现。5.2 无可变状态的训练步当模型没有可变状态时训练步进一步简化为jax.jit def train_step(state, inputs, labels): def loss_fn(params): outputs state.apply_fn({params: params}, inputs) loss xent_loss(outputs, labels) return loss loss, grads jax.value_and_grad(loss_fn)(state.params) new_state state.update(gradsgrads) return new_state, loss state flax.training.TrainState.create( apply_fnmodel.apply, paramsvariables[params], txtx, ) for batch in ds.as_numpy_iterator(): state, loss train_step(state, batch[image], batch[label])5.3 使用边界与断点保存FLIP 特别提醒在 Flax 训练循环中「每个 step 后用新状态更新TrainState数据类」是常见模式。flax.training.train_state中的简单方案可以扩展附加数据但不支持高级用例例如多个不同模型和/或多个优化器此时应 fork 该数据类并按其需求重新实现。与旧 API 的Optimizer抽象不同TrainState现在直接包含.params无需经由.optimizer中转。TrainState本身是标准 pytree因此可以直接交给 flax/training/checkpoints.py 中的save_checkpoint/restore_checkpoint进行断点保存与恢复——这正是 ImageNet 示例中save_checkpoint(state, workdir)的做法。六、旧 API 详解Optimizer 与 OptimizerDef为便于理解迁移前后的差异FLIP 回顾了旧 API 的实现方式。优化器通过继承OptimizerDef实现以下为 FLIP 中的flax/optim/momentum.py示意代码flax.struct.dataclass class _MomentumHyperParams: learning_rate: jnp.ndarray beta: jnp.ndarray flax.struct.dataclass class _MomentumParamState: momentum: np.ndarray class Momentum(flax.optim.OptimizerDef): def __init__(self, learning_rateNone, beta0.9): super().__init__( _MomentumHyperParams(learning_rate, beta) ) def init_param_state(self, param): return _MomentumParamState(jnp.zeros_like(param)) def apply_param_gradient(self, step, hyper_params, param, state, grad): del step assert hyper_params.learning_rate is not None new_momentum state.momentum * hyper_params.beta grad new_params param - hyper_params.learning_rate * new_momentum return new_params, _MomentumParamState(new_momentum)关键机制调用链用户代码调用Optimizer.apply_gradient()→ 内部调用OptimizerDef.apply_gradient()连同其他逻辑→ 最终调用由子类实现的OptimizerDef.apply_param_gradient()。init_param_state()与apply_param_gradient()会对 params/grads pytree 的每个叶子调用一次因此可以直接书写逐叶计算无需jax.tree_util.tree_map()。该接口定义于 pre-Linen 时代没有考虑variables中params与其他 collection 的区分。原 API 优雅之处在于只需传递一个 optimizer 对象——它包含了参数、优化器状态、优化器超参数以及对OptimizerDef的引用。6.1 旧训练步写法先由定义和参数 pytree 构造优化器optimizer_def flax.optim.Momentum(learning_rate0.1, beta0.9) optimizer optimizer_def.create(variables[params])然后在训练步中优化目标变量假设存在单个非 params 集合batch_statsdef make_train_step(apply_fn): jax.jit def train_step(optimizer, batch_stats, inputs, labels): def loss_fn(params): variables {params: params, batch_stats: batch_stats} logits, new_model_state apply_fn( variables, inputs, mutable[batch_stats]) loss xent_loss(logits, labels) return loss, new_model_state[batch_stats] (loss, new_batch_stats), grad jax.value_and_grad(loss_fn, has_auxTrue)( optimizer.target) lr get_learning_rate(step) new_optimizer optimizer.apply_gradient(grad, learning_ratelr) return new_optimizer, new_batch_stats, loss return train_step batch_stats variables[batch_stats] train_step make_train_step(model.apply) for step, batch in enumerate(ds) optimizer, batch_stats, loss train_step( optimizer, batch_stats, batch[image], batch[label])注意optimizer.apply_gradient()可以接受额外参数来更新超参数例如这里从独立的get_learning_rate()函数传入学习率——这种「学习率作为超参动态传入」的职责在 Optax 中被转移到了optax.scale_by_schedule等调度变换内部。对比可见旧 API 需要显式地「捞」出optimizer.target求梯度、再把学习率作为额外参数传入而 Optax 方案中优化器状态、参数与模型状态均显式参与训练步逻辑更直白。七、更新计划与仓库现状验证FLIP 给出了七步落地计划完成 FLIP 讨论定稿在 Optax 中添加等价性测试保证现有flax.optim优化器与对应optax优化器返回完全一致的值更新示例使用 Optax并验证在相同计算成本下达到相同的最终性能将缺失的优化器移植到 Optax例如 Adafactor并验证以上两点更新全部文档README、Flax guided tour、HOWTO 等只谈论 Optax 优化器编写从flax.optim迁移到 Optax 的过渡指南同时指向 Optax 的等价性测试与更新示例的 PR将flax.optim中的优化器标记为弃用。FLIP 同时列出了当时各示例的优化器对照表ExampleFlaxOptaxCommentsimagenetoptim.Momentumoptax.sgdDynamicScale 可保持不变mnistoptim.Momentumoptax.sgdnlp_seqoptim.Adamoptax.adamwpixelcnnoptim.Adamoptax.adamppooptim.Adamoptax.adamseq2seqoptim.Adamoptax.adamvaeoptim.Adamoptax.adamwmtoptim.Adamoptax.adamwFlax 的 Adam 实现有一个可选的权重衰减参数而 Optax 中带/不带权重衰减的 Adam 是两个不同别名optax.adamw/optax.adam。从当前仓库可以验证该计划已基本全部落地flax/目录下已不存在flax/optim/模块旧 API 已被移除全部官方示例都使用 Optaxexamples/mnist/train.py 用optax.sgd(learning_rate, momentum)NNX 接口下同样由nnx.Optimizer(model, optax.sgd(...), wrtnnx.Param)包裹examples/imagenet/train.py 用optax.sgd(..., momentum..., nesterovTrue)examples/wmt/train.py 用optax.adamw(...)examples/sst2/train.py 用optax.chain(...)组合flax/training/train_state.py 中的TrainState以optax.GradientTransformation为tx字段成为 Flax含 NNX 之外的传统 Linen 路径训练循环的公共底座。八、附录FLIP 全部代码片段的运行环境FLIP 附录提供了可直接运行上述代码片段的 Setup Code基于 MNIST 数据集的 Linen 模型import functools from typing import Callable, Sequence import jax import jax.numpy as jnp import flax import flax.linen as nn import tensorflow as tf import tensorflow_datasets as tfds def pp(features): return { image: tf.cast(features[image], tf.float32) / 255 - 0.5, label: features[label], } class Model(nn.Module): nn.compact def __call__(self, inputs): x inputs.reshape([inputs.shape[0], -1]) x nn.normalization.BatchNorm(True)(x) x nn.Dense(10)(x) x nn.log_softmax(x) return x def onehot(labels, num_classes, on_value1.0, off_value0.0): x (labels[..., None] jnp.arange(num_classes)[None]) x jax.lax.select( x, jnp.full(x.shape, on_value), jnp.full(x.shape, off_value)) return x.astype(jnp.float32) def xent_loss(logits, labels): return -jnp.sum( onehot(labels, num_classes10) * logits) / labels.size def get_learning_rate(step): return 0.1 model Model() rng jax.random.key(0) ds tfds.load(mnist)[train].take(160).map(pp).batch(16) batch next(iter(ds)) variables model.init(rng, jnp.array(batch[image][:1])) jax.tree_util.tree_map(jnp.shape, variables)该模型包含一个可变的batch_stats集合BatchNorm因此恰好能够同时演示「带 mutable 状态」与「无可变状态」两种训练步写法xent_loss采用 one-hot 与 log-softmax 输出计算交叉熵get_learning_rate返回恒定学习率 0.1。结语FLIP 1009 所确立的「以 Optax 梯度变换为核心、以TrainState为训练循环底座」的设计已成为 Flax 的主流实践参数更新从「封装在 Optimizer 对象内」转变为「显式tx.updateoptax.apply_updates」多优化器从MultiOptimizer走向optax.masked学习率调度从函数参数走向调度变换最终全部收敛到可直接断点保存的TrainStatepytree。建议读者结合本文与 flax/training/train_state.py、examples/imagenet/train.py、examples/wmt/train.py 三份源码对照阅读即可完整掌握从旧 API 迁移到 Optax 的全部细节。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考