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

Keras Sequential模型获取epoch后学习率报错:set_model()缺model参数

解决Keras中获取每个epoch学习率时的set_model()参数缺失错误

嘿,我来帮你搞定这个问题!你遇到的set_model() missing 1 required positional argument: 'model'错误,大概率是因为自定义回调的姿势不对——没有正确继承Keras的回调基类,导致框架没法自动把模型实例传递给回调函数。

下面是完整的解决方案,一步步来:

1. 正确定义学习率日志回调类

首先要继承keras.callbacks.Callback基类,这个基类会帮我们自动处理set_model()这类底层方法的调用,不用自己手动实现:

from tensorflow.keras.callbacks import Callback

class LearningRateLogger(Callback):
    def on_epoch_end(self, epoch, logs=None):
        # 获取当前优化器的学习率,转成numpy数值方便阅读
        current_lr = self.model.optimizer.lr.numpy()
        # 打印或者记录学习率,这里用print示例
        print(f"\nEpoch {epoch + 1}: 当前学习率 = {current_lr:.6f}")

2. 模型构建与训练时正确使用回调

接下来在你的模型构建、编译、训练流程中,实例化这个回调类,并传入model.fit()的callbacks参数里:

from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Conv2D
from tensorflow.keras.optimizers import Adam

# 假设你已经定义了IMG_HEIGHT和IMG_WIDTH
IMG_HEIGHT = 224
IMG_WIDTH = 224

# 构建你的Sequential模型
model = Sequential()
model.add(Conv2D(64, (5, 5), input_shape=(IMG_HEIGHT, IMG_WIDTH, 3), activation='relu'))
model.add(Conv2D(64, (5, 5), activation='relu'))
# 这里可以继续添加你需要的其他层(比如Pooling、Dense等)

# 初始化Adam优化器,设置初始学习率
optimizer = Adam(learning_rate=0.001)
model.compile(optimizer=optimizer, loss='categorical_crossentropy', metrics=['accuracy'])

# 实例化学习率日志回调
lr_logger = LearningRateLogger()

# 开始训练,把回调传入callbacks列表
model.fit(
    x=train_dataset,  # 你的训练数据
    epochs=10,
    callbacks=[lr_logger]
)

错误原因说明

你之前的错误是因为自定义的回调没有继承Callback基类,Keras在训练过程中会自动调用回调的set_model()方法来传递当前模型实例,但如果你的类没有继承基类,这个方法就不存在,或者参数不匹配,就会抛出set_model() missing 1 required positional argument: 'model'的错误。

另外,如果你的Adam优化器使用了动态学习率调度(比如ReduceLROnPlateau或者自定义的LearningRateSchedule),获取学习率的方式可以稍作调整:

  • 如果是用LearningRateSchedule,可以用optimizer.lr(epoch).numpy()来获取对应epoch的学习率
  • 如果是用ReduceLROnPlateau这类回调,上面的self.model.optimizer.lr.numpy()依然有效,因为它会实时读取优化器的当前学习率

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 06:33:03