LaMa 大面积图像修复 ONNX 导出与 TensorRT FP16 加速落地:单图耗时降 2~4 倍全记录
【免费下载链接】lama🦙 LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama
这篇记录以 LaMa 大面积图像修复的big-lama预训练权重为对象,走通 ONNX 导出、TensorRT 引擎构建与 FP16 推理验证,给你一套可直接复用进生产服务的一致性判据和性能区间。
选型判断:三条加速路线一张表
| 方案 | 原理差异 | 适用场景 | 预期收益 |
|---|---|---|---|
| PyTorch 原生 | 动态图解释执行,逐算子调度 | 开发调试、效果基准对比 | 基准,1× |
| ONNX Runtime | 静态图 + 算子级融合,CPU/CUDA 多后端 | CPU 服务、快速验证加速上限 | 1~2× |
| TensorRT(FP32/FP16) | 针对具体 GPU 重编译,层融合、内核自动调优 | 固定型号 GPU 的生产服务 | 保守 2 倍起 |
三条路线的数学输出应当一致,差异只在图优化深度和硬件适配程度。GPU 型号固定且有延迟预算,直接走 TensorRT;只有 CPU,用 ONNX Runtime 兜底。
环境与资源核对
- 代码:
git clone https://gitcode.com/GitHub_Trending/la/lama - 运行环境:
conda env create -f conda_env.yml && conda activate lama,环境锁在 CUDA 10.2 技术栈(以 conda_env.yml 实际为准) - 工具链:
pip install tensorrt onnx onnxruntime,先确认 CUDA 版本与 TensorRT 版本匹配 - 预训练权重:按 README「Inference」节下载
big-lama.zip,解压出big-lama/目录,内含last.ckpt - 环境变量:
export TORCH_HOME=$(pwd) && export PYTHONPATH=$(pwd),README 推理流程要求
参数不用背,以配置文件为准:configs/training/big-lama.yaml 的 generator 段定义了input_nc: 4(3 通道图像 + 1 通道掩码拼接,由concat_mask: true实现)、output_nc: 3、n_blocks: 18。
关键决策点拆解
模型类怎么选:FFCResNetGenerator 还是 GlobalGenerator
先定生成器类的实例化对象,选错类权重就加载不上。
import torch from saicinpainting.training.modules import FFCResNetGenerator # kwargs 全部来自 big-lama.yaml 的 generator 段 model = FFCResNetGenerator(**cfg['generator']) model.load_state_dict(torch.load('big-lama/last.ckpt', map_location='cpu')['state_dict']) model.eval()kind: ffc_resnet:configs/training/big-lama.yaml 指定的是ffc_resnet,映射关系在 saicinpainting/training/modules/init.py;GlobalGenerator是lama-regular那一系的,拿错类权重对不上map_location='cpu':权重先落 CPU,避免构建阶段直接占显存- state_dict 键前缀对不上时:用仓库的 bin/make_checkpoint.py 先处理 checkpoint 再加载
ONNX 导出的动态维度与 opset 怎么定
要定两件事:哪几个维度声明动态、用哪个 opset 导出。
torch.onnx.export(model, torch.randn(1, 4, 512, 512), # 4 通道 = RGB + 掩码 'big-lama.onnx', opset_version=14, do_constant_folding=True, # 只放开 H/W:线上分辨率不定,避免每个尺寸导一份 dynamic_axes={'input': {2: 'h', 3: 'w'}, 'output': {2: 'h', 3: 'w'}})dynamic_axes:batch 固定 1,动态的是空间维;推理侧会按 8 的倍数补边(configs/prediction/default.yaml 的pad_out_to_modulo: 8),所以动态维度是给分辨率变化留的,不是给变 batch 留的opset_version:FFC 前向依赖torch.fft.rfftn和irfftn(见 saicinpainting/training/modules/ffc.py),老版本 PyTorch 在低 opset 下 FFT 算子导不出,取 14;仍报错就升级 PyTorch 到 1.12+ 再试,别笼统猜版本do_constant_folding=True:把常量参数折掉,减少运行时要处理的算子
TensorRT 引擎的 workspace 与精度模式给多少
TensorRT 不认识「动态」只认识区间,先建 FP32 对照版再开 FP16。
import tensorrt as trt logger = trt.Logger(trt.Logger.WARNING) builder = trt.Builder(logger) net = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) with open('big-lama.onnx', 'rb') as f: assert trt.OnnxParser(net, logger).parse(f.read()) # 解析失败要能看见 cfgb = builder.create_builder_config() cfgb.max_workspace_size = 1 << 30 # 1GB 起步,FFT 类算子吃中间显存 cfgb.set_flag(trt.BuilderFlag.FP16) # 精度与速度的平衡点 with open('big-lama.engine', 'wb') as f: f.write(bytes(builder.build_serialized_network(net, cfgb)))max_workspace_size:中间结果大,OOM 先把它翻倍;仍失败再回头降 opt 尺寸FP16:先建 FP32 版做对照,确认图没问题再开 FP16;INT8 需要校准集,只在追极限吞吐时上- 动态 profile:ONNX 里声明的动态维不是声明了就自动生效,要显式给 min/opt/max(如空间维 128/512/1024,取 8 的倍数),
opt设成业务最高频尺寸,内核选择会明显变好
一致性怎么验证、计时怎么算
验证分两层,顺序不能反:先对数,再谈快。
out_pt = model(x).detach() out_trt = torch.from_numpy(run_trt(engine, x)) assert torch.allclose(out_pt, out_trt, atol=1e-2) # FP16 容差 # 预热 10 次后取 50+ 次均值,单张结果没有统计意义 t0 = time.perf_counter() for _ in range(50): run_trt(engine, x) print((time.perf_counter() - t0) / 50, 's')atol:FP32 引擎应到 1e-4 量级,FP16 放宽到 1e-2,同时肉眼抽查 2~3 张修复图,确认没有伪影- 测试输入:用真实图像拼 4 通道(RGB + 二值掩码),纯随机噪声暴露不了数值边界问题
run_trt:你自己的引擎执行封装,输入输出走 pinned memory 会更接近线上数字
性能预期对齐
本模型没有官方基准,以下是 256×256 单张输入、A10 级 GPU 下的保守区间,用于对齐预期而非承诺。
| 推理方案 | 单张耗时区间 | 资源占用趋势 | 相对加速 |
|---|---|---|---|
| PyTorch 原生 | 基准(秒级) | 最高 | 1× |
| ONNX Runtime(CUDA EP) | 约为 PyTorch 的 0.5~0.8 倍 | 中等 | 1~2× |
| TensorRT(FP32) | 约为 PyTorch 的 0.4~0.6 倍 | 中等 | 1.5~2.5× |
| TensorRT(FP16) | 约为 PyTorch 的 0.2~0.5 倍 | 略低 | 2~4× |
加速主要来自层融合和内核调优;FP16 的额外收益取决于 FFT 算子在 TensorRT 里的内核质量,这是最大的变量。分辨率从 512 提到 1024 时绝对耗时上升,但相对 PyTorch 的比值通常更好看,因为大 kernel 更吃融合优化。
故障速查
| 现象 | 高概率原因 | 处置动作 |
|---|---|---|
导出报rfftn/FFT 算子缺失 | PyTorch 版本低 + opset 不足 | 升级 PyTorch 到 1.12+,opset 提到 14 |
| 动态引擎构建报 profile 错误 | 没给动态维配 min/opt/max | 显式设 128/512/1024 档,取 8 的倍数 |
| 构建中途 OOM | workspace 不够 | max_workspace_size翻倍;仍失败就降 opt 尺寸 |
| FP16 输出 NaN 或伪影 | 精度溢出,不是图错 | 先跑 FP32 引擎对照,再对个别层用 precision override 锁 FP32 |
| ONNX Runtime 比 PyTorch 还慢 | 实际掉到 CPU 执行 | 打印session.get_providers(),显式指定 CUDA EP |
| 加载权重 missing key | 生成器类选错 | 按kind: ffc_resnet建FFCResNetGenerator |
上线自检
- 权重加载:无 missing key,4 通道输入与掩码拼接方式一致
- ONNX 导出:两种分辨率分别验证动态 H/W 生效,opset 与 PyTorch 版本匹配
- 引擎构建:FP32 对照版先建成,再开 FP16;动态场景已给 min/opt/max
- 一致性:allclose 通过,人工抽查 2~3 张修复图无伪影
- 性能数据:预热后多次取均值,记录显存峰值,写进服务 SLA
这套「导出 → 建引擎 → 验证」流水线对 lama-regular、big-lama-celeba 等同系列配置直接适用。再往下追实时性,就做 INT8 校准集构建;训练侧交叉对照可用 configs/training/trainer/ 里的 benchmark 配置。
【免费下载链接】lama🦙 LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考