搭完整的 MNIST 识别网络

🎉摘要:通过PyTorch一步步构建MNIST手写数字识别卷积神经网络,包含代码详解、网络结构展示、各层参数量统计,适合初学者实战学习深度学习基础。

将下面代码保存到 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


  

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号