资讯详情

资讯详情

基于TensorFlow 2的SRGAN超分辨率图像增强实战

简介基于TensorFlow 2.5与Keras搭建的SRGAN超分辨率生成对抗网络项目面向深度学习研究者与图像处理开发者用于从低分辨率图像恢复出纹理清晰的高分辨率图像同时支持自定义数据集参与训练可针对特定类型或风格图像强化细节并提升真实感。压缩包内共17个文件、总大小962KB包含4个Python脚本生成器与判别器结构、数据加载、预测流程等、4个Markdown说明文档、YAML超参数配置、txt操作指南、docx附赠资料以及示例图片基本覆盖数据准备、模型训练、推理验证全套流程。目前已有70人学习下载项目结构清晰各模块职责分明便于二次开发。使用者可借助完整源码快速理解SRGAN中对抗训练与感知损失的作用在此基础上调整网络参数或引入自有数据集打造更贴合业务场景的超分辨率方案。1. SRGAN 超分辨率项目为什么说它比传统插值更值得做同样是放大一张模糊图片双三次插值只是让像素变多细节依然是糊的而基于TensorFlow 2.x和Keras实现的SRGAN会让模型在训练中学会脑补出原本不存在的高频纹理比如头发丝、皮肤毛孔和叶片脉络。这种用生成对抗网络做超分辨率的技术项目落地时最吸引人的一点是它不依赖任何外部高清数据库你自己收集一批图片做成低分辨率-高分辨率配对样本就能训练出一个针对你数据场景的放大模型。适合的人群很明确手上有模糊图像想复原的开发者、做图像算法方向课题的人、以及想在毕业设计里展示一个能跑通也能讲清原理的深度学习项目。我接下来按网络搭建→数据制作→训练调参→排错的顺序把这个项目完整拆开讲。2. 生成器与判别器的Keras实现把SRGAN两个核心网络搭出来SRGAN的核心思想并不复杂一个生成器负责把低分辨率图放大并补出细节一个判别器负责判断这张图是真实高清图还是生成器伪造的。两者对抗训练生成器越来越会骗判别器越来越会辨最终生成器产出的图片在感知上接近真实高清图。用TensorFlow 2.x的Keras API做这件事网络结构的代码量其实不大真正花时间的是结构细节的取舍。2.1 生成器16个残差块与亚像素卷积上采样生成器在SRGAN原方案里由浅层特征提取16个残差块上采样模块构成。残差块的教学价值在于它让梯度能直接跨层回传训练更稳也避免了网络加深后细节被抹平。上采样部分我用的是亚像素卷积PixelShuffle也就是先通过卷积把通道数变成原来的4倍再把通道重新排列成空间像素相比直接反卷积棋盘伪影要少很多。import tensorflow as tf from tensorflow.keras import layers, Model def residual_block(x, filters64): shortcut x x layers.Conv2D(filters, 3, paddingsame)(x) x layers.BatchNormalization()(x) x layers.PReLU(shared_axes[1, 2])(x) x layers.Conv2D(filters, 3, paddingsame)(x) x layers.BatchNormalization()(x) return layers.Add()([shortcut, x]) def upsample_block(x, filters256, scale2): # 亚像素卷积卷积输出 shape 为 (h, w, filters*scale*scale) x layers.Conv2D(filters * scale * scale, 3, paddingsame)(x) x layers.BatchNormalization()(x) # depth_to_space 把通道重排到空间维度实现 2 倍放大 x layers.Lambda(lambda t: tf.nn.depth_to_space(t, scale))(x) x layers.PReLU(shared_axes[1, 2])(x) return x def build_generator(hr_size384, scale4): lr_size hr_size // scale lr_input layers.Input(shape(lr_size, lr_size, 3)) x layers.Conv2D(64, 9, paddingsame)(lr_input) x layers.PReLU(shared_axes[1, 2])(x) shortcut x for _ in range(16): x residual_block(x) x layers.Conv2D(64, 3, paddingsame)(x) x layers.BatchNormalization()(x) x layers.Add()([shortcut, x]) # 4 倍放大分两步每步 2 倍 x upsample_block(x, 256, scale2) x upsample_block(x, 256, scale2) sr_output layers.Conv2D(3, 9, paddingsame, activationtanh)(x) return Model(lr_input, sr_output, namesrgan_generator)代码里几个地方值得注意。第一所有卷积层的padding都设为same保证特征图尺寸不缩水残差连接要求输入输出shape完全一致。第二PReLU的shared_axes参数意味着同一个通道共享一个斜率参数这样做能让参数量小一些训练也稳定。第三上采样没有用UpSampling2D再跟卷积而是depth_to_space重排这是SRGAN原结构的核心做法。最后一个9x9的卷积把特征映射回RGB三通道activation用tanh是因为输入归一化到[-1,1]输出范围必须匹配。原始项目里hr_size也可以按你数据实际情况改成256但要保证能被4整除。2.2 判别器从VGG思路借来的下采样分类网络判别器的结构比生成器简单它只需要输出一个0到1之间的真实性分数。常见做法是参考VGG的堆叠思路用stride为2的卷积逐层减半图片分辨率同时通道数从64翻倍到512。激活函数用LeakyReLU斜率为0.2避免负区间梯度完全消失。def build_discriminator(hr_size384): hr_input layers.Input(shape(hr_size, hr_size, 3)) x layers.Conv2D(64, 3, paddingsame)(hr_input) x layers.LeakyReLU(alpha0.2)(x) filters 64 for i in range(7): if i % 2 1: filters * 2 x layers.Conv2D(filters, 3, strides2, paddingsame)(x) x layers.BatchNormalization()(x) x layers.LeakyReLU(alpha0.2)(x) x layers.Flatten()(x) x layers.Dropout(0.4)(x) x layers.Dense(1024)(x) x layers.LeakyReLU(alpha0.2)(x) x layers.Dense(1, activationsigmoid)(x) return Model(hr_input, x, namesrgan_discriminator)注意第一层卷积后面没有接BatchNormalization这是我踩过坑后留下的习惯GAN输入层直接接BN容易让训练早期的梯度不稳定尤其是当真实图和生成图分布差异大的时候。Dropout放在全连接层之前目的是给判别器加一点随机性防止它对训练集中某几张图过拟合从而过早压制生成器。判别器输入尺寸必须等于HR尺寸也就是384x384如果你的数据集没做中心裁剪至少要做resize否则shape不匹配直接报错。2.3 把生成器、判别器和VGG特征提取器组装成训练模型单独定义好两个网络还不够训练GAN需要把生成器和判别器接成两个训练步骤。另外SRGAN的损失里包含感知损失需要借助VGG19提取中间层特征所以还得把VGG19加载进来只取前几个卷积层的输出。def build_vgg_feature_extractor(): vgg tf.keras.applications.VGG19(include_topFalse, weightsimagenet) # block5_conv4 是 VGG19 第五个卷积块里第4次卷积的激活输出 feature_layer vgg.get_layer(block5_conv4).output return Model(vgg.input, feature_layer, namevgg_features) def build_combined(generator, discriminator, vgg, lr_size96, hr_size384): lr_input layers.Input(shape(lr_size, lr_size, 3)) hr_input layers.Input(shape(hr_size, hr_size, 3)) sr generator(lr_input) validity discriminator(sr) sr_features vgg(sr) hr_features vgg(hr_input) return Model([lr_input, hr_input], [validity, sr_features, hr_features], namecombined)这里的核心设计是复合模型同时输出三个东西判别器对生成图的真实性打分、生成图的VGG特征、真实图的VGG特征。这样设计是为了在训练生成器时一次前向计算出对抗损失和感知损失。VGG19需要把输入归一化到[-1,1]这正好和生成器tanh输出范围一致。注意VGG19的权重是在ImageNet上预训练好的如果你做的是灰度图像或医学影像建议保留预训练权重只把感知损失当作纹理特征距离来用这比随机初始化效果好很多。3. 自定义数据集准备从图片文件夹做出一对训练样本SRGAN项目的关键一环不是网络结构而是数据怎么变成低分辨率-高分辨率配对。原版SRGAN用的是ImageNet子集每张图随机裁剪成384x384的高分辨率块再通过bicubic下采样生成96x96的低分辨率图。你自己的数据集完全可以照这个套路走但有几个细节决定了训练效果。3.1 图片数据从哪来、怎么划分训练验证集数据来源决定模型最终适用的场景。比如你想修老照片就收集黑白或偏色的人像想放大商品图片就拍或下载白底产品图。总量不用贪多我做过的模拟项目X里用800张图训练效果已经可用但至少要保证500张以上并且内容足够多样化。图片不要全部来自同一个设备或同一个场景否则GAN会过拟合到那个场景的色调和纹理分布。划分方式是先分出10%作为验证集这部分图片在训练阶段完全不参与只在评估阶段跑一次生成器计算PSNR/SSIM或肉眼观察。剩下的90%参与训练然后用随机划分的方式建一个文件列表比直接遍历文件夹更可控。import os, random def split_dataset(image_dir, val_ratio0.1, seed42): paths [os.path.join(image_dir, f) for f in os.listdir(image_dir) if f.lower().endswith((.png, .jpg, .jpeg, .bmp))] random.Random(seed).shuffle(paths) val_num int(len(paths) * val_ratio) train_paths paths[val_num:] val_paths paths[:val_num] return train_paths, val_paths代码里的seed参数很重要保证每次运行划分结果一致方便复现。如果你后续想对比不同的损失函数或网络改动训练集和验证集必须固定否则指标变化说不清是模型变了还是数据变了。另一个容易踩的坑是图片里混入了RGBA格式的PNGdecode时要用channels3强制转成RGB否则某些带透明通道的图会让训练数据多出一个维度。3.2 低分辨率配对样本的生成bicubic下采样不是简单resizeSRGAN里低分辨率图不是独立找的而是从高清图主动降采样得到的。常见做法是把原始图先统一resize到一个稍大的尺寸再从中裁剪出HR块随后用bicubic插值缩小到LR尺寸。为什么不用最近邻或双线性因为SRGAN论文里就是用bicubic模拟真实低分辨率退化而且训练时生成的LR越接近你实际应用场景的退化方式效果越好。def load_pair(image_path, lr_size96, hr_size384): hr tf.io.read_file(image_path) hr tf.image.decode_image(hr, channels3, expand_animationsFalse) hr tf.image.resize(hr, (hr_size, hr_size), methodbicubic) # 先缩小再归一化保证 LR 和 HR 场景内容一致 lr tf.image.resize(hr, (lr_size, lr_size), methodbicubic) # 归一化到 [-1, 1]匹配生成器 tanh 输出 lr (lr / 127.5) - 1.0 hr (hr / 127.5) - 1.0 return lr, hr这里有个容易忽视的点先裁HR再生成LR实际上等价于对高清内容做降采样原图本身很小时强行resize到384x384会引入假细节模型学到的是插值痕迹而不是真实纹理。所以我一般会在数据准备阶段加一步如果原图短边小于hr_size就直接丢弃或者先放大到hr_size的1.2倍再裁剪。另外不要在归一化之后做resize因为插值算法是假设在非归一化空间工作的先resize再去均值才是稳定流程。3.3 tf.data管道裁剪、翻转与数据增强训练SRGAN吃显存很凶如果每张图都直接以384x384送入网络batch size稍微大一点就OOM。常见做法是离线把训练图切成多个384x384的小块或者在线上随机裁剪。线下裁剪的好处是数据管道的预处理压力小缺点是每张图的边缘区域被重复利用的概率低线上裁剪的好处是数据增强更灵活坏处是CPU预处理会成为瓶颈。def preprocess_for_train(image_path, lr_size96, hr_size384): hr tf.io.read_file(image_path) hr tf.image.decode_image(hr, channels3, expand_animationsFalse) hr tf.image.random_crop(hr, size(hr_size, hr_size, 3)) # 随机水平翻转相当于把训练数据量翻倍 hr tf.image.random_flip_left_right(hr) # 随机旋转90度的整数倍也属于常用增强 hr tf.image.rot90(hr, ktf.random.uniform((), 0, 4, dtypetf.int32)) lr tf.image.resize(hr, (lr_size, lr_size), methodbicubic) hr tf.cast(hr, tf.float32) / 127.5 - 1.0 lr tf.cast(lr, tf.float32) / 127.5 - 1.0 return lr, hr def build_tf_dataset(paths, batch_size16, lr_size96, hr_size384): ds tf.data.Dataset.from_tensor_slices(paths) ds ds.map(lambda p: preprocess_for_train(p, lr_size, hr_size), num_parallel_callstf.data.AUTOTUNE) ds ds.shuffle(256).batch(batch_size).prefetch(tf.data.AUTOTUNE) return ds几个参数说明shuffle(256)的buffer size不要小于batch size否则每次取出的数据分布不够随机prefetch(AUTOTUNE)让数据加载和GPU计算重叠可以明显提高训练吞吐num_parallel_calls按CPU核数设置8核以上机器用AUTOTUNE最省事。random_crop的前提是原图尺寸必须大于hr_size所以数据清洗那一步一定要过滤掉过小的图片。rot90的k是随机0到3的整数这一步对带有方向性的场景图很有效但对文字图像要慎用。4. 训练SRGAN多损失组合与完整训练流程网络搭好、数据就绪之后进入最核心的训练环节。SRGAN不是直接拿一个MSE损失训到底而是把内容损失、感知损失、对抗损失组合在一起。这也是超分领域很多优化策略的地基值得花篇幅把每个损失项讲清楚因为调参时它们彼此牵制。4.1 损失函数组合为什么MSE单独用会得到模糊结果如果只用MSE做损失生成器会倾向于输出所有可能高清图的平均值因为像素级误差最小化对应的解是数学期望而期望在视觉上就是模糊。SRGAN加入了两条修正路径一条是VGG特征空间的感知损失让生成图的语义特征接近真实图另一条是判别器给出的对抗损失逼迫生成图在真假难辨上做文章。我常用的组合是感知损失为主、对抗损失为辅具体权重按训练效果调整。# 假设已经实例化好的对象 d_optimizer tf.keras.optimizers.Adam(learning_rate1e-4, beta_10.9, beta_20.999) g_optimizer tf.keras.optimizers.Adam(learning_rate1e-4, beta_10.9, beta_20.999) bce tf.keras.losses.BinaryCrossentropy(from_logitsFalse) def generator_loss(sr_feat, hr_feat, validity, fake_labels, content_weight1e-3): # 感知损失VGG特征图的 L1/MSE perceptual_loss tf.reduce_mean(tf.square(sr_feat - hr_feat)) # 对抗损失生成图被判别为真的概率要尽可能大 adversarial_loss bce(fake_labels, validity) total perceptual_loss content_weight * adversarial_loss return total, perceptual_loss, adversarial_loss内容权重我一般设1e-3这是一个经过多次试验相对稳的值。如果对抗权重设得太大生成器训练早期就会输出颜色怪异的纹理因为判别器还没学会合理判断梯度方向是乱的设得太小对抗训练就形同虚设结果退化成纯感知损失训练。实际训练中我会打印两个loss的分量观察它们的比值是否在10:1到100:1之间偏离太远就调整content_weight。4.2 两阶段训练先预训练生成器再做对抗直接从头联合训练生成器和判别器很容易出现训练崩溃判别器loss迅速降到接近0生成器梯度消失之后无论怎么调都救不回来。SRGAN论文给了一个稳妥解法先只用MSE或感知损失训练生成器若干epoch等它能输出基本正确但不锐利的重建图再解锁判别器做对抗训练。这等同于给生成器一个先验起步点。tf.function def pretrain_step(lr_batch, hr_batch): with tf.GradientTape() as tape: sr generator(lr_batch, trainingTrue) loss tf.reduce_mean(tf.square(sr - hr_batch)) # 纯MSE预训练 grads tape.gradient(loss, generator.trainable_variables) g_optimizer.apply_gradients(zip(grads, generator.trainable_variables)) return loss tf.function def train_step(lr_batch, hr_batch): batch_size tf.shape(lr_batch)[0] real_labels tf.ones((batch_size, 1)) fake_labels tf.zeros((batch_size, 1)) # 训练判别器真实图和生成图都要判 with tf.GradientTape() as tape: sr generator(lr_batch, trainingTrue) real_validity discriminator(hr_batch, trainingTrue) fake_validity discriminator(sr, trainingTrue) d_loss bce(real_labels, real_validity) bce(fake_labels, fake_validity) grads tape.gradient(d_loss, discriminator.trainable_variables) d_optimizer.apply_gradients(zip(grads, discriminator.trainable_variables)) # 训练生成器结合感知损失和对抗损失 with tf.GradientTape() as tape: sr generator(lr_batch, trainingTrue) fake_validity discriminator(sr, trainingTrue) sr_feat vgg_extractor(sr) hr_feat vgg_extractor(hr_batch) total_g_loss, p_loss, a_loss generator_loss( sr_feat, hr_feat, fake_validity, tf.ones_like(fake_validity)) grads tape.gradient(total_g_loss, generator.trainable_variables) g_optimizer.apply_gradients(zip(grads, generator.trainable_variables)) return d_loss, total_g_loss, p_loss, a_loss注意判别器训练时真实图和生成图分别算一次BCE然后相加作为总loss。而生成器训练时fake_labels要置为1表示生成器希望判别器把生成图判为真这正是对抗博弈的直接表达。每个step里先更新判别器再更新生成器这个顺序不能乱否则同一步内生成器可能利用上一个版本的判别器输出往前跳。你要是用tf.function加速把tf.function放在train_step上面就行但第一次调用会有编译开销属于正常现象。4.3 学习率与训练超参数一张表看懂怎么设从实际项目经验看SRGAN的参数范围相对固定下面这份表格是我跑模拟项目X时反复验证过的起始配置。注意batch size受限于显存如果你只有6GB显存batch size必须降到8或4同时学习率可以适当调低。超参数 建议值 调整说明 LR尺寸 96x96 输入分辨率扩大则增加计算量 HR尺寸 384x384 4倍超分的标准设定 batch_size 8-16 16G显存可跑166G显存建议8 预训练轮数 2-3个epoch 看到重建图不糊就够 对抗训练轮数 50-100个epoch 以验证集视觉为准 生成器学习率 1e-4 如果震荡严重则降到5e-5 判别器学习率 1e-4 可以比生成器低一半 BatchNorm动量 0.9 数据分布变化大时提高动量还有两个细节值得说。第一判别器不要全程保持固定学习率训练中期如果它的loss长期在0.5以下说明判别器太强可以把它学习率降到生成器的十分之一给生成器喘息空间。第二BN层在训练和推理时的行为不一致生成器保存权重后推理时默认用训练时的统计量如果你发现训练图锐利而推理图发糊检查一下是否把training参数误传或者加载了旧的checkpoint。5. SRGAN训练避坑5个高频问题与排查路径GAN训练本来就是玄学重灾区SRGAN因为网络深、损失多翻车点更多。这里写5个我实际撞过的问题每条按现象、原因、解决三步走给相同困境的人一条能直接抄的排查路径。5.1 现象判别器Loss一路降到0生成器输出全是灰色噪声原因判别器太强生成器生成的图片被一眼识破梯度消失导致生成器不再更新。这种情况通常发生在跳过预训练直接跑对抗训练或者判别器学习率远大于生成器时。解决加载预训练好的生成器权重把判别器学习率降到1e-5并且给判别器输入加上标签平滑把真实标签从1改成0.9到1之间的随机值让它不要过度自信。real_labels tf.random.uniform((batch_size, 1), minval0.9, maxval1.0)5.2 现象生成图像发灰或者整体偏色原因输出层用tanh但输入数据错误地归一化到了[0,1]模型输出范围与标签范围不匹配训练时梯度虽然能传导但生成器很难跨越数值范围差异。解决确认LR和HR图都做了 /127.5-1.0同时检查数据加载时是否误把uint8类型直接参与计算。另一个隐蔽来源是VGG特征函数接收的输入不在[-1,1]导致感知损失恒为极大值生成器被迫向错误方向优化。5.3 现象输出图出现棋盘格纹理放大后像马赛克原因上采样模块的depth_to_space操作对通道排列敏感。如果你的特征图是channels_last布局TensorFlow默认卷积输出通道顺序必须是[r1,g1,b1,r2,g2,b2,...]这种按2x2块排列的方式一旦中间层用了channels_first布局重排出来的图像就会像棋盘。解决检查是否在某个环节用了tf.transpose改过轴顺序另外把上采样块里的卷积核从3x3改成5x5也能显著减弱高频伪影。5.4 现象16GB显存也OOM训练直接崩溃原因384x384的HR图加VGG19特征提取一张图的前向计算就占用不少显存。batch_size16同时计算生成器和判别器梯度时显存会达到峰值。解决两条路一起走一是梯度累积用小的batch_size多次前向累计梯度后再更新权重二是开启混合精度把float16计算交给支持它的GPU。tf.keras.mixed_precision.set_global_policy(mixed_float16)注意开启混合精度后判别器输出的sigmoid层必须保留float32计算否则精度损失会导致训练不稳定。我一般只对卷积层和全连接层启用自动混合精度关键损失计算处用tf.cast回到float32。5.5 现象PSNR指标涨了人眼看着还是糊原因PSNR和SSIM都建立在像素级误差或结构相似度之上它们对模糊的惩罚不够对轻微的颜色偏移却很敏感。SRGAN本来就是为感知真实感设计的所以用传统指标衡量会出现指标与肉眼倒挂。解决把评估指标换成LPIPS或者FID这一类感知指标没有外部库时至少做一个用户对比测试——让几个人在盲测中选更像真实照片的图比单看PSNR有说服力。6. 从能跑到效果好验证方法与大图推理技巧训练结束后的验证不能只看训练集重建效果要留出一部分验证集图片走完整前向过程。我习惯每个epoch结束后固定抽取4到6张验证图把LR、SR、HR三张图拼成一行保存下来直观观察纹理恢复程度。训练过程中随时打开这个对比图比盯着loss曲线更有用。定量指标层面如果环境允许装上lpips包算感知距离同时计算PSNR和SSIM。但更重要的是统计生成图是否有伪影这类视觉缺陷数量。我会把验证集里明显偏色、棋盘纹、白边明显的图截图归档写进项目说明里这样后续换版本训练可以拉同一批图对比。大图推理是另一个必须有预案的地方。训练用的HR尺寸是384x384但实际应用时往往要放大2000x1500的大图直接整图送入生成器很容易OOM。我一般用滑窗切片推理把大图切成384x384的块相邻块之间保留32像素重叠推理后去掉重叠边缘再做加权融合。def infer_large_image(model, large_lr, patch_size384, overlap32): h, w large_lr.shape[:2] step patch_size - overlap output tf.zeros((h * 4, w * 4, 3)) weight tf.zeros((h * 4, w * 4, 1)) for y in range(0, h, step): for x in range(0, w, step): patch large_lr[y:y patch_size, x:x patch_size] sr_patch model(tf.expand_dims(patch, 0))[0] y_out, x_out y * 4, x * 4 output[y_out:y_out patch_size * 4, x_out:x_out patch_size * 4] sr_patch weight[y_out:y_out patch_size * 4, x_out:x_out patch_size * 4] 1.0 return output / tf.maximum(weight, 1.0)重叠区域被多次推理累加后再除以累加权重相当于做了平均融合能有效避免块与块之间的接缝。这是我实际项目里最常用的一个函数它不需要额外依赖一张任意尺寸的图都能直接跑。另一个进阶操作是给生成器权重做指数移动平均EMA。GAN训练末期权重波动大直接保存某一时刻的权重可能抽到差的解。维护一个shadow变量每步把新权重按0.999的衰减系数合入推理时用shadow权重通常能拿到更稳的结果。写起来不复杂但收益很明显。TensorFlow官方没有直接封装需要自己写一个回调或在训练循环里维护。最后说一个习惯每次训练结束把最优权重的验证集效果截图、训练曲线、完整超参数都留存在项目目录里而不是只存一个.h5文件。跑超分辨率项目一半时间不是在调模型而是在对历史实验做复盘对比。有了这些过程记录你才能确定某个改动是变好了还是变坏了而不是凭感觉调参。这套TF2版的SRGAN项目我从数据准备到训练完成跑过不止一遍最深的体会是不要迷信某个loss值或者某个指标所有改动都要回到肉眼看起来是否真实这个起点上。希望帮到你。本文还有配套的精品资源点击获取
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →