TensorFlow 2.18.0中Keras动态更新学习率报错及正确实现方法
TensorFlow 2.18.0中动态更新Keras模型学习率的正确方法
你的代码报错是因为两个核心原因:一是K.set_value在TensorFlow 2.x后期版本已被废弃,二是模型编译后model.optimizer.learning_rate返回的是字符串类型的配置标识,不再是可直接操作的张量,导致K.set_value无法识别其数据类型。
以下是TF 2.18.0中动态更新学习率的几种正确方式:
1. 直接通过优化器变量赋值(手动即时更新)
直接调用优化器的learning_rate.assign()方法,这是最直接的手动更新方式:
import tensorflow as tf from tensorflow import keras # 定义模型与优化器 model = keras.models.Sequential([keras.layers.Dense(10)]) optimizer = keras.optimizers.SGD(learning_rate=0.01) model.compile(optimizer=optimizer, loss='mse') # 动态更新学习率为0.001 optimizer.learning_rate.assign(0.001) # 验证更新结果 print(f"当前学习率: {optimizer.learning_rate.numpy()}")
2. 使用内置学习率调度器(自动按规则更新)
如果需要根据训练轮次、验证指标等自动调整学习率,可以用Keras内置的调度器,比如LearningRateScheduler或ReduceLROnPlateau:
示例1:按轮次手动调整
import tensorflow as tf from tensorflow import keras def lr_scheduler(epoch, current_lr): # 训练10轮后学习率乘以0.1 if epoch > 10: return current_lr * 0.1 return current_lr model = keras.models.Sequential([keras.layers.Dense(10)]) model.compile(keras.optimizers.SGD(learning_rate=0.01), loss='mse') # 训练时传入调度器回调 model.fit( x_train, y_train, epochs=20, callbacks=[keras.callbacks.LearningRateScheduler(lr_scheduler)] )
示例2:根据验证loss自动降低
model = keras.models.Sequential([keras.layers.Dense(10)]) model.compile(keras.optimizers.SGD(learning_rate=0.01), loss='mse') # 当验证loss连续3轮不下降时,学习率乘以0.5 reduce_lr = keras.callbacks.ReduceLROnPlateau( monitor='val_loss', factor=0.5, patience=3, min_lr=0.0001 ) model.fit( x_train, y_train, validation_data=(x_val, y_val), epochs=20, callbacks=[reduce_lr] )
3. 自定义回调函数(灵活定制更新逻辑)
如果需要更复杂的更新规则(比如根据训练中的实时loss、准确率等调整),可以自定义回调函数:
import tensorflow as tf from tensorflow import keras class CustomLRScheduler(keras.callbacks.Callback): def on_epoch_end(self, epoch, logs=None): current_loss = logs.get('loss') if current_loss is not None and current_loss < 0.1: new_lr = 0.0001 self.model.optimizer.learning_rate.assign(new_lr) print(f"Epoch {epoch+1}: 学习率更新为 {new_lr}") model = keras.models.Sequential([keras.layers.Dense(10)]) model.compile(keras.optimizers.SGD(learning_rate=0.01), loss='mse') model.fit( x_train, y_train, epochs=20, callbacks=[CustomLRScheduler()] )
内容的提问来源于stack exchange,提问作者codebysumit
相关产品推荐
相关产品推荐

