ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

NAFNet图像去模糊实战:从环境搭建到推理优化的完整指南

NAFNet图像去模糊实战:从环境搭建到推理优化的完整指南 简介面向图像处理与深度学习开发者的一份NAFNet图像去模糊Python实现。NAFNet通过卷积层、残差块与注意力机制提取特征可有效恢复因拍摄移动或相机抖动导致的模糊图像。资源压缩包共20个文件、约11.98MB含6个png示例图、5个jpg输入图、4个xml配置、2个py核心脚本及gitignore、iml、md说明文档Python脚本涵盖模型搭建、训练与测试流程配合示例图片可快速验证去模糊效果。目前已有761人学习下载适合具备Python与深度学习基础的读者可用于算法研究、效果比较或实际图像修复场景。内容预览显示项目包含4k与普通分辨率两套处理脚本并附有PaddleGAN相关目录便于拓展图像复原方案。1. NAFNet 做图像去模糊去掉 ReLU指标反而更高NAFNetNonlinear Activation Free Network出自 ECCV 2022 的《Simple Baselines for Image Restoration》把残差块里必带的 ReLU 这类非线性激活全部换成 SimpleGate 和简化通道注意力SCA在 GoPro 去模糊基准上做到约 33.7 dB 的 PSNR模型只有 16 MB 量级却比一堆复杂结构更出效果。这个反直觉的结论让 NAFNet 成了图像去模糊任务最常用的 baseline。标题里的 zip 通常是完整工程网络结构定义、train/test 入口、yml 配置和说明文档都在里面。拿到手的固定流程是配好 Python 环境、摆对数据集、按 yml 起训练或直接加载权重推理。适合想快速用 Python 落地图像去模糊、又不想从零搭网络的工程师。下面按这个顺序讲重点放在参数含义和容易翻车的位置。2. NAFNet 的结构要点与 Python 环境搭建2.1 SimpleGate 与 SCA为什么没有非线性激活反而更强先理解两个核心组件后面调参才不盲目。传统 ResNet 块里 ReLU 负责引入非线性但负半轴被直接清零特征信息有损失。NAFNet 的 SimpleGate 把一个通道数为 2C 的特征沿通道维对半拆成两份逐元素相乘x1 * x2得到动态的、可学习的门控输出不依赖固定激活函数却保留非线性表达能力。SCA简化通道注意力则把流程压缩成全局平均池化 → 1x1 卷积 → sigmoid中间那层 ReLU 也去掉了。这样的 NAFBlock 由 LayerNorm、1x1 卷积、depthwise 3x3 卷积、SimpleGate 和 SCA 组成整个网络按编码器-解码器堆叠。看到配置里enc_blk_nums: [1, 1, 1, 28]时要知道前面三个数字是浅层块数最后一个 28 是 bottleneck 深度承担了绝大部分计算量也解释了为什么改这个数对训练速度影响最大。2.2 用 conda 建独立的 Python 环境与国内源加速官方代码基于 PyTorch 1.8 时代编写我实际跑下来验证最多的组合是 Python 3.8 PyTorch 1.13 pytorch-lightning 1.3.8PyTorch 2.x 也能跑通但个别接口要小改。Linux 服务器上常见做法是先装 Miniconda 再建环境而不是自己编译 PythonWindows 上同样用 conda 隔离避免污染系统 Python。还没装 conda 的话Linux 下用 wget 拉安装脚本后 bash 执行Windows 下载安装包两步装完再执行 conda init 让 shell 识别 conda 命令。建环境命令如下conda create -n nafnet python3.8 -y conda activate nafnet pip install torch1.13.1 torchvision0.14.1 --index-url https://download.pytorch.org/whl/cu117 pip install opencv-python numpy pyyaml tqdm pytorch-lightning1.3.8第一行创建名为 nafnet 的独立环境并指定 Python 3.8第二行激活第三行从 PyTorch 官方 wheel 源安装 CUDA 11.7 编译版的 torchcu117 是编译时用的 CUDA 版本和你机器的驱动版本不冲突第四行装图像读写、数值计算和训练框架。如果下载慢在 pip 命令末尾加-i https://pypi.tuna.tsinghua.edu.cn/simple走国内镜像源torch 这种大包还是建议从官方源拉减少校验失败的概率。提示zip 里自带一份改造过的 basicsr不要额外pip install basicsr否则 train.py 会优先 import pip 版出现 dataset 参数校验对不上的怪报错。zip 里的 requirements.txt 可以一次装完大部分依赖但版本可能偏老装完再核对 2.4 的表。装完用两条命令自检任何一条报错都能直接定位到缺什么python -c import torch; print(torch.__version__, torch.cuda.is_available()) python -c import cv2; print(cv2.__version__)第二句报ModuleNotFoundError: No module named cv2就是 opencv-python 没装上回到上面补装。vscode 配置 Python 环境时按 CtrlShiftP 调出命令面板选 Python: Select Interpreter 指向 nafnet 这个 conda 环境可以避免终端能跑、编辑器里却找不到 torch 这类解释器不一致问题。2.3 解压后先认准四个入口二次打包的目录结构可能略有出入但骨架通常有四块basicsr/ 是改造过的训练框架models/ 或 basicsr/models/archs/ 下能找到 nafnet_arch.py 网络定义options/ 按 train/ 和 test/ 分好 yml 配置根目录 train.py、test.py 是统一入口。拿到 zip 先确认这四个位置再读一遍 README 里对 Python 和 CUDA 的说明比直接跑、报错后再翻日志省时间。2.4 依赖核对表与 vscode 解释器选择组件推荐版本没装好时的表现Python3.83.10语法错误或依赖解析失败PyTorch1.13.x cu117import torch 报错或 CUDA 不可用pytorch-lightning1.3.8Trainer 属性缺失的 AttributeErroropencv-python4.xNo module named cv2项目自带 basicsr不要用 pip 版参数校验报错、找不到 dataset 类型版本不用逐字对齐只要和上表同主版本基本都能跑。真正要盯的是 torch 和 pytorch-lightning 的搭配这两个版本差太多时报错往往出现在训练循环内部而不是 import 阶段排查成本最高。3. 训练 NAFNet 去模糊模型数据目录与 yml 配置拆解3.1 GoPro 数据集的目录怎么摆训练配置里改得最多的就是数据路径。GoPro 基准包含 2103 对训练图和 1111 对测试图每对是一张模糊图和对应的清晰图basicsr 的 PairedImageDataset 按固定目录结构读图常见摆法如下datasets/GoPro/ ├── train/ │ ├── blur/ # 模糊输入 │ └── gt/ # 清晰参考图 └── test/ ├── blur/ └── gt/这里有个容易忽略的约定blur 和 gt 的文件名必须完全一致dataset 按文件名配对而不是按顺序后缀可以不同。目录建好后用ls datasets/GoPro/train/blur | head -5和ls datasets/GoPro/train/gt | head -5对比两边输出文件名能对上就可以进入下一步。如果数据是 lmdb 格式yml 里的 dataset 类型要换成对应的 lmdb 实现路径也要指到 .lmdb 目录两种格式别混用。3.2 width、enc_blk_nums、loss三个最核心的配置参数train.py 只接收一个 yml 参数所有超参数集中在里面。下面是最常见的 width32 配置核心片段network_g: type: NAFNet img_channel: 3 width: 32 middle_blk_num: 1 enc_blk_nums: [1, 1, 1, 28] dec_blk_nums: [1, 1, 1, 1] train: optim_type: AdamW lr: !!float 1e-3 scheduler: CosineAnnealingRestart total_iter: 300000 loss_type: charbonnier datasets: train: name: GoPro gt_folder: datasets/GoPro/train/gt lq_folder: datasets/GoPro/train/blur crop_size: 256 batch_size: 16width是基础通道数决定模型体量和指标上限width32 约 16 MBwidth64 计算量接近四倍PSNR 还能涨零点几 dB显存紧张先用 32。enc_blk_nums最后一个数字是 bottleneck 深度快速验证流程时改成 4 能省大量时间跑通再改回 28。crop_size是训练时随机裁剪的块大小配合随机翻转和 90 度旋转做数据增广16 GB 显存的单卡建议 batch_size 降到 48crop 保持 256两个一起缩会明显影响收敛质量。lr设为 1e-3 并配 CosineAnnealingRestartbasicsr 会在每个周期内把学习率从 1e-3 余弦衰减到接近零再跳回收敛过程比固定学习率平滑。charbonnier损失即sqrt(x^2 eps^2)对离群像素比 L2 鲁棒是像素级损失的首选。想让人眼看着更干净常见做法是在它基础上叠加 LPIPS 感知损失和 FFT 频域损失论文的三阶段训练就是这个思路第一次跑通流程只开 charbonnier 就够权重系数在 zip 里其他样例 yml 中能找到现成写法。3.3 启动训练、续训和指定显卡配置改好后一行命令启动训练python train.py -opt options/train/NAFNet/NAFNet-width32.yml多卡场景用CUDA_VISIBLE_DEVICES0,1 python -m torch.distributed.launch --nproc_per_node2 --master_port43289 train.py -opt ...master_port 随便指定一个没被占用的端口即可。训练中断不用从头再来basicsr 会把带 training_state 的目录实时写到 experiments/ 下续训命令python train.py -opt options/train/NAFNet/NAFNet-width32.yml \ --resume experiments/NAFNet-width32/training_state/latest注意--resume 指向的是 training_state 状态目录不是 .pth 模型文件给成模型路径会直接报找不到文件。训练结束后 experiments/ 下会生成 latest.pth 和按 validation PSNR 挑选的 best.pth做推理优先用 best.pth。3.4 训练日志里重点看什么日志会同时输出 iters、lr 和 loss。前几千 iter loss 掉得快是正常的真正要看的是每个 CosineAnnealingRestart 周期点之后、lr 跳回时 loss 还能不能创新低连续两三个周期都压不下去再考虑加宽网络或换损失组合。训练过程会定期跑 validation 并输出 PSNR/SSIM那才是最终要盯的指标loss 只是参考。如果 validation PSNR 一直不涨先回查数据配对和 crop、batch 的搭配而不是急着换模型。4. 用训练好的权重做推理验证脚本与单图测试4.1 先跑官方 test.py 拿到 PSNR权重下好后常见放置位置是 experiments/pretrained_models/然后在 test yml 里把pretrain_network_g指过去执行python test.py -opt options/test/NAFNet/NAFNet-width32.yml程序会遍历整个测试集每张图推理后和 GT 计算 PSNR/SSIM最后打印平均值输出图写在 results/ 下对应配置名的子目录里。结果和官方数字差得远时依次查三处权重对应的 width 是否和配置一致width64 权重塞进 width32 网络必然报 shape 错误test 数据路径是否真的指向 GT 而不是另一份模糊图预处理有没有做多余归一化网络输入约定是 01不是 0255也不是 ImageNet 的 mean/std 归一化。4.2 不依赖 train.py 的单张推理代码只想处理一张图时不必把整套 basicsr 跑起来直接构造网络再加权重即可这也是排查问题最快的方式import cv2 import numpy as np import torch from models.nafnet_arch import NAFNet def deblur(model_path, in_path, out_path, width32): model NAFNet( img_channel3, widthwidth, middle_blk_num1, enc_blk_nums[1, 1, 1, 28], dec_blk_nums[1, 1, 1, 1], ) state torch.load(model_path, map_locationcpu) if params in state: # basicsr 保存格式权重挂在 params 键下 state state[params] model.load_state_dict(state) model.eval().cuda() img cv2.imread(in_path) # BGR 顺序0~255 x cv2.cvtColor(img, cv2.COLOR_BGR2RGB) x torch.from_numpy(x.transpose(2, 0, 1)).float().div_(255.0) x x.unsqueeze(0).cuda() # (1, 3, H, W) with torch.no_grad(): pred model(x) pred pred.squeeze(0).permute(1, 2, 0).clamp_(0, 1).cpu().numpy() out cv2.cvtColor((pred * 255).astype(np.uint8), cv2.COLOR_RGB2BGR) cv2.imwrite(out_path, out) if __name__ __main__: deblur(NAFNet-GoPro-width32.pth, blur.png, sharp.png)逻辑说明模型按 width32 的配置构造构造参数必须和训练时一致torch.load 后先判断权重是否包在 params 键里这是 basicsr 的保存约定单独转换过的权重则可能直接是裸 state_dict兼容判断能省一次报错。输入侧做 BGR→RGB 和 HWC→CHW 转置再div_(255.0)缩放到 01原地除法不产生中间张量推理包在torch.no_grad()里省显存和耗时输出侧clamp_(0, 1)拉回合法区间后转回 BGR 写盘。如果 import 报模块找不到改成from basicsr.models.archs.nafnet_arch import NAFNet取决于 zip 里网络文件的实际层级。需要量化对比时补一个 PSNR 计算from skimage.metrics import peak_signal_noise_ratio print(peak_signal_noise_ratio(gt, pred, data_range255))gt 和 pred 都用 0255 的 uint8 读入data_range显式传 255否则函数按 dtype 推断会把结果算偏。4.3 shape 不匹配、偏色、发灰三个高频报错现象原因处理size mismatch / missing keyswidth 与权重不一致核对构造网络时传入的 width红蓝通道互换漏了 BGR/RGB 转换输入输出各做一次 cv2.cvtColor结果发灰发暗归一化两次或没 clamp01 只归一化一次输出 clamp 到 01CUDA out of memory图太大或 batch 过大batch 改 1或按 5.2 分块推理这四类问题占了推理阶段九成以上的报错而且前两个都不会让程序崩溃只是结果不对最容易耽误时间。建议先跑 4.2 的脚本验证一张已知清晰的图再批量处理。4.4 推理速度与精度的取舍width32 模型对 720p 图单张推理通常不到一秒耗时主要花在 bottleneck 那 28 个块的大分辨率卷积上。要提速就把模型和输入都转半精度model.half()后输入张量也.half()吞吐能再上一个台阶对指标敏感就保持 fp32。测 PSNR 时不要叠加任何图像增强预处理那会让对比失去公平性。5. 把 NAFNet 用到真实拍摄场景的三个技巧5.1 视频帧去模糊用半精度把吞吐提上去GoPro 训练出的模型直接吃视频帧是能用的但逐帧推理浪费了不少重复计算。常见做法是每帧独立推理外层用torch.cuda.amp.autocast()包住权重切半精度每秒能处理的帧数会有明显提升。帧率还是不够时把帧先缩到 720p 推理再放大回原分辨率视觉差异不大、速度能快好几倍要是画面出现帧间闪烁就要考虑加光流对齐或时序滤波那是单独的课题。5.2 大图分块推理避免一张大图挤爆显存超过 2K 的图直接塞进模型容易 OOM分块是通用解法把图切成 512x512 的块块间留 50% 重叠stride 取 256推理完用线性渐变融合重叠区消除接缝。切块时让每个块的宽高都是 4 的倍数和网络的步长对齐避免边缘出现黑边重叠比例低于 25% 时接缝会变明显50% 是效果和耗时的常见平衡点。5.3 用拉普拉斯方差快速验证去模糊效果不打开看图软件一个数字就能判断模型有没有生效。对同一张图去模糊前后各算一次灰度图的拉普拉斯方差数值越高代表边缘越锐利import cv2 g cv2.cvtColor(cv2.imread(sharp.png), cv2.COLOR_BGR2GRAY) print(cv2.Laplacian(g, cv2.CV_64F).var())对 1080p 真实照片清晰帧的方差通常在几百到上千明显模糊的帧往往只有几十到一百多。去模糊后这个值翻了 3 倍以上基本说明模型在起作用数值几乎没动先回查 4.3 的表格多半是通道顺序或归一化的问题而不是模型没训练好。真实拍摄场景没有 GT 图拿不到 PSNR 时去模糊前后各算一次这个值提升的倍数就能量化模型恢复了多少锐度信息。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进