NAFNet图像去模糊实战:从环境搭建到推理优化的完整指南
2026/9/14 2:39:58 网站建设 项目流程

简介:面向图像处理与深度学习开发者的一份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,指标反而更高

NAFNet(Nonlinear 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.8;PyTorch 2.x 也能跑通,但个别接口要小改。Linux 服务器上常见做法是先装 Miniconda 再建环境,而不是自己编译 Python;Windows 上同样用 conda 隔离,避免污染系统 Python。还没装 conda 的话,Linux 下用 wget 拉安装脚本后 bash 执行,Windows 下载安装包两步装完,再执行 conda init 让 shell 识别 conda 命令。建环境命令如下:

conda create -n nafnet python=3.8 -y conda activate nafnet pip install torch==1.13.1 torchvision==0.14.1 --index-url https://download.pytorch.org/whl/cu117 pip install opencv-python numpy pyyaml tqdm pytorch-lightning==1.3.8

第一行创建名为 nafnet 的独立环境并指定 Python 3.8;第二行激活;第三行从 PyTorch 官方 wheel 源安装 CUDA 11.7 编译版的 torch,cu117 是编译时用的 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 环境时按 Ctrl+Shift+P 调出命令面板,选 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.8~3.10语法错误或依赖解析失败
PyTorch1.13.x + cu117import torch 报错或 CUDA 不可用
pytorch-lightning1.3.8Trainer 属性缺失的 AttributeError
opencv-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 -5ls 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: 16

width是基础通道数,决定模型体量和指标上限:width32 约 16 MB,width64 计算量接近四倍,PSNR 还能涨零点几 dB,显存紧张先用 32。enc_blk_nums最后一个数字是 bottleneck 深度,快速验证流程时改成 4 能省大量时间,跑通再改回 28。crop_size是训练时随机裁剪的块大小,配合随机翻转和 90 度旋转做数据增广;16 GB 显存的单卡建议 batch_size 降到 4~8,crop 保持 256,两个一起缩会明显影响收敛质量。lr设为 1e-3 并配 CosineAnnealingRestart,basicsr 会在每个周期内把学习率从 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_DEVICES=0,1 python -m torch.distributed.launch --nproc_per_node=2 --master_port=43289 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 而不是另一份模糊图;预处理有没有做多余归一化,网络输入约定是 0~1,不是 0~255,也不是 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, width=32): model = NAFNet( img_channel=3, width=width, middle_blk_num=1, enc_blk_nums=[1, 1, 1, 28], dec_blk_nums=[1, 1, 1, 1], ) state = torch.load(model_path, map_location='cpu') 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)缩放到 0~1,原地除法不产生中间张量;推理包在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_range=255))

gt 和 pred 都用 0~255 的 uint8 读入,data_range显式传 255,否则函数按 dtype 推断会把结果算偏。

4.3 shape 不匹配、偏色、发灰:三个高频报错

现象原因处理
size mismatch / missing keyswidth 与权重不一致核对构造网络时传入的 width
红蓝通道互换漏了 BGR/RGB 转换输入输出各做一次 cv2.cvtColor
结果发灰发暗归一化两次或没 clamp0~1 只归一化一次,输出 clamp 到 0~1
CUDA 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 时,去模糊前后各算一次这个值,提升的倍数就能量化模型恢复了多少锐度信息。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询