如何在Keras 2.8自定义回调中获取model.fit()的参数值?
在Keras 2.8的自定义回调中获取model.fit()的参数值
你可以通过两种方式轻松获取model.fit()的参数:
方法1:手动将参数传递给回调实例
在初始化自定义回调时,直接传入fit()中设置的参数,这样在回调的任意方法里都能精准访问:
from tensorflow import keras class FitParamsCallback(keras.callbacks.Callback): def __init__(self, **fit_args): super().__init__() self.fit_args = fit_args def on_train_begin(self, logs=None): print("获取到的fit参数:") for param, value in self.fit_args.items(): print(f"{param}: {value}") # 使用示例 model = keras.Sequential([keras.layers.Dense(10, input_shape=(10,))]) model.compile(optimizer='adam', loss='mse') # 定义fit参数 fit_params = { "batch_size": 32, "epochs": 10, "validation_split": 0.2 } model.fit( x_train, y_train, **fit_params, callbacks=[FitParamsCallback(**fit_params)] )
这种方式完全由你控制要传递的参数,不受Keras内部逻辑限制。
方法2:利用回调内置的self.params属性
Keras 2.8中,回调函数的self.params字典会自动包含model.fit()的核心参数(比如epochs、batch_size、validation_split等),无需手动传递,直接在回调方法中访问即可:
from tensorflow import keras class FitParamsCallback(keras.callbacks.Callback): def on_train_begin(self, logs=None): # 打印所有可用参数 print("内置params中的fit参数:") for key, val in self.params.items(): print(f"{key}: {val}") # 提取特定参数 batch_size = self.params.get("batch_size") epochs = self.params.get("epochs") validation_split = self.params.get("validation_split") print(f"\n提取的关键参数:") print(f"batch_size: {batch_size}") print(f"epochs: {epochs}") print(f"validation_split: {validation_split}") # 使用示例 model.fit( x_train, y_train, batch_size=32, epochs=10, validation_split=0.2, callbacks=[FitParamsCallback()] )
注意事项
- 如果
model.fit()中某些参数使用默认值(比如未显式设置batch_size),self.params会自动填充Keras的默认值(比如默认batch_size=32)。 self.params还包含训练元数据(比如steps、verbose等),可按需提取。
内容的提问来源于stack exchange,提问作者Mehdi
相关产品推荐
相关产品推荐

