MLX 学习率调度器(Schedulers)完全指南:从 exponential_decay 到 warmup + cosine 组合实战
2026/9/10 23:36:39 网站建设 项目流程

MLX 学习率调度器(Schedulers)完全指南:从 exponential_decay 到 warmup + cosine 组合实战

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

导读

本文围绕 MLX(Apple silicon 上的数组框架)优化器模块内置的五个学习率调度器——cosine_decayexponential_decayjoin_scheduleslinear_schedulestep_decay——展开。调度器(scheduler)是一类以训练步数为输入、输出学习率数值的可调用对象,在 MLX 中它们以"可调度参数"的形式挂在优化器状态里,每次update自动更新。读完本文,你将掌握每个调度器的数学公式、参数语义、边界行为、与优化器的底层协作机制,并能独立实现"线性 warmup + 余弦退火"等业界常见的组合调度方案。

调度器 API 由 python/mlx/optimizers/schedulers.py 实现,并通过 python/mlx/optimizers/init.py 随mlx.optimizers一起导出;对应的单元测试集中在 python/tests/test_optimizers.py 的TestSchedulers类中。

一、先理解机制:调度器如何与优化器协作

在 MLX 中,调度器不是独立于优化器之外的概念,而是优化器构造参数的一种可选形态。任何接受learning_rate的优化器(SGD、Adam、AdamW、LARS、AdaGrad、RMSProp 等)都可以直接传入一个"返回学习率的函数"。

这一机制的落点在 python/mlx/optimizers/optimizers.py 的Optimizer._maybe_schedule方法(第 143–155 行):

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

关键行为:

  • 若传入的是可调用对象(Callable),它会被登记到优化器的_schedulers字典中,并立即以当前self.step(初始为 0)求值一次,结果存入self.state[name]
  • 若传入的是普通数值,则直接转成mx.array存入状态,不走调度逻辑。

真正驱动调度器"随时间变化"的是Optimizer.apply_gradients(第 85–109 行),在每次应用梯度前执行:

for param, scheduler in self._schedulers.items(): self.state[param] = scheduler(self.step) self.state["step"] = self.step + 1

也就是说:调度器在每个训练步都被调用一次,输入是当前步数(uint64 标量数组),输出直接覆盖状态中对应参数(如learning_rate),随后步数自增。这也解释了为什么优化器的learning_rate属性(第 136–137 行)只是一个从self.state["learning_rate"]读取的普通属性——调度后的值就在那里。

从 optimizers.py 的构造器签名可以看到,learning_rate的类型注解统一为Union[float, Callable[[mx.array], mx.array]],即"数值或返回数组的函数"两种形式任选其一。其余调度参数(如weight_decay)目前不支持调度,_schedulers机制目前主要服务于学习率。

二、五种内置调度器逐一详解

以下全部函数定义在 python/mlx/optimizers/schedulers.py 中,返回值都是一个形如schedule(step) -> mx.array的函数。

1. exponential_decay:指数衰减

def exponential_decay(init: float, decay_rate: float) -> Callable: def schedule(step): return init * decay_rate**step return schedule

参数

参数类型含义
initfloat初始值(第 0 步的学习率)
decay_ratefloat每个训练步的乘法衰减因子,通常取 (0, 1) 区间

数学形式lr(step) = init * decay_rate^step。每经过一步,学习率乘以一次decay_rate,因此decay_rate = 0.9意味着每一步学习率衰减为前一步的 90%,衰减非常快。

文档示例(取自源码 docstring):

lr_schedule = optim.exponential_decay(1e-1, 0.9) optimizer = optim.SGD(learning_rate=lr_schedule) optimizer.learning_rate # array(0.1, dtype=float32) for _ in range(5): optimizer.update({}, {}) optimizer.learning_rate # array(0.06561, dtype=float32)

验证:0.1 * 0.9^5 = 0.1 * 0.59049 = 0.059049?注意示例输出是 0.06561,对应0.9^5 = 0.59049不对——实际上0.1 * 0.9^4 = 0.1 * 0.6561 = 0.06561。这里的细节是:apply_gradients先按当前 step 求值再自增,且示例中update({}, {})传入空梯度树也能推进步数,因此第 5 次循环结束时 step 恰好为 5,但调度器在求值时用的是自增前的步数,最终展示的learning_rate是 step=4 时的值。测试 test_exponential_decay 直接调用lr_schedule(10)验证0.1 * 0.99**10,与数学公式严格一致。

适用场景:需要严格按步数等比退火的场景;注意它没有自然截断,训练步数越多学习率越趋近于 0,不会像 cosine 那样在decay_steps后保持恒定。

2. step_decay:阶梯衰减

def step_decay(init: float, decay_rate: float, step_size: int) -> Callable: if step_size < 1: raise ValueError(f"step_size must be greater than 0, but got {step_size}.") def schedule(step): return init * (decay_rate ** (step // step_size)) return schedule

参数

参数类型含义
initfloat初始值
decay_ratefloat每个"阶梯"的乘法衰减因子
step_sizeint每隔多少步衰减一次;必须大于 0,否则抛出ValueError

数学形式lr(step) = init * decay_rate^(floor(step / step_size))。学习率在整数个step_size内保持不变,到达边界时骤降,形成"台阶"曲线。

文档示例

lr_schedule = optim.step_decay(1e-1, 0.9, 10) optimizer = optim.SGD(learning_rate=lr_schedule) optimizer.learning_rate # array(0.1, dtype=float32) for _ in range(21): optimizer.update({}, {}) optimizer.learning_rate # array(0.081, dtype=float32)

验证:第 0–9 步学习率为0.1,第 10–19 步为0.1 * 0.9 = 0.09,第 20 步起为0.1 * 0.9^2 = 0.081。与示例输出一致。对应测试 test_step_decay 验证step_decay(1e-1, 0.9, 1000)(2500) == 0.1 * 0.9^2(因为2500 // 1000 = 2)。

3. cosine_decay:余弦衰减(含末端保持)

def cosine_decay(init: float, decay_steps: int, end: float = 0.0) -> Callable: if decay_steps < 1: raise ValueError(f"decay_steps must be greater than 0, but got {decay_steps}.") def schedule(step): s = mx.minimum(step, decay_steps) decay = 0.5 * (1.0 + mx.cos((math.pi / decay_steps) * s)) return end + decay * (init - end) return schedule

参数

参数类型含义
initfloat初始值(第 0 步)
decay_stepsint衰减的总步数,必须大于 0
endfloat衰减终点值,默认 0.0;超过decay_steps后学习率恒等于end

数学形式lr(step) = end + 0.5 * (1 + cos(π * min(step, decay_steps) / decay_steps)) * (init - end)。核心是mx.minimum(step, decay_steps)这一钳制——一旦step超过decay_steps,输入固定为decay_steps,余弦项恒为cos(π) = -1,学习率锁定在end。这是与exponential_decay的重要差异:cosine 是"有终点"的衰减

文档示例

lr_schedule = optim.cosine_decay(1e-1, 1000) optimizer = optim.SGD(learning_rate=lr_schedule) optimizer.learning_rate # array(0.1, dtype=float32) for _ in range(5): optimizer.update({}, {}) optimizer.learning_rate # array(0.0999961, dtype=float32)

验证:step=4 时0.5 * (1 + cos(4π/1000)) * 0.1 ≈ 0.099996,与输出吻合。测试 test_cosine_decay 覆盖了两个关键点:cosine_decay(0.1, 10)(4)严格等于公式值;cosine_decay(0.1, 10, 0.05)(20)恰好等于end=0.05,证明越界后保持终点值。注意该函数使用了mx.minimummx.cos,说明求值是惰性数组计算,返回的是 MLX 数组而非 Python 浮点。

典型用法:作为完整训练周期(epochs)总步数的退火,是 Transformer/大模型训练中最常见的"余弦退火到 0 或极小值"策略。

4. linear_schedule:线性调度(warmup 首选)

def linear_schedule(init: float, end: float, steps: int) -> Callable: if steps < 1: raise ValueError(f"steps must be greater than 0, but got {steps}.") def schedule(step): step = mx.minimum(step, steps) return step * ((end - init) / steps) + init return schedule

参数

参数类型含义
initfloat起始值
endfloat终点值
stepsint线性变化的步数,必须大于 0;超过steps后恒等于end

数学形式lr(step) = init + (end - init) * min(step, steps) / steps。学习率从init线性变化到end,到达后保持end不变(同样通过mx.minimum钳制实现)。

文档示例

lr_schedule = optim.linear_schedule(0, 1e-1, 100) optimizer = optim.Adam(learning_rate=lr_schedule) optimizer.learning_rate # array(0.0, dtype=float32) for _ in range(101): optimizer.update({}, {}) optimizer.learning_rate # array(0.1, dtype=float32)

验证:第 0 步为 0,第 100 步为0 + (0.1 - 0) * 100/100 = 0.1,此后保持 0.1。最经典的用法是 warmuplinear_schedule(0, target_lr, warmup_steps)让学习率在训练初期从 0 平滑爬升,避免大学习率破坏模型初期不稳定的权重。

5. join_schedules:多段调度拼接

def join_schedules(schedules: List[Callable], boundaries: List[int]) -> Callable: if len(schedules) == 0: raise ValueError("Must provide at least 1 schedule to join.") if len(schedules) != len(boundaries) + 1: raise ValueError(f"Received {len(boundaries)} boundaries but expected {len(schedules) - 1}.") def schedule(step): output = schedules0 for boundary, schedule in zip(boundaries, schedules[1:]): output = mx.where(step < boundary, output, schedule(step - boundary)) return output return schedule

参数

参数类型含义
scheduleslist(Callable)一段或多段子调度;第i+1段接收的步数是"距第i个边界以来的步数"
boundarieslist(int)长度为len(schedules) - 1的整数列表,标记各段切换的步数

约束与校验

  • schedules不能为空;
  • boundaries数量必须恰好等于len(schedules) - 1,否则抛ValueError(测试 test_schedule_joiner 专门验证了传错边界个数会报错)。

工作机制:逐段用mx.where(step < boundary, ...)做条件选择——step小于当前边界时保留前一段的输出,否则切换到后一段,且后一段的输入是step - boundary(即"从边界起重新计数")。这个设计非常关键:拼接后的每个子调度都从自己的 0 步开始,因此你无需为每一段手动平移时间轴。

文档示例(warmup 10 步 + 余弦退火 200 步):

linear = optim.linear_schedule(0, 1e-1, steps=10) cosine = optim.cosine_decay(1e-1, 200) lr_schedule = optim.join_schedules([linear, cosine], [10]) optimizer = optim.Adam(learning_rate=lr_schedule) optimizer.learning_rate # array(0.0, dtype=float32) for _ in range(12): optimizer.update({}, {}) optimizer.learning_rate # array(0.0999938, dtype=float32)

验证:第 10 步linear(10) = 0.1,第 11 步进入余弦段cosine(11 - 10) = cosine(1) ≈ 0.5*(1+cos(π/200))*0.1 ≈ 0.0999938,与输出吻合。这正是 warmup 衔接退火的衔接点行为。

三、实战:线性 warmup + 余弦退火组合调度

join_schedules最经典的工业级用法是"先线性升温、再余弦退火"。这一场景在仓库测试 test_linear_warmup_with_cosine_decay 中有完整的可验证用例:

warmup_schedule = opt.schedulers.linear_schedule(0.0, 1e-5, 100) cosine_schedule = opt.schedulers.cosine_decay(1e-5, 100) cos_with_warmup = opt.schedulers.join_schedules( [warmup_schedule, cosine_schedule], [101] )

测试断言:

  • cos_with_warmup(0) == 0.0:warmup 从 0 起步;
  • cos_with_warmup(101) ≈ 1e-5:第 101 步到达峰值学习率并切换进余弦段;
  • optim.Adam(learning_rate=cos_with_warmup)更新 100 步后学习率逼近1e-5,再更新 100 步后学习率符合1e-5 * 0.5 * (1 + cos(π * 200 / 10))——注意此处测试的decay_steps=10,即从峰值开始的第 100 步(全局第 201 步)已远超退火窗口,学习率迅速走向 0。

一个可直接复用的完整训练骨架(结合 docs/src/python/optimizers/optimizers.rst 中 SGD/Adam 的用法与 docs/src/python/optimizers/schedulers.rst 中的调度器):

import mlx.core as mx import mlx.optimizers as optim from mlx.nn import value_and_grad total_steps, warmup_steps = 1000, 100 peak_lr = 1e-3 # 阶段一:0 -> peak_lr(100 步线性爬升) warmup = optim.linear_schedule(0.0, peak_lr, warmup_steps) # 阶段二:peak_lr -> 0(剩余 900 步余弦退火) cosine = optim.cosine_decay(peak_lr, total_steps - warmup_steps) lr_schedule = optim.join_schedules([warmup, cosine], [warmup_steps]) optimizer = optim.AdamW(learning_rate=lr_schedule) def step(model, inputs, targets): loss, grads = value_and_grad(model, loss_fn)(model, inputs, targets) optimizer.update(model, grads) return loss for i in range(total_steps): loss = step(model, inputs, targets) if i % 100 == 0: print(f"step {i}: lr={optimizer.learning_rate.item():.6f}, loss={loss.item():.4f}")

注意边界语义:warmup_steps步完成后,linear_schedule已把学习率推到peak_lr;第warmup_steps + 1步起余弦段以step - warmup_steps重新计数,从峰值开始平滑下降。

四、进阶:把调度器与 mx.compile 一起使用

调度器天然适配 MLX 的图编译。由于apply_gradients中调度器的求值结果直接写入optimizer.state,你可以把整个"更新步"编译成图,输入输出都锚定在优化器状态上,从而获得端到端的编译收益。仓库测试 test_compile_with_schedule 给出了标准写法:

lr_schedule = opt.exponential_decay(1e-1, 0.9) optimizer = opt.SGD(learning_rate=lr_schedule) @partial(mx.compile, inputs=optimizer.state, outputs=optimizer.state) def update(): optimizer.update({}, {}) for step in range(5): update() # 每一步编译执行后,调度值都与公式严格一致 assert lr_schedule(step) == optimizer.learning_rate.item()

这里的关键点是:调度器内部完全由mx.minimummx.cosmx.where等 MLX 原语构成(见 python/mlx/optimizers/schedulers.py),是纯数组运算,因此可以被mx.compile捕获进计算图并随状态一起迭代;若调度器混入 Python 分支或外部状态,则无法享受这一优化。这也正是 MLX 调度器"返回数组而非浮点"设计的深层原因。

五、边界条件与错误处理速查

五个调度器在 python/mlx/optimizers/schedulers.py 中的参数校验规则汇总:

函数校验条件异常行为
step_decaystep_size < 1ValueError("step_size must be greater than 0, but got {step_size}.")
cosine_decaydecay_steps < 1ValueError("decay_steps must be greater than 0, but got {decay_steps}.")
linear_schedulesteps < 1ValueError("steps must be greater than 0, but got {steps}.")
join_schedulesschedules为空ValueError("Must provide at least 1 schedule to join.")
join_scheduleslen(boundaries) != len(schedules) - 1ValueError(f"Received {len(boundaries)} boundaries but expected {len(schedules) - 1}.")

exponential_decay不校验参数;decay_rate的建议取值在 (0, 1],decay_rate = 1时退化为恒定学习率。

末端行为对比(决定选型的关键):

调度器越界行为说明
exponential_decay无边界,持续指数趋零无自然终点
step_decay无边界,每step_size步骤降一次阶梯型,可长期运行
cosine_decaystep > decay_steps恒为end有明确终点
linear_schedulestep > steps恒为end有明确终点,适合 warmup
join_schedules由最后一段子调度决定分段切换,各段从 0 计数

六、验证与进一步阅读

如果你安装了 MLX,可以直接运行仓库测试来验证本文全部结论:

cd python python -m pytest tests/test_optimizers.py -k "Schedulers or schedule" -v

TestSchedulers类(python/tests/test_optimizers.py)覆盖了test_decay_lr(对所有优化器遍历验证调度生效)、test_step_decaytest_exponential_decaytest_cosine_decaytest_schedule_joinertest_linear_warmup_with_cosine_decaytest_compile_with_schedule共七个用例,是理解调度器行为最直接的参考。

调度器的官方文档入口为 docs/src/python/optimizers/schedulers.rst;与之配合的优化器完整 API 见 docs/src/python/optimizers/optimizers.rst 及 docs/src/python/optimizers/common_optimizers.rst。若想深入调度器之外的训练工具链,可继续阅读 docs/src/python/optimizers/optimizer.rst 了解Optimizer基类的状态管理与tree_*工具的使用方式。

小结:MLX 的调度器设计极简但内功扎实——五个纯函数覆盖了指数、阶梯、余弦、线性与分段拼接五种核心策略;与_maybe_schedule+apply_gradients的协作让调度值以状态形式流转于每个训练步;纯 MLX 原语实现则保证了调度逻辑可以无缝融入mx.compile的编译图中。无论你是训练小型模型还是大型 Transformer,都可以用这套 API 快速拼出符合论文标准的退火曲线。

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询