搭一个简单网络试试前向传播

🎉摘要:通过完整代码演示PyTorch下全连接层的矩阵运算、ReLU与Sigmoid激活函数效果,并构建一个784-128-64-10的简单网络模拟MNIST手写数字分类的前向传播过程,输出logits与Softmax概率,并统计模型参数量。

将下面代码保存到 chapter04_neural_network.py 文件,运行脚本演示 PyTorch 下全连接层、激活函数、完整神经网络的前向传播计算过程。

完整代码如下:

"""
第 4 章:理解神经网络 —— 搭一个简单网络试试前向传播
本脚本演示 PyTorch 下全连接层、激活函数、完整神经网络的前向传播计算过程
运行方式:python chapter04_neural_network.py
"""
import torch
import torch.nn as nn

# ====================== 1. 演示:全连接层做了什么 ======================
# nn.Linear就是全连接层(线性层),完成矩阵运算 Y = X @ W.T + b
print("=" * 50)
print("  1. 全连接层演示")
print("=" * 50)
# in_features输入特征数,out_features输出特征数;4个输入特征映射为2个输出特征
fc = nn.Linear(in_features=4, out_features=2)  # 4输入 → 2输出

# 打印线性层内部可学习参数:权重weight、偏置bias
# weight权重矩阵:输出维度 × 输入维度,矩阵用于和输入做矩阵乘法
print(f"权重 W 形状: {fc.weight.shape}")  # → (2, 4)  2个输出 × 4个输入
# bias偏置向量:每个输出通道对应一个偏置数值,运算时加到矩阵乘法结果上
print(f"偏置 b 形状: {fc.bias.shape}")    # → (2,)    2个输出各一个偏置

# 构造输入张量:1个样本,样本拥有4个特征
x = torch.tensor([[1.0, 2.0, 3.0, 4.0]])  # 1个样本,4个特征
# 调用全连接层,自动执行 y = x @ W^T + b 线性运算
y = fc(x)
print(f"输入: {x[0].tolist()}")
print(f"输出: {y[0].tolist()}( = x @ W^T + b )")
print()


# ====================== 2. 演示:激活函数的作用 ======================
# 激活函数给网络引入非线性;如果没有激活函数,无论多少层网络等价于单层线性变换
print("=" * 50)
print("  2. 激活函数演示")
print("=" * 50)

# 生成从-5到5,一共11个等间隔数字,包含正负,用来观察激活函数对正负值的处理
x = torch.linspace(-5, 5, steps=11)  # [-5, -4, ..., 4, 5]
print(f"输入:     {x.tolist()}")
# ReLU:小于0直接置0,大于0保持原值;最常用的非线性激活函数
print(f"ReLU后:   {torch.relu(x).tolist()}")
# Sigmoid:把任意实数压缩到0~1之间,早期激活函数,现在少用于隐藏层
print(f"Sigmoid后: {torch.sigmoid(x).tolist()}")
print()


# ====================== 3. 搭一个完整的简单网络 ======================
# 构建多层全连接神经网络,模拟手写数字MNIST分类任务的前向传播流程
print("=" * 50)
print("  3. 简单网络前向传播")
print("=" * 50)

class SimpleNet(nn.Module):
    """3 层全连接网络,用于演示前向传播
    任务:手写数字0~9分类;输入展平后的784维像素,输出10类得分
    继承nn.Module:PyTorch网络基类,自动管理参数、设备、保存加载等能力
    """
    def __init__(self):
        super().__init__()  # 必须调用父类构造函数,初始化module内部逻辑
        # 第一层全连接:输入784(MNIST图片28*28展平),映射到128隐藏神经元
        self.fc1 = nn.Linear(784, 128)  # 784 → 128
        # 第二层全连接:128维隐藏特征映射到64维隐藏特征
        self.fc2 = nn.Linear(128, 64)   # 128 → 64
        # 第三层输出层:64维隐藏特征映射到10维,对应数字0~9共10个类别
        self.fc3 = nn.Linear(64, 10)    # 64 → 10

    def forward(self, x):
        """
        前向传播函数:定义数据流过网络的计算逻辑
        x: 输入张量,shape [batch_size,784]
        return: logits原始分类得分 [batch_size,10]
        """
        # 注意:每一层后都跟一个 ReLU 激活函数,引入非线性,网络才有能力拟合复杂模式
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        x = self.fc3(x)          # 最后一层不跟激活,留给后面的 Softmax;输出logits原始得分
        return x

# 实例化神经网络模型
model = SimpleNet()
print(model)  # 打印网络结构,观察各层信息
print()

# 模拟一张"图片"输入:随机生成张量,batch=1,784个像素值,模拟一张展平后的MNIST图像
fake_image = torch.randn(1, 784)  # 随机生成 784 个像素值
# 执行前向传播,把数据送入网络,自动调用forward函数
output = model(fake_image)
print(f"网络输入形状: {fake_image.shape}")
print(f"网络输出形状: {output.shape}")
print(f"10个输出值(未过Softmax的原始得分,也叫 logits):")
# 遍历10个类别的原始得分,打印简易字符柱状图可视化分数大小
for i, score in enumerate(output[0].tolist()):
    bar = "█" * int(abs(score) * 5)
    print(f"  数字{i}: {score:7.3f} {bar}")

# 过 Softmax:把logits原始得分转换为0‑1之间概率,所有类别概率总和为1;dim=1代表在类别维度做运算
probs = torch.softmax(output, dim=1)
print(f"\nSoftmax 后(概率,加起来=1):")
for i, p in enumerate(probs[0].tolist()):
    print(f"  数字{i}: {p:.4f} ({p*100:.1f}%)")
print(f"  总和: {probs.sum():.4f}")

# 统计模型全部可训练参数量,p.numel()统计一个参数张量元素总个数,累加所有参数
total_params = sum(p.numel() for p in model.parameters())
print(f"\n网络总参数量: {total_params:,}")

运行代码,输出如下:

==================================================
  1. 全连接层演示
==================================================
权重 W 形状: torch.Size([2, 4])
偏置 b 形状: torch.Size([2])
输入: [1.0, 2.0, 3.0, 4.0]
输出: [-0.8024548292160034, -1.2125389575958252]( = x @ W^T + b )

==================================================
  2. 激活函数演示
==================================================
输入:     [-5.0, -4.0, -3.0, -2.0, -1.0, 0.0, 1.0, 2.0, 3.0, 4.0, 5.0]
ReLU后:   [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 5.0]
Sigmoid后: [0.006692850962281227, 0.01798621006309986, 0.04742587357759476,
0.11920291930437088, 0.2689414322376251, 0.5, 0.7310585975646973, 
0.8807970285415649, 0.9525741338729858, 0.9820137619972229, 0.9933071732521057]

==================================================
  3. 简单网络前向传播
==================================================
SimpleNet(
  (fc1): Linear(in_features=784, out_features=128, bias=True)
  (fc2): Linear(in_features=128, out_features=64, bias=True)
  (fc3): Linear(in_features=64, out_features=10, bias=True)
)

网络输入形状: torch.Size([1, 784])
网络输出形状: torch.Size([1, 10])
10个输出值(未过Softmax的原始得分,也叫 logits):
  数字0:  -0.270 █
  数字1:  -0.092 
  数字2:   0.184 
  数字3:   0.208 █
  数字4:  -0.036 
  数字5:  -0.203 █
  数字6:   0.106 
  数字7:  -0.120 
  数字8:  -0.095 
  数字9:   0.044 

Softmax 后(概率,加起来=1):
  数字0: 0.0775 (7.8%)
  数字1: 0.0926 (9.3%)
  数字2: 0.1221 (12.2%)
  数字3: 0.1251 (12.5%)
  数字4: 0.0981 (9.8%)
  数字5: 0.0829 (8.3%)
  数字6: 0.1130 (11.3%)
  数字7: 0.0901 (9.0%)
  数字8: 0.0924 (9.2%)
  数字9: 0.1062 (10.6%)
  总和: 1.0000

网络总参数量: 109,386


  

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