模板骨架:
所有自定义网络都遵循这套模板:
import torch from torch import nn # 1.定义网络类,继承nn.Module class MyNet(nn.Module): def __init__(self): super().__init__() # 必须调用父类构造函数 # --------在这里定义所有网络层/容器(Linear、ReLU、Sequential等)-------- # self.xxx = 层实例 def forward(self, x): # --------在这里写数据流动逻辑:输入x,经过各层计算,返回输出-------- # x = self.xxx(x) return x # 2.测试网络 if __name__ == '__main__': # 构造模拟输入数据 input_data = torch.randn(形状) net = MyNet() # 实例化网络 out = net(input_data) # 前向传播,自动调用forward print(out.shape)搭建神经网络示例:
from torch import nn import torch class FullyConnectedNet(nn.Module): def __init__(self): super().__init__() self.layer = nn.Sequential( nn.Flatten(), # [batch, 1, 28, 28] -> [batch, 784] nn.Linear(28 * 28, 512), # 784 个像素点映射到 512 个隐藏特征 nn.ReLU(), # 增加非线性表达能力 nn.Linear(512, 256), # 继续提取更紧凑的隐藏特征 nn.ReLU(), nn.Linear(256, 128), nn.ReLU(), nn.Linear(128, 10), # 模型输出的原始分数(分数越高,说明模型当前越偏向哪个候选项) ) def forward(self, x): return self.layer(x) if __name__ == '__main__': data = torch.randn(1,1,28,28) net = FullyConnectedNet() output = net(data) print(output)from torch import nn,nn是什么
from torch import nntorch是顶层大模块;nn是 torch 下面的一个子模块(module),全称torch.nn,专门用来搭建神经网络。
等价写法:
import torch nn = torch.nn # 完全一样torch.nn 里面装了什么
里面全是神经网络相关的类、工具:
- 网络层类:
nn.Linear、nn.Conv2d、nn.Flatten、nn.ReLU、nn.Sequential - 模型基类:
nn.Module—— 所有网络 / 层都继承它 - 损失函数类:
nn.CrossEntropyLoss、nn.MSELoss - 容器:
nn.Sequential、nn.ModuleList、nn.ModuleDict - 归一化、dropout、embedding`等等
语法拆解
from torch import nn # 从 torch 包中导入 nn 子模块,当前文件就可以直接写 nn.xxx,不用写 torch.nn.xxxnn:模块对象,不是类,不是函数。nn.Sequential:访问 nn 模块内部的Sequential类nn.Sequential(...):实例化这个类,得到一个网络层对象
层级关系梳理
torch (顶层包) └── nn 子模块(torch.nn) ├── nn.Module 【基类】 ├── nn.Sequential 【类,继承Module】 ├── nn.Linear 【全连接层类】 ├── nn.ReLU 【激活类】 └── nn.CrossEntropyLoss 损失类重点: 所有层
nn.Linear()、nn.ReLU()、nn.Sequential()都是实例化类,返回的实例全部继承自nn.Module。 只有继承nn.Module的对象,放到模型里面,参数才会被自动管理(net.parameters()、cuda、保存加载模型)。
容易混淆对比
| 代码 | 是什么 |
|---|---|
torch | 最顶层包 |
nn | torch.nn,神经网络子模块 |
nn.Module | 类,所有网络组件的父类 |
nn.Sequential | 类,Module 的子类 |
nn.Sequential(...) | 实例化,得到 Module 实例对象 |
self.layer = nn.Sequential(...) | self.layer是模型的属性,存这个 Module 实例 |
小坑提醒
self.layer = nn.Sequential # ❗错误!没有括号,只是把类本身赋值,没有创建对象 self.layer = nn.Sequential() # ✅加括号:实例化生成对象