如何在Keras Tuner中使用tf.keras.callbacks.ModelCheckpoint并关联试验信息
Keras Tuner中绑定试验标识的ModelCheckpoint配置方案
核心解决逻辑是利用Keras Tuner每个试验(trial)自带的唯一ID,动态为每个试验生成专属的检查点保存路径,实现检查点与试验记录的一一绑定,具体实现步骤如下:
1. 基础模型构建逻辑保留原有写法
和普通Keras Tuner调参的模型构建逻辑一致,无需修改:
import os import tensorflow as tf import keras_tuner as kt # 提前创建检查点保存目录 os.makedirs("./trial_checkpoints", exist_ok=True) def build_model(hp): model = tf.keras.Sequential([ tf.keras.layers.Dense( hp.Int("hidden_units", min_value=32, max_value=256, step=32), activation="relu" ), tf.keras.layers.Dense(10, activation="softmax") ]) model.compile( optimizer="adam", loss="sparse_categorical_crossentropy", metrics=["accuracy"] ) return model
2. 自定义Tuner类动态注入带试验标识的回调
重写run_trial方法,在每个试验启动时生成对应唯一路径的ModelCheckpoint回调,自动绑定试验ID、执行序号信息:
class TrialBoundTuner(kt.RandomSearch): # 可根据需要替换父类为kt.BayesianOptimization、kt.Hyperband等其他调优器 def run_trial(self, trial, *args, **kwargs): # 动态生成检查点保存路径,嵌入试验ID、执行序号、轮次、指标信息 checkpoint_path = "./trial_checkpoints/trial_{}_exec_{}_epoch_{{epoch:02d}}_val_acc_{{val_accuracy:.4f}}.h5".format( trial.trial_id, trial.execution_index ) # 初始化检查点回调 model_checkpoint_cb = tf.keras.callbacks.ModelCheckpoint( filepath=checkpoint_path, save_best_only=True, monitor="val_accuracy", verbose=1 ) # 将回调添加到回调列表中 kwargs["callbacks"] = kwargs.get("callbacks", []) + [model_checkpoint_cb] return super().run_trial(trial, *args, **kwargs)
3. 启动调参流程
和普通调参流程一致,无需额外传入ModelCheckpoint回调:
# 初始化自定义调优器 tuner = TrialBoundTuner( hypermodel=build_model, objective="val_accuracy", max_trials=10, executions_per_trial=2, directory="./tuner_logs", project_name="test_demo" ) # 启动搜索 (x_train, y_train), (x_val, y_val) = tf.keras.datasets.mnist.load_data() tuner.search( x_train, y_train, validation_data=(x_val, y_val), epochs=10, batch_size=32 )
效果说明
- 最终保存的检查点文件名形如
trial_0_exec_0_epoch_05_val_acc_0.9876.h5,可以直接通过文件名中的trial_id和Tuner生成的试验记录一一匹配 - 如需获取对应试验的最优检查点,调用
tuner.get_trials()获取所有试验对象,取trial_id字段匹配对应文件即可
内容的提问来源于stack exchange,提问作者Tiago Domingues
相关产品推荐
相关产品推荐

