将下面代码保存到 chapter05_build_network.py 文件,运行代码查看自己搭建的 MNIST 识别网络结构以及各层的参数信息。
代码如下:
"""
第 5 章:构建 CNN 网络 —— 搭完整的 MNIST 识别网络
运行方式:python chapter05_build_network.py
"""
import torch
import torch.nn as nn
class MNIST_CNN(nn.Module):
"""MNIST 手写数字识别 CNN
输入:batch张 1通道 28×28灰度手写数字图片
输出:batch×10 的logits得分,对应数字0~9,未做softmax
"""
def __init__(self):
# 调用父类nn.Module构造函数,注册网络层,是pytorch模型必须写的
super().__init__()
# ── 第1组:检测低级特征(边缘、角落、线条) ──
# Conv2d 参数:(输入通道, 输出通道, 卷积核大小, padding填充)
# in_channels=1:MNIST灰度图,单通道
# out_channels=32:使用32个3×3卷积核,输出32张特征图
# kernel_size=3:3×3卷积核,捕捉局部相邻像素
# padding=1:图片四周补1圈0,卷积前后宽高保持不变 28→28
self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1)
# MaxPool2d 最大池化,kernel_size=2:取2×2窗口内最大值作为输出
# 步长默认等于kernel_size,宽高减半:28×28 →14×14
# 作用:压缩特征图尺寸、降低计算量,让特征对微小位移更鲁棒
self.pool1 = nn.MaxPool2d(kernel_size=2)
# ── 第2组:检测高级特征(曲线、数字部件、局部形状) ──
# 输入通道32:接收上一层32张特征图;输出通道提升到64,提取更多高阶组合特征
# padding=1保证卷积后尺寸不变 14→14(因为前面已经池化过)
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
# 第二次池化,尺寸再次减半:14×14 →7×7
self.pool2 = nn.MaxPool2d(kernel_size=2)
# ── 全连接层:把卷积提取的特征映射为类别得分,完成最终分类 ──
# 64*7*7:conv2+pool2之后特征图维度(64,7,7),展平后向量长度
# 3136维特征压缩映射到128维中间特征向量
self.fc1 = nn.Linear(64 * 7 * 7, 128)
# 最后一层:128维特征映射到10个输出,对应0‑9十个数字的原始得分logits
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
"""
前向传播逻辑:数据流过网络的计算流程
:param x: 输入张量 x 形状: (batch_size, 1, 28, 28)
:return: logits原始分类得分,shape=(batch_size,10)
"""
# 第一层卷积+ReLU激活;ReLU引入非线性,网络才能学习复杂特征
# 输出shape:(batch, 32, 28, 28)
x = torch.relu(self.conv1(x))
# 最大池化下采样,宽高减半;输出shape:(batch, 32, 14, 14)
x = self.pool1(x)
# 第二层卷积提取高级特征 + ReLU非线性激活
# 输出shape:(batch, 64, 14, 14)
x = torch.relu(self.conv2(x))
# 第二次池化下采样;输出shape:(batch, 64, 7, 7)
x = self.pool2(x)
# 展平操作view:保留第0维batch,后面所有维度合并成一维
# x.size(0) 获取batch大小;-1自动计算剩余维度总大小64*7*7=3136
# 输出shape:(batch, 3136)
x = x.view(x.size(0), -1)
# 第一层全连接+ReLU,做特征融合变换;输出shape:(batch, 128)
x = torch.relu(self.fc1(x))
# 最后全连接输出10类原始得分logits,CrossEntropyLoss会内部做softmax,这里不用手动softmax
# 输出shape:(batch, 10)
x = self.fc2(x)
return x
if __name__ == "__main__":
# 实例化我们搭建好的CNN模型对象
model = MNIST_CNN()
# 打印模型结构,查看所有网络层
print(model)
print()
# 构造模拟输入:fake_batch 模拟 2 张MNIST图片,randn 生成随机正态数据,代替真实图片
# shape=(2, 1, 28, 28):batch_size=2,单通道,28高28宽
fake_batch = torch.randn(2, 1, 28, 28)
# 将模拟数据送入模型执行forward前向传播,得到预测输出
output = model(fake_batch)
print("─" * 50)
print(f"输入形状: {fake_batch.shape}")
print(f"输出形状: {output.shape}")
# detach()切断计算图,tolist()转为Python列表方便打印查看原始得分
print(f"输出内容(第1张图): {output[0].detach().tolist()}")
print()
# 统计模型全部参数的总数量
total = sum(p.numel() for p in model.parameters())
# 统计requires_grad=True的可训练参数,本网络全部参数都参与训练
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"总参数量: {total:,}")
print(f"可训练参数: {trainable:,}")
# 遍历命名参数,打印每一层权重的名称、shape、参数个数,方便分析参数量分布
print("\n各层参数量:")
for name, param in model.named_parameters():
print(f" {name:20s} shape={str(param.shape):30s} count={param.numel():,}")运行代码,输出如下:
MNIST_CNN(
(conv1): Conv2d(1, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(pool1): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
(conv2): Conv2d(32, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(pool2): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
(fc1): Linear(in_features=3136, out_features=128, bias=True)
(fc2): Linear(in_features=128, out_features=10, bias=True)
)
──────────────────────────────────────────────────
输入形状: torch.Size([2, 1, 28, 28])
输出形状: torch.Size([2, 10])
输出内容(第1张图): [-0.03671533614397049, -0.03836964815855026, 0.0932239443063736, 0.156317800283432, -0.21033373475074768,
0.07289294898509979, -0.018034493550658226, 0.18022002279758453, -0.1199900209903717, 0.16377055644989014]
总参数量: 421,642
可训练参数: 421,642
各层参数量:
conv1.weight shape=torch.Size([32, 1, 3, 3]) count=288
conv1.bias shape=torch.Size([32]) count=32
conv2.weight shape=torch.Size([64, 32, 3, 3]) count=18,432
conv2.bias shape=torch.Size([64]) count=64
fc1.weight shape=torch.Size([128, 3136]) count=401,408
fc1.bias shape=torch.Size([128]) count=128
fc2.weight shape=torch.Size([10, 128]) count=1,280
fc2.bias shape=torch.Size([10]) count=10