3Blue1Brown的手写数字图像识别实验复现
深度学习框架这里用到的是Pytorch。一、加载数据一import部分导入一些我们需要的库import torch from torch import nn #包括sigmoid() linear()等宝藏函数捏 from torch.utils.data import DataLoader #数据集加载器主要是为了把数据集切成batch且打乱 from torchvision import datasets,transforms from matplotlib.pyplot as plttorchvision.datasets是常用数据集导入就自动下载包括我们这次所用到的MNISTtorchvision.transforms是为了将原始图像转换成模型能用的格式PIL to float 看下文。二数据集与预处理tf transforms.ToTensor() ToTensor()在这有两个任务 类型转换MNIST原始数据是PIL (Python Imaging Library*其实就是一个图像数据库里的一个格式就是太古早了现在常用Pillow) 格式为0~255这里将其转换为Pytorch张量类型为浮点数。 2.数字缩放0~255缩放至0~1学过神经网络的你懂我意思。 train_datadatasets.MINIST(data,trainTrue,downloadTrue,tansformtf) #训练集 test_datadatasets.MINIST(data,trainFalse,downloadTrue,tansformtf) #测试集Anyway上古数据集官网打不开是正常的^^所以使用torchvision.MNIST下载是行不通的。所以这里用第二种方法其实就是国内的镜像网站⬇https://storage.googleapis.com/cvdf-datasets/mnist/tftransforms.ToTensor() datasets.MNIST.mirrors [https://storage.googleapis.com/cvdf-datasets/mnist/] train_data datasets.MNIST(data, trainTrue, downloadTrue, transformtf) test_datadatasets.MNIST(data,trainFalse,downloadTrue,transformtf)三DataLoader分批操作train_batchDataLoader(train_data,batch_size64,shuffleTrue) text_batchDataLoader(test_data,batch_size1000)训练集有6万张测试集有1万张。batch_size64是指以64张图片为一组分批一个历元epoch需要更新60000/64≈937.5向上取整938次每次取一小批这里是 64 张算梯度并更新一次。这叫小批量随机梯度下降mini-batch SGD是实际训练的标准做法。如果一次用全部6万张不仅占内存而且epoch只更新一次参数测试集无所谓这里1000是自设。shufflen/v.洗牌True意为打乱数据。走的路可能更长但梯度下降更快详见3Blue1Brown视频梯度下降那一部分。四查看一个样本img,label train_data[0] #数据集支持下标访问返回的是一个‘图片张量标签’的元组 print(img.shape,label) plt.imshow(img[0],cmapgray);plt.show() #这里不要害怕咋有那么多函数因为后面更多^^img.shape指的是图片的 [通道高宽]应输出为[1,28,28]img[0]从中取出第0个元素就是该通道内的像素即28×28的像素值是一个二维矩阵因为imshow需要二位灰度或三位彩色的数据所以才要把多余的通道维度去掉。ok数据加载成功进入训练。二、定义网络课程中将其定义为了一个784→16→16→10的网络接下来就让我们进行搭建吧。model nn.Sequential( nn.Flatten(), 把图片拉成一条向量flatten v.使变平使变薄摧毁夷平建筑、城镇或植物 例如:输入形状[N,1,28,28]→输出形状[N,16] N为batch大小64 nn.linear(784,16), #第一隐藏层的线性部分第一个16 nn.sigmoid(), #非线性激活 nn.linear(16,16), #第二个16 nn.sigmoid(), nn.linear(16,10), #输出层对应数字0~9 )nn.Sequential(adj.连续的按顺序的)把括号内的层按照顺序串起来上一层的输出就是下一层的输入向前传播。nn.Flatten()将[1,28,28]展平为784维向量所以[64,1,28,28]→[64,784]而非[50176]。nn.linear(78416)全连接层计算y xW^T b其中W权重是[16,784]16个神经元每个对784个像素各有一个权重 b偏置是[16]每个神经元一个。784*161612560。三、损失函数和优化器解决俩问题① 模型错的有多离谱损失函数②怎么改正模型优化器。我们常说的“梯度下降”是①和②之间的桥梁“反向传播”就是②优化器。代码量就两行但是要吃透loss_fn nn.CrossEntropyLoss() #损失函数 optimizer torch.optim.SGD(model.parameteers(),lr0.5) #优化器一损失函数nn.CrossEtropyLoss() 作用把模型的预测和真实标签比较输出一个数表示模型错得有多厉害。越小越好训练就是想办法让它变小。对应视频里的 cost function代价函数。它内部做了两件事*有一个点注意一下“CrossEtropyLoss”即“交叉熵损失”里面已经包含Softmax所以model内就不要加nn.Softmax了要不然模型白学习了Softmax把10个logits变成总和为1的概率分布取负对数根据真实标签取出对应的概率取负对数算出最终的Loss值。这里注意一下视频用的是平方误差MSE。分类任务里交叉熵更常用因为输出错得很离谱时梯度依然够大训练更快MSE 配 Sigmoid 在这种情况下梯度会变得很小学得慢。可以在后面做对照实验。二优化器torch.optim.SGD(...)作用拿到梯度之后负责更新参数。SGD 是随机梯度下降Stochastic Gradient Descent对应视频里沿着负梯度方向走一小步。权重/偏置原权重/偏置-学习率*梯度model.parameters()告诉优化器要更新哪些参数这里是网络里所有的权重和偏置共13002个lr0.5学习率即每一步走多远可以调整看区别。我这做了个对比学习率分别设了0.5和100结果发现初始损失相近这说明初始损失和学习率无关。而损失率的影响要在训练之后才能看到即第一轮前馈学习率不起作用。三分工关系角色负责model前向计算得到预测loss_fn衡量预测错多少得到损失值loss.backward()*反向传播算出每个参数的梯度optimizer.step()*根据梯度更新参数*下面将用到的函数*注意loss.backward()是计算梯度的即反向传播优化器只管实施。四、训练循环根据计算的梯度让损失真正地降低下来。是全部项目中最核心的部分主要分为两部分简建立评估函数、训练循环。一建立评估函数evaluate()评估的是准确率回答模型在没见过的测试集上一万张图里猜对了多少张返回值是一个0~1的小数。例如100张图猜对了95张就是0.9595%。训练靠的是损失检验看准确率。def evaluate(): model.eval() #同理为model.train()均为切换模式的意思有些层在训练和评估的时候行为不同切换是个好习惯 correct0 #计数器用来记录猜对的数量 with torch.no_grad(): #评估的时候不需要算梯度关掉主要是省内存 for x,y in test_batch: #x是图片y是真实标签在上文分批操作时将测试集分为了1000张一批共有1万张所以取10批 correct (model(x).argmax(1) y).sum().item() return correct /len (test_data)这份代码相对复杂除了注释之外这里主要解释for循环部分modelx每张图经过一次训练之后会得到10个得分数字0~9分别给出一个对应的得分.ardmax1对每张图找出得分最高的是第几个类别例如一个手写数字5对应最多的类别却是3那模型就认为这是数字3。在判断是否与标签相同后会得到1000个True/False将结果sum使用item()将结果从张量转换成数字累加到correct。最后return部分就是在计算并返回准确率。二训练循环主要完成任务前向算损失 → 清梯度 → 反向算梯度 → 更新参数for epoch in range (10): #epoch是把整个训练集完整过一遍10个epoch就是看10遍 model.eval() for x,y in train_batch: #每次取一个batch64一个epoch里会执行938次所以10个epoch总共更新参数约9380次 loss loss_fn(model(x),y) #向前传播 optimizer.zero_grad() #清空上一个batch留下的梯度如果不清backward会默认累加 loss.backward() #反向传播 optimizer.step() #更新参数 print(fepoch{epoch1},loss{loss.item():,4f},test acc {evaluate():.4f}) epoch{epoch1}当前第几轮loss{loss.item():,4f}当前损失值保留四位小数test acc {evaluate():.4f}调用评估函数显示准确率。这里需要注意一下zero_grad()只要放在backward()之前即可放在循环开头或者loss之后都行但不能放在backward()和step()之间否则梯度被清掉白算了。训练结果如下五、可视化恭喜你已经到可视化阶段了^^也恭喜我。胜利就在眼前其实还早最后我们需要知道模型在想什么分为三部分看预测、找错题、看第一层权重这也是我们要用plt包的原因。一看预测结果x,y next(iter(test_batch)) #取一批测试图 model.eval() with torch.no_grad(): pred model(x[:8]).argmax(1) #切片只预测前八张 fig,axes plt.subplots(1,8,figsize(12,2)) #这张图包含18个子图画布长乘宽122enumerate 同时给出下标 i 和格子 ax for i,ax in enumerate(axes): ax.imshow(x[1,0].cmapgray) ax.axis(off) #隐藏坐标轴 ax.set_title(fpred{pred[i].item()}\n true{y[i].item()}) plt.show()iter()即iteration迭代器目的是让DataLoader变成迭代器next()取出第一批二找出预测错的图with torch.no_grad(): #只预测不训练不需要记录梯度 pred_all model(x).argmax(1) wrong (pred_all !y).nonzero().flatten() #预测错的下标.nonzero是取了错了的图的下标但是由于是行列二维的所以需要使用flatten展平 print(这一批错了,len(wrong),张) #识别出来有多少张错的 fig,axesplt.subplot(2,8,figsize(14,4)) #28一共有16张图取错的图中的前16张图 for ax,idx in zip(axes.flat,wrong[:16]) #flat是指将28的格子拉成一排16个只取前面16个zip把两个序列一一配对 ax.imshow(x[idx,0],cmapgray) ax.axis(off) ax.set_title(fpred {pred_all[idx].item()}/true {y[idx].item()},fontsize8) #前者为这张图的预测值后者为这张图的真实值 plt.show本人朽木对这一串代码深感困惑幸得克劳德先生辅助以下面这个小例子为大家显化此代码与诸君共勉pred_all [7, 2, 1, 4, 4]y [7, 2, 1, 9, 4]· wrong [3] ← 第 3 张预测错了→ 取 x[3] 画出来标题 pred 4 / true 9最终结果如下图三看第一层的权重这一部分的目标是为了把第一层神经元各自学到的“权重”还原成16张28*28的小图看看他们长啥样。在此之前我们要先晓得权重图是个啥玩意儿第一层每个神经元对 784 个像素各有一个权重也就是784 个数。784 28×28所以可以把这 784 个数摆回 28×28 的方格里变成一张图。图上某个位置的颜色表示这个神经元对那个像素位置的态度蓝色正权重这个位置有笔画神经元就更兴奋红色负权重这个位置有笔画神经元就被抑制接近白色这个位置对它无所谓一共 16 个神经元所以有 16 张图。Wmodel[1].weight.datach() #形状是[16,784],W最后代表权重的张量 fig ,axes plt.subplots(2,8,figsize(14,4)) #创建画布 for i,ax in enumerate(axes.flat): #enumerate循环时同时给出编号和内容 ax.imshow(W[i].reshape(28,28),cmapRdBu) #.reshape意思是“反折”就是将展平的数字重新反折回28*28的格子RdBu 是红蓝色板 ax.axis(off) plt.show()model[1]需要记得我们modelnn.sequential当中第0层是Flatten第一层是linear[784,16]线性层我们取的就是它.weight就是这一层的权重形状[16,784]即16行每行784个数每一行属于一个神经元.detachv.拆下使分离把权重从带梯度记录的状态中分离出来变成纯数字的张量。参数默认带着梯度追踪信息直接拿去画图会报错所以先detach。最终结果如下好了看到这其实大家也累了主要是我累了古法编程真的蛮痛苦的QVQ这个手写数字图像识别的基础神经网络就到这里。起码到可视化这一步应对作业如果你有、入门神经网络、实操一个自己的项目应该都是可以的了。但是如果你想继续深入学习神经网络建议跟着我继续做对照试验我将在下一篇控制变量进行不断调整让大家查看不同超参数、学习率、激活函数、优化器变化会对我们准确率和速度产生什么其他的变化预告如下超参数学习率、隐藏层大小、epoch、batch_size学习率学习的快慢步长的大小激活函数Sigmoid VS ReLUSigmoid把值压到 01两端很平梯度很小网络深了学得慢ReLU负数变 0正数不变计算简单梯度不容易消失现代网络的默认选择优化器SGD VS AdamSGD所有参数用同一个学习率朴素Adam为每个参数自动调整步长通常收敛更快对学习率没那么敏感。然后我们就能回答这么几个问题学习率过大和过小各自的表现是什么ReLU 和 Sigmoid 谁收敛更快为什么隐藏层变大准确率和速度各发生了什么变化为什么实验要一次只改一个变量我们现在的结构是MLPMulti-Layer Perceptron多层感知机最基础的神经网络。后面我们也会将结构换成CNN卷积网络、RNN/LSTM/GRU、Transformer、自编码器Autoencoder and so on...欢迎大家点赞收藏、改错提问评论和私信我都会看^_^~

相关新闻

JavaScript 的原型与继承——对象之间是怎么扯上关系的

JavaScript 的原型与继承——对象之间是怎么扯上关系的

引言如果说作用域和闭包是 JS 的"空间"问题,那原型和继承就是 JS 的"血缘"问题。JS 没有传统面向对象语言里的类继承体系,它用的是另一套逻辑:原型链。这套逻辑让很多人初学时一头雾水,但一旦想通&#xff0c…

2026/10/11 1:52:42 阅读更多 →
枚举类型全解析:从代码可读性到状态机与设备枚举的编程之道

枚举类型全解析:从代码可读性到状态机与设备枚举的编程之道

枚举(enum)类型算是编程里最被低估的关键字了。很多人一提到 enum 就只想到“给数字起名字”,觉得它无非是让type 1变成type TYPE_A,值没变,只是顺眼了一点。我早先也这么觉得,直到后来在业务代码、算法题…

2026/10/11 1:52:42 阅读更多 →
WPF贝塞尔曲线绘制平滑折线图实战指南

WPF贝塞尔曲线绘制平滑折线图实战指南

简介:本资源是一个基于WPF与C#实现的贝塞尔曲线动态折线图可视化项目,面向.NET桌面开发初学者及图形学实践者,解决传统折线图缺乏平滑过渡与动态量程适配的问题。项目完整封装为RAR压缩包(65KB),共36个文件…

2026/10/11 1:51:41 阅读更多 →

最新新闻

自建Docker镜像仓库完整指南:从选型到落地的踩坑总结

自建Docker镜像仓库完整指南:从选型到落地的踩坑总结

在容器化落地走到一定规模之后,几乎每个团队都会遇到一个绕不开的基础设施问题:镜像仓库。项目标题就四个字“docker镜像仓库”,但真正动手自建过的人都知道,这四个字背后藏着选型、存储、安全、性能、运维一长串的决策链。这篇就…

2026/10/11 2:44:13 阅读更多 →
RoboMaster机器人硬件设计从电源到CAN总线再到电机驱动的排查指南

RoboMaster机器人硬件设计从电源到CAN总线再到电机驱动的排查指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/11 2:44:13 阅读更多 →
错误面板可优化清单记录

错误面板可优化清单记录

代码阅读总结 这是中望CAD插件里质检结果展示的 UserControl(QCResultUserControl),基于WinForm,核心功能: 两个构造:无参构造用于插件启动预创建控件;带dwg路径构造直接加载图纸质检数据UI&…

2026/10/11 2:44:13 阅读更多 →
C语言算法分析

C语言算法分析

本文通过讲解洛谷中珠心算测验的解题思路带编程小白了解C语言算法&#xff0c;同时会介绍一些函数知识和字符的使用方法&#xff0c;希望大家能够通过这篇文章学到更多编程知识&#xff0c;从而可以更好地运行代码。 一、函数名称及作用 <string.h> 常用函数 函数 …

2026/10/11 2:44:13 阅读更多 →
华为鸿蒙免费戒烟工具—小羊戒烟

华为鸿蒙免费戒烟工具—小羊戒烟

午饭刚放下筷子&#xff0c;手又往烟盒那边伸——饭后一支烟像按了开关。有时硬生生忍住了&#xff0c;过一会儿却忘自己撑过几回&#xff1b;周末回想&#xff0c;只剩“好像少抽了”&#xff0c;本周到底比上周少几支、省了多少&#xff0c;说不清。我想把抽了几支、忍住几次…

2026/10/11 2:44:13 阅读更多 →
ArcGIS属性表字段添加与编辑实战:类型选择、计算器及维护指南

ArcGIS属性表字段添加与编辑实战:类型选择、计算器及维护指南

1. 字段类型没选对&#xff0c;后面全是坑&#xff1a;先把数据需求想明白前天帮同事处理一份小区地块数据入库&#xff0c;忙活半小时后发现面积字段精度对不上&#xff0c;明明算好是123.45平方米&#xff0c;属性表里却挂着123.450000001。我问他当时添加字段选了什么类型&a…

2026/10/11 2:43:12 阅读更多 →

日新闻

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

简介&#xff1a;基于 ARIMA、LSTM、Transformer 等模型的流感时间序列预测 Python 源码&#xff0c;面向计算机相关专业课程设计与期末大作业学生&#xff0c;以及项目实战学习者。内容覆盖预处理、平稳性检验、定阶、残差分析、多模型对比预测的完整时序建模流程&#xff0c;…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程&#xff1a;键盘模拟输入实战——输入文本与模拟按键的区别 做影刀RPA自动化&#xff0c;十个新手有八个栽在"往输入框里填东西"这件事上&#xff1a;要么填不进去&#xff0c;要么填了一半&#xff0c;要么直接把原来内容追加在后面。这背后的根因&…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程&#xff1a;阅文起点小说数据采集实战——书籍信息与章节内容 1. 认识影刀&#xff1a;什么场景该用RPA采小说数据 起点中文网的页面结构相对稳定——分类榜单、书籍详情、章节内容三块独立页面&#xff0c;跳转链路清晰。这种场景非常适合影刀自动化&#x…

2026/10/11 0:00:27 阅读更多 →

周新闻

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

简介&#xff1a;基于 ARIMA、LSTM、Transformer 等模型的流感时间序列预测 Python 源码&#xff0c;面向计算机相关专业课程设计与期末大作业学生&#xff0c;以及项目实战学习者。内容覆盖预处理、平稳性检验、定阶、残差分析、多模型对比预测的完整时序建模流程&#xff0c;…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程&#xff1a;键盘模拟输入实战——输入文本与模拟按键的区别 做影刀RPA自动化&#xff0c;十个新手有八个栽在"往输入框里填东西"这件事上&#xff1a;要么填不进去&#xff0c;要么填了一半&#xff0c;要么直接把原来内容追加在后面。这背后的根因&…

2026/10/11 0:00:27 阅读更多 →
影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程&#xff1a;阅文起点小说数据采集实战——书籍信息与章节内容 1. 认识影刀&#xff1a;什么场景该用RPA采小说数据 起点中文网的页面结构相对稳定——分类榜单、书籍详情、章节内容三块独立页面&#xff0c;跳转链路清晰。这种场景非常适合影刀自动化&#x…

2026/10/11 0:00:27 阅读更多 →

月新闻

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/10 5:23:50 阅读更多 →
Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证

Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/9 21:32:20 阅读更多 →
黑夜航拍船只数据集训练YOLOV5模型全流程解析

黑夜航拍船只数据集训练YOLOV5模型全流程解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/10 10:38:42 阅读更多 →