FEATURED · 精选文章

MLX 优化器基类 Optimizer 深度解析:状态管理、梯度应用与学习率调度

发布时间 / 2026/9/10 12:47:44
来源 / 创域科博编辑部
栏目 / 资讯中心
MLX 优化器基类 Optimizer 深度解析:状态管理、梯度应用与学习率调度 MLX 优化器基类 Optimizer 深度解析状态管理、梯度应用与学习率调度【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlxMLX 的mlx.optimizers.Optimizer是所有内置优化器SGD、Adam、AdamW、Lion、Muon 等的基类它采用逐参数定义更新规则、对整个参数树统一应用的设计。本文以官方 API 文档 optimizer.rst 为主线结合 optimizers.py 的源码实现与 test_optimizers.py 的测试用例完整讲解state、init、update、apply_gradients等核心接口的内部机制、状态序列化方案以及如何通过调度器scheduler实现动态学习率。读完本文你将能够自行编写自定义优化器、正确管理优化器状态并将其无缝接入 MLX 的懒加载lazy evaluation训练循环。Optimizer 的设计定位逐参数更新 参数树应用从源码 docstring 可以确认Optimizer 基类的核心设计是The base class for all optimizers. It allows us to implement an optimizer on a per-parameter basis and apply it to a parameter tree.这意味着两件事逐参数per-parameter你只需实现单个参数的更新逻辑apply_single与单个参数的状态初始化逻辑init_single这两个方法在基类中均抛出NotImplementedError由子类覆写。整树应用apply to a parameter treeapply_gradients通过mlx.utils.tree_map把单参数更新函数广播到由 dict / list / tuple 嵌套构成的完整参数树上因此无论是nn.Module.trainable_parameters()返回的嵌套字典还是手工构造的{w: array}结构都能统一处理。Optimizer既可配合mlx.nn使用也可直接与纯mlx.core函数协同这一点在 optimizers.rst 中有明确说明。核心 API 逐一拆解optimizer.rst 列出了基类对外暴露的 1 个属性与 3 个方法下面结合源码逐一深入。state优化器状态字典property def state(self): The optimizers state dictionary. return self._state state.setter def state(self, state: dict): self._initialized False self._state statestate是优化器的内存保存所有需要跨 step 保留的张量。在__init__中基类把初始状态设置为self._state {step: mx.array(0, mx.uint64)}即每个优化器至少携带一个uint64类型的step计数器。step还有独立的只读属性property def step(self): return self.state[step]不同的内置优化器会在state中追加各自的动量 / 二阶矩累加器例如对应 optimizers.py 中各子类的init_single优化器state 中附加的键含义SGD / RMSprop / Adagradv动量SGD或梯度平方累加器AdaDeltav,u梯度平方与更新量平方的运行均值Adam / AdamW / Adamaxm,v一阶矩、二阶矩Adamax 为无穷范数Lionm动量Muonv动量Adafactorexp_avg_sq_row,exp_avg_sq_col≥2 维或exp_avg_sq1 维beta_1非空时还有exp_avg因子分解的二阶矩估计init 与 init_single显式初始化状态def init(self, parameters: dict): ... update_state(parameters, self._state) tree_map(lambda p, s: s or self.init_single(p, s), parameters, self._state) self._initialized Trueinit(parameters)递归地把参数树的结构与状态字典对齐新出现的参数键补上空 dict然后对每个参数调用子类实现的init_single(parameter, state)完成状态张量的分配通常用mx.zeros_like(parameter)。根据文档说明显式调用init是可选的优化器在第一次apply_gradients时会自动初始化。但在以下场景中显式初始化很有价值——你希望在第一次update之前就能访问Optimizer.state。官方示例 optimizer optim.SGD(learning_rate1e-1, momentum0.9) model nn.Linear(2, 2) optimizer.init(model.trainable_parameters()) optimizer.state.keys() dict_keys([step, learning_rate, weight, bias])注意一个细节init接受的参数是参数树而apply_gradients在首次调用时是用梯度树去隐式初始化状态的见下节源码。测试 test_optimizers.py 中的test_sgd同时验证了两种路径显式init(params)后state[v]全为 0隐式初始化下经过一次全 1 梯度的更新后state[v]恰好等于梯度本身。apply_gradients 与 apply_single梯度应用apply_gradients是状态机的中枢其源码逻辑依次为def apply_gradients(self, gradients: dict, parameters: dict): if not self._initialized: self.init(gradients) # Update any scheduled variables for param, scheduler in self._schedulers.items(): self.state[param] scheduler(self.step) # Increment the step self.state[step] self.step 1 # Apply the update return tree_map(self.apply_single, gradients, parameters, self.state)值得注意的要点首次调用自动初始化若未显式init这里用gradients树完成状态初始化调度变量在每步更新所有注册进_schedulers的参数典型如learning_rate都会以当前step为输入重新求值并写回statestep 先自增再使用apply_single内部读取self.step时已经是t1这保证 Adam 的bias_correction、Adafactor 的step**decay_rate等公式按第 t 步语义工作返回新参数树tree_map(self.apply_single, gradients, parameters, self.state)的返回值结构与gradients一致parameters允许是梯度的超集此时返回值仅覆盖梯度对应的键。apply_single(gradient, parameter, state)由子类实现是真正的数学更新。例如 SGD 的动量版本对应v self.momentum * state.get(v) ... state[v] v return parameter - self.learning_rate.astype(gradient.dtype) * update测试test_types_conserved验证了一个容易被忽略的约束所有优化器的更新都必须保持参数 dtype如float16参数更新后仍为float16因此各实现都会用self.learning_rate.astype(gradient.dtype)把学习率转换到梯度精度后再参与运算。update一行代码完成模型更新def update(self, model: Module, gradients: dict): model.update(self.apply_gradients(gradients, model))update是日常训练循环最常用的入口把apply_gradients(gradients, model)得到的更新后参数树直接通过nn.Module.update写回模型。其等价展开式即为文档所注明的model.update(opt.apply_gradients(grads, model))。learning_rate 属性与 _maybe_schedule可调度参数property def learning_rate(self): return self.state[learning_rate] learning_rate.setter def learning_rate(self, learning_rate: Union[float, mx.array]): self.state[learning_rate] mx.array(learning_rate)学习率被存进state因此可以随时读写甚至在编译后的函数之间修改见测试test_update_lr_compiled。子类构造时通过_maybe_schedule注册学习率def _maybe_schedule(self, name, param): if isinstance(param, Callable): self._schedulers[name] param parameter param(self.step) else: parameter mx.array(param) self.state[name] parameter由此可以理解 optimizers.rst 中保存与加载一节的判断规则——凡是能传入 callable即可被调度的参数都会进入 optimizer state学习率如此而 Adam 的betas、eps这类固定超参不会入 state因此序列化状态时它们不会随 checkpoint 保存。状态的生命周期懒执行与 mx.evalMLX 采用懒执行lazy evaluation。调用optimizer.update(model, grads)只是把更新运算加入计算图此时还没有任何数值计算发生。必须在合适的时机调用mx.eval同时求值模型参数与优化器状态。官方在 optimizers.rst 给出的标准训练循环骨架如下model MLP(num_layers, train_images.shape[-1], hidden_dim, num_classes) mx.eval(model.parameters()) loss_and_grad_fn nn.value_and_grad(model, loss_fn) optimizer optim.SGD(learning_ratelearning_rate) for e in range(num_epochs): for X, y in batch_iterate(batch_size, train_images, train_labels): loss, grads loss_and_grad_fn(model, X, y) # Update the model with the gradients. So far no computation has happened. optimizer.update(model, grads) # Compute the new parameters but also the optimizer state. mx.eval(model.parameters(), optimizer.state)要点mx.eval的入参同时包含model.parameters()与optimizer.state。如果只 eval 参数而忽略optimizer.state动量 / 二阶矩等状态张量仍是未求值的惰性节点会随着训练步数增加而累积出极长的计算图导致内存膨胀与性能劣化。学习率调度器把 callable 交给 learning_rate调度器是learning_rate接受 callable 的直接应用。mlx.optimizers.schedulers模块schedulers.py提供了 5 个工厂函数其签名与行为如下调度器签名行为exponential_decay(init, decay_rate)init * decay_rate**stepstep_decay(init, decay_rate, step_size)每step_size步乘一次decay_ratecosine_decay(init, decay_steps, end0.0)余弦衰减step decay_steps后恒定于endlinear_schedule(init, end, steps)线性插值超过steps后恒为endjoin_schedules(schedules, boundaries)按边界拼接多个调度后一段的步数从边界处重新计数用法示例来自 schedulers.py 的 docstringlr_schedule optim.cosine_decay(1e-1, 1000) optimizer optim.SGD(learning_ratelr_schedule)调度器本质是Callable[[mx.array], mx.array]在每次apply_gradients中由_schedulers以当前step求值并写回state[learning_rate]因此optimizer.learning_rate属性在每一步后都会反映最新值。测试test_decay_lr用step_decay(1e-1, 0.9, 1)遍历全部优化器验证了这一点test_compile_with_schedule还证明调度器可以与mx.compile协同把optimizer.state声明为编译函数的inputs/outputs即可在编译图内更新调度值。join_schedules是构造warmup 衰减组合的常用手段warmup optim.linear_schedule(0.0, 1e-5, 100) cosine optim.cosine_decay(1e-5, 100) lr_schedule optim.join_schedules([warmup, cosine], [101]) optimizer optim.Adam(learning_ratelr_schedule)对应测试test_linear_warmup_with_cosine_decay精确校验了该组合在各 step 上的取值。MultiOptimizer一个模型、多个优化器MultiOptimizer允许按参数路径 / 形状对模型参数做分流让不同参数使用不同优化器optimizers.py。构造规则optimizers优化器列表filters谓词列表数量必须为len(optimizers) - 1每个谓词接收(name, weight)并返回布尔值列表中的最后一个优化器是兜底fallback不配谓词。optimizer opt.MultiOptimizer( [opt.Adam(learning_rate0.001), opt.SGD(learning_rate0.1)], [lambda name, weight: weight.ndim 1], )实现上_split_dictionary用tree_flatten展平梯度树按filters依次匹配并归入对应优化器再tree_unflatten还原。其state结构为{states: [各子优化器 state, ...]}learning_rate属性代理到第一个优化器setter 则广播到全部子优化器。测试test_multi_optimizer验证了分流正确性高维权重归 Adam、其余归 SGDtest_multi_optimizer_with_parameterless_layers则验证了含无参数层如 ReLU的模型也能正常更新。保存与加载优化器状态MLX 官方的序列化方案是保存optimizer.state加载时重建优化器再整体设置state。完整示例optimizers.rstimport mlx.core as mx from mlx.utils import tree_flatten, tree_unflatten import mlx.optimizers as optim optimizer optim.Adam(learning_rate1e-2) # Perform some updates with the optimizer model {w: mx.zeros((5, 5))} grads {w: mx.ones((5, 5))} optimizer.update(model, grads) # Save the state state tree_flatten(optimizer.state, destination{}) mx.save_safetensors(optimizer.safetensors, state) # Later on, for example when loading from a checkpoint, # recreate the optimizer and load the state optimizer optim.Adam(learning_rate1e-2) state tree_unflatten(mx.load(optimizer.safetensors)) optimizer.state state这里的关键是optimizer.state state会走 state 的 setter把_initialized重置为False、替换内部_state——因此加载后的优化器可以直接继续update而无需再次init。测试test_init_from_state走的就是这条往返路径tree_flatten→ 重建优化器 →tree_unflatten→ 赋值 state → 正常 update。务必牢记 optimizers.rst 中特别说明的边界并非所有配置参数都在 state 中。例如 Adam 的betas、eps不随 state 保存判断规则是能被调度的参数才进入 state学习率在、betas/eps不在。因此恢复训练时必须用与保存时一致的超参重新构造优化器再覆盖其 state。梯度裁剪clip_grad_normmlx.optimizers.clip_grad_norm是训练稳定性工具按全局范数裁剪梯度。其源码逻辑optimizers.pynorm_squared tree_reduce(lambda acc, g: acc g.square().sum(), grads, 0.0) total_norm mx.sqrt(norm_squared) normalizer mx.minimum(max_norm / (total_norm 1e-6), 1.0) clipped_grads tree_map(lambda g: g * normalizer, grads) return clipped_grads, total_norm入参grads为梯度字典max_norm为允许的全局范数上限max_norm 0会抛ValueError返回(裁剪后梯度, 原始总范数)当梯度范数未超限时normalizer为 1.0梯度原样返回。测试test_clip_grad_norm验证了三种情形小梯度不裁剪结果与输入逐元素相等、大梯度裁剪后全局范数收敛到max_norm、以及负max_norm报错。典型用法是在optimizer.update之前对梯度做裁剪grads, _ clip_grad_norm(grads, max_norm1.0) optimizer.update(model, grads)常见陷阱与最佳实践附测试佐证综合源码与测试以下实践值得固化始终 eval 优化器状态懒执行下训练循环必须mx.eval(model.parameters(), optimizer.state)否则状态张量成为悬空惰性节点。不要手工修改累积器Optimizer.update内部依赖state[step]的自增语义apply_gradients在调度与 Adam 偏置校正中都会读取step绕过update/apply_gradients直接改 state 会破坏一致性。epsilon 必须 0RMSprop / Adagrad / AdaDelta这些优化器的累加器从零开始eps 0时零梯度会算出0/0产生 NaN 参数因此源码显式校验并抛出ValueError测试test_epsilon_validation。保持 dtype优化器会按梯度 dtype 转换学习率但不会隐式改参数类型混合精度场景下参数与梯度 dtype 不一致时需自行保证。weight decay 不改写输入测试test_weight_decay_keeps_inputs_unchanged确认SGD / Adafactor 的 weight decay 通过构造新张量实现不会原地写入调用方传入的grads或params。与mx.compile配合测试test_compiled_optimizer展示了两种编译模式——把optim.state作为编译函数的输入输出纯函数式或用partial(mx.compile, inputs[model.state, optim.state], outputs[...])保留原地风格impure两者结果与未编译版一致可用于大幅降低每步 Python 开销。训练中断恢复务必用 safetensors 持久化optimizer.state并保证重建时传入与 checkpoint 一致的超参betas、eps、weight_decay等不进 state。如需查阅全部内置优化器的参数签名与数学公式可继续阅读 common_optimizers.rst涵盖 SGD、RMSprop、Adagrad、Adafactor、AdaDelta、Adam、AdamW、Adamax、Lion、MultiOptimizer、Muon与 schedulers.rst这些类的实现细节统一沉淀在 optimizers.py 与 schedulers.py 中test_optimizers.py 则为每个行为提供了可复现的数值验证。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻