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

如何在文件夹式图像数据集上应用GridSearchCV进行超参数调优?

自定义CNN的GridSearch超参数调优实现方案

要给你的CNN做超参数调优,得先把Keras模型适配到scikit-learn工具链,再定义参数搜索范围,最后执行网格搜索。以下是完整修改方案:

1. 补充导入必要依赖

除原有库外,需导入适配工具和网格搜索模块:

import numpy as np
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Conv2D, MaxPool2D, Flatten, Dense
from tensorflow.keras.preprocessing.image import ImageDataGenerator
from tensorflow.keras.wrappers.scikit_learn import KerasClassifier
from sklearn.model_selection import GridSearchCV

2. 重构模型为可配置函数

将固定参数的模型改成函数形式,让超参数可动态传入:

def build_cnn(filters=16, kernel_size=(3,3), dense_units=64, optimizer='adam'):
    imageSize = [101,168,3]
    model = Sequential()
    model.add(Conv2D(filters=filters, kernel_size=kernel_size, input_shape=imageSize, 
                     activation="relu", padding="same"))
    model.add(MaxPool2D(strides=2, pool_size=(2,2)))
    model.add(Flatten())
    model.add(Dense(dense_units, activation="relu"))
    model.add(Dense(1, activation="sigmoid"))
    # 必须添加编译步骤,网格搜索需要编译好的模型
    model.compile(optimizer=optimizer, loss='binary_crossentropy', metrics=['accuracy'])
    return model

3. 加载并转换数据集为数组格式

GridSearchCV需要X(特征)和y(标签)格式的数据,需将生成器的所有批次合并为数组:

trainDirectory = "../Images/DATA3/Training"
testDirectory = "../Images/DATA3/Test"

trainingGenerator = ImageDataGenerator(rescale = 1./255,
                                      shear_range = 0.2,
                                      zoom_range = 0.1,
                                      horizontal_flip = True)
testingGenerator = ImageDataGenerator(rescale = 1./255)

trainingSet = trainingGenerator.flow_from_directory(trainDirectory,
                                                   target_size = (101, 168),
                                                   batch_size = 16,
                                                   class_mode = 'binary',
                                                   shuffle=False)  # 关闭打乱保证合并顺序正确

testingSet = testingGenerator.flow_from_directory(testDirectory,
                                                 target_size = (101,168),
                                                 batch_size = 16,
                                                 class_mode = "binary",
                                                 shuffle=False)

# 提取生成器所有数据的函数
def extract_full_data(generator):
    data_batches = []
    label_batches = []
    for _ in range(len(generator)):
        batch_data, batch_labels = next(generator)
        data_batches.append(batch_data)
        label_batches.append(batch_labels)
    return np.concatenate(data_batches), np.concatenate(label_batches)

train_data, train_labels = extract_full_data(trainingSet)
test_data, test_labels = extract_full_data(testingSet)

4. 定义超参数搜索网格

列出要调优的参数和候选值,新手建议选少量候选值避免搜索时间过长:

param_grid = {
    'filters': [16, 32, 64],        # 卷积核数量
    'kernel_size': [(3,3), (5,5)],  # 卷积核尺寸
    'dense_units': [32, 64, 128],   # 全连接层神经元数
    'optimizer': ['adam', 'sgd'],   # 优化器
    'epochs': [10, 20]              # 训练轮数
}

5. 执行网格搜索

把Keras模型包装成scikit-learn兼容的分类器,初始化并执行搜索:

# 包装模型,verbose=0表示训练时不输出日志
model_wrapper = KerasClassifier(build_fn=build_cnn, verbose=0)

# 初始化网格搜索,cv=3表示3折交叉验证,用准确率评估
grid_search = GridSearchCV(estimator=model_wrapper, param_grid=param_grid,
                           cv=3, scoring='accuracy', verbose=2)

# 开始搜索
grid_result = grid_search.fit(train_data, train_labels)

6. 查看结果并评估最佳模型

搜索完成后输出最优参数,再用最优模型测试测试集:

# 打印最佳结果
print(f"交叉验证最佳准确率: {grid_result.best_score_:.4f}")
print(f"最佳参数组合: {grid_result.best_params_}")

# 获取最佳模型并评估测试集
best_model = grid_result.best_estimator_.model
test_loss, test_acc = best_model.evaluate(test_data, test_labels)
print(f"测试集准确率: {test_acc:.4f}")

注意事项

  • 若数据集过大,加载全部数据会导致内存溢出,可改用生成器配合自定义训练逻辑,但新手先从小数据集测试更稳妥。
  • 超参数候选值不宜过多,否则搜索组合量会指数级增长,训练时间大幅延长。
  • 可扩展更多调优参数,比如Dropout层的丢弃率、优化器的学习率(需单独配置优化器并加入参数网格)。

内容的提问来源于stack exchange,提问作者Anna TOes

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 15:30:20