ARTICLE DETAIL

资讯详情

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

MXNet 学习率调度器(lr_scheduler)完全指南:从阶梯衰减到余弦退火的实现与实战

MXNet 学习率调度器(lr_scheduler)完全指南:从阶梯衰减到余弦退火的实现与实战 MXNet 学习率调度器lr_scheduler完全指南从阶梯衰减到余弦退火的实现与实战【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址: https://gitcode.com/gh_mirrors/mxnet1/mxnet导读学习率learning rate是深度学习中最重要的超参数之一其调度策略直接决定模型能否稳定收敛、能否跳出局部最优。MXNet 在 python/mxnet/lr_scheduler.py 中提供了一套完整的mxnet.lr_scheduler模块内置阶梯衰减、多项式衰减、余弦退火三种主流策略并统一支持 warmup 预热机制。本文将以该模块源码为主体逐一讲解每个调度器的数学公式、参数含义、与 Optimizer / Gluon Trainer 的集成方式并结合仓库测试用例给出可复制的实战示例帮助你按训练阶段精细控制学习率。一、模块总览调度器是什么在 MXNet 中lr_scheduler学习率调度器是一个根据已执行的更新次数num_update返回新学习率的可调用对象。模块顶部注释python/mxnet/lr_scheduler.py只有一句话Scheduling learning rate但其实现涵盖了四个具体类调度器类衰减策略适用场景FactorScheduler每 n 步乘以固定因子简单固定间隔衰减MultiFactorScheduler在指定的多个步数点乘以因子里程碑式阶梯衰减最常用PolyScheduler多项式函数衰减长周期训练、渐进收敛CosineScheduler余弦函数平滑衰减现代 CNN 训练的平滑退火此外还有基类LRScheduler它本身不实现衰减逻辑__call__直接抛出NotImplementedError但统一实现了warmup 预热机制供所有子类复用。二、基类 LRScheduler统一预热机制warmup所有调度器都继承自LRSchedulerpython/mxnet/lr_scheduler.py其构造函数统一接受四个参数LRScheduler(base_lr0.01, warmup_steps0, warmup_begin_lr0, warmup_modelinear)参数说明base_lr : float初始学习率默认0.01。这也是 warmup 结束后衰减正式开始的起点源码中记为warmup_final_lr。warmup_steps : int预热步数默认0即不预热。必须是int类型源码用assert isinstance(warmup_steps, int)强制校验且不能为负数。warmup_begin_lr : float预热起始学习率默认0。源码规定warmup_begin_lr必须小于等于base_lr否则抛出ValueError(Base lr has to be higher than warmup_begin_lr)。warmup_mode : str预热模式仅支持linear与constant两种其他值会抛出ValueError。get_warmup_lr(num_update)是预热的核心实现python/mxnet/lr_scheduler.pylinear 模式学习率从warmup_begin_lr出发按固定增量线性爬升到base_lrlr warmup_begin_lr (base_lr - warmup_begin_lr) * num_update / warmup_stepsconstant 模式在整个预热期内保持warmup_begin_lr不变直到预热结束才跳到base_lr。预热机制的价值在于训练初期模型参数随机初始化梯度较大且方向不稳定若直接用较大学习率容易震荡甚至发散先用小学习率热身再过渡到正式学习率能显著提升训练的稳定性。基类的__call__(num_update)中明确规定了num_update的语义它是所有权重中被更新次数的最大值num_update max([k_i for all i])即optimizer.update(i, weight_i)被调用次数的上界。各子类基于此约定判断当前处于哪个衰减阶段。三、FactorScheduler按固定步长等比衰减FactorSchedulerpython/mxnet/lr_scheduler.py是最基础的阶梯调度器其数学公式为lr base_lr * factor ^ floor(num_update / step)构造参数step : int每隔多少个更新步改变一次学习率必须 ≥ 1否则抛ValueError。factor : float衰减因子默认1不衰减。源码强制要求factor 1.0否则抛ValueError(Factor must be no more than 1 to make lr reduce)。stop_factor_lr : float学习率下限默认1e-8。当衰减后的学习率低于该值时学习率被钳制在stop_factor_lr并记录日志now learning rate arrived at ... will not change in the future。其余base_lr、warmup_*参数继承自基类。实现要点__call__python/mxnet/lr_scheduler.py若num_update warmup_steps直接返回get_warmup_lr(num_update)使用while而非if处理衰减——源码注释明确说明这是为了支持通过load_epoch恢复断点续训的场景当一次调用跨越多个step区间时需要连续多次乘以factor才能回到正确状态。实战示例与 tests/python/unittest/test_optimizer.py 中的test_factor_scheduler一致import mxnet as mx sched mx.lr_scheduler.FactorScheduler( step100, factor0.1, stop_factor_lr1e-4, base_lr1, warmup_steps20, warmup_begin_lr0.1, warmup_modeconstant ) print(sched(0)) # 0.1 预热期constant 模式保持 0.1 print(sched(10)) # 0.1 预热期内 print(sched(21)) # 1.0 预热结束回到 base_lr print(sched(101)) # 0.1 第一次衰减1 * 0.1 print(sched(201)) # 0.01 第二次衰减 print(sched(1000)) # 0.0001被 stop_factor_lr 钳制四、MultiFactorScheduler里程碑式多段衰减MultiFactorSchedulerpython/mxnet/lr_scheduler.py是工程中最常用的调度器它在指定的多个更新步数点各衰减一次若存在 k 使得 step[k] num_update step[k1]则 lr base_lr * factor^(k1)构造参数step : list of int里程碑步数列表必须是非空且严格递增的整数列表源码分别用assert isinstance(step, list) and len(step) 1与逐元素检查保证。每个元素必须 ≥ 1且后一个必须大于前一个否则抛ValueError。factor : float每次里程碑处的衰减因子同样要求 ≤ 1.0。base_lr、warmup_*同上。实现要点__call__python/mxnet/lr_scheduler.py同样使用while循环遍历step列表用cur_step_ind记录已触发的里程碑个数一旦num_update超过某个里程碑就将base_lr乘以factor并记录日志Change learning rate to %0.5e。实战示例对应 tests/python/unittest/test_optimizer.py 的test_multifactor_schedulersched mx.lr_scheduler.MultiFactorScheduler( step[15, 25], factor0.1, base_lr0.1, warmup_steps10, warmup_begin_lr0.05, warmup_modelinear ) print(sched(0)) # 0.05预热起点 print(sched(5)) # 0.075线性预热中段(0.1-0.05)*5/10 0.05 print(sched(15)) # 0.1预热结束尚未触发第一个里程碑 print(sched(16)) # 0.01触发 step15第一次衰减 print(sched(26)) # 0.001触发 step25第二次衰减 print(sched(100)) # 0.001保持在经典 CNN 训练如 ResNet中MultiFactorScheduler常配合每 30/60/80 epoch 衰减 0.1的策略使用是大多数 MXNet 图像分类例程如 example/image-classification/train_imagenet.py的默认选择。五、PolyScheduler多项式衰减PolySchedulerpython/mxnet/lr_scheduler.py按多项式函数在max_update步内把学习率从base_lr平滑过渡到final_lr若 num_update max_update lr final_lr (base_lr - final_lr) * (1 - (num_update - warmup_steps) / max_steps)^pwr 否则lr final_lr其中max_steps max_update - warmup_steps即去除预热步数后的有效衰减区间。构造参数max_update : int衰减到达最终学习率所需的最大更新步数必须是正整数源码assert isinstance(max_update, int)且 1时抛错。base_lr : float起始学习率默认0.01。pwr : int多项式幂次默认2。幂次越大前期衰减越慢、后期衰减越快。final_lr : float所有步数结束后的最终学习率默认0。warmup_*同上。实现要点__call__python/mxnet/lr_scheduler.py构造时把base_lr存入独立的base_lr_orig避免衰减过程污染初始值num_update max_update时按公式计算否则返回final_lr由于pow计算代码在超出范围时直接返回上一次的base_lr即final_lr。验证示例对应test_poly_schedulertests/python/unittest/test_optimizer.pysched mx.lr_scheduler.PolyScheduler( max_update1000, base_lr3, pwr2, final_lr0, warmup_steps100, warmup_begin_lr0, warmup_modelinear ) print(sched(0)) # 0预热起点 print(sched(50)) # 1.5线性预热中段 print(sched(100)) # 3预热结束回到 base_lr print(sched(500)) # 1.6二次多项式衰减中段 print(sched(1000)) # 0到达 final_lr六、CosineScheduler余弦退火CosineSchedulerpython/mxnet/lr_scheduler.py采用余弦函数实现平滑退火公式为若 num_update max_update lr final_lr (base_lr - final_lr) * (1 cos(pi * (num_update - warmup_steps) / max_steps)) / 2 否则lr final_lr该曲线从base_lr出发先缓慢、再加速、最后又缓慢地逼近final_lr相比阶梯衰减更平滑近年来在图像分类、目标检测等任务中被广泛采用。构造参数与PolyScheduler几乎一致max_update正整数必须显式传入、base_lr默认 0.01、final_lr默认 0外加warmup_*三件套。验证示例对应test_cosine_schedulertests/python/unittest/test_optimizer.py该用例特意验证了不带 warmup的场景sched mx.lr_scheduler.CosineScheduler(max_update1000, base_lr3, final_lr0.1) print(sched(0)) # 3.0等于 base_lr print(sched(250)) # 约 1.55曲线中点附近 print(sched(1000)) # 0.1到达 final_lr七、与 Optimizer、Gluon Trainer 的集成调度器本身不会自动运行它必须挂载到优化器上由优化器在每次更新参数前调用。集成方式有两种。7.1 与符号式/命令式 Optimizer 集成在 python/mxnet/optimizer/optimizer.py 中Optimizer构造函数接受lr_scheduler参数。其内部逻辑python/mxnet/optimizer/optimizer.py值得注意若lr_scheduler与learning_rate都为Nonelearning_rate默认取0.01若两者同时传入且值不一致会打印UserWarning(learning rate from lr_scheduler has been overwritten by learning_rate in optimizer.)并用learning_rate覆盖调度器的base_lr每次update时优化器通过self.lr_scheduler(self.num_update)取回当前学习率python/mxnet/optimizer/optimizer.py。对应的单元测试见 tests/python/unittest/test_optimizer.pytest_learning_rate验证了调度器base_lr与优化器学习率的联动关系test_learning_rate_expect_user_warning则验证了调度器与 learning_rate 冲突会触发 UserWarning的行为。7.2 与 Gluon Trainer 集成在 Gluon 命令式训练中更常见的做法是把调度器直接传给gluon.Trainerpython/mxnet/gluon/trainer.py 的 docstring 明确列出optimizer.lr_scheduler是 Trainer 可接受的配置项import mxnet as mx from mxnet import gluon, autograd net gluon.nn.Sequential() net.add(gluon.nn.Dense(10)) net.initialize() scheduler mx.lr_scheduler.MultiFactorScheduler( step[30, 60, 90], factor0.1, base_lr0.1 ) trainer gluon.Trainer(net.collect_params(), optimizersgd, optimizer_params{lr_scheduler: scheduler}) # 训练循环内无需手动改学习率Trainer 内部按 num_update 自动查询 for epoch in range(100): for batch in train_data: with autograd.record(): loss compute_loss(net(batch.data)) loss.backward() trainer.step(batch_size)Trainer 还提供learning_rate属性python/mxnet/gluon/trainer.py和set_learning_rate(lr)方法python/mxnet/gluon/trainer.py可在训练中途手动覆盖学习率调度器会从更新后的num_update起继续生效。相关行为在 tests/python/unittest/test_gluon_trainer.py 中有直接验证。八、选择调度器的实践建议综合源码实现与仓库用例给出如下选型建议阶梯式训练分类任务常规做法首选MultiFactorScheduler在总 epoch 的 1/2、3/4 附近设置里程碑factor取 0.1配合固定的base_lr即可稳定收敛。需要精确控制总训练预算使用PolyScheduler或CosineScheduler将max_update设为预估总步数epoch 数 × 每 epoch 步数两者都能保证在训练结束时学习率恰好降到final_lr适合对比实验与论文复现。大 batch 或 transformer 类模型务必开启 warmupwarmup_steps通常设为总步数的 5%10%warmup_mode优先选linear以平滑过渡。断点续训由于FactorScheduler与MultiFactorScheduler使用while而非if推进状态通过load_epoch恢复训练时学习率仍能正确落在当前阶段无需额外处理。九、深入阅读调度器全部源码python/mxnet/lr_scheduler.py优化器与调度器的耦合逻辑python/mxnet/optimizer/optimizer.pyGluon Trainer 的接入方式python/mxnet/gluon/trainer.py各调度器数值行为的单元测试tests/python/unittest/test_optimizer.pyGluon 场景下的调度器测试tests/python/unittest/test_gluon_trainer.py模块在包中的导出位置python/mxnet/init.py如需在真实任务中验证调度器效果可参考仓库内的完整训练脚本如 example/image-classification/train_mnist.py 与 example/image-classification/train_imagenet.py其中都展示了调度器与数据迭代、模型训练循环的组合用法。【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址: https://gitcode.com/gh_mirrors/mxnet1/mxnet创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表