11#!/usr/bin/env python
22# coding: utf-8
33
4- # 导入必要的库
5- # 导入操作系统模块,用于环境变量设置、路径管理等系统相关操作
4+ # =============================================================================
5+ # 基于 PyTorch 的卷积神经网络(CNN)实现
6+ # 数据集:MNIST 手写数字识别(0-9,共10类)
7+ # 网络结构:2个卷积层 + 2个全连接层 + Dropout正则化
8+ # =============================================================================
9+
10+ # 导入操作系统模块,用于路径管理和环境变量设置
611import os
712
8- # 导入 NumPy,用于高效的数值计算(如矩阵、向量操作 )
13+ # 导入 NumPy,用于高效的数值计算(矩阵、向量操作等 )
914import numpy as np
1015
1116# 导入 PyTorch 主库
1217import torch
1318
14- # 从 PyTorch 中导入神经网络模块(构建模型的基础类和层 )
19+ # 导入神经网络模块(构建模型的基础类和各类网络层 )
1520import torch .nn as nn
1621
17- # 导入常用的函数接口模块,包括激活函数、损失函数等
22+ # 导入函数接口模块,包含激活函数、损失函数等常用操作
1823import torch .nn .functional as F
1924
20- # 从 PyTorch 中导入自动求导模块,用于构建支持梯度的变量
21- from torch .autograd import Variable
22-
23- # 导入数据处理模块,包括 Dataset 封装与 DataLoader 批处理等功能
25+ # 导入数据处理模块,用于封装数据集和批量加载
2426import torch .utils .data as Data
2527
26- # 导入 torchvision 库,它包含了常用的计算机视觉数据集、模型结构和图像处理工具
28+ # 导入 torchvision,包含常用视觉数据集、模型和图像处理工具
2729import torchvision
2830
29- # 设置超参数。超参数(Hyperparameters)是机器学习模型在训练前需要手动设定(或通过算法优化)的配置参数,
30- # 它们不直接从数据中学习,而是控制模型的整体行为和性能。
31- learning_rate = 1e-4 # 学习率:控制参数更新步长
32- keep_prob_rate = 0.7 # Dropout保留神经元的比例:防止过拟合
33- max_epoch = 3 # 训练的总轮数
34- BATCH_SIZE = 50 # 每批训练数据的大小为50:影响内存使用和训练稳定性
31+ # =============================================================================
32+ # 超参数设置
33+ # 超参数:训练前手动指定的配置参数,控制模型行为,不从数据中学习
34+ # =============================================================================
35+ LEARNING_RATE = 1e-4 # 学习率:控制每次参数更新的步长,过大容易震荡,过小收敛慢
36+ KEEP_PROB_RATE = 0.7 # Dropout 保留率:训练时随机保留70%的神经元,防止过拟合
37+ MAX_EPOCH = 3 # 训练轮数:整个数据集被遍历的次数
38+ BATCH_SIZE = 50 # 批大小:每次迭代使用的样本数量,影响内存占用和训练稳定性
3539
36- # 检查是否需要下载 MNIST 数据集
40+ # =============================================================================
41+ # 数据集加载
42+ # =============================================================================
43+
44+ # 检查本地是否已存在 MNIST 数据集,若不存在则自动下载
3745DOWNLOAD_MNIST = False
38- if not (os .path .exists ('./mnist/' )) or not os .listdir ('./mnist/' ):
39- # 如果不存在 mnist 目录或者目录为空,则需要下载
40- DOWNLOAD_MNIST = True
46+ if not os .path .exists ('./mnist/' ) or not os .listdir ('./mnist/' ):
47+ DOWNLOAD_MNIST = True # 目录不存在或为空时,标记为需要下载
4148
4249# 加载训练数据集
4350train_data = torchvision .datasets .MNIST (
44- root = './mnist/' , # 数据集保存路径
45- train = True , # 加载训练集(False则加载测试集)
46- transform = torchvision .transforms .ToTensor (), # 将PIL图像转换为 Tensor 并归一化到 [0,1]
47- download = DOWNLOAD_MNIST # 如果需要则下载
51+ root = './mnist/' , # 数据集本地存储路径
52+ train = True , # True= 加载训练集(60000张),False=加载测试集
53+ transform = torchvision .transforms .ToTensor (), # 将 PIL 图像转为 Tensor,并自动归一化到 [0,1]
54+ download = DOWNLOAD_MNIST # 是否需要从网络下载
4855)
4956
50- # 创建数据加载器,用于批量加载数据
57+ # 创建训练数据加载器(支持批量读取、数据打乱)
5158train_loader = Data .DataLoader (
5259 dataset = train_data , # 使用的数据集
53- batch_size = BATCH_SIZE , # 每批数据量
54- shuffle = True # 是否在每个epoch打乱数据顺序(重要!避免模型学习到顺序信息)
60+ batch_size = BATCH_SIZE , # 每批加载的样本数
61+ shuffle = True # 每个 epoch 开始前打乱数据,避免模型学到顺序规律
5562)
5663
57- # 加载测试数据集(不用于训练,仅用于评估)
58- # torchvision.datasets.MNIST用于加载 MNIST 数据集
59- # root='./mnist/'指定数据集的存储路径
60- # train=False表示加载测试集(而不是训练集)
61- test_data = torchvision .datasets .MNIST (root = './mnist/' , train = False )
62- # 预处理测试数据:转换为 Variable(旧版PyTorch自动求导机制) ,调整维度(原始MNIST是28x28,需要变为1x28x28),转换为FloatTensor类型,归一化到[0,1]范围(/255.),只取前500个样本
63- test_x = Variable (torch .unsqueeze (test_data .test_data , dim = 1 ), volatile = True ).type (torch .FloatTensor )[:500 ]/ 255.
64- # 获取测试集的标签(前500个),并转换为 numpy 数组
65- test_y = test_data .test_labels [:500 ].numpy ()
66-
67- # 定义CNN模型
64+ # 加载测试数据集(仅用于评估,不参与训练)
65+ test_data = torchvision .datasets .MNIST (
66+ root = './mnist/' ,
67+ train = False # False 表示加载测试集(10000张)
68+ )
69+
70+ # 预处理测试数据:
71+ # 1. unsqueeze(dim=1):增加通道维度,从 [N,28,28] 变为 [N,1,28,28](灰度图通道数为1)
72+ # 2. .float() / 255.:转为浮点类型并归一化到 [0,1](像素值原为0-255的整数)
73+ # 3. [:500]:只取前500张用于快速评估
74+ test_x = torch .unsqueeze (test_data .data , dim = 1 ).float ()[:500 ] / 255.
75+
76+ # 获取前500个测试样本的真实标签,转为 numpy 数组(用于准确率计算)
77+ test_y = test_data .targets [:500 ].numpy ()
78+
79+ # =============================================================================
80+ # CNN 模型定义
81+ # 网络结构:
82+ # 输入(1x28x28)
83+ # → 卷积层1(32个3x3卷积核) + BN + ReLU + 最大池化 → (32x14x14)
84+ # → 卷积层2(64个3x3卷积核×2) + BN + ReLU + 最大池化 → (64x7x7)
85+ # → 展平 → 全连接层1(3136→1024) + ReLU + Dropout
86+ # → 全连接层2(1024→10) → 输出10类预测值
87+ # =============================================================================
6888class CNN (nn .Module ):
6989 def __init__ (self ):
70- super (CNN , self ).__init__ () # 调用父类构造函数
71- """
72- 设计特点:
73- 使用小尺寸卷积核(3x3)保留更多局部特征
74- - 批量归一化(BN)加速训练收敛
75- - 最大池化逐步降低空间维度
76- - Dropout层防止过拟合
77- """
78- # 第一个卷积层
90+ super (CNN , self ).__init__ () # 调用父类 nn.Module 的构造函数,必须调用
91+
92+ # ---------- 第一卷积块 ----------
93+ # 输入:1通道(灰度图),28x28
94+ # 输出:32通道,14x14(经过池化后尺寸减半)
7995 self .conv1 = nn .Sequential (
80- nn .Conv2d (1 , 32 , kernel_size = 3 , stride = 1 , padding = 1 ), # 3x3卷积核
81- nn .BatchNorm2d (32 ), # 添加批量归一化
82- nn .ReLU (), # ReLU激活函数,引入非线性,ReLU 函数的公式为 f(x) = max(0, x),可以将负值置为0。
83- nn .MaxPool2d (2 ) # 最大池化,减小特征图尺寸
96+ # 卷积层:输入1通道 → 输出32通道,3x3卷积核,padding=1保持特征图尺寸不变
97+ nn .Conv2d (in_channels = 1 , out_channels = 32 , kernel_size = 3 , stride = 1 , padding = 1 ),
98+ # 批量归一化:对每批数据做归一化,加速训练收敛,提高稳定性
99+ nn .BatchNorm2d (32 ),
100+ # ReLU激活函数:f(x)=max(0,x),引入非线性,使网络能学习复杂特征
101+ nn .ReLU (),
102+ # 最大池化:2x2窗口取最大值,特征图尺寸从28x28变为14x14
103+ nn .MaxPool2d (kernel_size = 2 )
84104 )
85-
86- # 第二个卷积层
105+
106+ # ---------- 第二卷积块 ----------
107+ # 输入:32通道,14x14
108+ # 输出:64通道,7x7(经过池化后尺寸减半)
87109 self .conv2 = nn .Sequential (
88- nn .Conv2d (32 , 64 , kernel_size = 3 , stride = 1 , padding = 1 ), # 3x3卷积核
89- nn .BatchNorm2d (64 ), # 添加批量归一化
90- nn .ReLU (), # ReLU激活函数
91- nn .Conv2d (64 , 64 , kernel_size = 3 , stride = 1 , padding = 1 ), # 增加一层3x3卷积
92- nn .BatchNorm2d (64 ), # 批量归一化,加速训练并提高模型稳定性
93- nn .ReLU (), # ReLU激活函数,引入非线性变换
94- nn .MaxPool2d (2 ) # 最大池化,减小特征图尺寸
110+ # 第一个卷积:32通道 → 64通道,提取更丰富的特征
111+ nn .Conv2d (in_channels = 32 , out_channels = 64 , kernel_size = 3 , stride = 1 , padding = 1 ),
112+ nn .BatchNorm2d (64 ),
113+ nn .ReLU (),
114+ # 第二个卷积:64通道 → 64通道(深度叠加,增强特征提取能力)
115+ nn .Conv2d (in_channels = 64 , out_channels = 64 , kernel_size = 3 , stride = 1 , padding = 1 ),
116+ nn .BatchNorm2d (64 ),
117+ nn .ReLU (),
118+ # 最大池化:特征图尺寸从14x14变为7x7
119+ nn .MaxPool2d (kernel_size = 2 )
95120 )
96-
97- # 第一个全连接层:输入是7*7*64=3136(两次池化后图像尺寸变为7x7),输出1024维
98- self .out1 = nn .Linear (7 * 7 * 64 , 1024 , bias = True )
99-
100- # Dropout层:训练时随机丢弃神经元,防止过拟合
101- self .dropout = nn .Dropout (keep_prob_rate )
102-
103- # 第二个全连接层:1024维输入,10维输出(对应10个数字类别)
104- self .out2 = nn .Linear (1024 , 10 , bias = True )
105-
106- #定义了一个神经网络的前向传播过程,进行特征提取和分类预测
107- def forward (self , x ):
108- x = self .conv1 (x ) # 第一卷积层特征提取,输入 -> 卷积 -> 激活 (ReLU由self.conv1定义)
109- x = self .conv2 (x ) # 第二卷积层特征提取,特征图 -> 卷积 -> 激活
110- x = x .view (x .size (0 ), - 1 ) # 展平张量:保留批量维度,合并其他所有维度
111- out1 = self .out1 (x ) # 第一个全连接层 + 激活函数,线性变换: [B, in_features] -> [B, hidden_features]
112- out1 = F .relu (out1 ) # 应用ReLU激活函数引入非线性
113- out1 = self .dropout (out1 ) # 应用dropout正则化,随机丢弃部分神经元输出
114- out2 = self .out2 (out1 ) # 将上一层输出out1传递给当前层self.out2进行处理
115- return out2
116-
117- # 测试函数 - 评估模型在测试集上的准确率
118- def test (cnn ):
119- global prediction # 使用全局变量prediction保存预测结果
120-
121- # 模型预测:输入测试数据,得到原始输出logits(未归一化的预测值)
122- y_pre = cnn (test_x )
123-
124- # 计算softmax概率分布(将logits转换为概率值,dim=1表示对类别维度做归一化)
125- y_prob = F .softmax (y_pre , dim = 1 )
126-
127- # 获取预测类别:找到每个样本概率最大的类别索引
128- # torch.max返回(最大值, 最大值的索引)
129- _ , pre_index = torch .max (y_prob , 1 )
130-
131- # 调整张量形状为1维向量(例如从[N,1]变为[N])
132- pre_index = pre_index .view (- 1 )
133-
134- # 将预测结果从PyTorch张量转换为numpy数组
135- prediction = pre_index .data .numpy ()
136-
137- # 计算正确预测的数量(预测值与真实标签test_y比较)
138- correct = np .sum (prediction == test_y )
139-
140- # 返回准确率(假设测试集共500个样本)
141- return correct / 500.0
142121
122+ # ---------- 全连接层 ----------
123+ # 展平后的尺寸:64通道 × 7 × 7 = 3136
124+ # 全连接层1:3136 → 1024,进行高层特征整合
125+ self .fc1 = nn .Linear (64 * 7 * 7 , 1024 , bias = True )
143126
127+ # Dropout 正则化层:训练时随机"关闭"部分神经元,防止过拟合
128+ # p=1-KEEP_PROB_RATE 表示随机丢弃的比例(这里是30%)
129+ self .dropout = nn .Dropout (p = 1 - KEEP_PROB_RATE )
130+
131+ # 全连接层2(输出层):1024 → 10,对应10个数字类别(0-9)
132+ self .fc2 = nn .Linear (1024 , 10 , bias = True )
133+
134+ def forward (self , x ):
135+ """
136+ 前向传播:定义数据从输入到输出的计算流程
137+ 参数:
138+ x: 输入张量,形状为 [batch_size, 1, 28, 28]
139+ 返回:
140+ 输出张量,形状为 [batch_size, 10](10个类别的原始分数)
141+ """
142+ x = self .conv1 (x ) # 第一卷积块:[B,1,28,28] → [B,32,14,14]
143+ x = self .conv2 (x ) # 第二卷积块:[B,32,14,14] → [B,64,7,7]
144+ x = x .view (x .size (0 ), - 1 ) # 展平:[B,64,7,7] → [B,3136],保留batch维度
145+ x = self .fc1 (x ) # 全连接1:[B,3136] → [B,1024]
146+ x = F .relu (x ) # ReLU 激活,引入非线性
147+ x = self .dropout (x ) # Dropout 正则化(仅在训练模式下生效)
148+ x = self .fc2 (x ) # 全连接2(输出层):[B,1024] → [B,10]
149+ return x
150+
151+
152+ # =============================================================================
153+ # 测试函数:评估模型在测试集上的准确率
154+ # =============================================================================
155+ def evaluate (cnn ):
156+ """
157+ 评估模型准确率
158+ 参数:
159+ cnn: 训练中的 CNN 模型
160+ 返回:
161+ accuracy: 测试集准确率(0~1之间的浮点数)
162+ """
163+ cnn .eval () # 切换为评估模式(关闭 Dropout 和 BatchNorm 的训练行为)
164+
165+ with torch .no_grad (): # 关闭梯度计算,节省内存,加快推理速度
166+ # 前向传播得到原始输出(logits,未经归一化的预测分数)
167+ y_pre = cnn (test_x )
168+
169+ # 获取预测类别:找每个样本10个类别中分数最高的索引
170+ # torch.max 返回 (最大值, 最大值索引),我们只需要索引
171+ _ , pre_index = torch .max (y_pre , dim = 1 )
172+
173+ # 转为 numpy 数组,与真实标签 test_y 比较
174+ prediction = pre_index .numpy ()
175+
176+ # 计算准确率:预测正确的样本数 / 总样本数
177+ correct = np .sum (prediction == test_y )
178+ accuracy = correct / len (test_y )
179+
180+ cnn .train () # 切回训练模式
181+ return accuracy
182+
183+
184+ # =============================================================================
144185# 训练函数
186+ # =============================================================================
145187def train (cnn ):
146- # 使用Adam优化器,学习率为learning_rate,并添加L2正则化(weight_decay)
147- optimizer = torch .optim .Adam (cnn .parameters (), lr = learning_rate , weight_decay = 1e-4 )
148- # 使用交叉熵损失函数
188+ """
189+ 训练 CNN 模型
190+ 参数:
191+ cnn: 待训练的 CNN 模型
192+ """
193+ # Adam 优化器:自适应学习率优化算法,结合了动量和自适应学习率
194+ # weight_decay=1e-4 是 L2 正则化系数,防止参数过大导致过拟合
195+ optimizer = torch .optim .Adam (cnn .parameters (), lr = LEARNING_RATE , weight_decay = 1e-4 )
196+
197+ # 交叉熵损失函数:适用于多分类任务,内部包含 Softmax 计算
149198 loss_func = nn .CrossEntropyLoss ()
150199
151- # 训练max_epoch轮
152- for epoch in range (max_epoch ):
153- # 遍历训练数据加载器
200+ print ("开始训练..." )
201+ print (f"超参数:学习率={ LEARNING_RATE } , Dropout保留率={ KEEP_PROB_RATE } , "
202+ f"训练轮数={ MAX_EPOCH } , 批大小={ BATCH_SIZE } " )
203+ print ("=" * 60 )
204+
205+ for epoch in range (MAX_EPOCH ):
206+ print (f"\n 第 { epoch + 1 } /{ MAX_EPOCH } 轮训练开始" )
207+
154208 for step , (x_ , y_ ) in enumerate (train_loader ):
155- # 将数据转换为Variable(自动求导需要)
156- x , y = Variable (x_ ), Variable (y_ )
157- output = cnn (x ) # 前向传播得到预测结果
158- loss = loss_func (output , y ) # 计算损失
159- optimizer .zero_grad (set_to_none = True ) # 清空模型参数的梯度缓存,set_to_none=True可减少内存占用
160- loss .backward () # 反向传播计算梯度
161- optimizer .step () # 更新参数
162-
163- # 每20个batch打印一次测试准确率
164- if step != 0 and step % 20 == 0 : # 跳过初始训练前的测试(通常初始准确率无意义)
165- print ("=" * 10 , step , "=" * 5 , "=" * 5 , "测试准确率: " , test (cnn ), "=" * 10 )
166- # step != 0: 跳过初始训练前的测试(通常初始准确率无意义)
209+ # 前向传播:将输入数据传入模型,得到预测结果
210+ output = cnn (x_ )
211+
212+ # 计算损失:预测结果与真实标签之间的差距
213+ loss = loss_func (output , y_ )
214+
215+ # 反向传播三步骤:
216+ optimizer .zero_grad () # 1. 清空上一步的梯度(否则梯度会累加)
217+ loss .backward () # 2. 反向传播,计算各参数的梯度
218+ optimizer .step () # 3. 根据梯度更新参数
219+
220+ # 每隔20个batch打印一次测试准确率
221+ if step != 0 and step % 20 == 0 :
222+ accuracy = evaluate (cnn )
223+ print (f" Epoch { epoch + 1 } | Step { step :4d} | "
224+ f"Loss: { loss .item ():.4f} | 测试准确率: { accuracy :.4f} ({ accuracy * 100 :.2f} %)" )
225+
226+ print ("\n " + "=" * 60 )
227+ print ("训练完成!" )
228+ final_accuracy = evaluate (cnn )
229+ print (f"最终测试准确率:{ final_accuracy :.4f} ({ final_accuracy * 100 :.2f} %)" )
230+
231+
232+ # =============================================================================
167233# 主程序入口
234+ # =============================================================================
168235if __name__ == '__main__' :
169- cnn = CNN () # 创建CNN实例
170- train (cnn ) # 开始训练
236+ # 创建 CNN 模型实例
237+ cnn = CNN ()
238+
239+ # 打印模型结构,方便了解网络层次
240+ print ("模型结构:" )
241+ print (cnn )
242+ print ("=" * 60 )
243+
244+ # 开始训练
245+ train (cnn )
0 commit comments