从图像去雾到PyTorch实战:手把手教你复现AOD-Net模型(附完整代码与数据集处理技巧)
从图像去雾到PyTorch实战:手把手教你复现AOD-Net模型(附完整代码与数据集处理技巧)
当你在雾天拍摄的照片总是灰蒙蒙的,是否想过用AI技术让画面重获清晰?AOD-Net(All-in-One Dehazing Network)作为图像去雾领域的经典模型,以其端到端的简洁架构和出色的去雾效果备受关注。本文将带你从零开始,用PyTorch完整复现这个模型,不仅会深入解析其独特设计,还会分享数据处理、训练优化的实战技巧,让你真正掌握图像去雾的核心技术。
1. 理解AOD-Net:为什么它如此特别
在传统图像去雾方法中,大气散射模型(Atmospheric Scattering Model)是理论基础,其数学表达式通常为:
I(x) = J(x)t(x) + A(1 - t(x))
其中:
I(x):观测到的有雾图像J(x):待恢复的无雾图像t(x):透射率图A:大气光值
AOD-Net的创新之处在于将传统模型中的 t(x) 和 A 整合为一个统一的参数 K(x) ,重构公式为:
J(x) = K(x)I(x) - K(x) + b
这种设计带来了三个显著优势:
- 端到端学习 :无需分别估计透射率和大气光,直接学习映射关系
- 物理可解释性 :保持与传统模型的数学联系
- 计算高效 :网络结构轻量,适合实时应用
表:AOD-Net与传统去雾方法对比
| 特性 | AOD-Net | 传统方法 |
|---|---|---|
| 是否需要分步估计 | 否 | 是 |
| 参数数量 | 约1.6K | 通常更多 |
| 推理速度 | 快 | 较慢 |
| 物理可解释性 | 保留 | 强 |
2. 数据准备:处理配对图像的关键技巧
NYU-Depth数据集是去雾任务常用的基准数据,包含室内场景的无雾/有雾图像对。处理这类配对数据需要特别注意以下要点:
2.1 数据集结构解析
典型结构如下:
dataset/
├── original_images/ # 无雾图像
│ └── NYU2_1.jpg
├── training_images/ # 有雾图像
│ └── NYU2_1_alpha_0.9_beta_0.1.jpg
└── splits/ # 数据划分
文件名匹配规则:
- 无雾图像:
NYU2_1.jpg - 对应有雾图像:
NYU2_1_*.jpg(不同雾浓度)
2.2 高效数据加载方案
使用自定义Dataset类可以优化加载流程:
class DehazeDataset(Dataset):
def __init__(self, clean_dir, haze_dir, transform=None):
self.clean_paths = sorted(glob(f"{clean_dir}/*.jpg"))
self.haze_dict = defaultdict(list)
for haze_path in glob(f"{haze_dir}/*.jpg"):
base_name = "_".join(haze_path.split("_")[:2])
self.haze_dict[base_name].append(haze_path)
self.transform = transform
def __len__(self):
return len(self.clean_paths)
def __getitem__(self, idx):
clean_path = self.clean_paths[idx]
base_name = Path(clean_path).stem
haze_path = random.choice(self.haze_dict[base_name])
clean_img = Image.open(clean_path).convert("RGB")
haze_img = Image.open(haze_path).convert("RGB")
if self.transform:
clean_img = self.transform(clean_img)
haze_img = self.transform(haze_img)
return haze_img, clean_img
提示:使用
defaultdict存储雾图路径可以显著提升配对效率,避免每次遍历整个目录
关键优化点:
- 内存映射 :对于大型数据集,使用
lmdb格式存储 - 在线增强 :添加随机裁剪、翻转等增强操作
- 批量预取 :设置
num_workers>0和pin_memory=True加速数据加载
3. 模型构建:逐层解析AOD-Net架构
AOD-Net的核心是由5个卷积层和3个特征拼接操作组成的轻量网络。下面我们拆解每个组件的设计意图:
3.1 基础卷积模块
class AODNet(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Sequential(
nn.Conv2d(3, 3, kernel_size=1, stride=1, padding=0),
nn.ReLU()
)
self.conv2 = nn.Sequential(
nn.Conv2d(3, 3, kernel_size=3, stride=1, padding=1),
nn.ReLU()
)
# 其余卷积层类似...
设计要点 :
- 第一层使用1x1卷积:学习通道间的线性组合
- 后续采用3x3/5x7卷积:捕获局部雾浓度变化
- 保持输入输出同尺寸:通过
padding=same实现
3.2 特征拼接与融合
def forward(self, x):
conv1_out = self.conv1(x) # [b,3,h,w]
conv2_out = self.conv2(conv1_out)
concat1 = torch.cat([conv1_out, conv2_out], dim=1) # [b,6,h,w]
conv3_out = self.conv3(concat1)
concat2 = torch.cat([conv2_out, conv3_out], dim=1)
conv4_out = self.conv4(concat2)
concat3 = torch.cat([conv3_out, conv4_out], dim=1)
conv5_out = self.conv5(concat3)
return self.final_fusion(x, conv5_out)
多尺度特征融合的优势 :
- 浅层特征保留细节纹理
- 深层特征编码雾浓度信息
- 拼接操作实现信息流动
3.3 最终融合层
def final_fusion(self, input_img, k_map):
""" 实现J(x) = K(x)I(x) - K(x) + b """
return torch.relu(k_map * input_img - k_map + 1.0)
注意:最后的ReLU激活确保输出像素值在[0,1]范围内
4. 训练技巧:从损失函数到超参调优
4.1 损失函数选择
除了基础的MSE损失,推荐尝试组合损失:
class CompositeLoss(nn.Module):
def __init__(self):
super().__init__()
self.mse = nn.MSELoss()
self.ssim = SSIMLoss() # 结构相似度损失
def forward(self, pred, target):
return 0.7*self.mse(pred, target) + 0.3*self.ssim(pred, target)
表:不同损失函数效果对比
| 损失类型 | PSNR ↑ | SSIM ↑ | 训练稳定性 |
|---|---|---|---|
| MSE | 22.1 | 0.83 | 高 |
| MAE | 21.8 | 0.82 | 中 |
| 混合损失 | 23.4 | 0.87 | 较高 |
4.2 学习率调度策略
使用 CosineAnnealingLR 配合热身阶段:
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
scheduler = torch.optim.lr_scheduler.SequentialLR(
optimizer,
[
torch.optim.lr_scheduler.LinearLR(
optimizer, start_factor=0.1, total_iters=5
),
torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=20
)
],
milestones=[5]
)
4.3 GPU内存优化技巧
当处理高分辨率图像时:
- 使用梯度累积:
for i, (haze, clean) in enumerate(dataloader):
pred = model(haze)
loss = criterion(pred, clean) / accumulation_steps
loss.backward()
if (i+1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
- 启用混合精度训练:
scaler = torch.cuda.amp.GradScaler()
with torch.cuda.amp.autocast():
pred = model(haze)
loss = criterion(pred, clean)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
5. 模型部署与效果可视化
5.1 使用Netron分析模型
保存可视化模型:
model_scripted = torch.jit.script(model)
torch.jit.save(model_scripted, "aodnet_scripted.pt")
然后通过Netron打开生成的 .pt 文件,可以清晰看到:
- 各层输入输出维度
- 参数数量统计
- 计算图结构
5.2 推理测试代码
def dehaze_image(model, img_path, device="cuda"):
img = Image.open(img_path).convert("RGB")
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
])
img_tensor = transform(img).unsqueeze(0).to(device)
with torch.no_grad():
output = model(img_tensor)
output_img = output.squeeze().cpu().permute(1,2,0).numpy()
output_img = np.clip(output_img*255, 0, 255).astype(np.uint8)
return Image.fromarray(output_img)
5.3 效果对比示例
实际测试中发现,模型在以下场景表现最佳:
- 均匀薄雾场景
- 室内人造光环境
- 中等分辨率图像(480p-720p)
而对于浓雾或极端光照条件,建议:
- 增加更多样化的训练数据
- 调整网络深度
- 尝试后处理方法(如直方图均衡化)
更多推荐




所有评论(0)