LaMa ONNX导出与TensorRT推理加速指南
【免费下载链接】lama🦙 LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama
LaMa(WACV 2022)是一个基于傅里叶卷积的大掩码修复模型,解决"高分辨率图像里大面积缺失区域怎么补得自然"的问题。它原生 PyTorch 推理偏慢,上生产通常要走两步:先做 LaMa ONNX 导出,再构建 TensorRT 引擎榨出硬件性能。本文带你走完整条链路:核对权重、导出、建引擎、数值验收。
推理慢、GPU空转:先定位瓶颈在哪
PyTorch推理里多花了什么
big-lama 含 18 个残差块,其中 FFC(傅里叶滤波卷积)块要做rfftn正变换和irfftn逆变换,实现见 ffc.py。PyTorch eager 模式逐算子发 kernel,复数拆分、频域卷积、BatchNorm 之间没有融合,输入一到 1024x1024 以上,单张图延迟就很难压下来。这就是部署优化的空间来源。
部署三级链路:ckpt → ONNX → TensorRT引擎
先分清三种格式各自的角色,再决定在哪一层投入:
| 格式 | 角色 | 关键特性 |
|---|---|---|
| last.ckpt | 训练权重 | 只含 state_dict |
| ONNX | 跨框架交换格式 | 算子图,不绑推理后端 |
| TRT engine | 硬件专用二进制 | 算子已融合调优,绑 GPU 型号 |
注意最后一点:engine 不能跨机器搬。在 A 卡上构建的引擎,换 B 卡要重建,这是部署事故最常见的根源。
LaMa ONNX导出前:核对配置里的关键参数
big-lama.yaml的四个关键数字
构建模型前只读 configs/training/big-lama.yaml 这一段就够了:
input_nc: 4:3 通道图像 + 1 通道掩码拼接output_nc: 3:修复后的 RGBngf: 64、n_downsampling: 3:特征图下采样 8 倍,所以推理输入要补齐到 8 的倍数(预测默认配置 里就是pad_out_to_modulo: 8)n_blocks: 18、add_out_act: sigmoid:18 个残差块,输出直接落在 [0,1],省掉一次归一化
FFC块:导出麻烦的根源
FFC 的原理是"到频域里做卷积":空间特征 FFT 后在频域乘系数,再逆变换回来。它给导出带来两个坎:
torch.fft.rfftn / irfftn依赖 ONNX opset 17 才有的 RFFT2D / IRFFT2D 算子;- 块内部走
torch.complex的实部虚部拆分,导出器覆盖不全时最容易在这里报"算子不支持"。
权重在哪
预训练权重是big-lama/last.ckpt单文件(big-lama.zip 解压所得,下载地址见 README),取其中的state_dict。
LaMa ONNX导出:三步走与两个易错点
第一步:搭导出环境、构建生成器
先 clone 仓库并了解依赖:
git clone https://gitcode.com/GitHub_Trending/la/lama cd lama && cat conda_env.yml仓库自带环境锁在 python 3.6.13(conda_env.yml 第 118 行),FFT 的 ONNX 导出器在旧 torch 上不稳定。建议单独建一个 torch ≥ 1.13(推荐 2.x)的 venv 做导出。然后构建模型:
from saicinpainting.training.modules.pix2pixhd import GlobalGenerator model = GlobalGenerator( input_nc=4, output_nc=3, ngf=64, n_downsampling=3, n_blocks=18, ffc_positions=list(range(18)), ffc_kwargs=dict(ratio_gin=0.75, enable_lfu=False), ) ckpt = torch.load("big-lama/last.ckpt", map_location="cpu") model.load_state_dict(ckpt["state_dict"], strict=True) model.eval()ratio_gin: 0.75和enable_lfu: false直接取自配置的resnet_conv_kwargs段,别改。如果load_state_dict报 key 不匹配,说明 FFC 块位置与权重对不上,拿 ckpt 的 key 名逐一比对ffc_positions。
第二步:opset必须写17而不是12 ⚠️
dummy = torch.randn(1, 4, 512, 512) torch.onnx.export( model, dummy, "big-lama.onnx", opset_version=17, do_constant_folding=True, input_names=["input"], output_names=["output"], dynamic_axes={"input": {2: "h", 3: "w"}, "output": {2: "h", 3: "w"}}, )第一个易错点就在这里:不少教程沿用opset_version=12,而 opset 12 没有 FFT 算子,rfftn会直接导出失败。
第三步:静态还是动态尺寸,看你的服务形态
上面的dynamic_axes让高宽在推理期可变,但代价是 TensorRT 构建时要给 min/opt/max 三档 profile,内核调优效果打折。如果你的服务固定只处理 512 和 1024 两档,更稳的做法是导出静态 512 模型,或给动态模型指定两档 opt shape。
上图是仓库里分割图生成掩码的示例,修复输入始终是"图像 + 掩码"拼成的 4 通道张量,导出时的 dummy 输入也必须是这个形状。
TensorRT引擎构建:FP16开关放在哪
解析器与构建配置
logger = trt.Logger(trt.Logger.WARNING) builder = trt.Builder(logger) network = builder.create_network( 1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser = trt.OnnxParser(network, logger) with open("big-lama.onnx", "rb") as f: parser.parse(f.read()) config = builder.create_builder_config() config.set_flag(trt.BuilderFlag.FP16) config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30) engine = builder.build_serialized_network(network, config) open("big-lama.engine", "wb").write(engine)FP16 开关就是BuilderFlag.FP16;工作空间用set_memory_pool_limit(旧版 API 的max_workspace_size已废弃)。FFT 算子需要 TensorRT 8.0 起提供的 ONNX FFT 插件,版本低了会在 parser 阶段报错。
为什么工作空间给1GB
FFC 的频域中间特征是实部、虚部拼接的复数张量,通道数翻倍;1024x1024 输入下激活占用不小。工作空间给小了构建直接失败或被迫回退低效内核,1GB 是这类模型的常见起点,不够再往上加。
FP16掉点怎么验证
别默认"FP16 基本不掉点",做两级对比:
- ONNX Runtime(FP32)输出 vs PyTorch 输出:期望最大像素差在 1e-3 量级;
- TRT 引擎(FP16)输出 vs ORT 输出:差异会放大。修复任务对高频纹理敏感,建议再对一批固定测试图算 LPIPS 与 SSIM(仓库有现成实现,见 losses 目录),和 PyTorch 基线比。
上图是仓库中掩码生成流程的内存 profile,展示部署监控应有的图表形态:预热冲高后回落,随后平稳无泄漏。
上线验收:数值、速度、内存三道关
数值:两级对比留档
import onnxruntime as ort sess = ort.InferenceSession("big-lama.onnx") out = sess.run(None, {"input": dummy.cpu().numpy()})[0] print(np.abs(out - model(dummy).numpy()).max()) # 期望 1e-3 量级引擎输出再和 ORT 输出做一遍同样的对比,两组 max abs diff 都记进验收单。
速度:必须标注三个测试条件
for _ in range(10): # 预热 run_once() torch.cuda.synchronize() t0 = time.time() for _ in range(50): run_once() print(f"avg {(time.time()-t0)/50*1e3:.1f} ms")写任何性能数字都要带三样:输入尺寸(如 512x512)、GPU 型号、精度模式(FP16)。参考文章给出的量级是 ONNX Runtime 约 1.5-2 倍于 PyTorch、TensorRT 约 2-5 倍(未标注测试条件),只有按上面协议在你自己硬件上复测的数字才可用。
落地检查清单
导出前
- 模型构建参数与 big-lama.yaml 一致(n_blocks=18、ratio_gin=0.75)
load_state_dict(strict=True)通过,无缺失 keyopset_version=17,导出日志无"Unsupported"字样- ORT 与 PyTorch 最大差在 1e-3 量级
引擎构建后
- 构建日志无 FP16 层回退告警,或已确认可接受
- 引擎 vs ORT 的差值与 LPIPS 已记录
- 速度表带全三条件:输入尺寸 / GPU / 精度
- engine 文件旁注明目标 GPU 型号与驱动版本
【免费下载链接】lama🦙 LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考