You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow中如何在批次训练结束后调整学习率

学习率搜索回调实现方案

核心实现逻辑

  • 自定义回调类继承对应深度学习框架的基础回调类
  • 初始化方法中定义三个核心参数:初始学习率、学习率增长系数(需设置为1.05即可实现每次提升5%的需求)、累计训练批次计数
  • 在on_train_batch_end钩子中每触发一次就依次执行:批次计数+1、计算新学习率=初始学习率 * (增长系数 ** 累计批次)、将新学习率赋值给优化器
  • 可新增两个列表同步记录每一步的学习率和对应损失,后续用于绘制学习率-损失曲线定位最优学习率

PyTorch Lightning 版本实现代码

from pytorch_lightning.callbacks import Callback

class LRFinderCallback(Callback):
    def __init__(self, init_lr: float = 1e-7, growth_factor: float = 1.05):
        self.init_lr = init_lr
        self.growth_factor = growth_factor
        self.batch_count = 0
        # 存储历史数据用于后续分析
        self.lr_history = []
        self.loss_history = []

    def on_train_start(self, trainer, pl_module):
        # 训练启动时先将优化器学习率设置为初始值
        for param_group in trainer.optimizers[0].param_groups:
            param_group['lr'] = self.init_lr

    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
        self.batch_count += 1
        # 计算更新后的学习率
        new_lr = self.init_lr * (self.growth_factor ** self.batch_count)
        # 写入优化器生效
        for param_group in trainer.optimizers[0].param_groups:
            param_group['lr'] = new_lr
        
        # 记录历史数据
        self.lr_history.append(new_lr)
        self.loss_history.append(outputs["loss"].item())

Keras/TensorFlow 版本实现代码

from tensorflow.keras.callbacks import Callback

class LRFinderCallback(Callback):
    def __init__(self, init_lr: float = 1e-7, growth_factor: float = 1.05):
        super().__init__()
        self.init_lr = init_lr
        self.growth_factor = growth_factor
        self.batch_count = 0
        self.lr_history = []
        self.loss_history = []

    def on_train_begin(self, logs=None):
        self.model.optimizer.learning_rate.assign(self.init_lr)

    def on_train_batch_end(self, batch, logs=None):
        self.batch_count += 1
        new_lr = self.init_lr * (self.growth_factor ** self.batch_count)
        self.model.optimizer.learning_rate.assign(new_lr)
        
        self.lr_history.append(new_lr)
        self.loss_history.append(logs["loss"])

使用说明

  • 初始化回调实例时可按需调整初始学习率,建议设置为1e-8~1e-7的极小值避免一开始就超出最优学习率区间
  • 将回调实例传入训练器的callbacks参数列表后启动训练即可
  • 学习率搜索无需训练完整epoch,通常训练100200个batch、学习率涨到110区间即可停止,避免损失爆炸浪费资源
  • 分析最优学习率时建议对学习率取对数坐标绘制曲线,选择损失下降速度最快的区间对应值作为最优学习率

内容的提问来源于stack exchange,提问作者aitor

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.28 03:27:03