共计 3828 个字符,预计需要花费 10 分钟才能阅读完成。
背景痛点:传统图像去噪的局限性
传统图像去噪算法如 BM3D(Block-Matching 3D)和 NLM(Non-Local Means)在简单噪声场景下表现尚可,但在复杂噪声或高动态范围图像中往往力不从心。这些方法主要依赖手工设计的特征和统计假设,比如 BM3D 利用图像块相似性进行协同滤波,NLM 通过像素邻域相似性加权平均去噪。它们的局限性主要体现在三个方面:

- 噪声假设单一:传统方法通常假设噪声服从高斯分布,而真实场景噪声往往复杂多变(如泊松 - 高斯混合噪声)
- 细节保留不足:在强噪声下容易过度平滑,导致纹理细节丢失
- 参数敏感:滤波窗口大小、相似性阈值等参数需要针对不同场景反复调整
技术对比:扩散模型 vs GAN/VAE
| 维度 | 扩散模型 | GAN | VAE |
|---|---|---|---|
| 训练稳定性 | 高(分步训练) | 低(模式坍塌风险) | 中等(后验坍缩风险) |
| 生成质量 | 极高(渐进式生成) | 高(依赖架构设计) | 中等(模糊倾向) |
| 计算成本 | 高(多步迭代) | 中等 | 低 |
| 模式覆盖 | 全面(理论上界明确) | 部分(判别器限制) | 保守(KL 约束) |
| 收敛速度 | 慢 | 快(对抗训练) | 快 |
核心实现:AA 去噪扩散模型
前向过程(加噪)
扩散模型通过逐步添加高斯噪声破坏图像,定义噪声调度函数 β_t:
import torch
def linear_beta_schedule(timesteps, beta_start=1e-4, beta_end=2e-2):
return torch.linspace(beta_start, beta_end, timesteps)
class ForwardProcess:
def __init__(self, timesteps=1000):
self.betas = linear_beta_schedule(timesteps)
self.alphas = 1. - self.betas
self.alpha_bars = torch.cumprod(self.alphas, dim=0)
def q_sample(self, x0, t, noise=None):
"""
前向扩散过程:q(x_t | x_0)
输入:x0: 原始图像 [B,C,H,W]
t: 时间步 [B,]
返回:加噪后的图像
"""
if noise is None:
noise = torch.randn_like(x0)
sqrt_alpha_bar = self.alpha_bars[t].sqrt().view(-1,1,1,1)
sqrt_one_minus_alpha_bar = (1 - self.alpha_bars[t]).sqrt().view(-1,1,1,1)
return sqrt_alpha_bar * x0 + sqrt_one_minus_alpha_bar * noise
反向过程(去噪)
关键是通过神经网络预测噪声:
import torch.nn as nn
class ResidualBlock(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.conv = nn.Sequential(nn.Conv2d(in_channels, in_channels, 3, padding=1),
nn.GroupNorm(8, in_channels),
nn.SiLU(),
nn.Conv2d(in_channels, in_channels, 3, padding=1),
nn.GroupNorm(8, in_channels)
)
def forward(self, x):
return x + self.conv(x)
class DenoiseModel(nn.Module):
def __init__(self, in_channels=3, hidden_dims=[64,128,256]):
super().__init__()
# 时间步嵌入
self.time_embed = nn.Sequential(nn.Linear(1, hidden_dims[0]),
nn.SiLU(),
nn.Linear(hidden_dims[0], hidden_dims[0])
)
# 编码器
self.encoder = nn.ModuleList([
nn.Sequential(nn.Conv2d(in_channels, hidden_dims[0], 3, padding=1),
ResidualBlock(hidden_dims[0])
)
])
# 解码器
self.decoder = nn.ModuleList([
nn.Sequential(ResidualBlock(hidden_dims[0]),
nn.Conv2d(hidden_dims[0], in_channels, 3, padding=1)
)
])
def forward(self, x, t):
# t 形状转换 [B,] -> [B,1] -> [B,D]
t_emb = self.time_embed(t.float().unsqueeze(-1))
# 编码过程
for layer in self.encoder:
x = layer(x)
# 加入时间信息
x = x + t_emb.view(-1, t_emb.shape[1], 1, 1)
# 解码过程
for layer in self.decoder:
x = layer(x)
return x
训练循环
from torch.utils.data import DataLoader
from torchvision.datasets import ImageFolder
def train(model, dataloader, optimizer, device, epochs=100):
forward_process = ForwardProcess()
model.train()
for epoch in range(epochs):
for batch_idx, (x0, _) in enumerate(dataloader):
x0 = x0.to(device)
batch_size = x0.shape[0]
# 随机采样时间步
t = torch.randint(0, forward_process.timesteps, (batch_size,), device=device)
# 前向加噪
noise = torch.randn_like(x0)
xt = forward_process.q_sample(x0, t, noise)
# 预测噪声
pred_noise = model(xt, t)
# 计算损失
loss = nn.MSELoss()(pred_noise, noise)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
# 数据加载示例
dataset = ImageFolder("path/to/images", transform=transforms.ToTensor())
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
model = DenoiseModel().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
train(model, dataloader, optimizer, device)
性能考量
显存优化技巧
-
梯度检查点:
from torch.utils.checkpoint import checkpoint # 在模型 forward 中替换 x = checkpoint(layer, x) # 代替直接调用 layer(x) -
混合精度训练:
scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): pred_noise = model(xt, t) loss = nn.MSELoss()(pred_noise, noise) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
推理时质量 - 速度权衡
- 步数缩减:从 1000 步减少到 50-100 步(需调整 β 调度)
- 蒸馏技术:训练学生模型模仿多步去噪过程
- 隐空间加速:在低维空间执行扩散过程
避坑指南
扩散步数选择
- 低噪声图像:50-200 步足够
- 高噪声 / 医学图像:建议 500-1000 步
- 可通过信噪比 (SNR) 分析确定临界步数:
def find_optimal_steps(alpha_bars, target_snr=0.01): snr = alpha_bars / (1 - alpha_bars) return (snr > target_snr).sum().item()
常见训练问题
- 梯度爆炸:
- 解决方案:梯度裁剪(
nn.utils.clip_grad_norm_(model.parameters(), 1.0)) -
检查时间步嵌入是否归一化
-
模式坍缩:
- 现象:生成结果多样性不足
-
解决方法:增加噪声调度方差,使用余弦调度替代线性调度
-
生成模糊:
- 检查损失函数(推荐使用 L1+L2 混合损失)
- 增加模型容量(更多残差块)
开放性问题
扩散模型在图像修复领域展现出强大潜力,但仍有若干值得探索的方向:
- 如何结合 CLIP 等跨模态模型实现文本引导的智能修复?
- 能否设计动态扩散步数,对简单区域快速收敛,复杂区域精细处理?
- 在计算资源受限的边缘设备上,如何实现实时去噪(<50ms 延迟)?
期待读者在实践中发现更多创新应用,也欢迎分享你们遇到的独特挑战和解决方案。
正文完
