Skip to content

필사 모드: 人人可懂的 AI 第5篇 — 用47万参数的 U-Net 给黑白照片上色,以及颜色为何发灰

中文
0%
정확도 0%
💡 왼쪽 원문을 읽으면서 오른쪽에 따라 써보세요. Tab 키로 힌트를 받을 수 있습니다.

引言 — 当输出本身是图像

此前的输出都是文本:第1篇是句子,第3篇是词,第4篇是字幕。这次 输出本身就是图像

任务是上色。输入黑白照片,返回彩色照片。

输入: 32×32 黑白 (1 通道)
输出: 32×32 彩色 (3 通道)

不需要另外找数据。把 CIFAR-10 的彩色图转成黑白, 输入与标签的配对就自动生成了

rgb = torch.stack([...])                                    # 原始彩色 = 标签
gray = (rgb * torch.tensor([0.299, 0.587, 0.114])
        [None, :, None, None]).sum(1, keepdim=True)         # 转黑白 = 输入

权重 0.299, 0.587, 0.114 是标准亮度公式,反映了人眼对绿色最敏感、对蓝色最迟钝。这样构造的数据叫 自监督学习 (self-supervised):没人打标注,标签却存在。

U-Net — skip connection 运送的东西

上色有个棘手的要求。 要知道"是什么"才能定颜色,同时还要知道"在哪里"。 认出是青蛙才会涂绿,但绿色越过青蛙的轮廓就毁了。

普通的编码器-解码器很难同时满足两者,因为压缩图像时获得了语义却丢失了位置。U-Net 用 skip connection 解决这一点。

class TinyUNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.e1 = block(1, 32); self.e2 = block(32, 64); self.e3 = block(64, 128)
        self.d2 = block(128 + 64, 64); self.d1 = block(64 + 32, 32)
        self.out = nn.Conv2d(32, 3, 1)

    def forward(self, x):
        s1 = self.e1(x)                       # 32x32  — 最精细
        s2 = self.e2(F.max_pool2d(s1, 2))     # 16x16
        h  = self.e3(F.max_pool2d(s2, 2))     # 8x8    — 最抽象
        h = self.d2(torch.cat([F.interpolate(h, scale_factor=2), s2], 1))
        h = self.d1(torch.cat([F.interpolate(h, scale_factor=2), s1], 1))
        return torch.sigmoid(self.out(h))

关键是 torch.cat([上采样后的h, s2], 1)。解码器在恢复分辨率时, 把同分辨率的编码器输出原样取来拼接

  • 从下方上来的 h —— 分辨率低,但携带"这是青蛙"这类语义
  • 从旁边来的 s2 —— 未经压缩,边缘与纹理完好

解码器同时看这两路。语义来自深层通路,位置来自 skip 通路。通道数写成 128 + 64 正因如此。

sigmoid 收尾也是有意的:输出必须是 0 到 1 之间的 RGB 值,用激活函数把范围钉住。

训练

opt = torch.optim.AdamW(model.parameters(), lr=2e-3)

while not run.over_budget():          # 12分钟预算
    ix = torch.randint(0, len(rgb), (128,))
    x, y = gray[ix].to(dev), rgb[ix].to(dev)
    loss = F.l1_loss(model(x), y)
    opt.zero_grad(); loss.backward(); opt.step()

损失用的是 L1 (平均绝对误差),即直接度量预测 RGB 与真实 RGB 之差的回归损失。这个选择决定了结果的性格,后面再展开。

[   14.1s] pairs=30000
[   14.2s] params {"params": 472323, "res": 32}
[  722.0s] SUMMARY {"steps": 62136, "final_l1": 0.0231}

62,136 步,最终 L1 为 0.023。像素值在 0 到 1 之间,也就是 平均误差 2.3% 。单看数字非常漂亮。

结果 — 形状对了,颜色发灰

每行是两组三张,顺序为 输入黑白 / 模型预测 / 真实彩色

CIFAR-10 上色结果 —— 每三张一组,从左到右依次为输入黑白、模型预测、真实彩色

诚实地读一读。

做得好的 —— 形状被精确保留。猫的毛、船的甲板结构、青蛙的腿、汽车的窗框都没有糊。这是 skip connection 起了作用。颜色的 方向 也大体正确:青蛙周围偏绿,船周围偏蓝,猫偏棕。

没做好的 —— 颜色整体 发灰 。标签里鲜艳的红色汽车,预测里成了暗灰褐。标签里红色船体和蓝色海面,预测里成了浑浊的灰蓝。不少结果接近棕褐色调。

颜色为何发灰 —— L1 损失的本性

这不是因为模型小,而是 选了 L1 的结果

设想一张黑白汽车照片。这辆车可能是红的、蓝的,也可能是白的。仅凭黑白信息无法判定。换句话说, 这是一个有多个正确答案的问题

此时 L1 对模型说:把与标签的绝对误差降到最小。在多种答案都可能时,遵守这条指令最安全的做法是什么?

押所有可能性的中间值。

押红色而真值是蓝色,罚分很重。押灰色则无论真值是什么,罚分都被压在中等水平。训练越久,模型越向这个安全策略收敛,产物就是 低饱和度的颜色

这就是 L1 为 0.023 这样漂亮的数字与发灰的颜色能够并存的原因。 L1 恰恰把交代给它的事做好了,只是交代给它的并不是我们想要的。

真实的上色模型怎么做

这个问题广为人知,解法也分成几类。

把颜色变成分类问题 —— Zhang 等人 2016 年的工作把色彩空间分成 313 个区间,转而预测每个像素属于哪个区间。分类允许把概率分给多个候选,因此不会被平均抹平;采样时挑一个鲜艳的即可。

加入对抗损失 —— GAN 的判别器判断"这张图像不像真的"。发灰的颜色不像真实照片,会被罚分。L1 负责形状,判别器负责鲜艳度。

使用感知损失 —— 不直接比较像素值,而在预训练网络的特征空间里比较,更接近人对相似性的判断。

三者有个共同点: 把损失设计成"逃向平均不再划算"

小结

项目
参数量472,323 (0.47M) —— 系列最小
训练时间722.0秒
步数62,136
最终 L10.0231
数据CIFAR-10 三万对 (自监督)

47 万参数,在保持形状完好的前提下完成了上色。skip connection 让细节存活,自监督的方式让标注成本为零。

而本篇真正的教训在发灰的那一侧。 损失函数是用来定义目标的,不是用来达成目标的工具。 选定 L1 的那一刻,"接近平均的答案更安全"这条规则也随之确立,模型忠实地遵守了它。若结果不合心意,先要回头审视的不是模型,而是 你让它最小化的是什么

下一篇正面处理这个问题 —— 扩散模型 。它的结构能在多个答案皆可能时挑出一个,而不是逃向平均。

🧠 理解度自测

1. 没有 skip connection,上色结果会有什么不同?

形状会糊。压缩到 8×8 的过程中边缘与纹理信息被丢弃,没有 skip connection 解码器就无从恢复。颜色也许大体不错,但轮廓会烂、细节会消失。"语义走深层通路、位置走 skip 通路"的分工就此瓦解。

2. L1 损失 0.023 是个好数字,颜色为什么还发灰?

上色是多解问题:黑白的汽车可能是红也可能是蓝。L1 要求最小化绝对误差,所以在多种可能并存时,押中间值最安全 —— 押鲜艳色一旦押错罚分很重,而灰色对任何真值都只挨中等罚分。L1 达成了它自己的目标,只是那目标不是我们的。

3. 为什么把这种训练称作自监督学习?

因为没人打标注,标签却存在。用亮度公式把 CIFAR-10 的彩色图转成黑白便得到输入,原始彩色就是标签。零标注成本即可造出三万对样本。

4. 把颜色改成 313 个区间的分类问题,为什么能缓解饱和度问题?

分类的输出是各区间上的概率分布。红蓝都合理时,可以同时给两个区间高概率,采样时再鲜艳地挑其中一个。回归必须给出单一数值,只能落在两色之间的灰;分类没有这个强制。

参考资料

현재 단락 (1/73)

此前的输出都是文本:第1篇是句子,第3篇是词,第4篇是字幕。这次 **输出本身就是图像** 。

작성 글자: 0원문 글자: 3,490작성 단락: 0/73