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

如何在Keras Tuner中兼顾模型精度与训练、推理速度?

解决方案:兼顾精度与模型速度的超参数调优

1. 自定义复合优化目标,平衡精度与模型复杂度

Hyperband默认仅优化单一精度指标,我们可以将验证精度与模型参数总量结合成加权目标,让调优过程自动偏好高效模型。

代码修改示例:

def model_builder(hp):
    # 原模型构建逻辑保留
    model = keras.Sequential([
        tf.keras.layers.Conv2D(hp.Int('conv1filter', min_value=32, max_value=512*3, step=512/2),
                               hp.Int('conv1kernal', min_value=2, max_value=20, step=2), 
                               padding="same", 
                               activation="relu", 
                               input_shape=(14,8,8)),
        tf.keras.layers.BatchNormalization(axis=-1, momentum=0.99, epsilon=1e-05),
        tf.keras.layers.Conv2D(hp.Int('conv2filter', min_value=32, max_value=512*3, step=512/2),
                               hp.Int('conv2kernal', min_value=2, max_value=20, step=2), 
                               padding="same", 
                               activation="relu"),
        tf.keras.layers.BatchNormalization(axis=-1, momentum=0.99, epsilon=1e-05),
        layers.Flatten(),
        tf.keras.layers.Dense(hp.Int('dense1', min_value=32, max_value=512, step=32), activation='relu'),
        tf.keras.layers.Dense(hp.Int('dense2', min_value=32, max_value=512, step=32),  activation='relu'),
        tf.keras.layers.Dense(hp.Int('dense3', min_value=32, max_value=512, step=32), activation='relu'),
        tf.keras.layers.Dense(1, activation='tanh'),
    ])
    
    # 计算模型参数总数并作为超参数记录
    total_params = model.count_params()
    hp.set_hparam('total_params', total_params)
    
    hp_learning_rate = hp.Choice('learning_rate', values=[1e-2, 1e-3, 1e-4])
    model.compile(optimizer=keras.optimizers.Adam(learning_rate=hp_learning_rate),
                loss='mean_absolute_error',
                metrics=['accuracy'])

    return model

# 自定义加权目标:90%权重给验证精度,10%权重给参数数量(越小越好)
tuner = kt.Hyperband(
    model_builder,
    objective=kt.Objective('val_accuracy', direction='max') * 0.9 + kt.Objective('total_params', direction='min') * 0.1,
    max_epochs=10,
    overwrite=True,
    directory='my_dir30',
    project_name='intro_to_kt30'
)

可根据需求调整权重比例:如果速度优先级更高,可增大total_params的权重。

2. 约束超参数搜索空间,从源头限制模型规模

直接缩小超参数的取值范围,避免生成过于庞大的模型:

# 缩小Conv2D过滤器数量上限,步长调整为64
tf.keras.layers.Conv2D(hp.Int('conv1filter', min_value=32, max_value=256, step=64),
                       hp.Int('conv1kernal', min_value=2, max_value=8, step=2),  # 卷积核上限从20降至8
                       padding="same", 
                       activation="relu", 
                       input_shape=(14,8,8)),
# 同理修改第二层卷积
tf.keras.layers.Conv2D(hp.Int('conv2filter', min_value=32, max_value=256, step=64),
                       hp.Int('conv2kernal', min_value=2, max_value=8, step=2), 
                       padding="same", 
                       activation="relu"),
# 缩小全连接层神经元上限
tf.keras.layers.Dense(hp.Int('dense1', min_value=32, max_value=256, step=32), activation='relu'),
tf.keras.layers.Dense(hp.Int('dense2', min_value=32, max_value=256, step=32),  activation='relu'),
tf.keras.layers.Dense(hp.Int('dense3', min_value=32, max_value=128, step=32), activation='relu'),

3. 后处理筛选:从候选模型中选精度达标且速度最快的

Hyperband结束后,遍历topN候选模型,评估推理速度,筛选出精度满足阈值且速度最优的模型:

# 获取top10精度的候选超参数
all_hps = tuner.get_best_hyperparameters(num_trials=10)
# 设定精度阈值(取最优模型精度的95%作为达标线)
top_val_acc = tuner.get_best_models()[0].evaluate(x_val, y_val)[1]
acc_threshold = top_val_acc * 0.95

best_model = None
min_infer_time = float('inf')

for hp in all_hps:
    model = tuner.hypermodel.build(hp)
    # 验证精度是否达标
    val_loss, val_acc = model.evaluate(x_val, y_val, verbose=0)
    if val_acc >= acc_threshold:
        # 评估推理速度(用100个样本测试批量推理耗时)
        import time
        start = time.time()
        model.predict(x_val[:100], verbose=0)
        infer_time = time.time() - start
        
        if infer_time < min_infer_time:
            min_infer_time = infer_time
            best_model = model

# 评估并保存最优模型
eval_result = best_model.evaluate(x_test, y_test)
print("[test loss, test accuracy]:", eval_result)
best_model.save('/notebooks/saved_model/my_model')

内容的提问来源于stack exchange,提问作者Lukas Taylor

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 03:32:15