尧图网站建设 尧图网络
  • 首页
  • 关于我们
  • 服务项目
  • 案例展示
  • 建站流程
  • 资讯中心
  • 联系我们
首页/资讯中心/详情

T1-实现mnist手写数字识别

T1-实现mnist手写数字识别
📅 发布时间:2026/8/1 9:19:02

● 🍨 本文为🔗365天深度学习训练营中的学习记录博客
● 🍖 原作者:K同学啊

一、前期准备

1.设置GPU

import tensorflow as tf gpus = tf.config.list_physical_devices("GPU") if gpus: gpu0 = gpus[0] #如果有多个GPU,仅使用第0个GPU tf.config.experimental.set_memory_growth(gpu0, True) #设置GPU显存用量按需使用 tf.config.set_visible_devices([gpu0],"GPU") print(gpus)

2.导入数据

一种方法是可以直接下载

from tensorflow.keras import datasets, layers, models import matplotlib.pyplot as plt # 导入mnist数据,依次分别为训练集图片、训练集标签、测试集图片、测试集标签 (train_images, train_labels), (test_images, test_labels) = datasets.mnist.load_data()

如果无法连接下载数据可以先下载好数据集在本地加载

import os import gzip import numpy as np def load_mnist_local(path): files = [ 'train-images-idx3-ubyte.gz', 'train-labels-idx1-ubyte.gz', 't10k-images-idx3-ubyte.gz', 't10k-labels-idx1-ubyte.gz' ] with gzip.open(os.path.join(path, files[0]), 'rb') as f: x_train = np.frombuffer(f.read(), np.uint8, offset=16).reshape(-1,28,28) with gzip.open(os.path.join(path, files[1]), 'rb') as f: y_train = np.frombuffer(f.read(), np.uint8, offset=8) with gzip.open(os.path.join(path, files[2]), 'rb') as f: x_test = np.frombuffer(f.read(), np.uint8, offset=16).reshape(-1,28,28) with gzip.open(os.path.join(path, files[3]), 'rb') as f: y_test = np.frombuffer(f.read(), np.uint8, offset=8) return (x_train, y_train), (x_test, y_test) # 这里改成数据集所在位置 data_path = r"C:\Users\asus\Desktop\T1\data\MNIST\raw" (train_images, train_labels), (test_images, test_labels) = load_mnist_local(data_path)

3.归一化

# 将像素的值标准化至0到1的区间内。(对于灰度图片来说,每个像素最大值是255,每个像素最小值是0,也就是直接除以255就可以完成归一化。) train_images, test_images = train_images / 255.0, test_images / 255.0 # 查看数据维数信息 train_images.shape,test_images.shape,train_labels.shape,test_labels.shape

4.数据可视化

# 将数据集前20个图片数据可视化显示 # 进行图像大小为20宽、10长的绘图(单位为英寸inch) plt.figure(figsize=(20,10)) # 遍历MNIST数据集下标数值0~49 for i in range(20): # 将整个figure分成2行10列,绘制第i+1个子图。 plt.subplot(2,10,i+1) # 设置不显示x轴刻度 plt.xticks([]) # 设置不显示y轴刻度 plt.yticks([]) # 设置不显示子图网格线 plt.grid(False) # 图像展示,cmap为颜色图谱,"plt.cm.binary"为matplotlib.cm中的色表 plt.imshow(train_images[i], cmap=plt.cm.binary) # 设置x轴标签显示为图片对应的数字 plt.xlabel(train_labels[i]) # 显示图片 plt.show()

5.调整图片格式

train_images = train_images.reshape((60000, 28, 28, 1)) test_images = test_images.reshape((10000, 28, 28, 1)) train_images.shape,test_images.shape,train_labels.shape,test_labels.shape

二、训练模型

1.构建CNN模型

model = models.Sequential([ # 设置二维卷积层1,设置32个3*3卷积核,activation参数将激活函数设置为ReLu函数,input_shape参数将图层的输入形状设置为(28, 28, 1) # ReLu函数作为激活励函数可以增强判定函数和整个神经网络的非线性特性,而本身并不会改变卷积层 # 相比其它函数来说,ReLU函数更受青睐,这是因为它可以将神经网络的训练速度提升数倍,而并不会对模型的泛化准确度造成显著影响。 layers.Conv2D(32, (3, 3), activation='relu', input_shape=(28, 28, 1)), #池化层1,2*2采样 layers.MaxPooling2D((2, 2)), # 设置二维卷积层2,设置64个3*3卷积核,activation参数将激活函数设置为ReLu函数 layers.Conv2D(64, (3, 3), activation='relu'), #池化层2,2*2采样 layers.MaxPooling2D((2, 2)), layers.Flatten(), #Flatten层,连接卷积层与全连接层 layers.Dense(64, activation='relu'), #全连接层,特征进一步提取,64为输出空间的维数,activation参数将激活函数设置为ReLu函数 layers.Dense(10) #输出层,输出预期结果,10为输出空间的维数 ]) # 打印网络结构 model.summary()

2.编译模型

model.compile( optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=['accuracy'])

3.模型训练

history = model.fit( train_images, train_labels, epochs=10, validation_data=(test_images, test_labels))

import matplotlib.pyplot as plt #隐藏警告 import warnings warnings.filterwarnings("ignore") plt.rcParams['font.sans-serif'] = ['SimHei'] plt.rcParams['axes.unicode_minus'] = False plt.rcParams['figure.dpi'] = 100 from datetime import datetime current_time = datetime.now() epochs_range = range(10) plt.figure(figsize=(12, 3)) plt.subplot(1, 2, 1) plt.plot(epochs_range, train_acc, label='Training Accuracy') plt.plot(epochs_range, test_acc, label='Test Accuracy') plt.legend(loc='lower right') plt.title('Training and Validation Accuracy') plt.xlabel(current_time) plt.subplot(1, 2, 2) plt.plot(epochs_range, train_loss, label='Training Loss') plt.plot(epochs_range, test_loss, label='Test Loss') plt.legend(loc='upper right') plt.title('Training and Validation Loss') plt.show()

三、模型预测

plt.imshow(test_images[1])

pre = model.predict(test_images) # 对所有测试图片进行预测 pre[1] # 输出第一张图片的预测结果

第三项为29.019821,说明预测结果为2,和图片一致。

个人总结:安装tensorflow2的GPU版本时,一定要注意各个库的版本协调,我因为在学深度学习之前就已经安装了python3.12,anaconda和CUDA12.6,python和CUDA的版本都过高无法适配tensorflow2,所以在anaconda里建立了一个pyhon3.9的虚拟环境,在该环境内下载了CUDA11.2和cudnn8.1,最后安装tensorflow2.10。本周使用tensorflow实现mnist手写数字识别,使用的是最简单的CNN模型 LeNet-5,由输入层、卷积层1,池化层1,卷积层2,池化层2,flatten层,全连接层,输出层按序构成。这里的tensorflow2(Keras)是高阶API的写法,能够快速搭建出CNN模型。

相关新闻

  • HMC998A,DC~22GHz 2W 超宽带功率放大器
  • 【MDX】 Markdown 和 JSX 融合
  • 桌面多功能交互终端:硬件调试与多协议配置实践指南

最新新闻

  • 西点培训哪家最靠谱,2026靠谱西点培训学校真实口碑榜,避坑不踩雷 - 工业推荐榜
  • ESP32-S3单路继电器模块:从硬件设计到智能控制节点的完整开发指南
  • 雅思同义词替换技巧:characteristic、feature、property的精准使用
  • 深入掌握pip配置:自定义路径、切换镜像源与离线安装whl文件
  • Unity3D从入门到精通:结构化学习路径与核心模块实战指南
  • 公众号投票怎么做?海投票小程序2026免费搭建流程分享 - 微信投票小程序

日新闻

  • ClickHouse版本管理深度实战:4步构建零风险升级与回滚体系
  • Java 23 种设计模式:从踩坑到精通 | 番外:责任链模式 —— 物流审批流程实战
  • 华硕笔记本性能解放指南:G-Helper轻量级控制工具全面解析

周新闻

  • 大连理工大学与东京大学联手打造的“主动型AI助手“
  • 170.2026年国家级科研瓶颈:超精密单点金刚石切削(SPDT)光学表面生成
  • SongBloom:革命性歌曲生成框架深度解析——如何通过交织自回归与扩散模型创作完整音乐

月新闻

  • ClickHouse版本管理深度实战:4步构建零风险升级与回滚体系
  • Java 23 种设计模式:从踩坑到精通 | 番外:责任链模式 —— 物流审批流程实战
  • 华硕笔记本性能解放指南:G-Helper轻量级控制工具全面解析

关于尧图

  • 公司简介
  • 团队介绍
  • 企业文化
  • 荣誉资质

服务项目

  • 定制开发
  • 电商建站
  • UI 设计
  • 运维服务

快速链接

  • 案例展示
  • 建站流程
  • 常见问题
  • 资讯中心

联系方式

  • 📍北京市朝阳区互联网产业园 A 座 10 层
  • 📞400-888-8888
  • ✉️contact@rkmt.cn
  • 🕐周一至周日 9:00-21:00

© 2024 北京尧图网络科技有限公司 版权所有 | 京 ICP 备 XXXXXXXX 号