简介:这份资源面向希望上手图像自动着色、了解深度先验应用的 Python 开发者与计算机视觉学习者,核心是两套预训练着色模型:eccv16 与 siggraph17,可在实时用户引导下为黑白照片上色。包内共 23 个文件,以 py 脚本与 pyc 缓存为主,辅以 jpg、jpeg、png 示例图,另有 license、md 说明与 txt 依赖清单,压缩包约 4.47MB,结构紧凑、开箱即用。已有 680 人学习下载。读者可借助 demo_release.py 直接运行推理,观察从 Lab 空间转换、256×256 缩放、着色再融合回全分辨率并转回 RGB 的完整预处理与后处理链路;colorizers 模块封装了模型加载与调用方式,示例图与输出结果便于对照验证。对于想复现经典着色网络、理解深度先验在着色任务中作用,或以此为基线做二次开发的读者,这份代码提供了清晰可读的参考实现与可运行入口。
1. 彩色图像着色:从灰度图到彩色图的深度神经网络方案
手里有一批老照片、医学影像或者监控截图,全是灰度的,想批量上色又不想一张张丢进在线工具里等半天——这个场景下,用深度神经网络做自动着色(Colorization)就是最直接的解法。它的核心思路不复杂:把灰度图当作 L 通道输入,让网络去预测对应的 a、b 两个色度通道,拼回 Lab 空间再转 RGB,就得到彩色图。整套流程用 Python 就能跑通,代码量不大,但选型、损失函数和数据准备这几步决定了最终效果是「能看」还是「翻车」。这篇笔记面向想自己动手复现的开发者,从环境搭建、数据准备、模型搭建到训练调参和避坑,一步步拆开讲,新手能跟着跑,熟手能看到参数边界和常见陷阱。
2. 自动着色的技术选型:为什么是 Lab 空间加卷积神经网络
2.1 着色问题的本质是回归而非分类
灰度图着色的数学本质是:给定亮度 L,预测色度 a 和 b。这是一个一对多的映射问题——同一张灰度图,天空可以是蓝色也可以是橙色,草地可以是绿色也可以是枯黄色。正因为存在多种合理答案,如果用逐像素的 L2 损失去训练,网络会倾向于输出所有可能颜色的均值,结果就是饱和度极低、灰蒙蒙的「安全色」。
常见做法是把 a、b 通道离散化成网格(比如量化成 313 个色块),把回归问题转成分类问题,再用 softmax 输出概率分布。推理时取概率最高的色块或者做退火均值。这样训练更稳定,颜色也更鲜艳。我一般会推荐这种方式,尤其是数据量不大的时候。
2.2 为什么选 Lab 而不是 RGB 或 HSV
RGB 三个通道高度耦合,直接预测 RGB 意味着网络要同时学亮度和颜色,训练难度大。HSV 的色相 H 在低饱和度区域不稳定,灰色像素的 H 值几乎是噪声。Lab 空间把亮度(L)和色度(a、b)解耦,灰度图直接就是 L 通道,网络只需要专注预测 a、b,输入输出关系清晰。这是目前主流着色方案的标准做法。
2.3 网络结构:编码器-解码器加跳跃连接
基础结构用编码器-解码器(Encoder-Decoder)就够了。编码器逐层下采样提取语义特征,解码器逐层上采样恢复空间分辨率。中间加跳跃连接(Skip Connection)把浅层的高频细节传到深层,避免上采样后边缘模糊。编码器可以用几层卷积加池化,解码器用转置卷积或上采样加卷积。如果追求更好的效果,可以把编码器换成预训练的分类网络(比如 ResNet 的前几层),利用 ImageNet 上学到的语义特征来指导着色——网络见过「天空」「草地」「人脸」这些概念后,上色会合理得多。
提示:编码器用预训练权重时,注意输入通道数要改成 1(只接 L 通道),或者把 L 复制成 3 通道再送入。前者需要改第一层卷积,后者不用改结构但计算量稍大。
3. 用 Python 跑通着色模型:环境、数据与训练代码
3.1 环境搭建与依赖安装
先确认 Python 版本,建议 3.8 及以上。核心依赖就几个:PyTorch、NumPy、Pillow、scikit-image。安装命令如下:
pip install torch torchvision numpy pillow scikit-image如果要用 GPU 训练,去 PyTorch 官网对照 CUDA 版本选对应的安装命令。装完后验证一下:
import torch print(torch.__version__) print(torch.cuda.is_available()) # 有 GPU 应返回 True逻辑说明:torch.cuda.is_available()返回 False 时,检查显卡驱动和 CUDA 版本是否匹配。CPU 也能跑,只是训练时间会从几小时变成几天。
3.2 数据准备:把彩色图转成 Lab 并提取 L 通道
训练数据就是一批彩色图片。用 scikit-image 做 RGB 到 Lab 的转换:
import numpy as np from skimage import color, io, transform def load_and_preprocess(img_path, size=256): """读取彩色图,转 Lab,返回 L 通道和 ab 通道""" rgb = io.imread(img_path) # 统一尺寸 rgb = transform.resize(rgb, (size, size), anti_aliasing=True) # RGB 转 Lab,注意 skimage 要求输入为 [0,1] 浮点 lab = color.rgb2lab(rgb) L = lab[:, :, 0] # 亮度通道,范围约 [0, 100] ab = lab[:, :, 1:] # 色度通道,范围约 [-128, 127] # 归一化到 [-1, 1],方便网络训练 L_norm = L / 50.0 - 1.0 ab_norm = ab / 128.0 return L_norm, ab_norm参数说明:size控制输入分辨率,256×256 是常见起点,显存够可以上 512。L / 50.0 - 1.0把 L 从 [0,100] 映射到 [-1,1]。ab / 128.0把色度压到 [-1,1]。归一化这步不能省,否则损失值波动大,收敛慢。
3.3 模型定义:一个最小可用的着色网络
import torch import torch.nn as nn class ColorNet(nn.Module): def __init__(self): super().__init__() # 编码器:输入 1 通道(L),逐层下采样 self.encoder = nn.Sequential( nn.Conv2d(1, 64, 3, stride=2, padding=1), # 256 -> 128 nn.ReLU(inplace=True), nn.Conv2d(64, 128, 3, stride=2, padding=1), # 128 -> 64 nn.ReLU(inplace=True), nn.Conv2d(128, 256, 3, stride=2, padding=1), # 64 -> 32 nn.ReLU(inplace=True), ) # 解码器:逐层上采样,输出 2 通道(a, b) self.decoder = nn.Sequential( nn.ConvTranspose2d(256, 128, 3, stride=2, padding=1, output_padding=1), nn.ReLU(inplace=True), nn.ConvTranspose2d(128, 64, 3, stride=2, padding=1, output_padding=1), nn.ReLU(inplace=True), nn.ConvTranspose2d(64, 2, 3, stride=2, padding=1, output_padding=1), nn.Tanh() # 输出范围 [-1, 1],对应归一化后的 ab ) def forward(self, L): feat = self.encoder(L) ab = self.decoder(feat) return ab逻辑说明:编码器三层卷积把 256×256 降到 32×32,解码器再升回 256×256。output_padding=1是为了让输出尺寸和输入对齐。最后一层用Tanh把输出限制在 [-1,1],和前面 ab 的归一化范围一致。
参数说明:卷积核统一用 3×3,通道数 64→128→256 是常见配置。显存不够就把通道数减半。stride=2代替池化做下采样,保留更多信息。
3.4 训练循环与损失函数
from torch.utils.data import DataLoader, Dataset import os class ColorDataset(Dataset): def __init__(self, img_dir, size=256): self.files = [os.path.join(img_dir, f) for f in os.listdir(img_dir) if f.lower().endswith(('.jpg', '.png', '.jpeg'))] self.size = size def __len__(self): return len(self.files) def __getitem__(self, idx): L, ab = load_and_preprocess(self.files[idx], self.size) # 增加通道维度: (H,W) -> (1,H,W) 和 (2,H,W) L = torch.tensor(L, dtype=torch.float32).unsqueeze(0) ab = torch.tensor(ab, dtype=torch.float32).permute(2, 0, 1) return L, ab # 训练 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = ColorNet().to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) criterion = nn.MSELoss() dataset = ColorDataset('path/to/your/images') loader = DataLoader(dataset, batch_size=16, shuffle=True, num_workers=2) for epoch in range(50): total_loss = 0 for L, ab in loader: L, ab = L.to(device), ab.to(device) pred_ab = model(L) loss = criterion(pred_ab, ab) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() print(f'Epoch {epoch+1}, Loss: {total_loss/len(loader):.4f}')逻辑说明:ColorDataset负责读图、转 Lab、归一化、转 tensor。训练循环就是标准的 PyTorch 流程。损失函数先用 MSE 跑通,后面再换更好的。
参数说明:batch_size=16是 8GB 显存的保守值,显存大可加到 32 或 64。lr=1e-3是 Adam 的常用起点,loss 震荡就降到 1e-4。num_workers设成 CPU 核心数的一半左右,太多反而拖慢。
3.5 推理:把预测的 ab 拼回彩色图
from skimage import color def colorize(model, img_path, size=256): model.eval() L_norm, _ = load_and_preprocess(img_path, size) L_tensor = torch.tensor(L_norm, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device) with torch.no_grad(): ab_pred = model(L_tensor).squeeze(0).permute(1, 2, 0).cpu().numpy() # 反归一化 L_orig = (L_norm + 1.0) * 50.0 ab_orig = ab_pred * 128.0 lab = np.concatenate([L_orig[:, :, None], ab_orig], axis=2) rgb = color.lab2rgb(lab) return (rgb * 255).astype(np.uint8)逻辑说明:推理时不需要计算梯度,用torch.no_grad()省显存。预测出的 ab 反归一化后和原始 L 拼成 Lab,再转 RGB。注意lab2rgb输出是 [0,1] 浮点,乘 255 转成 uint8 才能保存。
参数说明:推理的size要和训练时一致,否则网络看到的分布不匹配,颜色会偏。如果原图不是正方形,先做中心裁剪或 padding 再 resize。
4. 着色效果翻车的五个坑:从灰蒙蒙到颜色溢出
4.1 输出全是灰色或低饱和度
现象:训练 loss 降下去了,但推理结果几乎看不出颜色,整体灰蒙蒙。
原因:用 MSE 损失直接回归 ab 值时,网络学到的是条件均值。对于一张图里某个像素可能是蓝也可能是绿的情况,均值就是灰色。
解决:改用分类方案。把 ab 空间量化成 313 个色块(这是常见做法),网络输出每个色块的概率,损失用交叉熵。推理时取概率最高的色块对应的 ab 值,或者用温度参数做 softmax 退火。颜色饱和度会明显提升。
4.2 颜色溢出到相邻区域
现象:天空的蓝色渗到了建筑物边缘,或者人物衣服的颜色糊到了背景上。
原因:编码器下采样太狠,空间信息丢失严重,解码器上采样时无法精确恢复边界。另外,感受野太大也会导致颜色「漏」到不该去的地方。
解决:加跳跃连接,把编码器浅层的高分辨率特征直接拼到解码器对应层。或者用 U-Net 结构,它的跳跃连接是标配。另一个办法是减小下采样倍数,比如只下采样到 1/4 而不是 1/8。
4.3 训练 loss 不下降或震荡剧烈
现象:loss 在某个值附近来回跳,或者一直不降。
原因:学习率太大、batch size 太小、数据归一化没做对,或者网络输出没有做值域限制。
解决:先检查归一化。L 必须在 [-1,1],ab 也必须在 [-1,1]。然后降学习率,从 1e-3 降到 1e-4 试试。batch size 至少 8,太小梯度噪声大。最后确认最后一层有没有 Tanh,没有的话输出可能跑到几百,loss 直接爆炸。
4.4 推理时颜色和训练时不一致
现象:训练集上的图着色正常,换一张新图颜色就偏了。
原因:推理时的预处理和训练时不一致。比如训练用了 anti_aliasing 的 resize,推理用了最近邻;或者训练时做了归一化,推理忘了。
解决:把预处理封装成一个函数,训练和推理都调同一个。检查 resize 方法、归一化参数、通道顺序是否完全一致。这个坑很隐蔽,血泪经验是写个单元测试对比训练和推理的预处理输出。
4.5 显存不够导致训练中断
现象:跑几个 batch 就报 CUDA out of memory。
原因:输入分辨率太高、batch size 太大、模型通道数太多,或者没及时释放中间变量。
解决:先把 batch size 降到 4 试试。不够就把输入从 256 降到 128。还不够就砍通道数,把 256 改成 128。另外,训练循环里用del及时删掉不用的中间 tensor,配合torch.cuda.empty_cache()。如果这些都不行,用梯度累积模拟大 batch。
5. 进阶技巧:用感知损失和类别平衡把颜色做自然
基础版跑通后,想让颜色更自然、更符合人类审美,有两个方向值得试。
感知损失(Perceptual Loss)。MSE 只关心像素值差异,不关心颜色看起来是否合理。感知损失的做法是把预测图和真实图都送进一个预训练的分类网络(比如 VGG),取中间某层的特征图,算它们的 L2 距离。这样损失函数约束的是「语义特征要像」,而不是「每个像素要一模一样」。颜色会更符合语义——天空像天空,草地像草地。实现上,把 VGG 的前几层冻结,接在着色网络后面,训练时只更新着色网络的参数。
import torchvision.models as models vgg = models.vgg16(pretrained=True).features[:16].to(device).eval() for p in vgg.parameters(): p.requires_grad = False def perceptual_loss(pred_rgb, target_rgb): # pred_rgb, target_rgb: (B, 3, H, W), 范围 [0,1] feat_pred = vgg(pred_rgb) feat_target = vgg(target_rgb) return nn.functional.mse_loss(feat_pred, feat_target)参数说明:features[:16]取的是 VGG 前几层,感受野小,关注局部纹理和颜色。取太深层会丢失空间信息。权重上,感知损失和 MSE 按 1:10 或 1:100 混合,具体看效果调。
类别平衡(Class Rebalancing)。ab 空间量化成 313 个色块后,分布极不均匀——灰色、棕色、蓝色占了大头,红色、紫色、亮绿色很少。直接训练的话,网络会偏向高频色块,稀有颜色永远预测不出来。做法是给每个色块按出现频率的倒数加权,频率越低权重越高。这样网络会被迫关注稀有颜色,输出更多样化。
# 假设 class_freq 是 313 个色块的频率统计 weights = 1.0 / (class_freq + 1e-6) weights = weights / weights.sum() * len(weights) # 归一化 criterion = nn.CrossEntropyLoss(weight=torch.tensor(weights, dtype=torch.float32).to(device))参数说明:1e-6防止除零。归一化让权重均值为 1,避免整体 loss 尺度变化太大。权重太极端会导致训练不稳定,可以开根号或取对数缓和一下。
验证方法:除了看 loss,更直观的是固定几张测试图,每个 epoch 存一次着色结果,拼成网格图观察颜色变化。另外可以算 PSNR 和 SSIM,但这两个指标和人类感知的相关性有限,只能做参考。真正靠谱的还是人眼看。
我自己踩过的坑是:一开始只盯着 loss 调,结果 loss 很低但颜色一塌糊涂。后来养成习惯,每训几个 epoch 就导出一批测试图的着色结果,肉眼过一遍。颜色偏了、溢出了、灰了,一眼就能看出来,比看数字快得多。希望帮到你。
本文还有配套的精品资源,点击获取