☰
LaMa 大面积图像修复 ONNX 导出与 TensorRT FP16 加速落地:单图耗时降 2~4 倍全记录
2026/9/29 3:05:17 网站建设 项目流程

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 的倍数
构建中途 OOMworkspace 不够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

上线自检

  1. 权重加载:无 missing key,4 通道输入与掩码拼接方式一致
  2. ONNX 导出:两种分辨率分别验证动态 H/W 生效,opset 与 PyTorch 版本匹配
  3. 引擎构建:FP32 对照版先建成,再开 FP16;动态场景已给 min/opt/max
  4. 一致性:allclose 通过,人工抽查 2~3 张修复图无伪影
  5. 性能数据:预热后多次取均值,记录显存峰值,写进服务 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),仅供参考

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

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

立即咨询