残差连接必须显式对齐形状,dense连接须用torch.cat(dim=1)拼接通道;混合使用时需确保残差路径取自block输入、dense包含所有前序特征;统一用nn.identity()作恒等映射,避免lambda导致导出或优化问题。

残差连接必须显式相加,不能依赖自动广播
PyTorch 的 torch.nn.Module 不会自动处理张量形状不匹配时的残差相加——哪怕两个张量语义上“应该对齐”。常见错误是主干输出通道数(如 64)和跳跃路径输入通道数(如 32)不一致,直接写 x + residual 会报 RuntimeError: The size of tensor a (32) must match tensor b (64) at non-singleton dimension 1。
实操建议:
- 始终检查
residual.shape == x.shape,尤其在 stride=2 的下采样残差块中; - 用
nn.Conv2d或nn.Sequential(nn.AvgPool2d(2), nn.Conv2d(in_c, out_c, 1))对齐跳跃路径的尺寸与通道; - 避免在
forward中用F.interpolate动态缩放——它引入额外计算且易漏梯度; - 标准 ResNet 实现中,
downsample参数就是为解决此问题而设的,别绕开它。
Dense 连接要用 torch.cat 拼接,但得管好维度和内存
DenseNet 的核心是把前面所有层的输出沿 channel 维度拼接(dim=1),不是堆叠或相加。错用 torch.stack 会导致新增维度,后续卷积直接崩;漏掉 dim=1 参数则默认拼在第 0 维(batch),结果 shape 彻底错乱。
实操建议:
- 所有中间特征存入 list 后,统一用
torch.cat(features, dim=1); - 每层输出通道数要提前规划(如 growth_rate=32),否则拼接后 channel 爆涨,显存瞬间告急;
- 不要在循环里反复
cat:先收集再拼一次,比out = torch.cat([out, new_feat], dim=1)每次都新建 tensor 快得多; - 如果某层输出是
(N, C, H, W),而跳跃来的特征是(N, C', H', W'),必须先插值或卷积对齐空间尺寸——Dense 连接不豁免 shape 检查。
混合残差 + Dense 时,顺序和命名容易混淆
有人想“每个 dense 层内部加残差”,结果把 residual 接在 BN-ReLU-Conv 前,导致 ReLU 把负残差截断,破坏恒等映射性质;也有人把 dense 拼接后的 tensor 直接喂给下一个残差块,却忘了拼接输出通道数已变,downsample 配置没同步更新。
实操建议:
- 残差结构必须满足“变换前 vs 变换后”可相加,即
residual应取自该 block 输入,而非某子模块输出; - Dense 连接的“前序特征”应包含本 block 输入 + 所有上游 dense 输出,别漏掉初始输入;
- 给变量起名带语义:比如
dense_feats、res_input,别全叫x; - 在 forward 开头加
assert x.dim() == 4和assert x.is_contiguous(),能早暴露 view / transpose 引发的隐性 bug。
nn.Identity 是最轻量的残差占位符,别手写 lambda x: x
当某层不需要实际变换,只为了保持模块接口统一(比如 backbone 中某些 stage 可选是否加残差),用 lambda x: x 或 def identity(x): return x 看似简单,但 PyTorch JIT 和 ONNX 导出时可能无法追踪,训练中也可能被优化掉。
实操建议:
- 一律用
nn.Identity()作为恒等映射模块,它是torch.nn官方支持的轻量级实现; - 它参与
model.modules()遍历,方便统一 freeze 或替换; - 搭配
nn.Sequential使用时不会引入额外参数或计算,print(list(model.parameters()))能验证为空; - 注意它不处理 shape 变更——如果真需要“pass-through but reshape”,就得自己写个 wrapper,
Identity不背这个锅。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











