RDN(Residual Dense Network,残差密集网络)凭借密集连接和特征复用能力,在超分辨率、去噪等视觉任务中表现出色。但当网络堆叠到几十个残差密集块(RDB)时,训练常常出现损失下降缓慢、浅层参数几乎不更新等典型症状,这就是梯度消失问题。本文从梯度传播的数学机制出发,重点讲解两种工程上最有效的手段:残差缩放与权重初始化,并给出完整代码。

为什么RDN容易出现梯度消失
RDN的核心结构是残差密集块,每个块内部包含多层卷积,每层的输入都拼接了之前所有层的输出特征图。以一个8层的RDB为例,第l层的输入通道数是l乘以增长率,这种设计虽然增强了特征复用,但也让反向传播路径变得非常复杂。
从链式法则来看,损失函数对浅层参数的梯度需要经过大量乘法运算才能传回去。假设每一层的局部梯度幅值平均为0.8,经过20层传播后梯度幅值只剩0.8的20次方,约0.01,浅层基本学不到东西。反之如果局部梯度普遍大于1,梯度会指数级膨胀,造成梯度爆炸。RDB中的1x1卷积压缩层和局部门控机制(类似SE模块)进一步拉长了链路,问题会更明显。
实际训练中的典型表现包括:浅层卷积核的梯度范数接近于零、损失曲线在前几个epoch后长期停滞、加权初始化改变后训练结果差异巨大。如果你观察到这些现象,基本可以确认是梯度消失而非数据或学习率的问题。
方案一:残差缩放的原理与实现
残差缩放的思想很直接:残差分支学到的特征先乘以一个小于1的缩放因子,再与主干特征相加。这样每经过一个残差块,主干特征的幅值增长是受控的,前向传播的数值不会逐层放大,反向传播时梯度也不会因为特征幅值过大而失真。
用公式表达就是:输出等于输入加上beta乘以F(x),其中F(x)是残差分支的映射,beta通常取0.1到0.2。这个思路与ResNet中Stochastic Depth、Fixup等技巧一脉相承,但在RDN这种密集连接结构里尤其重要,因为每个RDB输出的特征还要继续参与后续拼接,不缩放的话特征幅值会滚雪球式增长。
下面是带残差缩放的残差密集块的PyTorch实现:
import torch
import torch.nn as nn
import torch.nn.functional as F
class ResidualDenseBlock(nn.Module):
def __init__(self, nf=64, gc=32, n_layers=6, res_scale=0.2):
super().__init__()
self.res_scale = res_scale
self.convs = nn.ModuleList()
for i in range(n_layers):
in_ch = nf + i * gc
self.convs.append(nn.Conv2d(in_ch, gc, 3, padding=1))
def forward(self, x):
feats = [x]
for conv in self.convs:
y = conv(torch.cat(feats, dim=1))
feats.append(F.leaky_relu(y, negative_slope=0.2))
# 残差缩放:控制特征幅值增长,稳定梯度传播
return x + self.res_scale * feats[-1]
需要注意一个细节:缩放因子如果用除法形式写在初始化里(即直接把卷积权重除以scale),与在forward里乘以res_scale在数学上等价,但后者更直观,也方便后续调参。ESRGAN的官方实现采用的正是0.2这个默认值,训练深层生成器时非常稳定。
res_scale的选择需要权衡:取值太小(如0.05)会让残差分支贡献过弱,网络需要更多epoch才能拟合;取值太大(如0.5)则失去缩放意义,深层网络仍可能不稳定。经验上0.1到0.25是安全区间,层数越多取值应越保守。一个实用的判断方法是训练几个epoch后打印各层激活的标准差,正常情况下应该维持在0.5到2之间,逐层不应明显递增。
方案二:正确的权重初始化策略
不当的初始化是梯度消失的另一大源头。如果权重初始方差过小,前向信号逐层衰减,反向梯度同样衰减;方差过大则激活值饱和(对ReLU类激活尤为明显),梯度又会趋近于零。RDN中通道数随层数变化,统一的初始化方式往往不合适。
对于使用Leaky ReLU的卷积层,推荐He初始化(在PyTorch中对应kaiming_normal_,并指定nonlinearity为leaky_relu)。He初始化将方差设为fan_in除以2(针对ReLU的一半有效区域),保证前向传播中激活方差不衰减。关键在于初始化要针对每层的实际输入通道数动态计算,而不是用一个固定标准差:
def init_weights(module, scale=0.1):
if isinstance(module, nn.Conv2d):
nn.init.kaiming_normal_(module.weight, a=0.0, mode='fan_in', nonlinearity='leaky_relu')
if module.bias is not None:
nn.init.constant_(module.bias, 0.0)
model = ResidualDenseBlock()
model.apply(init_weights)
# 如需进一步抑制初始输出幅值,可额外缩放权重
for m in model.modules():
if isinstance(m, nn.Conv2d):
m.weight.data.mul_(scale)
上面代码里的额外乘以scale是一种简化版的自归一化手段,效果类似残差缩放,但它是静态的。更严谨的做法是结合缩放因子修正方差:如果残差输出要乘以beta,那么卷积权重的初始化方差应除以beta的平方,让缩放后的输出方差依然保持在1附近。这种推导来自Fixup论文的思想,适用于不引入BatchNorm的深层网络(RDN通常就是无归一化层的结构,因为超分任务中BatchNorm会引入伪影)。
还要提醒一个常见误区:不少人在整个网络上统一使用std等于0.01的随机初始化,这对单层网络勉强可行,对几十层的RDN是灾难性的。标准差0.01意味着输出方差约为输入的0.0001倍,信号不到十层就消失了。务必按fan_in计算初始化方差。
联合调优与训练监控建议
残差缩放和权重初始化并不是二选一,两者作用在不同环节:初始化决定训练起点的信号分布,残差缩放决定训练过程中的动态平衡。实际工程中建议先固定res_scale为0.2,用He初始化跑通训练,再根据梯度统计微调。
监控梯度是否健康,可以在每个epoch结束后注册钩子统计各模块梯度的L2范数:
stats = {}
def hook(module, grad_in, grad_out):
stats[module._get_name()] = grad_out[0].norm().item()
for name, m in model.named_modules():
if isinstance(m, nn.Conv2d):
m.register_backward_hook(hook)
loss.backward()
for k, v in stats.items():
print(k, round(v, 6))
健康的表现是:浅层与深层的梯度范数在同一数量级内(相差不超过一个数量级)。如果浅层范数持续比深层小两个数量级以上,优先减小res_scale或检查激活函数选择;如果梯度范数整体剧烈震荡,考虑加梯度裁剪(clip_grad_norm_设为0.5到1.0)配合上述两项手段。
最后总结一下实践要点:RDB层数超过10层时务必启用残差缩放,取值0.1到0.2起步;初始化统一采用按fan_in计算的kaiming方法;避免在超分任务的网络中使用BatchNorm;训练初期打开梯度统计确认浅层参数在被有效更新。做到这几点,即使几十层的RDN也能稳定收敛,训练速度和最终指标都会有明显改善。