使用Keras SKLearnClassifier包装器拟合MNIST数据遇报错求助
Keras SKLearnClassifier 网格搜索报错解决
问题场景
尝试用Keras的SKLearnClassifier包装器结合sklearn开展网格搜索与交叉验证,但模型无法正常运行;直接调用模型构建函数build_model拟合时却一切正常。
问题代码
def build_model(X, y, n_neurons: List[str], learning_rate: float): model = keras.models.Sequential() model.add(keras.Input(shape=(28*28,))) model.add(keras.layers.Dense(n_neurons[0], activation="relu")) model.add(keras.layers.Dense(n_neurons[1], activation="relu")) model.add(keras.layers.Dense(10, activation="softmax")) optimizer = keras.optimizers.SGD(learning_rate=learning_rate) model.compile(loss="sparse_categorical_crossentropy", optimizer=optimizer, metrics=["accuracy"]) return model sk_train = X_train.reshape((X_train.shape[0],X_train.shape[1]*X_train.shape[2])) sk_val = X_val.reshape((X_val.shape[0],X_val.shape[1]*X_val.shape[2])) model = keras.wrappers.SKLearnClassifier(model=build_model, model_kwargs={ "n_neurons": [300, 100], "learning_rate": 3e-4 }) model.fit(sk_train, y_train, epochs=30, validation_data=(sk_val, y_val))
报错信息
ValueError: Argument `output` must have rank (ndim) `target.ndim - 1`. Received: target.shape=(None, 10), output.shape=(None, 10)
问题分析与修复方案
核心原因
SKLearnClassifier默认会对目标变量y做one-hot编码,但你的模型使用的sparse_categorical_crossentropy损失函数,要求目标是整数形式的类别索引而非one-hot向量。直接调用build_model时,传入的是原始整数索引的y_train,所以正常;但包装后y被自动转成2维one-hot向量,和模型输出维度(2维)冲突,触发报错。
修复方式(二选一)
方式1:修改损失函数适配one-hot目标
把损失函数换成categorical_crossentropy,同时删除build_model中多余的X、y参数(包装器不会传递训练数据给构建函数),修正n_neurons的类型标注:
def build_model(n_neurons: List[int], learning_rate: float): model = keras.models.Sequential() model.add(keras.Input(shape=(28*28,))) model.add(keras.layers.Dense(n_neurons[0], activation="relu")) model.add(keras.layers.Dense(n_neurons[1], activation="relu")) model.add(keras.layers.Dense(10, activation="softmax")) optimizer = keras.optimizers.SGD(learning_rate=learning_rate) model.compile(loss="categorical_crossentropy", optimizer=optimizer, metrics=["accuracy"]) return model
方式2:禁用自动one-hot编码
给SKLearnClassifier添加y_int=True参数,让它保留原始整数索引的目标变量,同时同样需要删除build_model中多余的X、y参数:
model = keras.wrappers.SKLearnClassifier( model=build_model, model_kwargs={"n_neurons": [300, 100], "learning_rate": 3e-4}, y_int=True )
验证修复
修复后直接调用构建函数测试:
dummy_model = build_model([300, 100], learning_rate=3e-4) dummy_model.summary() dummy_model.fit(sk_train, y_train, epochs=30, validation_data=(sk_val, y_val))
确认正常运行后,再用SKLearnClassifier包装即可开展网格搜索与交叉验证。
内容的提问来源于stack exchange,提问作者Jesus Diaz Rivero
相关产品推荐
相关产品推荐

