TensorFlow 2.0自定义训练循环中如何查看当前学习率?
获取TensorFlow 2.x自定义训练循环中的当前学习率
在TensorFlow 2.x的自定义训练循环里,要获取优化器的当前学习率,分两种场景处理就很清晰:
固定学习率场景
如果你的优化器用的是固定值学习率(比如初始化时直接传learning_rate=0.001),直接访问优化器的learning_rate属性就能拿到当前值:
current_lr = optimizer.learning_rate.numpy()
你可以把这段代码加到训练循环里,比如在apply_gradients操作之后打印或者记录这个值。
动态学习率场景
如果用了学习率调度器(比如ExponentialDecay、PiecewiseConstantDecay这类动态调整策略),此时optimizer.learning_rate是一个学习率调度对象,需要传入当前迭代步数(也就是optimizer.iterations)来计算实时学习率:
current_lr = optimizer.learning_rate(tf.cast(optimizer.iterations, tf.float32)).numpy()
这里要把迭代步数转成浮点型,因为大多数学习率调度器的输入要求是浮点数类型。
适配两种场景的示例代码
结合你给出的训练循环示例,我们可以加入自动判断逻辑,不管是固定还是动态学习率都能正确获取:
def train_step(images, labels): with tf.GradientTape() as tape: predictions = model(images) loss = loss_object(labels, predictions) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) # 自动判断并获取当前学习率 if isinstance(optimizer.learning_rate, tf.keras.optimizers.schedules.LearningRateSchedule): current_lr = optimizer.learning_rate(tf.cast(optimizer.iterations, tf.float32)).numpy() else: current_lr = optimizer.learning_rate.numpy() print(f"当前训练步数 {optimizer.iterations.numpy()},学习率: {current_lr:.6f}") train_loss(loss) train_accuracy(labels, predictions)
这样每次执行train_step时,都会打印当前迭代步数和对应的学习率,方便你监控训练过程。
内容的提问来源于stack exchange,提问作者yun
相关产品推荐
相关产品推荐

