ESRGAN自发表以来一直是图像超分辨率领域的明星模型,它凭借RRDB(Residual-in-Residual Dense Block)结构和相对GAN损失,生成出的纹理细节常常比SRCNN、EDSR这类传统方法细腻得多。不过用过的人基本都遇到过同一个问题:输出图像上会出现一些不自然的纹理,比如规律排列的网格纹、突兀的色块、人物皮肤上凭空多出来的斑点。这些就是典型的GAN伪影。伪影的来源并不单一,但归纳起来主要集中在两个方面:网络结构本身(深度、感受野、残差连接方式)和训练过程的不稳定。本文重点讨论前者,也就是通过网络深度调整和残差连接优化来压制伪影。

一、伪影从哪里来:先搞清RRDB结构的副作用
ESRGAN的生成器由23个RRDB堆叠而成,每个RRDB内部包含3个稠密残差块,每个稠密残差块又有5层卷积,层层相加、层层拼接。这种设计的初衷是扩大感受野、增强特征复用,但也带来了副作用。第一,过深的网络让低频信息和高频噪声在传递过程中不断混合,当判别器把某个噪声模式误判为“真实纹理”时,生成器会变本加厉地放大它,于是你在图像上看到重复出现的纹路。第二,残差分支直接以系数1叠加到主干上,训练前期梯度抖动剧烈,权重容易冲进不稳定的区间,输出画面随之出现色块状失真。
还有一个容易被忽视的因素是beta系数。原始ESRGAN在RRDB的残差连接上乘了一个0.2的缩放因子,用来削弱残差分支的影响、稳定训练。很多复现者在魔改网络时不小心把这个系数去掉或者改成了1,结果伪影明显加重。理解这一点很重要:残差缩放不是可有可无的trick,而是控制信息注入强度的阀门。
可以先做个简单实验验证:分别用beta=0.2和beta=1.0训练同样步数,观察中间输出的图像。beta=1.0的版本通常在第几千步就出现明显噪点,而beta=0.2的版本过渡平滑得多。这个现象直接说明残差连接的强度与伪影高度相关。
二、网络深度调整:不是越深越好
很多复现者的第一反应是“加深网络,细节更多”。但在GAN框架下,加深生成器等于扩大判别器与生成器的博弈空间,一旦二者失衡,伪影会成倍增加。实践中有两个调整方向。
第一个方向是直接减少RRDB数量。对于人脸、动漫这类结构规整的领域,23个RRDB往往过深,16个甚至12个就足够,输出反而更干净。第二个方向是保留深度但引入渐进式训练:先在浅层配置下训练若干epoch,收敛后再逐级添加残差块继续训练。这样网络不需要在一开始就面对巨大的参数空间,梯度更新更平稳。
下面是PyTorch下的残差块数量配置和渐进式加载示例:
import torch
import torch.nn as nn
class RRDB(nn.Module):
def __init__(self, nf=64, gc=32):
super().__init__()
self.conv1 = nn.Conv2d(nf, gc, 3, 1, 1)
self.conv2 = nn.Conv2d(nf + gc, gc, 3, 1, 1)
self.conv3 = nn.Conv2d(nf + 2 * gc, gc, 3, 1, 1)
self.conv4 = nn.Conv2d(nf + 3 * gc, gc, 3, 1, 1)
self.conv5 = nn.Conv2d(nf + 4 * gc, nf, 3, 1, 1)
self.lrelu = nn.LeakyReLU(0.2, inplace=True)
self.beta = 0.2 # 残差缩放系数,压制伪影的关键
def forward(self, x):
x1 = self.lrelu(self.conv1(x))
x2 = self.lrelu(self.conv2(torch.cat((x, x1), 1)))
x3 = self.lrelu(self.conv3(torch.cat((x, x1, x2), 1)))
x4 = self.lrelu(self.conv4(torch.cat((x, x1, x2, x3), 1)))
x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1))
return x + x5 * self.beta
class RRDBNet(nn.Module):
def __init__(self, num_block=16): # 从23降到16,减少伪影
super().__init__()
self.conv_first = nn.Conv2d(3, 64, 3, 1, 1)
self.body = nn.Sequential(*[RRDB() for _ in range(num_block)])
self.conv_body = nn.Conv2d(64, 64, 3, 1, 1)
self.conv_up1 = nn.Conv2d(64, 64, 3, 1, 1)
self.conv_up2 = nn.Conv2d(64, 64, 3, 1, 1)
self.conv_hr = nn.Conv2d(64, 64, 3, 1, 1)
self.conv_last = nn.Conv2d(64, 3, 3, 1, 1)
self.lrelu = nn.LeakyReLU(0.2, inplace=True)
def forward(self, x):
fea = self.conv_first(x)
body_fea = self.conv_body(self.body(fea))
fea = fea + body_fea # 长跳连接,稳定低频信息
fea = self.lrelu(self.conv_up1(
torch.nn.functional.interpolate(fea, scale_factor=2, mode='nearest')))
fea = self.lrelu(self.conv_up2(
torch.nn.functional.interpolate(fea, scale_factor=2, mode='nearest')))
out = self.conv_last(self.lrelu(self.conv_hr(fea)))
return out
注意代码里的两处细节:RRDB内部的beta保留了0.2,主干上的长跳连接(fea + body_fea)也建议保留。长跳连接让原始低分辨率特征绕过整段深层网络直接参与重建,图像的整体结构信息不至于被深层变换扭曲,这对抑制大面积色块伪影非常有效。
深度调整之后记得同步调整学习率。残差块减少后参数量下降,原来1e-4的生成器学习率可以保持,但判别器学习率建议略降到4e-5左右,避免判别器过强把生成器逼进“投机取巧”的状态——那正是伪影滋生的典型环境。
三、残差连接进阶优化:稠密连接与局部特征融合
单纯调beta和深度能解决大部分问题,但如果你的训练数据噪声较多或者目标是4倍放大,还可以对残差连接做进一步改造。核心思路是:不要让每个残差分支都直接跳到输出,而是在中间增加融合层,让网络自己学习各层特征的权重。
具体做法有两种。第一种是局部特征融合(LFFL思路):在稠密块内部,每次拼接特征后先过一个1x1卷积压缩通道,再进入下一层。1x1卷积起到了信息筛选作用,噪声特征在传递中被衰减,伪影自然减少。第二种是加权残差:把固定的beta换成一个可学习的参数,初始化为0.2,训练过程中让模型自己决定残差强度,通常收敛后这个值会稳定在0.1到0.3之间,说明网络确实偏好温和的残差注入。
class RRDBLearnable(nn.Module):
def __init__(self, nf=64, gc=32):
super().__init__()
self.convs = nn.ModuleList([
nn.Conv2d(nf + i * gc, gc, 3, 1, 1) for i in range(4)
])
# 可学习的残差缩放系数,初始0.2,训练中自适应调整
self.beta = nn.Parameter(torch.tensor(0.2))
self.lrelu = nn.LeakyReLU(0.2, inplace=True)
def forward(self, x):
feats = [x]
for conv in self.convs:
feats.append(self.lrelu(conv(torch.cat(feats, 1))))
return x + feats[-1] * self.beta
可学习beta的写法很简单,但效果实测不错。需要注意的是,不要把它初始化得太大(比如0.5以上),否则训练初期就容易出现不稳定,等于回到固定beta=1的老问题。
除了结构层面,训练配置上还有两个搭配建议。一是给判别器加谱归一化,把每层卷积的谱范数约束住,判别器输出的梯度更平滑,生成器收到的信号更稳定,重复网格纹这类伪影会明显减少。二是采用TTUR(两个时间尺度的学习率更新),生成器慢一点、判别器快一点,维持博弈平衡。伪影本质上大多是博弈失衡的产物,结构优化负责降低失衡的概率,训练策略负责在失衡时快速纠正。
最后给一个排查顺序的建议:先确认beta系数是否保留,再检查数据集是否干净(脏数据同样会教出带伪影的模型),然后尝试减少残差块数量,最后才考虑可学习beta和谱归一化这类进阶手段。逐项验证比一次性改一堆参数更容易定位问题根源。