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

《Python机器学习手册》MNIST卷积网络代码运行报错求助

问题根因
  • 首次报错:全局设置了channels_first(即NCHW数据格式,形状为[样本数, 通道数, 高度, 宽度]),但CPU版本的Keras/TensorFlow默认最大池化操作仅支持channels_last(即NHWC数据格式,形状为[样本数, 高度, 宽度, 通道数]),因此触发报错。
  • 修改后负维度报错:调整时仅修改了层的默认数据格式为NHWC,没有同步修改数据reshape逻辑和输入形状定义,导致实际输入的形状是[?,1,28,28],把通道数1放到了NHWC要求的高度维度位置,高度仅为1的前提下使用5x5的卷积核无padding计算,输出维度为负数,触发报错。
修复方案(CPU运行优先选择)

统一使用CPU原生支持的NHWC格式,对齐所有数据和层的格式配置,仅需修改3处代码:

  1. 替换全局图像格式配置为channels_last
  2. 修改训练/测试数据的reshape维度顺序
  3. 调整第一层卷积的输入形状定义

修改后完整可运行代码如下:

import numpy as np
from keras.datasets import mnist
from keras.models import Sequential
from keras.layers import Dense,Dropout,Flatten
from keras.layers.convolutional import Conv2D, MaxPooling2D
from keras.utils import np_utils
from keras import backend as K

# 修改点1:改为通道最后格式,适配CPU默认支持
K.set_image_data_format("channels_last")
np.random.seed(0)

channels=1
height=28
width=28

(data_train,target_train),(data_test,target_test)=mnist.load_data()

# 修改点2:reshape顺序改为(样本数, 高度, 宽度, 通道数)
data_train=data_train.reshape(data_train.shape[0], height, width, channels)
data_test=data_test.reshape(data_test.shape[0], height, width, channels)

features_train=data_train/255
features_test=data_test/255

target_train=np_utils.to_categorical(target_train)
target_test=np_utils.to_categorical(target_test)
number_of_classes=target_test.shape[1]

net=Sequential()

# 修改点3:输入形状改为(高度, 宽度, 通道数)
net.add(Conv2D(filters=64, kernel_size=(5,5), input_shape=(height, width, channels), activation="relu" ))
net.add(MaxPooling2D(pool_size=(2,2)))
net.add(Dropout(0.5))
net.add(Flatten())
net.add(Dense(128,activation="relu"))
net.add(Dropout(0.5))
net.add(Dense(number_of_classes,activation="softmax"))
net.compile(loss="categorical_crossentropy", optimizer="rmsprop",metrics=["accuracy"])
net.fit(features_train,target_train,epochs=2, verbose=1, batch_size=1000,validation_data=(features_test,target_test))

修改后运行即可正常训练,2轮验证集准确率可达97%以上。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 16:06:04