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

训练手写数字识别DNN遇输入维度错误:期望2维却得到(60000,28,28)

解决MNIST手写数字识别DNN的输入维度不匹配问题

嘿,我一眼就揪出问题啦——你在训练模型的时候传错了输入数据!咱们来拆解下:

错误根源

你已经辛辛苦苦把原始的28x28图像(x_train,形状(60000,28,28))扁平化处理成了train_images(形状(60000,784)),也把标签转成了one-hot编码的train_labels,但在调用model.fit()的时候,你却还是传了原始的x_train和y_train!这就导致模型收到的输入和它预期的(784,)维度完全不匹配,自然报错了。

另外还有两个小问题得顺便修正:

  • 你用了binary_crossentropy损失函数,但MNIST是10分类任务,应该用categorical_crossentropy(因为标签已经转成one-hot了)
  • 你把训练数据当成了验证数据,这完全起不到验证模型泛化能力的作用,建议换成测试集test_images和test_labels

修正后的完整代码

import matplotlib.pyplot as plt
import keras
from keras import optimizers
from keras.models import Sequential
from keras.layers import Dense, Activation, Flatten
from keras.datasets import mnist

# 加载数据
(x_train, y_train), (x_test, y_test) = mnist.load_data()

# 数据预处理:扁平化+归一化
train_images = x_train.reshape(60000, 784)
test_images = x_test.reshape(10000, 784)
train_images = train_images.astype('float32') / 255
test_images = test_images.astype('float32') / 255

# 标签转one-hot编码
train_labels = keras.utils.to_categorical(y_train, 10)
test_labels = keras.utils.to_categorical(y_test, 10)

# 构建模型
model = Sequential()
model.add(Dense(512, activation="relu", input_shape=(784,)))
for x in range(0, 10):
    model.add(Dense(512, activation="relu"))
model.add(Dense(10, activation="softmax"))

model.summary()

# 编译模型:修正损失函数
model.compile(optimizer="rmsprop", loss="categorical_crossentropy", metrics=['accuracy'])

# 训练模型:修正输入数据和验证数据
model.fit(
    train_images, train_labels, 
    epochs=100, verbose=2, 
    shuffle=True, 
    validation_data=(test_images, test_labels), 
    steps_per_epoch=10, 
    validation_steps=10, 
    validation_freq=1
)

关键修改点说明

  1. 输入数据替换:把model.fit()里的x_train换成train_images,y_train换成train_labels,确保输入维度和模型第一层的input_shape=(784,)匹配
  2. 损失函数修正:用categorical_crossentropy适配多分类任务,binary_crossentropy是给二分类用的,不适合10分类场景
  3. 验证数据替换:用测试集做验证,能真实反映模型在未见过的数据上的表现,而不是用训练数据自欺欺人

这样修改后,模型就能正常启动训练啦!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 06:42:08