Python神经网络网格搜索参数定义报错learning_rate不合法如何解决
问题核心原因
报错及参数不生效是两个问题导致的:
- 你需要搜索的
units、learning_rate等自定义参数没有在create_network的入参中声明,KerasRegressor无法将网格搜索传入的参数映射到模型构建逻辑 - 学习率是优化器实例的属性,不能直接和字符串格式的优化器名称绑定,需要先实例化优化器再配置学习率
修复方案
第一步:改写模型构建函数
将所有待搜索的参数都作为create_network的入参,函数内部逻辑使用传入的参数值,不要硬编码固定值:
import tensorflow as tf from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense from tensorflow.keras.wrappers.scikit_learn import KerasRegressor from sklearn.model_selection import GridSearchCV def create_network(optimizer='rmsprop', units=16, learning_rate=0.001, leaky_relu_alpha=0.3): network = Sequential() # 输入层+第一层隐藏层,units用传入的参数 network.add(Dense(units = units, activation = tf.keras.layers.LeakyReLU(alpha=leaky_relu_alpha))) # 第二层隐藏层 network.add(Dense(units = units, activation = tf.keras.layers.LeakyReLU(alpha=leaky_relu_alpha))) # 第三层隐藏层 network.add(Dense(units = units, activation = tf.keras.layers.LeakyReLU(alpha=leaky_relu_alpha))) # 输出层 network.add(Dense(units = 1)) # 实例化优化器并配置学习率 opt = tf.keras.optimizers.get(optimizer) opt.learning_rate = learning_rate # 编译模型 network.compile(optimizer = opt, loss = 'mean_squared_error', metrics=['mae', tf.keras.metrics.RootMeanSquaredError()]) return network
第二步:网格搜索逻辑兼容扩展
原有网格搜索逻辑不需要大幅修改,后续如果要新增搜索参数,只需要做两个操作:
- 在
create_network的入参中添加对应参数,设置默认值,函数内部用到该配置的位置替换为参数变量 - 在
hyperparameters字典中添加该参数的候选值列表即可
修改后的网格搜索完整代码:
# 封装Keras模型 ann = KerasRegressor(build_fn=create_network, verbose=0) # 超参数搜索空间,后续新增参数直接往这个字典里加就行 epoch_values = [10, 25, 50, 100, 150, 200] batches = [10, 20, 30, 40, 50, 100, 1000] optimizers = ['rmsprop', 'adam', 'SGD'] neurons = [16, 32, 64, 128, 256] lr_values = [0.001, 0.01, 0.1, 0.2, 0.3] # 示例:后续要加leaky_relu的alpha搜索,直接加变量再放到hyperparameters里即可 # alpha_values = [0.1, 0.2, 0.3, 0.5] hyperparameters = dict( optimizer=optimizers, epochs=epoch_values, batch_size=batches, units=neurons, learning_rate=lr_values # 新增参数直接加在这里就行: leaky_relu_alpha=alpha_values ) # 网格搜索实例化、训练 grid = GridSearchCV(estimator=ann, cv=5, param_grid=hyperparameters) grid_result = grid.fit(X, y)
内容的提问来源于stack exchange,提问作者Joehat
相关产品推荐
相关产品推荐

