colorization-pytorch 的 SIGGRAPHGenerator 网络逐层拆解:4x4 子采样与残差跳连如何实现快速上色 colorization-pytorch 的 SIGGRAPHGenerator 网络逐层拆解4x4 子采样与残差跳连如何实现快速上色【免费下载链接】colorization-pytorchPyTorch reimplementation of Interactive Deep Colorization项目地址: https://gitcode.com/gh_mirrors/co/colorization-pytorchcolorization-pytorch 是 SIGGRAPH 2017 经典工作《Real-Time User-Guided Image Colorization with Learned Deep Priors》的 PyTorch 复现核心功能是把灰度图像上色为彩色图并支持交互式上色用户只需点几下颜色提示hint模型就能实时把颜色扩散到整张图。整个项目的灵魂是定义在models/networks.py中的SIGGRAPHGenerator生成网络——它用 4x4 子采样/上采样块把计算量压到 1/8 分辨率再用残差跳连把边缘细节送回收码器从而做到既准又快。本文带你逐层读懂它。先看效果交互式上色的整体演示上图演示了完整流程左侧灰度图用户在中部涂抹少量彩色提示网络实时补全其余区域的颜色。能做到实时全靠生成网络本身够轻量。网络总览一个沙漏形的编码器-解码器SIGGRAPHGenerator 接收 3 路输入并沿通道维拼接输入含义来源input_A灰度图Lab 空间的 L 通道1 通道用户原图input_B颜色提示AB 通道2 通道用户涂抹处有值其余为 0mask_B提示掩码1 通道标记哪些像素给了颜色提示前向流程见models/networks.py中forward方法约 347 行起可以概括为一个沙漏阶段代码模块通道数变化输出分辨率作用编码 1model13 → 641/1低层特征编码 2model264 → 1281/2降采样 特征增强编码 3model3128 → 2561/4语义特征编码 4model4256 → 5121/8深层特征瓶颈model5/model6/model7512 → 5121/8空洞卷积扩大感受野解码 1model8upmodel8512 → 2561/44x4 上采样 跳连分类头model_class256 → 5291/4输出颜色分布解码 2model9upmodel9256 → 1281/24x4 上采样 跳连解码 3model10upmodel10128 → 1281/14x4 上采样 跳连回归头model_out128 → 21/1输出 AB 颜色值可以看到全网络只有一个 1/8 分辨率的瓶颈这是快速上色的第一大功臣。编码器4x4 子采样三步降到 1/8 编码器每步都是两块 3x3 卷积 ReLU BatchNorm的组合通道逐层翻倍64 → 128 → 256 → 512。与常见 U-Net 用 stride2 卷积下采样不同这里的降采样发生在卷积之后用一行张量切片完成conv2_2 self.model2(conv1_2[:, :, ::2, ::2])[:, :, ::2, ::2]相当于对 4x4 的块做 2x2 最近邻子采样把空间分辨率减半。这样做有两个好处零参数、零计算下采样不引入任何可学习权重比 stride2 卷积省算力无混叠设计先卷积充分提取特征再丢弃像素信息损失更小。编码器连续三次这样操作1/1 → 1/2 → 1/4 → 1/8输入提示信号也在这一路被读懂——提示区域的颜色与掩码信息随特征一起流入深层。瓶颈层空洞卷积扩大感受野model5与model6使用dilation2 的 3x3 卷积步长 1、padding2在不降低分辨率、不增加计算量的前提下把感受野翻倍。为什么要这么做颜色扩散本质上是一个让远处像素看到提示的问题一个提示点要影响几米之外的区域就需要足够大的感受野。两次空洞卷积让 1/8 分辨率下的 512 通道特征能覆盖相当大的原图范围model7则把 512 通道特征整理好作为瓶颈输出。解码器4x4 转置卷积上采样 残差跳连解码端是本文重点。三个上采样模块model8up、model9up、model10up全部是同一个套路nn.ConvTranspose2d(512, 256, kernel_size4, stride2, padding1)4x4 卷积核、步长 2每次把分辨率放大 2 倍恰好与编码器的子采样对称。而残差跳连的接法很特别——不是 U-Net 式的通道拼接而是加一下conv8_up self.model8up(conv7_3) self.model3short8(conv3_3)上采样特征先与编码器同分辨率的特征相加跳连侧只经过一层轻量 3x3 卷积model3short8/model2short9/model1short10做通道对齐开销极小。 为什么跳连必不可少解码器从 1/8 的模糊特征放大回全分辨率单靠上采样无法恢复边缘、纹理等高频细节。编码器在 1/4、1/2、1/1 分辨率处保存的特征conv3_3、conv2_2、conv1_2恰好携带这些信息逐层加回去后上色结果才能保持清晰的轮廓和边界。双输出头颜色分布 回归颜色网络最终返回两个结果forward末尾的(out_class, out_reg)分类输出out_classmodel_class是一层 1x1 卷积把 256 通道映射到529 类。529 23 × 23对应把 A、B 两个颜色通道各量化成 23 个等级参数ab_max110、ab_quant10见options/base_options.py。它表示这个像素可能是哪种颜色的概率分布在 1/4 分辨率计算再由pix2pix_model.py中的upsample4放大回全图用于交叉熵损失与熵的可视化。回归输出out_regmodel_out是一层 1x1 卷积128 → 2 通道加 Tanh直接输出全分辨率的 AB 颜色值用于 L1 回归损失。训练时还有一个巧思classificationTrue分支里解码上采样路径的输入会调用detach()停止梯度让分类损失只训练编码器主干、回归损失只精修解码器两个阶段互不干扰。这正是项目两阶段训练策略的网络层配合先纯分类训练自动上色--classification再注入颜色提示做回归微调学习信任用户提示。完整命令就在scripts/train_siggraph.sh里分siggraph_class→siggraph_reg→siggraph_reg2三段逐步调低学习率。实测效果提示越多上色越准跑test_sweep.py可以复现论文 Figure 6随机揭示不同数量的 6x6 颜色提示块统计上色结果的 PSNR。项目自带参考曲线从图中可以看到即使没有任何提示Auto 档纯自动上色PyTorch 复现模型也能达到约 26 dB提示点从 1 个增加到 500 个PSNR 稳步升到 33 dB 以上。这条平滑上升的曲线正说明编码器读懂了提示、跳连保住了细节——两者缺一不可。小结SIGGRAPHGenerator 快的 4 个原因设计代码位置收益2x2 切片式子采样无参数下采样models/networks.py的forward比 stride2 卷积省算力瓶颈只在 1/8 分辨率model4~model7深层计算量降到约 1/64跳连用相加而非拼接model8up等三处通道对齐开销极小1x1 卷积输出头 1/4 分辨率分类头model_class、model_out最贵的一层只算 529 类理解了这套4x4 子采样 残差跳连 双头输出的组合拳你就掌握了 colorization-pytorch 从灰度图到实时交互式上色的全部秘密。想动手验证可以按scripts/train_siggraph.sh复现两阶段训练或用test.py加载checkpoints/下的预训练权重查看上色结果。【免费下载链接】colorization-pytorchPyTorch reimplementation of Interactive Deep Colorization项目地址: https://gitcode.com/gh_mirrors/co/colorization-pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考