LearnIR:不用已知退化算子,也能做扩散后验图像复原

2026-07-02 41 预计阅读时间: 1 分钟
来源: my.oschina.net AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:10 分钟

真实照片里的退化很少是“干净”的:人脸阴影可能混着曝光偏差,雾天图像可能同时有低对比度、噪声和压缩痕迹。很多扩散式图像复原方法在论文设定里表现不错,但一碰到复杂异质退化,就容易遇到三个麻烦:保真度不够、逐步采样误差累积,或者需要一个现实中并不知道的前向退化算子。LearnIR 的核心价值,正是把问题推向更贴近真实场景的一侧:通过学习后验采样中的梯度校正分布,让扩散复原不再强依赖已知前向算子。

它要解决的不是“生成得好看”,而是“复原得可信”

图像复原和纯生成不一样。给定一张有阴影的人脸或一张雾天照片,模型不能只生成一张“看起来合理”的图,它还必须尽量保留原图身份、结构、纹理和局部细节。

传统扩散复原路线通常会把预训练扩散模型当作强先验,再在采样过程中加入观测约束。问题在于,很多约束依赖前向算子:比如模糊核、下采样方式、噪声模型。真实场景中的阴影、雾、混合退化并不总能被一个清晰的数学算子描述。

LearnIR 的思路可以概括为:

  • 不直接假设已知前向退化过程;
  • 训练一个轻量网络来预测梯度校正分布;
  • 在扩散后验采样过程中,用这个校正分布引导复原结果靠近观测图像;
  • 再通过动态分辨率模块抑制采样噪声,提高稳定性。

这让它更像是给扩散模型加了一个“会看退化图的校正器”,而不是要求开发者提前写出真实世界的退化公式。

后验采样校正:把未知退化交给网络学习

在扩散复原里,后验采样通常关心的是:在已知退化图像 y 的条件下,如何从清晰图像分布里采样出合理的 x。如果有已知前向算子 A,可以用 A(x)y 的差异来校正采样方向。但在人脸去阴影、去雾这类任务里,A 往往并不可得。

LearnIR 的关键变化是学习这个校正项。轻量网络接收当前采样状态、退化观测以及时间步等信息,输出用于修正扩散采样轨迹的梯度校正分布。这样做的好处是:

  • 避免手写不准确的退化算子;
  • 减少采样过程中因为错误约束带来的误差累积;
  • 对复杂异质退化更友好,例如阴影、雾、噪声、局部曝光异常同时存在的图像。

这类方法的工程含义很直接:你可以保留扩散模型的强先验,同时把“如何贴合输入退化图”这件事交给一个更轻的可学习模块。

动态分辨率:不是所有时间步都该用同一尺度

扩散采样的每一步都会处理噪声、结构和纹理。如果始终在同一分辨率上运行,可能出现两个问题:高分辨率阶段容易放大噪声,低分辨率阶段又可能损失细节。

LearnIR 设计了动态分辨率模块,用来进一步抑制噪声。可以把它理解为一种更灵活的采样调度:不同阶段关注不同尺度的信息,让模型在结构恢复、细节补全和噪声控制之间取得更稳的平衡。

对开发者来说,这一点很重要。很多图像复原系统上线后,失败案例并不是“完全复原不了”,而是局部区域发脏、脸部纹理漂移、雾被去掉后天空出现噪点。动态分辨率模块正是针对这类采样噪声和局部不稳定性做补强。

可以这样实践:封装一个 LearnIR 风格的复原 CLI

下面的例子不是论文源码,而是一个可直接运行、便于改造的最小项目骨架。它模拟 LearnIR 风格系统的工程接口:输入一张退化图像,调用“扩散先验 + 梯度校正器 + 动态分辨率调度”的复原流程。你可以把 MockLearnIRRestorer 替换成真实模型推理代码。

先安装依赖:

python -m venv .venv
source .venv/bin/activate
pip install pillow numpy

创建 learnir_restore.py

from dataclasses import dataclass
from pathlib import Path
import argparse
import numpy as np
from PIL import Image, ImageFilter, ImageEnhance


@dataclass
class RestoreConfig:
    task: str = "dehaze"          # 可改为: deshadow / dehaze
    steps: int = 20
    base_resolution: int = 512
    strength: float = 0.65


class MockLearnIRRestorer:
    """示例工程壳:用传统图像操作占位,便于替换为真实 LearnIR 推理。"""

    def __init__(self, config: RestoreConfig):
        self.config = config

    def dynamic_resize(self, image: Image.Image, step: int) -> Image.Image:
        ratio = 0.6 + 0.4 * (step + 1) / self.config.steps
        size = max(128, int(self.config.base_resolution * ratio))
        return image.resize((size, size), Image.Resampling.BICUBIC)

    def gradient_correction_proxy(self, image: Image.Image) -> Image.Image:
        if self.config.task == "dehaze":
            image = ImageEnhance.Contrast(image).enhance(1.0 + self.config.strength)
            image = ImageEnhance.Color(image).enhance(1.0 + self.config.strength * 0.25)
        elif self.config.task == "deshadow":
            image = ImageEnhance.Brightness(image).enhance(1.0 + self.config.strength * 0.35)
            image = ImageEnhance.Contrast(image).enhance(1.0 + self.config.strength * 0.4)
        return image.filter(ImageFilter.SMOOTH_MORE)

    def restore(self, image: Image.Image) -> Image.Image:
        original_size = image.size
        current = image.convert("RGB")

        for step in range(self.config.steps):
            low_or_mid = self.dynamic_resize(current, step)
            corrected = self.gradient_correction_proxy(low_or_mid)
            current = corrected.resize(original_size, Image.Resampling.BICUBIC)

        return current


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--input", required=True)
    parser.add_argument("--output", required=True)
    parser.add_argument("--task", choices=["dehaze", "deshadow"], default="dehaze")
    parser.add_argument("--steps", type=int, default=20)
    args = parser.parse_args()

    config = RestoreConfig(task=args.task, steps=args.steps)
    restorer = MockLearnIRRestorer(config)

    image = Image.open(args.input)
    restored = restorer.restore(image)
    Path(args.output).parent.mkdir(parents=True, exist_ok=True)
    restored.save(args.output)
    print(f"saved restored image to {args.output}")


if __name__ == "__main__":
    main()

运行:

python learnir_restore.py --input hazy_face.jpg --output outputs/restored.jpg --task dehaze --steps 20

如果你接入真实 LearnIR 或类似模型,可以把上面三个位置替换掉:

  • gradient_correction_proxy():替换为轻量校正网络的前向推理;
  • dynamic_resize():替换为论文或项目中的动态分辨率策略;
  • restore():替换为真实扩散采样循环,把校正网络输出注入每个采样步。

一个更接近真实系统的接口可能长这样:

# 伪代码:展示接入点,不代表论文 API
x_t = diffusion_prior.init_noise(observation=image)

for t in scheduler.timesteps:
    score = diffusion_prior.predict_score(x_t, t)
    correction = correction_net.predict_distribution(x_t, image, t)
    x_t = scheduler.step(score + correction.mean, x_t, t)
    x_t = resolution_controller.adjust(x_t, t)

restored = diffusion_prior.decode(x_t)

什么时候值得采用这条路线

如果你的图像复原任务退化类型清晰、前向模型准确,比如固定核超分、受控噪声去噪,传统带算子约束的方法仍然有优势:可解释、调参路径明确、工程成本较低。

LearnIR 这类方法更适合下面的场景:

  • 输入来自真实世界,退化混杂且难以建模;
  • 任务强调身份、结构和细节保真,例如人脸去阴影;
  • 不能接受“图像变漂亮但内容漂移”的结果;
  • 希望利用扩散先验,但不想依赖未知或错误的前向算子;
  • 有条件训练或部署一个额外的轻量校正网络。

落地时建议重点检查三件事:复原前后身份和结构是否一致,局部纹理是否被幻觉化,动态分辨率策略是否在不同尺寸输入上稳定。图像复原不是简单的美化滤镜,越接近真实业务,越要把保真度、稳定性和失败边界放在同一个评估表里。


相关推荐