GridSearchCV使用KerasClassifier时activation参数无效报错问题
问题修复方案
报错原因
KerasClassifier 本身仅支持识别自身构造方法内置的参数(如batch_size、epochs等),你自定义的模型构建函数MultiPerceptron的入参(activation、units、kernel_initializer等)属于透传的自定义参数,无法被GridSearchCV直接识别传递。
另外代码中存在一处笔误:MultiPerceptron定义时默认loss参数为binary_cross_entropy,多了一个下划线,和Keras内置的损失函数名binary_crossentropy不匹配,后续也会触发报错。
修复方法(二选一即可)
方法1:给自定义参数加嵌套前缀
按照scikit-learn嵌套参数的传递规则,所有要传给build_fn的参数,在param_grid的键名前加build_fn__前缀(双下划线)即可,修改后的param代码如下:
param = {'batch_size': [10, 30], 'epochs': [50, 100], 'build_fn__optimizer': ['adam', 'sgd'], 'build_fn__loss': ['binary_crossentropy', 'hinge'], 'build_fn__kernel_initializer': ['random_uniform', 'normal'], 'build_fn__activation': ['relu', 'tanh'], 'build_fn__units': [16, 8]}
同时修正MultiPerceptron的默认loss参数:
def MultiPerceptron(optimizer = 'adam', loss = 'binary_crossentropy', kernel_initializer = 'random_uniform', activation = 'relu', units = 16): # 剩余代码保持不变
方法2:显式声明自定义参数
初始化KerasClassifier时,把所有自定义的模型参数都显式作为构造参数传入(可以直接用默认值占位),param_grid不需要做任何修改:
classifier = KerasClassifier( build_fn = MultiPerceptron, validation_split = 0.1, validation_batch_size = 50, # 显式声明所有自定义参数 optimizer = 'adam', loss = 'binary_crossentropy', kernel_initializer = 'random_uniform', activation = 'relu', units = 16 )
补充说明
如果你使用的是新版SciKeras(TensorFlow官方已弃用tf.keras.wrappers.scikit_learn下的KerasClassifier,迁移到SciKeras库),不需要加前缀,直接保证param_grid的键名和构建函数的入参名完全一致即可正常运行。
内容的提问来源于stack exchange,提问作者Murilo
相关产品推荐
相关产品推荐

