torch.nn 提供带状态的网络模块:有些保存可学习参数,有些保存统计量,也有些只实现固定变换。构建模型时,先确定每一步的数据形状与状态,再选择层;一个很长的类名清单不能代替这些约定。
本文依据 PyTorch 2.8,将原笔记的池化、激活、损失、归一化、序列模块与容器重新分组,并纠正混入内部类型和错误维度的条目。下面两个完整程序可独立在 CPU 运行。
1. Module、参数和容器
跳转到“1. Module、参数和容器”继承 nn.Module 时,先调用 super().__init__(),将子模块作为属性或放入注册容器,并实现 forward()。使用 model(x) 调用模型,可以保留模块调用过程中的钩子等机制,不直接把 model.forward(x) 当作通用替代。Module 文档
| 类型 | 负责什么 | 容易遗漏的条件 |
|---|---|---|
Module | 管理子模块、参数、缓冲区与训练模式 | 普通 Tensor 属性不会仅因是 Tensor 就自动成为参数 |
Parameter | 表示可被模型登记的参数 | 优化器必须实际收到它,才会执行相应更新 |
register_buffer | 登记需要跟随模型设备等操作的非参数状态 | 持久缓冲区通常进入 state_dict,不作为优化器参数 |
Sequential | 顺序组合模块 | 适合前一模块输出直接交给下一模块的链路 |
ModuleList、ModuleDict | 登记列表或字典中的子模块 | 不自动替你定义前向执行顺序,应在 forward 中调用 |
ParameterList、ParameterDict | 登记一组参数 | 与保存子模块的容器不同 |
把多个 Linear 仅放入普通 Python list,不会让父模型自动递归发现它们。其参数可能缺席于 parameters()、设备迁移和权重保存;应使用合适的注册方式。ModuleList
2. 常见空间与形状模块
跳转到“2. 常见空间与形状模块”| 模块族 | 作用 | 关键边界 |
|---|---|---|
Conv1d/2d/3d | 在一个、两个或三个空间/序列轴上应用卷积 | 维数名称不包含批次和通道;检查通道数、步幅、填充与分组 |
ConvTranspose1d/2d/3d | 与卷积线性算子的转置相关,可改变空间大小 | 不是把一般卷积的丢失信息恢复出来的数学逆 |
Linear、Bilinear | 最后一个特征轴上的仿射映射、双输入双线性映射 | Linear 自身不包含 ReLU 等非线性激活 |
AvgPool1d/2d/3d、MaxPool1d/2d/3d | 按核与步幅进行平均或最大值汇聚 | 会丢失信息;输出尺寸不等于张量的维数减少 |
AdaptiveAvgPool1d/2d/3d、AdaptiveMaxPool1d/2d/3d | 直接指定输出空间尺寸 | 通道数不变;区域可能重叠 |
FractionalMaxPool2d/3d、LPPool1d/2d | 随机区域最大池化、Lp 范数形式的池化 | 按各自接口核对输出尺寸或比例、随机性与范数参数 |
MaxUnpool1d/2d/3d | 结合最大池化索引将保留值放回相应位置 | 其他位置补零,无法恢复已被丢弃的输入值 |
Identity | 原样返回输入 | 不是一份独立副本 |
Flatten、Unflatten | 合并指定轴、将指定轴拆为给定尺寸 | nn.Flatten() 默认保留第 0 轴;Unflatten 不记忆过去形状 |
ConstantPad1d/2d/3d | 在边界填常数 | 填充次序与空间轴需要核对 |
ReflectionPad1d/2d/3d、ReplicationPad1d/2d/3d | 反射边界、重复边缘 | 反射并非简单重复边缘,且有输入尺寸限制 |
原表把 MaxUnpool 写成完整逆过程,把 Identity 写成复制,并把 Flatten 概括为一维输出,这些都会导致错误的数据流判断。对照 MaxUnpool2d、ConvTranspose2d、Identity、Flatten。
自适应平均池化的区域推导与可微实现见算子笔记。
3. 激活函数按数值行为理解
跳转到“3. 激活函数按数值行为理解”| 组别 | 原笔记中的接口 | 读参数时的重点 |
|---|---|---|
| 分段线性或截断 | ReLU、ReLU6、LeakyReLU、RReLU、Hardtanh | 负半轴斜率、截断范围,以及训练时是否随机 |
| 指数分支 | ELU、CELU、SELU | 负半轴行为和参数条件;SELU 的自归一化需要网络设置配合 |
| 平滑门控 | GELU、SiLU、Mish、Hardswish | 精确函数与近似形式、数值与计算成本 |
| 有界输出 | Sigmoid、Hardsigmoid、Tanh、Softsign | 输出范围与饱和;不自动等于适当的损失或校准概率 |
| 指数归一化 | Softmax、Softmin | 明确 dim;Softmin 对应 softmax(-x) |
| 平滑/阈值收缩 | Softplus、Softshrink、Hardshrink、Tanhshrink | Softshrink 与 Hardshrink 均需区分正阈值、负阈值及中间区间 |
| 二分门控 | GLU | 沿指定轴拆成两半,计算 a * sigmoid(b),因此该轴缩为一半 |
GLU 本身不含可学习权重;前面的投影层通常负责生成两半特征。LeakyReLU 允许负侧梯度,不代表整个深层网络就不会出现梯度消失。GLU、激活函数目录
4. 归一化先看在哪些轴计算统计量
跳转到“4. 归一化先看在哪些轴计算统计量”| 模块 | 常见输入 | 统计范围与模式 |
|---|---|---|
BatchNorm1d | [N, C] 或 [N, C, L] | 每通道跨批次及相应长度位置统计 |
BatchNorm2d | [N, C, H, W],四维 | 每通道跨 N, H, W 统计 |
BatchNorm3d | [N, C, D, H, W],五维 | 每通道跨批次及三个空间轴统计 |
LayerNorm(normalized_shape) | 任意满足末尾形状的输入 | 对指定的末尾若干轴统计,不一定是整个样本所有特征 |
GroupNorm(groups, channels) | [N, C, ...] | 每个样本内按通道组和空间位置统计,通道数须能整除组数 |
InstanceNorm1d/2d/3d | 按对应接口的批次/非批次形状 | 每个样本每个通道的空间统计;默认不跟踪运行统计量 |
SyncBatchNorm | 分布式训练的特定配置 | 同步进程组内的统计量,不是随便放到多个设备就自动同步 |
BatchNorm 默认在训练时更新运行均值与方差,在 eval() 时使用保存的统计量;设置 track_running_stats=False 会改变这条规则。LayerNorm 和 GroupNorm 通常都使用当前输入统计,不能照抄 BatchNorm 的运行均值解释。BatchNorm2d、LayerNorm、GroupNorm
选层时同时考虑统计轴和有效样本数。GroupNorm 不跨批次求统计,不等于任何退化的组大小或输入形状都有效。
5. 一个从图像到分类损失的完整模型
跳转到“5. 一个从图像到分类损失的完整模型”模型把 [N, 3, H, W] 图像变成 [N, 3] 类别 logits。池化到 1×1 后,Flatten(1) 将通道维留作特征;最后不接 Softmax,直接交给交叉熵。
import torchfrom torch import nn
class SmallClassifier(nn.Module): def __init__(self): super().__init__() self.register_buffer("input_scale", torch.tensor(255.0)) self.features = nn.Sequential( nn.Conv2d(3, 4, kernel_size=3, padding=1), nn.BatchNorm2d(4), nn.ReLU(), nn.AdaptiveAvgPool2d(1), nn.Flatten(1), ) self.dropout = nn.Dropout(p=0.25) self.head = nn.Linear(4, 3)
def forward(self, x): features = self.features(x / self.input_scale) return self.head(self.dropout(features))
torch.manual_seed(5)model = SmallClassifier()assert "head.weight" in dict(model.named_parameters())assert "input_scale" not in dict(model.named_parameters())assert "input_scale" in model.state_dict()assert "features.1.running_mean" in model.state_dict()images = torch.rand(4, 3, 9, 7) * 255.0labels = torch.tensor([0, 1, 2, 1], dtype=torch.long)optimizer = torch.optim.SGD(model.parameters(), lr=0.05)
model.train()optimizer.zero_grad(set_to_none=True)logits = model(images)assert logits.shape == (4, 3)loss = nn.CrossEntropyLoss()(logits, labels)assert torch.isfinite(loss)loss.backward()assert all(p.grad is not None and torch.isfinite(p.grad).all() for p in model.parameters())before_update = model.head.weight.detach().clone()optimizer.step()assert not torch.equal(before_update, model.head.weight)
model.eval()with torch.inference_mode(): first = model(images) second = model(images) torch.testing.assert_close(first, second)assert not first.requires_gradprint("模块登记、形状、一次参数更新与评估模式检查通过")model.eval() 改变模块模式,不会自动关闭梯度计算;no_grad() 和 inference_mode() 控制自动求导行为,也不会自动把模型设为评估模式。Dropout 在训练时随机置零并缩放保留项,在评估时是恒等操作;BatchNorm 的行为还取决于上述统计配置。Dropout、Module 模式
这里的一次更新只验证数据流与参数管理,不是模型准确率实验。
6. 损失函数与相似度分开选择
跳转到“6. 损失函数与相似度分开选择”| 问题 | 相关接口 | 输入或标签约定 |
|---|---|---|
| 二分类、多标签 | BCELoss、BCEWithLogitsLoss、MultiLabelSoftMarginLoss | 前者接概率,带 Logits 的接未变换分数;核对形状和浮点目标 |
| 互斥分类 | CrossEntropyLoss、NLLLoss | 分别接 logits、对数概率;类别维与目标编号范围必须匹配 |
| 数值回归 | MSELoss、L1Loss、HuberLoss、SmoothL1Loss | MSE 比 L1 更放大大误差;Huber/SmoothL1 的阈值和缩放须核对 |
| 分布与计数 | KLDivLoss、GaussianNLLLoss、PoissonNLLLoss | 分别核对 log-space、正方差、计数率或 log-rate 等约定 |
| 无对齐序列标签 | CTCLoss | 输入通常为对数概率,还需要输入与目标长度、blank 编号 |
| 排名与嵌入 | CosineEmbeddingLoss、HingeEmbeddingLoss、MarginRankingLoss、TripletMarginLoss、TripletMarginWithDistanceLoss | 先明确正负目标、间隔和距离函数 |
| 间隔分类 | SoftMarginLoss、MultiMarginLoss、MultiLabelMarginLoss | SoftMarginLoss 是目标为 -1/+1 的二类 logistic 损失,不是普通多类交叉熵 |
| 大类别集合 | AdaptiveLogSoftmaxWithLoss | 针对类别频率与分组的专门设计,不是所有分类的默认替代 |
| 相似度测量 | CosineSimilarity | 计算相似度,本身不等于带目标和归约的训练损失 |
NLLLoss2d 在该版本中是兼容旧名,新内容直接使用支持相应空间维度的 NLLLoss。具体公式、稳定性与形状检查见损失函数。SoftMarginLoss、损失目录
7. 序列、嵌入、重排与并行
跳转到“7. 序列、嵌入、重排与并行”| 模块族 | 用途 | 边界 |
|---|---|---|
RNN、GRU、LSTM | 处理一段序列 | 核对 batch_first;隐藏状态的维度不能简单跟着它交换 |
RNNCell、GRUCell、LSTMCell | 处理一个时间步 | 序列循环与状态传递由调用者组织;基类 RNNBase/RNNCellBase 不作为日常首选入口 |
Embedding、EmbeddingBag | 根据离散索引查找向量、对一组向量聚合 | 索引范围、padding 与各组边界需明确 |
MultiheadAttention、Transformer | 注意力与完整编解码结构 | 嵌入维数、头数、mask 和批次/序列轴要匹配 |
TransformerEncoder/Decoder 及其 Layer | 组合编码器或解码器层 | Layer 与多个层的容器不同;不能仅按名字替换 |
Upsample、UpsamplingBilinear2d、UpsamplingNearest2d | 通过插值改变空间尺寸 | 插值模式、输出大小、align_corners 条件不能混用 |
PixelShuffle、PixelUnshuffle | 在通道和空间维之间重排 | 需要满足倍率对应的整除条件,不凭空增加元素 |
ChannelShuffle | 将通道按组重新排列 | 不改变通道总数 |
Unfold、Fold | 提取局部块、将局部块累加回空间位置 | 重叠位置会累加,不是无条件互逆 |
Dropout1d/2d/3d | 在对应特征结构上做 dropout | 与逐元素 Dropout 的掩码粒度不同 |
AlphaDropout、FeatureAlphaDropout | 配合自归一化网络的 dropout 形式 | 应结合 SELU 等条件理解均值/方差设计 |
DataParallel、DistributedDataParallel | 数据并行训练包装 | 涉及设备、进程、采样器和通信配置;本文 CPU 例子未验证分布式行为 |
接口与参数入口见 RNN、LSTM、Transformer 与 nn 分类目录。
8. 用小例子纠正“复制”和“逆操作”的误读
跳转到“8. 用小例子纠正“复制”和“逆操作”的误读”import torchfrom torch import nn
class RegisteredLayers(nn.Module): def __init__(self): super().__init__() self.layers = nn.ModuleList([nn.Linear(2, 2), nn.Linear(2, 1)])
def forward(self, x): for layer in self.layers: x = layer(x) return x
model = RegisteredLayers()assert "layers.0.weight" in dict(model.named_parameters())assert model(torch.zeros(3, 2)).shape == (3, 1)
x = torch.arange(24, dtype=torch.float64).reshape(2, 3, 4)assert nn.Identity()(x) is xflattened = nn.Flatten()(x)assert flattened.shape == (2, 12)torch.testing.assert_close(nn.Unflatten(1, (3, 4))(flattened), x)
# 最大池化只留下4;反池化按索引放回,不能恢复1、2、3。image = torch.tensor([[[[1., 2.], [3., 4.]]]])pooled, indices = nn.MaxPool2d(2, return_indices=True)(image)restored = nn.MaxUnpool2d(2)(pooled, indices, output_size=image.shape)torch.testing.assert_close(restored, torch.tensor([[[[0., 0.], [0., 4.]]]]))assert not torch.equal(restored, image)
# 两个不同输入被同一有损卷积映射到同一输出,无法仅凭输出唯一恢复。conv = nn.Conv2d(1, 1, kernel_size=1, stride=2, bias=False)with torch.no_grad(): conv.weight.fill_(1)a = torch.tensor([[[[1., 2.], [3., 4.]]]])b = torch.tensor([[[[1., 9.], [8., 7.]]]])assert not torch.equal(a, b)torch.testing.assert_close(conv(a), conv(b))print("注册容器、Identity、Flatten和信息丢失反例检查通过")原表末尾的 Block、Capsule、ClassType、ConcreteModuleType、ConcreteModuleTypeBuilder,以及 ScriptClass、ScriptClassFunction、ScriptDict、ScriptFunction、ScriptList、ScriptMethod、ScriptModule、ScriptModuleSerializer、StaticModule 等名称,不能混称为 torch.nn 的普通网络层。有些属于其他命名空间、脚本系统或内部实现;按实际导入位置和公开文档查询,不能根据名字猜测功能。旧 Container 也不作为新的组合方式,使用 Module、Sequential 或明确的注册容器。
原笔记的 utils 是工具命名空间,参数初始化、梯度处理等工具应按具体函数查询。完成层选择之后,再用张量基础检查轴、类型、设备和共享关系,比单纯增加层数更能排除基础错误。