使用Keras运行MNIST分类代码出现Negative dimension size错误如何解决
报错原因与修复方案
核心原因
你的报错本质是数据格式、卷积层参数与Keras默认配置不匹配,结合了多个拼写/参数错误共同导致:
- TensorFlow后端默认使用
NHWC(样本数、高度、宽度、通道数)的数据排列格式,你把单通道放在了第二维度,和默认格式错位 - 卷积层参数写法属于Keras 1.x版本的旧语法,你错误把卷积步长设为了3,直接导致卷积计算时维度为负
- 存在多处变量名、类名大小写拼写错误
- 使用了已废弃的训练参数
具体修改点
- 修改数据reshape逻辑,把通道维度放在最后,适配NHWC格式:
原代码X_train = X_train.reshape(X_train.shape[0], 1, 28, 28)改为X_train = X_train.reshape(X_train.shape[0], 28, 28, 1),测试集同理修改 - 修正
Sequential导入的大小写,将from keras.models import sequential改为from keras.models import Sequential - 修正卷积层参数写法,新版本
Conv2D(Convolution2D的别名,更常用)的卷积核大小用tuple传入:Convolution2D(32, (3, 3), activation='relu', input_shape=(28,28,1)),第二层卷积同理修改 - 修正训练时的变量名大小写,把
model.fit(X_train, Y_train改为model.fit(X_train, y_train,评估时的Y_test改为y_test - 把废弃参数
nb_epoch改为epochs
修复后完整代码
import numpy as np np.random.seed(123) # 可复现设置 from keras.models import Sequential from keras.layers import Dense, Dropout, Activation, Flatten from keras.layers import Conv2D, MaxPooling2D from keras.utils import np_utils from keras.datasets import mnist # 加载数据集 (X_train, y_train), (X_test, y_test) = mnist.load_data() # 数据预处理:适配NHWC格式,通道放最后 X_train = X_train.reshape(X_train.shape[0], 28, 28, 1) X_test = X_test.reshape(X_test.shape[0], 28, 28, 1) X_train = X_train.astype('float32') X_test = X_test.astype('float32') X_train /= 255 X_test /= 255 # 标签独热编码 y_train = np_utils.to_categorical(y_train, 10) y_test = np_utils.to_categorical(y_test, 10) # 搭建模型 model = Sequential() model.add(Conv2D(32, (3, 3), activation='relu', input_shape=(28,28,1))) model.add(Conv2D(32, (3, 3), activation='relu')) model.add(MaxPooling2D(pool_size=(2,2))) model.add(Dropout(0.25)) model.add(Flatten()) model.add(Dense(128, activation='relu')) model.add(Dropout(0.5)) model.add(Dense(10, activation='softmax')) # 编译与训练 model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy']) model.fit(X_train, y_train, batch_size=32, epochs=10, verbose=1) score = model.evaluate(X_test, y_test, verbose=0) print(f"测试集准确率:{score[1]:.4f}")
验证效果
运行修复后代码,训练10轮后测试集准确率可以达到99%以上,无维度报错。
内容的提问来源于stack exchange,提问作者Tuba Umer
相关产品推荐
相关产品推荐

