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

Keras输入矩阵适配问题:MNIST手写数字识别模型报错排查

问题分析与解决

你的代码存在三个核心问题,导致输入28×28矩阵时模型报错:

1. 全连接层输入形状不匹配

Dense层默认对输入的最后一维进行全连接运算。你错误地在input_shape中包含了batch维度(None代表batch),且没有先将28×28的矩阵扁平化,导致全连接层无法正确处理二维特征矩阵。

正确处理方式:

  • 输入形状只保留单样本维度,设为(28, 28)
  • 在第一个Dense层前添加Flatten层,将28×28矩阵转为784维向量,适配全连接层的输入要求

2. 分类任务损失函数错误

手写数字识别是多分类任务,你使用了回归任务的mae(平均绝对误差)损失函数,这会导致模型训练逻辑完全偏离目标。应根据标签类型选择:

  • 若标签是整数(如你的代码中train_labels是0-9的整数),用sparse_categorical_crossentropy
  • 若标签是one-hot编码,用categorical_crossentropy

3. 数据标准化逻辑错误

你分别计算训练集、验证集、测试集的均值和标准差进行标准化,这会引入数据泄露(验证集和测试集的统计信息不应参与训练过程)。正确做法是仅用训练集的均值和标准差标准化所有数据集。


修正后的完整代码
import pickle
import gzip
import numpy as np
from tensorflow import keras
from tensorflow.keras import layers
import tensorflow as tf

def load_data():
    f = gzip.open("./data/mnist.pkl.gz")
    training_data, validation_data, test_data = pickle.load(f, encoding="latin1")
    return training_data, validation_data, test_data

training_data, validation_data, test_data = load_data()    

# 转换为28×28矩阵
train_images = training_data[0].reshape(len(training_data[0]), 28, 28)
# 仅用训练集计算标准化统计量
train_mean = train_images.mean()
train_std = train_images.std()
# 标准化所有数据集
train_images = (train_images - train_mean) / train_std
train_labels = training_data[1]

validation_images = validation_data[0].reshape(len(validation_data[0]), 28, 28)
validation_images = (validation_images - train_mean) / train_std
validation_labels = validation_data[1]

test_images = test_data[0].reshape(len(test_data[0]), 28, 28)
test_images = (test_images - train_mean) / train_std
test_labels = test_data[1]

# 修正模型结构
model = keras.Sequential([
    layers.Input(shape=(28, 28)),  # 输入形状为单样本的维度:28×28
    layers.Flatten(),  # 将二维矩阵扁平为784维向量
    layers.Dense(512, activation='relu'),
    layers.Dropout(rate=0.01),
    layers.Dense(512, activation='relu'),
    layers.Dense(512, activation='relu'),
    layers.Dense(512, activation='relu'),
    layers.Dense(10, activation='softmax'),
])

# 修正损失函数并添加准确率指标
model.compile(
    optimizer='adam',
    loss='sparse_categorical_crossentropy',
    metrics=['accuracy']
)

history = model.fit(
    train_images, train_labels,
    validation_data=(validation_images, validation_labels),
    batch_size=256,
    epochs=20,
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 05:17:03