尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

基于CNN与U-Net的图像着色实战:从Lab空间到313类分类

发布时间:2026/9/24 18:31:18

资讯中心
01
ARTICLE

基于CNN与U-Net的图像着色实战:从Lab空间到313类分类

基于CNN与U-Net的图像着色实战:从Lab空间到313类分类
简介一份基于深度学习CNN网络的图像着色Python源码包配套完整可运行的训练与推理脚本面向计算机、人工智能、数据科学等专业的在校学生、教师及企业开发者适用于课程设计、毕业设计或入门进阶学习。源码包含ECCV16与SIGGRAPH17两套经典着色模型的实现涵盖模型定义、特征提取、颜色空间转换等核心模块并提供示例图片与运行结果图便于对比不同模型的着色效果。压缩包共40个文件以Python脚本14个py为主辅以测试图片12个png、编译缓存10个pyc、说明文档2个txt及参考文献2个pdf整体约17.52MB目录结构清晰便于按模块阅读与二次开发。目前已有50人学习浏览。资源适合希望深入理解CNN图像着色原理、快速上手深度学习项目或在此基础上DIY功能的读者源码注释与模块划分有助于二次开发与算法调优。1. 拿到图像着色源码包时先想清楚一个问题拿到「基于深度学习CNN网络实现图像着色python源码.zip」这个包第一反应通常是灰度图补色有什么好做的直到你把一张旧照片丢进去才意识到着色的输出不是唯一解同一张黑白人像头发可能是深棕也可能是纯黑天空可能是青灰也可能是日落橙黄。深度学习 CNN 在这里做的事是把「预测颜色」重新定义成「从亮度推断一个最合理的色度分布」。这个方向适合两类人一类是刚学完 CNN、ImageNet 分类练得没新鲜感、想找能跑通且能持续调优的实战项目的入门者另一类是需要批量给历史影像、老电影帧或监控灰度图做着色的落地工程师。前者借它把数据管线、模型结构、训练调参整条链路走通后者拿它当生产工具用。下面按我自己的落地习惯把原理、实现、训练、排查和进阶评估完整讲一遍。2. 为什么着色要放弃 RGBLab 空间、U-Net 与 313 类分类2.1 为什么选 CNN 而不是 Transformer 或 RNN图像着色本质上是稠密预测任务每个像素都要输出一个颜色。RNN 擅长序列建模硬把图像按行展开喂进去全局结构和局部纹理都会被打散计算效率也差这个任务基本没人用它。Transformer 这两年很火能做长程依赖但图像着色最关键的信号恰恰是局部的——皮肤材质、树叶纹理、金属反光这些都是由邻近像素的亮度关系决定的。CNN 的局部感受野和权重共享天然匹配这种信号结构而且在几十万张图的规模下卷积网络收敛更稳、显存占用更可控。你如果已经在用cnn explainer 离线包这类工具做可视化调试更会发现 CNN 的每一层特征响应都能直观看到这对排查模型为什么把草地着成蓝色非常有帮助。我的建议是先拿 CNN 把基线做扎实再去折腾 Transformer 结构否则一旦出问题你根本分不清是数据问题还是注意力机制的问题。2.2 Lab 色彩空间把「预测颜色」变成「预测偏移」很多第一次做着色的人会直接想输入灰度图单通道输出 RGB 三个通道回归一下就完事。这个方案基本必翻车。原因是 RGB 三个通道高度相关亮度信息全部混在里面模型既要预测亮度又要预测色度而亮度在灰度输入里已经写死了模型很容易学成「复制亮度、色度取平均」最后输出一片灰蓝色调的图。常见做法是转换到 Lab 色彩空间。L 通道表示亮度a 通道表示绿色到品红的偏移b 通道表示蓝色到黄色的偏移。这样输入是 L灰度图就是 L模型只需要预测 a 和 b 两个偏移量。原来的问题从「凭空生成整套颜色」变成「在给定亮度下猜色度偏移」约束更清晰模型也更好学。这也是这个标题下的源码包几乎统一采用的设计。实现时有个坑OpenCV 的cvtColor转 Lab 后数值范围和 scikit-image 不一样前者 a、b 会有 128 的偏移后者 a、b 大致在 -128 到 127 之间。如果不统一量化 bin 和归一化参数全得重做。我习惯用 scikit-image语义更直白。训练时输入 L 统一除以 100 归一化到 0~1 区间推理时同样处理这一步必须严格一致否则模型看到的数据分布完全不同。2.3 回归改分类313 个 bin 与类别重加权那能不能直接回归 a、b 两个通道能但效果差。MSE 回归对多解问题天然不友好一个灰度像素可以对应「红色毛衣」也可以对应「蓝色毛衣」回归模型会把所有可能答案求平均红色和蓝色平均出来是灰紫色。这是均值回归的数学宿命不是模型结构能救的。主流做法是把着色转成分类问题。把 ab 平面量化成一个个离散的 bin每个像素的标签变成「这 313 个 bin 里的哪一个」。模型输出的每一层特征图有 313 个通道每个通道代表该像素选某个 bin 的概率。分类不需要在连续空间里求平均而是可以学出一个多峰分布配合后面的期望解码能保留颜色选择的多样性。这里还有一个关键的细节颜色分布不均衡。自然图像里棕色、灰色、深绿色的像素占比极高纯金黄、荧光橙这些鲜艳颜色占比极低。如果不做处理模型学到的就是一个「永远输出保守颜色」的分布。常用做法是按每个 bin 在训练集里的统计频率 p 做重加权权重约等于(1-λ)/p λ/θλ 取 0.5θ 是类别数的倒数。写代码时通常把这组权重提前算好存成文件训练时直接传给 CrossEntropyLoss后面第 4 章会给出完整实现。3. 从零搭一个图像着色项目数据管线与 U-Net 实现3.1 数据管线读图、转 Lab、量化标签拿到这个方向的项目我一般会把代码拆成四个模块数据加载、模型定义、训练脚本、推理脚本。先写数据管线。核心逻辑是读 RGB 图转 Lab取 L 通道做输入把 a、b 通道量化成 bin 索引做标签。import numpy as np import torch from torch.utils.data import Dataset from skimage import color, io from skimage.transform import resize # 预设 ab 空间量化网格范围 [-110, 110]步长 10共 23x23 个候选点 # 去掉离原点过远的无色区域保留 313 个有效 bin def build_color_bins(): grid np.arange(-110, 111, 10) a_grid, b_grid np.meshgrid(grid, grid) points np.stack([a_grid.ravel(), b_grid.ravel()], axis1).astype(np.float32) radius np.sqrt((points ** 2).sum(axis1)) valid radius 110 # 半径截断保证 bin 落在真实可感知的色域内 return points[valid] # 形状约为 (313, 2) class ColorizationDataset(Dataset): def __init__(self, image_paths, size224, binsNone): self.paths image_paths self.size size self.bins build_color_bins() if bins is None else bins self.bins_t torch.from_numpy(self.bins) # (K, 2) def __len__(self): return len(self.paths) def _quantize_ab(self, ab): # ab: (H, W, 2)把每个像素映射到最近 bin 的索引 h, w ab.shape[:2] flat ab.reshape(-1, 2) # (N, 2) diff flat[:, None, :] - self.bins[None, :, :] # 广播求差 dist (diff ** 2).sum(axis2) # (N, K) return dist.argmin(axis1).reshape(h, w) # 每个像素的 bin 索引 def __getitem__(self, idx): rgb io.imread(self.paths[idx]) if rgb.ndim 2: rgb color.gray2rgb(rgb) rgb resize(rgb, (self.size, self.size), anti_aliasingTrue) lab color.rgb2lab(rgb).astype(np.float32) l_ch lab[:, :, 0] / 100.0 # 归一化到 [0, 1]训练推理必须一致 ab lab[:, :, 1:] # 原始 a、b 范围 label self._quantize_ab(ab) # (H, W) 的标签图 l_t torch.from_numpy(l_ch).unsqueeze(0).float() label_t torch.from_numpy(label).long() return l_t, label_t这段代码的量化逻辑要重点讲一下。build_color_bins在 ab 平面上按 10 的步长撒网格点再用半径 110 截断目的是把网格限制在真实照片会出现的色域内——ab 空间离原点太远的地方对应的是显示器都很难准确还原的极端颜色样本量也极少留它们做分类只会浪费类别数。如果截断半径设得太大bin 数会远超 313最后几层卷积的通道数跟着膨胀显存压力变大设得太小则会丢掉高饱和度的鲜艳颜色输出发灰。_quantize_ab用全量距离计算做最近邻匹配视觉上蛮直观但要注意它对每个像素都算了 313 个距离224×224 分辨率下就是 5 万个像素乘 313 个 bin训练时每张图都要做一次。数据量大的话建议预先把标签存成 npy 文件避免每个 epoch 重复计算。因为 ab 范围是固定的量化结果不随训练变化属于「一次计算、到处使用」的重复劳动。3.2 U-Net 编码器堆卷积还是借用预训练权重模型结构我选了 U-Net 的变体。编码器逐级下采样提取语义信息解码器逐级上采样恢复分辨率中间用跳跃连接把高分辨率细节传给解码器。这个结构在图像翻译类任务里几乎是默认选择。编码器可以自己堆卷积也可以直接拿 VGG16 的预训练权重初始化。我一般会给两个选项如果机器能联网下载权重用预训练 VGG16 前几层训练收敛明显更快如果环境受限就自建一个 5 层卷积编码器效果差一些但完全可控。import torch import torch.nn as nn class Encoder(nn.Module): def __init__(self, in_ch1): super().__init__() # 每个 block: 两次卷积 BN ReLU随后 stride2 下采样 self.block1 self._make_block(in_ch, 64) self.block2 self._make_block(64, 128) self.block3 self._make_block(128, 256) self.block4 self._make_block(256, 512) self.block5 self._make_block(512, 512) self.pool nn.MaxPool2d(2) def _make_block(self, in_c, out_c): return nn.Sequential( nn.Conv2d(in_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), nn.Conv2d(out_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), ) def forward(self, x): f1 self.block1(x) # 224 - 64 p1 self.pool(f1) # 112 f2 self.block2(p1) # 128 p2 self.pool(f2) # 56 f3 self.block3(p2) # 256 p3 self.pool(f3) # 28 f4 self.block4(p3) # 512 p4 self.pool(f4) # 14 f5 self.block5(p4) # 512 p5 self.pool(f5) # 7 return p5, [f1, f2, f3, f4]为什么输入用 1 通道而不用三通道灰度复制预训练权重方案里确实要把灰度图复制三次再输入因为 ImageNet 预训练的第一层卷积接受三通道输入。但自建编码器没有这个约束1 通道输入更省内存第一层也不需要做 3 到 1 的通道适配。要注意的是如果以后想切换预训练方案输入层要改成in_ch3其余结构不用动。每个 block 内部塞两次卷积而不是一次是为了在每次下采样前先扩大感受野提取更丰富的局部模式。通道数从 64 翻倍到 512符合「空间分辨率减半、通道数翻倍」的经典设计原则既能压缩信息又不会让特征图过薄。输入 224 分辨率经过 5 次下采样变成 7×7512 个通道这是整条网络里语义最浓缩的地方解码器要在这基础上逐级恢复细节。3.3 解码器与跳跃连接输出 313 通道热力图解码器负责把 7×7 的特征图逐步放大回 224×224。每一步先上采样再把编码器对应层的特征图拼过来这样解码器在恢复颜色时既能看见全局语义又能看见原始边缘和纹理。拼接操作是沿通道维度的torch.cat不是加法简单加法会损失信息拼接让解码器自己学会怎么融合粗细粒度特征。class Decoder(nn.Module): def __init__(self, out_ch313): super().__init__() # 上采样 跳跃连接逐级恢复分辨率 self.up1 self._make_up(512 512, 512) # 7 - 14 self.up2 self._make_up(512 256, 256) # 14 - 28 self.up3 self._make_up(256 128, 128) # 28 - 56 self.up4 self._make_up(128 64, 64) # 56 - 112 self.up5 self._make_up(64 64, 64) # 112 - 224 self.head nn.Conv2d(64, out_ch, 3, padding1) def _make_up(self, in_c, out_c): return nn.Sequential( nn.Upsample(scale_factor2, modebilinear, align_cornersTrue), nn.Conv2d(in_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), ) def forward(self, x, skips): x self.up1(x) # 7x7 - 14x14 x torch.cat([x, skips[3]], dim1) # 拼上编码器 block4 的特征 x self.up2(x) x torch.cat([x, skips[2]], dim1) x self.up3(x) x torch.cat([x, skips[1]], dim1) x self.up4(x) x torch.cat([x, skips[0]], dim1) x self.up5(x) return self.head(x) # 输出 (B, 313, 224, 224)头部的输出通道数 313 直接对应 bin 数量。每个空间位置有 313 个 logits经过 softmax 就得到该像素选择每个 bin 的概率分布。这里有个很容易忽略的点Upsample用的align_cornersTrue不能省默认的align_cornersFalse会对齐方式有细微差异特征图尺寸在奇数分辨率下可能和编码器端差 1 个像素拼接时直接报维度错误。我自己踩过这个坑报错信息是Sizes of tensors must match光看堆栈很难一眼定位到是这里的问题。跳跃连接前的通道数也要细心算。比如up2这层上采样后通道数是 512拼上编码器 block3 的 256实际输入是 768所以_make_up的第一个参数必须是 768。我把编码器和解码器分开定义就是为了让每个阶段的通道拼接关系一目了然真要改结构时不容易出错。3.4 推理模块把概率分布转回一张彩色图训练完模型推理时不能直接拿 logits 的 argmax 当颜色。argmax 会让每个像素只选概率最大的 bin相邻像素可能选到不同的鲜艳颜色结果全是椒盐噪点。常用做法是对 313 个 bin 的 ab 坐标做期望E[a] Σ p_i * a_iE[b] Σ p_i * b_i。期望天然做了平滑高概率的颜色贡献大低概率的只是轻微拉偏。import torch import numpy as np from skimage import color def decode_ab(logits, bins, temperature0.38): # logits: (1, 313, H, W)bins: (313, 2) probs torch.softmax(logits / temperature, dim1) # 温度缩放后取 softmax probs probs.permute(0, 2, 3, 1) # (1, H, W, 313) ab_flat torch.tensor(bins, dtypetorch.float32) a torch.sum(probs * ab_flat[:, 0].view(1, 1, 1, -1), dim-1) b torch.sum(probs * ab_flat[:, 1].view(1, 1, 1, -1), dim-1) return a, b # (1, H, W) def lab_to_rgb(l_ch, a, b, saturation1.0): l l_ch * 100.0 # 还原归一化前的 L 范围 lab torch.stack([l, a * saturation, b * saturation], dim-1) lab_np lab.squeeze(0).cpu().numpy() rgb color.lab2rgb(lab_np) # 自动裁剪到合法范围 return (rgb * 255).astype(np.uint8)decode_ab里的temperature是整条推理链路上最值得调的一个参数后面第 6 章会专门展开。简单说temperature 小于 1 会让分布更尖锐颜色更鲜艳更大胆大于 1 会让分布更平滑颜色更保守更灰。0.38 是论文作者验证过的经验值但不同数据集上差不少建议在自己的验证集上扫一遍。代码里a * saturation是对颜色饱和度的后处理saturation 大于 1 会推高色彩强度小于 1 会压向灰色。我一般默认 1.0如果模型已经够鲜艳就不再动。注意lab2rgb输出范围是 0 到 1转回 0 到 255 的 uint8 时别丢了.astype否则后面存图会报警告甚至写不进文件。4. 训练配置与调参损失函数、学习率、Batch Size 怎么设4.1 损失函数多分类交叉熵与平衡权重训练标签是每个像素的 bin 索引自然用多分类交叉熵。但如果直接把nn.CrossEntropyLoss()套上去模型会严重偏向高频出现的颜色。前面提过的类别重加权在这里落地权重数组要在训练前统计一次不要每个 batch 都算。import torch import torch.nn as nn def compute_class_weights(labels_all, num_classes313, lam0.5): # labels_all: 所有训练标签拼成的长数组统计每个 bin 的全局频率 counts np.bincount(labels_all, minlengthnum_classes).astype(np.float32) prob counts / counts.sum() theta 1.0 / num_classes # lambda 控制均匀先验的占比避免稀有颜色权重过大 weight (1 - lam) / (prob 1e-6) lam * theta weight weight / weight.mean() # 归一化保持损失数值尺度稳定 return torch.from_numpy(weight).float() weight compute_class_weights(all_train_labels) loss_fn nn.CrossEntropyLoss(weightweight)lam取 0.5 是常见做法它在「完全按频率反比加权」和「完全均匀加权」之间取平衡。如果把lam调到 0稀有颜色权重会极高模型看到棕色像素几乎不更新反而容易在鲜艳区域产生色块噪点调大lam则重加权效果消失回到保守灰。weight / weight.mean()这步很多人会漏不归一化的话整体损失数值会被放大几十倍学习率效果完全失真。注意这里统计权重用的是所有训练像素的 bin 索引数组而不是每张图单独统计。每张图的颜色分布差很多单图统计的权重噪声极大。我通常会在第一次跑数据管线时把所有 label 存成 npy顺手就把权重算好一劳永逸。4.2 优化器与学习率策略优化器选型没什么悬念Adam 在这个任务上比 SGD 省心得多不用费劲调动量参数。学习率我一般从 1e-3 开始batch size 大的时候可以适当上调。训练分辨率用 224×224 是性价比比较高的选择太低边缘细节不够太高显存撑不住。配置项推荐值说明优化器Adambetas 默认 (0.9, 0.999)不需要动初始学习率1e-3batch size 32 时bs 越大 lr 可适度调大学习率策略CosineAnnealing总 epoch 30~50T_max 设成总 epoch 数Batch Size16 或 32取决于显存8GB 显存建议 16输入分辨率224×224兼顾细节与训练速度梯度裁剪max_norm1.0防止分类头梯度爆炸训练时每次迭代要做的步骤很常规前向拿到 logits算损失反向传播优化器更新。但有个细节是标签必须在 GPU 上和 logits 对齐量化后的标签是(B, H, W)的整数张量logits 是(B, 313, H, W)CrossEntropyLoss 内部要求 logits 在第一维是类别数标签形状刚好匹配不需要额外squeeze或permute。for epoch in range(epochs): model.train() for l_t, label_t in train_loader: l_t l_t.to(device) label_t label_t.to(device) logits model(l_t) # (B, 313, H, W) loss loss_fn(logits, label_t) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step()如果显存只够 batch size 8又想让模型看到更大 batch 的统计量常见做法是梯度累积。写一个accum_steps变量每accum_steps步做一次optimizer.step()。注意梯度累积时zero_grad要放在累积结束之后否则梯度被清零就白攒了。混合精度训练可以用torch.cuda.amp.GradScaler对 313 通道这种大头输出能省接近一半显存训练时间也能缩短 30% 左右值得加。4.3 数据增强的边界什么能抖什么不能抖数据增强在这个任务里是双刃剑。随机水平翻转、随机裁剪、随机缩放都可以放心用它们不改变颜色语义人穿红色衣服翻转后还是红色。但颜色抖动类增强一定要慎用。torchvision.transforms.ColorJitter会直接改 RGB 的亮度、对比度、饱和度和色相这套操作放在灰度图任务里等于伪造信息——图像明明只有亮度输入模型却可能学到「这个 RGB 值对应某个颜色」的捷径。更隐蔽的坑是有些实现先对 RGB 图做增强再转 Lab这样 a、b 通道都被污染了标签也跟着错。安全的做法是只对亮度通道做轻微扰动。比如把输入 L 乘一个 0.9 到 1.1 的随机系数模拟不同曝光条件下的灰度图。注意这个扰动只能加在输入上不能加在标签的 a、b 上。我做数据增强时一般把 L 通道单独拿出来处理标签量化走独立的纯函数两者不互相干扰。4.4 验证指标PSNR 之外的判断方法图像着色是典型的不适定问题灰姑娘的裙子可以是蓝色也可以是粉色PSNR、SSIM 这类逐像素指标根本反映不了「颜色合不合理」。真正有用的验证手段是三个第一可视化对比每训练 500 步从验证集抽 8 张图把输入灰度图、模型输出、原图拼成一行存到 TensorBoard肉眼直接看颜色是否符合语义第二FID用预训练 Inception 网络提取真实图和生成图的特征分布算两个高斯分布的 Wasserstein 距离FID 越低说明生成图像整体风格越接近真实照片第三色彩分布统计画一张模型输出图的 ab 散点图如果所有点堆在原点附近说明模型又回到灰色均值解如果散点覆盖整个 bin 空间说明输出颜色多样性足够。我见过不少人直接拿 PSNR 调参盯着这个数字调了几天发现模型越调越灰——因为 PSNR 是逐像素误差最安全的答案就是平均色跟模型是不是真的生成合理颜色毫无关系。这个图上位的评判原则是PSNR 只能作为训练曲线平滑度的参考真正决定模型能不能用的判断方式是人眼投票和 FID 趋势。5. 图像着色排查指南五个最常见的翻车现场5.1 现象训练 Loss 下降但输出一片灰训练曲线降得很漂亮Loss 从 5 点多降到 2 点多可视化一看整张图灰蒙蒙的几乎分不清颜色差异。原因最常见的是推理时没做温度缩放直接用温度 1.0 求期望。类别分布本身是长尾的灰色和棕色的概率远高于鲜艳颜色期望运算把所有低概率的鲜艳颜色一平均就只剩灰色的残影。另一个常见原因是decode_ab里对 a、b 做了归一化但推理时忘了还原输出范围错位颜色被压扁。解决先把 temperature 调到 0.38 左右再试。还不行就检查训练时输入 L 的归一化和推理时是否一致比如训练时除以了 100推理时也必须是 100。调试的小技巧是直接打印模型输出的 ab 均值如果训练集真实图像的 ab 均值接近 0 是正常的但输出的 ab 标准差如果远小于真实图的 ab 标准差基本就是温度或归一化的问题。5.2 现象物体边缘有彩色晕染物体边缘一圈五颜六色的光晕像是把彩色的蛋糕抹在剪影外圈。原因这个基本就是跳跃连接和上采样之间没有对齐。Upsample和MaxPool2d在某些分辨率下输出尺寸差 1 个像素torch.cat能正常执行是因为 PyTorch 做了自动的维度匹配不会实际会直接报维度错误。如果没报错又出现了晕染通常是因为编码器靠后的 feature map 分辨率太低边缘位置被卷积核污染解码器在恢复时把模糊的边界特征放大继承了。解决检查所有 concat 位置的 shape打印出来逐一比对收敛到完全一致为止。同时在解码器的每个上采样后面加一层 3×3 卷积而不是 1×1让网络有参数去修正上采样带来的锯齿。边缘晕染如果特别严重也可以给输入图像加一个轻微的边缘权重图让损失在边缘像素上有更高的权重这是后置手段结构修正才是治本。5.3 现象Loss 正常下降但颜色饱和度不足训练一切正常Loss 稳步降温度调小了颜色还是不够鲜艳输出像蒙了一层灰纱。原因类别重加权没生效。检查CrossEntropyLoss是否真的收到了你计算好的权重数组。我遇到过权重算好了但训练脚本里loss_fn nn.CrossEntropyLoss()在compute_class_weights之前初始化导致 loss_fn 一直用的是默认等权版本。这种问题在代码里完全不报错只能靠日志确认权重实际参与计算。解决训练脚本里初始化 loss_fn 之前先确认 weight 数组的 shape 和数据类型weight.dtype必须是 float32元素最小值不能是 0某些 bin 在训练集里一次都没出现bincount会给出 0权重那里要加一个1e-6的平滑项。顺便检查一下compute_class_weights里的prob 1e-6是不是加在了除法之前加错位置权重会直接爆掉。5.4 现象显存爆炸Batch Size 起不来224×224 分辨率下 Batch Size 设 32训练刚开始就报 CUDA out of memory。原因313 通道的分类头是显存大户。中间层 feature map 是 313 通道乘以 56×56 或 28×28 的空间尺寸一张图的激活值就有百万级别。再加上 Adam 优化器要存两份动量状态显存占用直接翻倍。解决三个优化方向按顺序做。第一Batch Size 降到 8配合梯度累积模拟更大的 batch第二开混合精度用GradScaler把前向和反向计算降成 float16能省一半显存第三改模型结构分类头只在 28×28 分辨率上输出也就是说up3之后不要继续上采样到 224而是直接在这个分辨率接分类头后处理阶段用插值把 ab 图放大回原分辨率。这么做会损失一些边缘细节但显存占用能降非常多我遇到显存紧张时通常先动这一刀。5.5 现象平滑区域出现彩色噪点天空、白墙、柏油路面这种本应是均匀颜色的区域着出来的颜色像打翻的调色盘满是彩色斑点。原因分类分布熵太高。天空这种区域亮度变化很小模型对颜色的置信度本就不高如果直接用 argmax 选 bin相邻像素各自选到不同颜色自然全是噪点。另一个推动因素是 bin 步长太大10 的步长让相邻 bin 之间颜色跳变明显。解决推理时默认就用期望解码而不是 argmax期望天然有平滑效果。如果还有噪点对概率分布的空间平滑一下——常见做法是先对 313 通道的 softmax 输出做一次 3×3 的平均池化让空间上相邻像素的类别分布趋于一致再走期望解码。这个操作在 PyTorch 里就是一句F.avg_pool2d(probs, 3, 1, 1)基本不耗时但效果立竿见影。也可以在后处理用引导滤波以输入 L 通道为引导图对 ab 通道做边缘保持平滑这招能压住大部分色点。6. 进阶玩法与验证方法温度系数、FID 与 ONNX 导出训练收敛后图像着色的产品化还有一个关键参数调优过程。上面反复提到的 temperature 值得在这单独说透。我习惯在验证集上固定取 300 张图分别用T0.2、0.38、0.6、0.8跑一轮推理对比人眼观感和 FID 数值。T 越小softmax 分布越尖模型更敢于选高概率的鲜艳颜色但错误的概率也跟着被放大可能出现青色人脸T 越大分布越平每类概率被压向均匀输出就越灰。这个参数没有全局最优值同一个模型在不同数据分布上的最佳 T 能差出 0.2 以上值得花半小时扫一遍。另一个值得做的验证是 FID它比 PSNR 更能反映「生成图看起来像不像真实照片」。具体做法是拿 1000 张真实彩色图和对应的 1000 张模型着色图分别过预训练 InceptionV3 取倒数第二层特征算特征分布之间的 Frechet 距离。FID 小于 40 说明整体风格已经比较接近真实照片再往下调参空间不大如果 FID 在 70 以上说明颜色风格明显偏离真实分布优先回头查类别重加权和数据增强。最后落到部署。模型要上线给业务方调用通常要导出成 ONNX。这里有两个注意点一是导出时输入输出的动态轴要显式声明否则固定成 224×224 之后换分辨率就得重新导出二是 313 通道的输出要在后处理层处理好再导出可以把 softmax、期望解码和饱和度调整全部封装进torch.nn.Module里导出成一个端到端的模型业务方拿到的就是一个「灰度图进、彩色图出」的黑盒省去他们理解 bin 量化的成本。class ColorizationWrapper(torch.nn.Module): def __init__(self, model, bins_tensor): super().__init__() self.model model self.register_buffer(bins, bins_tensor) def forward(self, gray): logits self.model(gray) probs torch.softmax(logits / 0.38, dim1) ab torch.einsum(bchw,cd-bdhw, probs, self.bins) return torch.cat([gray, ab], dim1) # 输出 L、a、b 拼接外部再转 RGB我自己现在的习惯是所有调参对比之前先把 temperature 定死再谈模型结构变化和数据增强得失否则两个变量一起动结果差点都说不清是谁的功劳。每次训练完跑一轮 300 张图的 FID 和人眼抽样把数值记在实验表格里这样迭代了十几版之后才能准确知道每一处改动实际带来了什么。着色这个任务没有唯一正确答案衡量标准永远是「看起来合不合理」所以验证方法比训练技巧更值得投入时间。希望帮到你。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。