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

TensorFlow教程中学习率的更新及衰减实现方式咨询

TensorFlow中学习率更新的执行逻辑详解

嘿,我来帮你拆解清楚这种学习率更新的实现方式~你看到的这个代码片段是TensorFlow 1.x时代非常经典的动态调整学习率的手动实现方式,核心是通过变量赋值操作在运行时修改学习率,具体执行逻辑可以分成这几个步骤:

1. 初始化学习率变量

self._lr = tf.Variable(0.0, trainable=False)

这里把学习率定义成一个tf.Variable,但关键是设置了trainable=False——这意味着这个变量不会被优化器当作模型参数来更新,完全由我们手动控制它的取值。

2. 绑定优化器与学习率

optimizer = tf.train.GradientDescentOptimizer(self._lr)

优化器初始化时传入这个学习率变量,这样后续每次执行训练操作时,优化器都会读取self._lr的当前值来计算梯度更新量。

3. 定义学习率更新操作

self._new_lr = tf.placeholder(tf.float32, shape=[], name="new_learning_rate")
self._lr_update = tf.assign(self._lr, self._new_lr)
  • self._new_lr是一个占位符(placeholder),用来接收外部传入的新学习率值;
  • tf.assign()会创建一个赋值操作,把占位符传入的新值覆盖到self._lr变量中。这个操作不会自动执行,必须显式通过会话(session)调用才会生效。

4. 训练时的执行流程

在实际训练循环中,你需要按以下步骤操作:

  • 首先初始化所有变量,包括这个学习率变量;
  • 先设置初始学习率:通过session.run(self._lr_update, feed_dict={self._new_lr: 初始学习率})把初始值写入self._lr;
  • 每轮训练(或每N步训练)后,根据你的衰减规则(比如按epoch衰减、按步数衰减,或者根据验证集精度调整)计算出新的学习率;
  • 调用session.run(self._lr_update, feed_dict={self._new_lr: 新学习率})执行更新操作,此时self._lr的取值就被替换成新值了;
  • 后续再执行训练操作(比如session.run(train_op))时,优化器就会使用更新后的学习率来调整模型参数。

完整示例片段

# 初始化计算图
self._lr = tf.Variable(0.0, trainable=False)
loss = ... # 定义你的损失函数
optimizer = tf.train.GradientDescentOptimizer(self._lr)
train_op = optimizer.minimize(loss)

# 定义学习率更新操作
self._new_lr = tf.placeholder(tf.float32, shape=[], name="new_learning_rate")
self._lr_update = tf.assign(self._lr, self._new_lr)

# 开始训练
initial_lr = 0.1
decay_epochs = 5
decay_rate = 0.9
total_epochs = 30

with tf.Session() as sess:
    sess.run(tf.global_variables_initializer())
    # 设置初始学习率
    sess.run(self._lr_update, feed_dict={self._new_lr: initial_lr})
    
    for epoch in range(total_epochs):
        # 执行一轮训练
        _, current_loss = sess.run([train_op, loss], feed_dict={...})
        
        # 每5轮衰减一次学习率
        if (epoch + 1) % decay_epochs == 0:
            current_lr = sess.run(self._lr)
            new_lr = current_lr * decay_rate
            sess.run(self._lr_update, feed_dict={self._new_lr: new_lr})
            print(f"Epoch {epoch+1}: 学习率从 {current_lr:.4f} 更新为 {new_lr:.4f}")

这种方式的好处是灵活性极高——你可以自定义任何复杂的学习率衰减规则,只要在运行时计算出新值传入即可,不需要重新构建计算图。不过在TensorFlow 2.x中,官方更推荐使用tf.keras.optimizers.schedules里的内置衰减策略(比如ExponentialDecay),但你看到的这种手动更新方式是理解TF1.x计算图模型的典型案例。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 12:11:31