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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 07:06:03