各位做视觉算法和图像处理的朋友,大家好。
在真实场景中,天气退化往往不是单一类型。户外摄像头可能同时遇到雨雾交加,自动驾驶车辆会在雪天遭遇低照度,监控系统在夜间又得处理模糊与噪声。过去我们习惯针对每种退化训练一个专用模型,但在实际落地时,这种做法维护成本高、部署体积大、泛化能力也有限。近年来,All-in-One Weather Restoration(一体化天气恢复)逐渐成为研究热点,目标是让一个模型处理雨、雪、雾、低光等多种天气退化。
今天这篇文章,围绕Efficient All-in-One Weather Restoration using Spectral Harmonization这篇论文的核心思路展开,讲清楚它解决什么问题、频谱协调是怎么工作的、整体网络如何设计,并给出基于 PyTorch 的核心模块实现思路,帮助大家快速理解并复现这类方法。
文章适合对图像恢复、底层视觉、深度学习模型设计有一定基础的开发者阅读。如果你刚刚接触图像去雨、去雾、去雪等任务,本文也会先补充必要的基础概念,再逐步进入核心原理和代码分析。
1. 背景与核心概念
1.1 什么是天气退化图像恢复
天气退化图像恢复,是底层计算机视觉中的一个经典研究方向。它的输入是一张被雨、雪、雾、低光、噪声等环境因素污染的照片,输出是一张恢复清晰纹理、正确颜色和结构细节的干净图像。
单任务模型在过去十年里取得了不错的成绩,比如专门去雨的 DerainNet、专门去雾的 AOD-Net、专门去雪的 DesnowNet。这些模型在各自的任务上表现出色,但一旦遇到训练时没见过的退化组合,效果就会明显下降。
这是因为单任务模型把退化类型当成固定不变的假设。现实世界并不是这样:一场雨夹雪天气会同时包含雨和雪,雾天往往伴随低对比度和色彩偏移,夜间图像还会叠加低光噪声。如果线下维护多个模型,推理时还要先判断当前输入属于哪类退化,流程复杂且容易误判。
1.2 All-in-One 一体化恢复的提出
All-in-One Weather Restoration 的思路很直接:用一个统一的模型,接受多种天气退化输入,输出干净恢复结果。模型不再关心输入是哪一种退化,而是直接学习“退化图像到干净图像”的通用映射。
这种设置有以下几点优势:
- 部署简化:模型体积缩小,只需要维护一个权重文件。
- 推理高效:不需要先做退化分类,避免级联错误。
- 泛化提升:模型有机会学到不同退化之间的共享特征和差异特征。
但“一个模型处理所有退化”听起来简单,做起来却不容易。不同退化的物理模型差异很大:
- 雨滴和雨线是局部高频叠加。
- 雾气是全局低频散射,还伴随颜色偏移。
- 雪点形状多样、大小不一。
- 低光图像整体暗部噪声严重。
如果用简单卷积去混洗所有这些特征,网络很容易在特征空间发生冲突。关键问题就变成了:如何设计一种机制,让不同退化特征既能共享通用表示,又能保持各自特有的结构信息?
1.3 Spectral Harmonization 频谱协调
Spectral Harmonization 是这篇论文给出的答案。这里的 Spectral 指图像的频率域表示,也就是把图像从空间域通过傅里叶变换转换到频域,得到幅度谱和相位谱。
在图像恢复任务中,频率域有一个非常重要的特性:
- 幅度谱(Amplitude Spectrum)记录像素强度的整体分布,决定图像全局风格、亮度和对比度,与天气退化类型高度相关。
- 相位谱(Phase Spectrum)记录结构的空间位置关系,决定边缘、纹理和物体轮廓,更多与内容相关。
频谱协调的核心思想,是在频域上对不同退化分支的特征进行重新校准与融合,而不是直接在空间域做简单的加法或拼接。这样做的直接好处是:网络可以在频率分量级别上识别出哪些信息属于退化伪影、哪些信息属于干净结构,从而避免特征冲突。
在进入代码实现之前,我们先从任务定义和现有方法两个角度做更细致的分析。
2. 任务定义与现有方法分析
2.1 数学形式化
将天气退化图像恢复看成一个监督学习问题,训练集可以表示为:
{ (X_r, Y), (X_s, Y), (X_f, Y), ... }其中 Y 是干净参考图,X_r、X_s、X_f 分别代表雨、雪、雾等退化输入。目标函数通常写为:
L = || F(X) - Y ||_1 + L_perceptual + L_frequency第一项是逐像素重建损失,第二项是感知损失,第三项是频域损失。
对 All-in-One 模型来说,X 不再标注具体退化类型,模型直接学习条件分布 P(Y|X)。这个任务的难点在于,不同退化共享图像内容和结构信息,但退化的频域模式完全不同。
2.2 现有 All-in-One 方法分类
目前常见的 All-in-One 天气恢复方法有两大类。
第一类是动态卷积和稀疏门控方法。
这类方法通过可学习的路由模块,让网络自动为不同输入分配不同的卷积核或网络分支。优点是灵活性高,缺点是路由模块本身需要额外监督或容易陷入局部最优。
第二类是特征解耦方法。
这类方法试图把特征分解为“内容无关分量”和“退化相关分量”,只对退化部分做处理。优点是原理清晰,缺点是解耦不彻底时,容易丢失纹理细节。
Spectral Harmonization 更接近第二类,但它不是在空间域解耦,而是借助傅里叶变换,在频域中对多分支特征做协调和融合。这样做在数学上更自然,因为不同天气退化在频域的分布差异,比在空间域更明显且更容易分离。
2.3 为什么要引入频域处理
空间域卷积的感受野有限,无法轻松获取全局频带信息。虽然可以通过堆叠更多卷积层扩大感受野,但带来的参数和计算开销都不小。
频域处理的好处在于:
- 全局感受野:傅里叶变换天然覆盖整张图像,幅度谱和相位谱直接反映全局统计特性。
- 退化解耦友好:雨、雪、雾在频带上表现出明显差异,便于网络学习和分离。
- 计算效率高:FPT(傅里叶变换) 运算本身可以通过 FFT 算法高效实现,加上频域特征图尺寸天然具备压缩特性。
Spectral Harmonization 正是利用这些性质,对多分支编码器输出的特征进行频域对齐和融合。
3. 核心方法拆解:整体架构与频谱协调原理
这一节我们从宏观到微观逐步拆解。
3.1 整体架构
论文提出的网络整体结构可以概括为:
输入图像 │ ▼ ┌─────────────────────────────┐ │ Multi-Branch Encoder │ │ 主干 + N 个退化感知分支 │ └─────────────────────────────┘ │ 多组特征 ▼ ┌─────────────────────────────┐ │ Spectral Harmonization │ │ 频域协调模块 │ └─────────────────────────────┘ │ 协调后特征 ▼ ┌─────────────────────────────┐ │ Decoder + Reconstruction │ └─────────────────────────────┘ │ ▼ 干净图像大致流程是:
- 输入退化图像先经过主干编码器提取基础特征。
- 多个退化感知分支分别处理不同退化类型的专属信息。
- 特征送入 Spectral Harmonization 模块,在频域将多分支特征协调融合。
- 解码器把融合特征重建为干净图像。
这种设计保留了多任务学习中“共享-差异”的优点,同时通过频域协调抑制了特征冲突。
3.2 频域中的幅度谱和相位谱处理
以二维图像特征 X 为例,其形状为 (B, C, H, W),其中 B 是批量大小,C 是通道数,H 和 W 是空间尺寸。
对每个通道做二维傅里叶变换:
F(u, v) = FFT2D(x(h, w))然后可以分解为幅度和相位:
A(u, v) = |F(u, v)| # 幅度谱 P(u, v) = angle(F(u, v)) # 相位谱Spectral Harmonization 的核心不是简单地把不同分支的幅度谱相加,而是设计门控或注意力机制,对不同频段的重要性加权。
在天气退化中:
- 雨线是高频突发信号,表现为幅度谱中的高频突起。
- 雾气是全局平滑效果,主要集中在低频段。
- 雪花在频域中会出现离散冲击点。
如果直接在空间域叠加特征,网络需要很深的层级才能区分这些细节。而在频域,只需要对相应频段施加不同权重,就能快速实现退化抑制或信息增强。
3.3 Spectral Harmonization 模块设计思路
整个模块可以拆成三个子步骤。
第一步,把多个分支提取的特征分别做 FFT,得到各自的幅度谱和相位谱。
第二步,将幅度谱拼接或相加,送入一个小型卷积网络或者注意力网络,生成一个协调后的融合幅度谱。相位谱可以保留最为详细的某一个分支,也可以全部拼接后卷积重建。
第三步,将融合后的幅度谱与处理后的相位谱组合,做逆傅里叶变换(IFFT),回到空间域,得到协调后的重建特征。
这里的逻辑是:内容结构主要由相位谱承载,而退化类型和退化强度主要由幅度谱承载。因此,幅度谱更适合跨分支融合和校准,相位谱则更适合保留细节。
从行业落地角度来看,这类方法的工程设计也很值得参考。不同任务的模型往往在推理框架、算子库和硬件调度上有所差异,频域模型在 TensorRT、ONNX Runtime 等部署环境中要注意 FFT 算子是否支持,部分嵌入式芯片对 FFT 的加速并不友好。这一点我们放在后面的工程建议部分展开。
4. 核心代码实现:基于 PyTorch 的简化版本
下面给出核心模块的 PyTorch 实现思路。这里说明一下,这是根据论文思路整理的简化示例,用于理解网络结构和频域处理流程,实际训练和推理请以论文官方代码为准。
4.1 项目结构建议
在动手写代码之前,我们先规划好项目目录:
weather_restoration/ ├── models/ │ ├── __init__.py │ ├── encoder.py # 多分支编码器 │ ├── harmonization.py # 频谱协调模块 │ └── decoder.py # 解码器 ├── losses/ │ ├── __init__.py │ └── restoration_loss.py ├── data/ │ └── dataset.py ├── train.py └── test.py下面的代码片段均以该目录结构为参考。
4.2 频域处理基础函数
在实现频谱和谐调之前,先封装傅里叶变换的处理函数。
# 文件路径:models/harmonization.py import torch import torch.nn as nn import torch.fft def to_frequency(x): """ 将空间域特征转换到频域,返回幅度谱和相位谱。 输入 x 形状:(B, C, H, W) """ fft_result = torch.fft.fft2(x, norm="backward") fft_shifted = torch.fft.fftshift(fft_result) amplitude = torch.abs(fft_shifted) phase = torch.angle(fft_shifted) return amplitude, phase def from_frequency(amplitude, phase): """ 将幅度谱和相位谱组合并转换回空间域。 """ complex_tensor = amplitude * torch.cos(phase) + 1j * amplitude * torch.sin(phase) fft_ishifted = torch.fft.ifftshift(complex_tensor) x = torch.fft.ifft2(fft_ishifted, norm="backward") x = torch.real(x) return x这里使用了torch.fft.fft2和torch.fft.ifft2。fftshift的作用是把零频移到中心,方便后续对不同频段施加注意力。
需要注意,FFT 在 GPU 上计算时,不同尺寸的特征图性能差异较大,建议在实现时固定特征图尺寸,避免在动态尺寸上频繁计算 FFT。
4.3 Spectral Harmonization 模块实现
下面实现核心的频谱协调模块。
# 文件路径:models/harmonization.py class SpectralHarmonization(nn.Module): """ 频谱协调模块:将多个分支的特征在频域进行协调融合。 输入: features: list[Tensor],每个分支的输出特征,形状 (B, C, H, W) 输出: 协调后的特征,形状 (B, C, H, W) """ def __init__(self, channels, num_branches=3): super().__init__() self.num_branches = num_branches # 用于生成幅度谱的融合权重 self.amp_fuse = nn.Sequential( nn.Conv2d(channels * num_branches, channels, kernel_size=1, bias=False), nn.ReLU(inplace=True), nn.Conv2d(channels, channels, kernel_size=1, bias=False), nn.Sigmoid(), ) # 用于处理相位谱的小型卷积 self.phase_fuse = nn.Sequential( nn.Conv2d(channels * num_branches, channels, kernel_size=1, bias=False), nn.ReLU(inplace=True), nn.Conv2d(channels, channels, kernel_size=1, bias=False), ) # 频域残差校准,增加模型表达能力 self.refine = nn.Sequential( nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False), nn.ReLU(inplace=True), nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False), ) def forward(self, features): assert len(features) == self.num_branches amplitudes = [] phases = [] for feat in features: amp, phase = to_frequency(feat) amplitudes.append(amp) phases.append(phase) # 将所有分支的幅度谱在通道维度拼接 amp_cat = torch.cat(amplitudes, dim=1) phase_cat = torch.cat(phases, dim=1) # 通过注意力生成融合幅度谱 amp_weight = self.amp_fuse(amp_cat) amp_fused = sum(amp * w for amp, w in zip(amplitudes, torch.chunk(amp_weight, self.num_branches, dim=1))) # 通过卷积融合相位谱 phase_fused = self.phase_fuse(phase_cat) # 重建空间域特征 fused = from_frequency(amp_fused, phase_fused) # 空间域小规模残差校准 fused = fused + self.refine(fused) return fused这段代码有几点说明:
第一,amp_fuse生成的是逐通道的注意力权重,它决定了每个分支的幅度谱在融合中占多大比例。这个设计非常关键,因为不同退化类型的幅度谱特征差异较大,如果直接取均值,会削弱每个分支中最独特的退化信息。
第二,phase_fuse采用拼接后卷积的方式。因为相位谱对内容结构至关重要,直接用可学习的卷积网络来生成融合相位,可以保留更多结构细节。
第三,refine在空间域加了小规模校准,让模块在频域融合之后,还能在空间域做局部修正,避免逆变换带来的轻微伪影。
4.4 多分支编码器示例
下面给出一个简单的多分支编码器。为了稳住篇幅,这里不堆叠很深的网络,主要以理解结构为目标。
# 文件路径:models/encoder.py import torch import torch.nn as nn class BasicBlock(nn.Module): def __init__(self, in_ch, out_ch, stride=1): super().__init__() self.conv1 = nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=stride, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(out_ch) self.relu = nn.ReLU(inplace=True) self.conv2 = nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(out_ch) if stride != 1 or in_ch != out_ch: self.shortcut = nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=1, stride=stride, bias=False), nn.BatchNorm2d(out_ch), ) else: self.shortcut = nn.Identity() def forward(self, x): identity = self.shortcut(x) out = self.conv1(x) out = self.bn1(out) out = self.relu(out) out = self.conv2(out) out = self.bn2(out) out = out + identity out = self.relu(out) return out class MultiBranchEncoder(nn.Module): """ 多分支编码器:一个主干 + 多个退化感知分支。 分支数量根据任务设定,这里默认 3 个分支。 """ def __init__(self, in_ch=3, base_ch=32, num_branches=3): super().__init__() self.num_branches = num_branches # 主干特征提取 self.stem = nn.Sequential( nn.Conv2d(in_ch, base_ch, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(base_ch), nn.ReLU(inplace=True), ) # 多分支特征提取 self.branches = nn.ModuleList() for _ in range(num_branches): branch = nn.Sequential( BasicBlock(base_ch, base_ch), BasicBlock(base_ch, base_ch), ) self.branches.append(branch) def forward(self, x): stem_feat = self.stem(x) branch_feats = [] for branch in self.branches: branch_feats.append(branch(stem_feat)) return stem_feat, branch_feats这个编码器里,主干负责提取通用浅层特征,分支负责提取不同退化的差异特征。实际论文中每个分支内部会有更精细的结构,但整体逻辑是一样的。
4.5 解码器与整体网络
解码器的作用是把协调后的特征逐步恢复成原图尺寸。这里提供最简版本。
# 文件路径:models/decoder.py import torch import torch.nn as nn class SimpleDecoder(nn.Module): def __init__(self, in_ch=32, out_ch=3): super().__init__() self.conv1 = nn.Conv2d(in_ch, in_ch * 2, kernel_size=3, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(in_ch * 2) self.relu = nn.ReLU(inplace=True) self.conv2 = nn.Conv2d(in_ch * 2, in_ch, kernel_size=3, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(in_ch) self.last = nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1) def forward(self, x): x = self.relu(self.bn1(self.conv1(x))) x = self.relu(self.bn2(self.conv2(x))) x = self.last(x) return x之后把编码器、协调模块、解码器拼在一起,就构成了一个完整的简化版 All-in-One 模型。
# 文件路径:models/__init__.py import torch import torch.nn as nn from .encoder import MultiBranchEncoder from .harmonization import SpectralHarmonization from .decoder import SimpleDecoder class AllInOneRestorationNet(nn.Module): def __init__(self, in_ch=3, base_ch=32, num_branches=3): super().__init__() self.encoder = MultiBranchEncoder(in_ch=in_ch, base_ch=base_ch, num_branches=num_branches) self.harmonization = SpectralHarmonization(channels=base_ch, num_branches=num_branches) self.decoder = SimpleDecoder(in_ch=base_ch, out_ch=in_ch) def forward(self, x): stem_feat, branch_feats = self.encoder(x) # 将主干特征也作为一路输入,增强全局信息 branch_feats = [stem_feat] + branch_feats fused_feat = self.harmonization(branch_feats) out = self.decoder(fused_feat) # 残差学习 return x + out这里最后的x + out是典型的残差学习思路,让网络学习的是“退化残差”,也就是输入与干净图像之间的差距,而不是直接回归整张干净图像。这样能显著降低优化难度,尤其在低频分量上。
4.6 损失函数设计
图像恢复任务通常不会只用单一损失。下面给出一个综合损失函数示例。
# 文件路径:losses/restoration_loss.py import torch import torch.nn as nn import torch.nn.functional as F import torchvision.models as models class PerceptualLoss(nn.Module): def __init__(self, device="cuda"): super().__init__() vgg = models.vgg16(pretrained=True).features[:16].to(device).eval() for param in vgg.parameters(): param.requires_grad = False self.vgg = vgg self.criterion = nn.L1Loss() def forward(self, pred, target): pred_feat = self.vgg(pred) target_feat = self.vgg(target) return self.criterion(pred_feat, target_feat) class RestorationLoss(nn.Module): def __init__(self, lambda_rec=1.0, lambda_per=0.1, lambda_freq=0.1, device="cuda"): super().__init__() self.lambda_rec = lambda_rec self.lambda_per = lambda_per self.lambda_freq = lambda_freq self.perceptual_loss = PerceptualLoss(device=device) def forward(self, pred, target): # 像素级重建损失 rec_loss = F.l1_loss(pred, target) # 感知损失 per_loss = self.perceptual_loss(pred, target) # 频域损失,直接用 FFT 幅度谱做 L1 pred_fft = torch.fft.fft2(pred) target_fft = torch.fft.fft2(target) freq_loss = F.l1_loss(torch.abs(pred_fft), torch.abs(target_fft)) total = ( self.lambda_rec * rec_loss + self.lambda_per * per_loss + self.lambda_freq * freq_loss ) return total频域损失的作用是让模型在频率分量上贴近干净图像,这在天气恢复中非常有效,因为很多退化伪影在频域上很有辨识度。
5. 训练与评估方法
5.1 数据准备
All-in-One 天气恢复的数据集通常由多组子集构成:
- 雨图对,例如 Rain100L、Rain100H。
- 雾图对,例如 RESIDE 的合成子集。
- 雪图对,例如 Snow100K。
训练时,通常将多个数据集的退化图像混合成一个批次。这个过程中不需要告诉模型当前输入是什么退化类型,模型直接学习映射关系。
这里有个细节值得注意:如果不同数据集的干净参考图差异过大,比如风格不一致、亮度分布不同,模型可能会在它们之间“折中”,导致每个子任务的效果都受到影响。实际操作时,可以先统计各数据集的均值和方差,必要时对图像做归一化或色彩匹配。
5.2 优化策略
论文方法通常采用 AdamW 优化器,初始学习率设置在 1e-4 左右。训练过程中使用余弦退火或者阶段性衰减,可以有效提升收敛稳定性。
一个常用的经验是:
前 10 个 epoch 只训练像素重建损失。 10 个 epoch 之后逐渐加入感知损失和频域损失。这样做是因为网络在前期需要快速拟合图像结构,感知损失过早加入反而会干扰稳定收敛。
5.3 评估指标
图像恢复任务最常用的两个指标是 PSNR(峰值信噪比)和 SSIM(结构相似性)。
PSNR 计算公式为:
PSNR = 10 * log10(MAX^2 / MSE)其中 MAX 是像素最大值,一般为 255。PSNR 越高,说明像素级误差越小。
SSIM 从亮度、对比度和结构三个维度衡量两幅图像的相似性,取值范围为 -1 到 1,越高越好。
除了数值指标,主观视觉效果也很重要。很多情况下 PSNR 很高但纹理过于平滑,这就要求结合感知质量来评估。
6. 实验分析与方法优势解读
6.1 核心实验结果方向
根据论文公布的整体设计,这类模型在多个天气恢复数据集上都能取得具有竞争力的效果。相比单任务模型,它在跨任务泛化上有明显优势;相比早期 All-in-One 方法,频谱协调机制带来了更稳定的一致性能提升。
重要的不是具体数字,而是方法的行为特性:
优点一:特征冲突被明显抑制。
多分支特征在频域协调后,融合特征避免了空间域直接相加时的互相干扰。
优点二:小模型也能保持不错的效果。
因为频域操作天然具有全局信息编码能力,模型不需要依靠很深的网络来扩展感受野,所以参数效率相对较高。
优点三:训练收敛更稳定。
频域损失与空间域损失互补,网络优化目标更平滑,不容易出现某一阶段退化类型完全学不会的情况。
6.2 与其他 All-in-One 方法对比
如果用一张表来概括当前 All-in-One 方法的特点,可以这样看:
| 方法类型 | 代表作思路 | 优点 | 缺点 |
|---|---|---|---|
| 动态卷积路由 | 根据输入动态选择卷积核 | 灵活、适配性强 | 路由不稳定、额外开销 |
| 空间域特征解耦 | 分离内容特征与退化特征 | 原理清晰 | 解耦不彻底时细节丢失 |
| 频谱协调方法 | 在频域对齐多分支特征 | 全局频带建模、退化分离自然 | FFT 算子部署需额外处理 |
SpectrAl Harmonization 属于第三类,且从实现来看,它并不完全抛弃空间域操作,而是在频域融合后加空间域残差校准,兼顾全局和局部信息。
7. 常见问题与改进方向
7.1 常见问题排查表
在使用或复现这类方法时,你可能会遇到以下问题。
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 训练 Loss 下降但指标不涨 | 感知损失权重过大,模型过度平滑 | 降低感知损失权重,增大像素损失权重 |
| 雨图恢复好但雾图偏色 | 数据集亮度分布不一致 | 对雾图子集做亮度归一化,或增加色彩损失 |
| FFT 运算显存溢出 | 特征图分辨率太大 | 先下采样到固定尺寸再做频域融合,或使用分块 FFT |
| 多分支退化分类不明确 | 分支初始化相同,梯度更新一致 | 对不同分支的初始权重做差异化设置 |
| 测试时出现网格伪影 | fftshift/ifftshift 使用不一致 | 检查正变换和逆变换是否成对使用 |
| 相位谱融合后结构模糊 | 相位谱被平均化 | 相位融合改用带注意力的卷积,保留主分支相位信息 |
7.2 改进方向
结合工程实践,这个方向还有不少值得探索的地方。
改进一:引入退化类型感知的混合监督。
虽然 All-in-One 模型不强制要求退化类型标签,但如果数据集中存在轻量级标签,可以用来辅助训练一个辅助分类头,对频域注意力权重进行弱监督约束。
改进二:在频域中建模模糊核和雨线方向。
雨的局部方向性在频域中其实表现为特定方向的高频能量分布,如果能显式建模方向信息,去雨效果还有提升空间。
改进三:把频域协调扩展为多尺度频域协调。
目前单层 FFT 关注的是整体频率分布,实际上不同尺度上的退化特征是不同的。可以考虑对多个尺度的特征分别做频谱协调,再逐层融合。
改进四:轻量化部署。
FFT 算子在某些边缘设备上缺乏硬件加速。可以考虑用可学习的频域滤波核来近似全局频带操作,或者将 FFT 换成 DCT(离散余弦变换),因为 DCT 在 JPEG 等传统图像处理链路中已有大量优化实现。
改进五:扩展到视频输入。
单帧天气恢复已经比较成熟,但视频输入中还存在时间一致性抖动问题。频域方法如果能结合时序信息,在频域中对时间轴上的退化分量做一致性约束,会是一个很有价值的方向。
8. 最佳实践与工程建议
8.1 数据处理建议
- 训练时不要直接混洗所有数据,可以按照“每个 batch 内混合所有退化类型”的策略组织数据。这样一个 batch 内网络能同时看到多种退化,强迫模型学习通用表示。
- 数据增强时,随机裁剪尺寸建议控制在 128 到 256 之间。尺寸太小,FFT 计算的高频信息不足;尺寸太大,显存压力增大。
- 对雨、雪等局部退化类型,可以加入随机遮挡和混合增强,提升模型对退化区域位置变化的适应能力。
8.2 训练稳定性建议
- 初期阶段冻结频域模块的参数,只训练空间域网络。等空间域网络稳定后,再联合调优频域模块。
- 梯度裁剪对这类模型很有用。频域操作的梯度尺度往往和空间域不同,直接使用整体梯度裁剪能避免训练震荡。
- 使用混合精度训练时要小心 FFT 的数值稳定性。建议先在 FP32 下验证整个前向和反向流程,确认无误后再切换到 AMP。
8.3 工程部署建议
- 在服务端推理时,如果使用 TensorRT,要确认当前版本是否支持
torch.fft导出后的算子。若不支持,可以将频域模块替换为自定义的周期卷积或者可学习的固定频域滤波核。 - 在移动端和边缘端部署时,优先选择小尺寸输入配合下采样频域模块。例如把特征图先降采样到原图一半,再做频谱融合,能显著减少推理时延。
- 模型导出前,务必固定输入尺寸。ONNX 的 FFT 算子对动态尺寸支持不稳定。
8.4 模型设计建议
- 分支数量不是越多越好。实际使用时,二到四个分支已经能覆盖雨、雾、雪、低光等常见场景。过多分支不仅增加参数,还可能引入噪声分支。
- 主干网络建议采用残差连接。频域融合模块已经引入了一定的计算复杂度,主干加残差结构可以降低训练难度,防止梯度消失。
- 对每路分支特征做 L2 归一化后再进入频谱协调模块,可以在一定程度上减小不同分支特征尺度不一致带来的影响。
9. 个人理解与总结
Spectral Harmonization 这类方法给 All-in-One 天气恢复带来的最大启发是:当多个任务共享一个模型时,与其在空间域做复杂的特征路由,不如回到信号处理的本源,在频域中重新思考特征之间的关系。
频率域的幅度谱和相位谱,天然把“退化类型”和“图像内容”分开。幅度谱承载了更多的全局风格和退化信息,适合跨分支融合;相位谱承载了更多的结构和纹理细节,适合保留和精修。这个观察不仅适用于天气恢复,对于低光增强、超分辨率、图像去噪等任务也有借鉴意义。
从论文到工程落地,还有一段路要走。FFT 在训练框架中很好用,但在推理框架里算子支持度参差不齐。如果要在实际产品中使用这类方法,我的建议是:
- 先跑通论文官方代码,在公开数据集上验证效果。
- 把模型切换到自己的业务数据上测试,重点关注雾天色彩偏移和雨天纹理模糊这两类常见问题。
- 在部署阶段,优先用静态输入尺寸导出模型,并检查推理框架对 FFT 算子的支持情况。
- 如果 FFT 算子缺失,可以先用 DCT 替代,或者把频域模块替换为几个并联的全局平均池化分支来近似全局频带信息。
整体来说,All-in-One 天气恢复目前已经到了一个比较成熟的阶段,Spectral Harmonization 这类方法让“一个模型处理多种天气退化”变得更加可行。如果你正在做相关方向的研究或开发,不妨从频域这个视角重新审视已有的特征融合模块,也许能打开新的思路。
如果你对本文中的代码或思路有任何疑问,欢迎在评论区留言交流。后续我也会继续写一些关于频域图像处理和底层视觉实战的内容,感兴趣的可以保持关注。