手写数字识别:MNIST 数据集处理与模型训练
手写数字识别
获取手写数字数据集
from tensorflow.keras.datasets import mnist
import numpy as np
import tensorflow as tf
# 加载 MNIST 数据集
(train_images, train_labels), (test_images, test_labels) = mnist.load_data()
# 训练集数据维度的调整:N H W C
train_images = np.reshape(train_images,
(train_images.shape[0], train_images.shape[1], train_images.shape[2], 1))
# 测试集数据维度的调整:N H W C
test_images = np.reshape(test_images,
(test_images.shape[0], test_images.shape[1], test_images.shape[2], 1))
# 定义获取训练数据的函数
def get_train(size):
# 随机生成要抽样的样本的索引
index = np.random.randint(0, np.shape(train_images)[0], size)
# 将这些数据resize成224*224大小
resized_images = tf.image.resize_with_pad(train_images[index], 224, 224)
# 返回抽取的
return resized_images.numpy(), train_labels[index]
# 定义获取测试数据的函数
def get_test(size):
# 随机生成要抽样的样本的索引
index = np.random.randint(0, np.shape(test_images)[0], size)
# 将这些数据resize成224*224大小
resized_images = tf.image.resize_with_pad(test_images[index], 224, 224)
# 返回抽样的测试样本
return resized_images.numpy(), test_labels[index]
# 获取训练样本和测试样本
train_images, train_labels = get_train(256)
test_images, test_labels = get_test(128)
# 假设 net 是你的模型
# ...
# 指定优化器,损失函数和评价指标
optimizer = tf.keras.optimizers.SGD(learning_rate=0.01, momentum=0.0)
net.compile(optimizer=optimizer, loss='sparse_categorical_crossentropy', metrics=['accuracy'])
# 模型训练:指定训练数据,batchsize,epoch,验证集
net.fit(train_images, train_labels, batch_size=128, epochs=3, verbose=1, validation_split=0.1)
# 指定测试数据
net.evaluate(test_images, test_labels, verbose=1)
代码说明:
- 导入库: 导入必要的库,包括
tensorflow.keras.datasets用于加载 MNIST 数据集、numpy用于数组操作和tensorflow用于模型构建和训练。 - 加载数据集: 使用
mnist.load_data()加载 MNIST 数据集,并分别获取训练集和测试集的图像和标签。 - 数据预处理:
- 调整图像维度:将图像数据从 (N, H, W) 调整为 (N, H, W, C) 的格式,其中 C 为通道数,这里为 1(灰度图像)。
- 随机抽样:定义
get_train()和get_test()函数,用于从训练集和测试集中随机抽样数据。 - 图像缩放:使用
tf.image.resize_with_pad()将图像缩放至 224*224 大小。
- 模型训练:
- 定义优化器、损失函数和评价指标:使用
tf.keras.optimizers.SGD()定义随机梯度下降优化器,使用'sparse_categorical_crossentropy'定义损失函数,使用'accuracy'定义评价指标。 - 训练模型:使用
net.fit()训练模型,指定训练数据、batch 大小、epoch 数、验证集比例等参数。 - 评估模型:使用
net.evaluate()对模型进行评估,指定测试数据。
- 定义优化器、损失函数和评价指标:使用
注意: 以上代码仅提供参考,具体实现需要根据你的模型和需求进行调整。
常见错误:
- 模型未定义:请确保在代码中定义了
net模型。 - 数据维度不匹配:请确保训练集和测试集的维度与模型输入的维度匹配。
- 未导入相关模块:请确保你已导入
tensorflow.keras.datasets、numpy和tensorflow模块。 - 变量未定义:请确保在代码中使用了所有必要的变量。
原文地址: https://www.cveoy.top/t/topic/pdyo 著作权归作者所有。请勿转载和采集!