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

如何将MNIST训练图像从(60000,28,28)转换为(60000,16,16)适配Keras全连接层?

解决MNIST图像从28x28缩放到16x16的问题

嘿,我懂你想把MNIST的28×28图像改成16×16再用全连接网络训练的需求——直接reshape肯定不行,那只是打乱像素排列,不是真正的图像缩放。咱们只需要在数据预处理阶段加一步图像尺寸调整,就能搞定!

核心思路

要改变图像的空间尺寸,得用专门的图像缩放工具,比如TensorFlow自带的tf.image.resize(不用额外装库,和你现有的Keras环境兼容)。步骤大概是:

  1. 给图像加上通道维度(因为缩放函数需要4D张量:样本数×高×宽×通道数)
  2. 把图像缩放到16×16
  3. 去掉多余的通道维度,回到3D张量
  4. 调整后续展平的维度和模型输入形状

修改后的完整代码

import keras
import numpy as np
import mnist
import tensorflow as tf  # 新增:导入TensorFlow用于图像缩放
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense
from tensorflow.keras.utils import to_categorical

train_images = mnist.train_images()
train_labels = mnist.train_labels()
test_images = mnist.test_images()
test_labels = mnist.test_labels()

# --- 新增:图像缩放步骤 ---
# 给图像添加通道维度(从(60000,28,28)变成(60000,28,28,1))
train_images = np.expand_dims(train_images, axis=-1)
test_images = np.expand_dims(test_images, axis=-1)

# 缩放到16×16,返回的是float32类型张量,转成numpy数组
train_images = tf.image.resize(train_images, [16, 16]).numpy()
test_images = tf.image.resize(test_images, [16, 16]).numpy()

# 去掉通道维度(回到(60000,16,16))
train_images = np.squeeze(train_images, axis=-1)
test_images = np.squeeze(test_images, axis=-1)
# --- 缩放步骤结束 ---

# 归一化图像(注意现在图像已经是float32,归一化逻辑不变)
train_images = (train_images / 255) - 0.5
test_images = (test_images / 255) - 0.5
print(train_images.shape)  # 现在应该输出(60000, 16, 16)
print(test_images.shape)   # 输出(10000, 16, 16)

# 展平图像:16×16=256,所以维度改成(-1, 256)
train_images = train_images.reshape((-1, 256))
test_images = test_images.reshape((-1, 256))
print(train_images.shape)  # 输出(60000, 256)
print(test_images.shape)   # 输出(10000, 256)

# 构建模型:输入形状改成(256,)
model = Sequential([
  Dense(10, activation='softmax', input_shape=(256,)),
])

# 编译、训练、评估、保存逻辑不变
model.compile(
  optimizer='adam',
  loss='categorical_crossentropy',
  metrics=['accuracy'],
)

model.fit(
  train_images,
  to_categorical(train_labels),
  epochs=5,
  batch_size=32,
)

model.evaluate(
  test_images,
  to_categorical(test_labels)
)

model.save_weights('model.h5')

关键修改点说明

  • 不能用reshape直接改尺寸:reshape只是重新排列元素顺序,不会改变图像的视觉内容,而tf.image.resize会通过插值算法(默认是双线性插值)真正调整图像的空间大小。
  • 通道维度的处理:TensorFlow的图像缩放函数要求输入是4D张量(包含通道),所以我们先加通道、缩放、再去掉通道,回到你需要的3D形状。
  • 模型输入调整:因为图像变成了16×16,展平后是256个特征,所以模型的input_shape要从(784,)改成(256,),否则会报形状不匹配的错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 18:37:44