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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 08:48:17