ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

LaMa 图像修复模型推理提速指南:ONNX 导出 + TensorRT 加速完整流程

LaMa 图像修复模型推理提速指南:ONNX 导出 + TensorRT 加速完整流程 LaMa 图像修复模型推理提速指南ONNX 导出 TensorRT 加速完整流程【免费下载链接】lama LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lamaLaMaLarge Mask Inpainting是一个用傅里叶卷积解决大面积图像修复的开源项目但它的 PyTorch 原生推理在批量场景下偏慢。这篇文章带你走一遍推理优化的完整链路先把权重导出为 ONNX再用 TensorRT 构建 GPU 加速引擎最后给出精度、速度和显存之间的取舍建议帮助新手直接落地。先搞清楚要优化什么图像修复的输入是一对文件一张原图 一张蒙版白色区域表示这里被挖掉了需要补。下图就是 LaMa 用来生成大蒙版素材的分割示例可以看到蒙版可以覆盖画面中很大一块区域这正是 LaMa 相对小洞修复模型的主场痛点在哪LaMa 的生成器不小——18 个残差块、多组傅里叶卷积处理一张 512×512 的图PyTorch 原生推理通常需要 1 秒上下。单张可接受但一次跑几万个样本比如离线处理一批产品图就会非常耗时。优化路线不换模型结构只换推理后端。PyTorch → ONNX Runtime → TensorRT精度基本可控在毫厘级别速度可提升到 2~5 倍ONNX Runtime 大约 1.5~2 倍TensorRT 在 GPU 上更明显。模型结构与关键参数速览LaMa 的生成器是一个下采样 → 瓶颈 → 上采样的编码器结构实现在 pix2pixhd 模块 的GlobalGenerator里。你可以把它理解成一条传送带图像先被逐步压小细节变少、信息变浓缩在瓶颈层做核心补画再逐步放大回原尺寸。big-lama 的关键参数都写在 big-lama.yaml 中参数值含义input_nc4RGB 三通道 1 通道蒙版拼接output_nc3修复后的 RGBngf64基础特征通道数n_downsampling3下采样级数特征图缩小 8 倍n_blocks18瓶颈层残差块数量FFC 位置第 9、10 个残差块每个位置放 1 个傅里叶卷积块频域权重 0.75其中 FFCFourier Filter Convolution实现在 ffc.py是修复质量好的关键普通卷积只能看小范围邻居FFC 把一部分特征转到频域全局处理相当于让模型一眼看全图。但 FFT 类算子也是后面导出时最可能出兼容问题的地方先记住这一点。另一个容易踩的坑是输入尺寸。模型里有 8 级缩放3 级下采样 3 级上采样 7×7 首尾卷积所以送入前通常要把图按 8 的倍数做填充——项目的 推理配置 里pad_out_to_modulo: 8就是干这个的。导出前的环境准备环境按三步走即可git clone https://gitcode.com/GitHub_Trending/la/lama cd lama conda env create -f conda_env.yml conda activate lama另外两件事装 TensorRTpip install tensorrt建议选与当前 CUDA 版本匹配的 wheel版本不匹配会直接 import 报错。准备预训练权重把官方的 big-lama 权重包下载解压到项目目录导出脚本会用到big-lama/last.ckpt。五步完成 ONNX 导出第 1 步按配置构建模型并加载权重。注意从 配置 里取resnet_conv_kwargs.ratio_gin0.75传给 FFC 参数漏掉它权重就对不上model GlobalGenerator(input_nc4, output_nc3, ngf64, n_downsampling3, n_blocks18, ffc_positions[9, 10], ffc_kwargs{ratio_g: 0.75}) ckpt torch.load(big-lama/last.ckpt, map_locationcpu) model.load_state_dict(ckpt[state_dict], strictFalse) model.eval()第 2 步造一个示例输入并导出。输入是图像蒙版拼接出的 4 通道张量dynamic_axes声明高宽两维可变这样同一个 ONNX 文件能接受不同分辨率dummy torch.randn(1, 4, 512, 512) torch.onnx.export(model, dummy, big-lama.onnx, opset_version12, input_names[input], output_names[output], dynamic_axes{input: {2: h, 3: w}, output: {2: h, 3: w}})第 3 步导出后验证。先用 ONNX Runtime 跑一遍 dummy 输入和 PyTorch 的输出对比 L2 差值正常应在 1e-3 量级内差值异常说明权重没对齐或激活函数处理有误。第 4 步处理 opset 与算子兼容性。FFC 块内部的 FFT 算子不是所有后端都支持ONNX Runtime一般没问题。TensorRTOnnxParser部分版本解析FFT/rFFT节点会报错可尝试调低 opset11/12或在 parser 日志里定位具体节点。第 5 步实在绕不过时的兜底方案。换用纯空域卷积的生成器配置重新导出例如 ffc_resnet 生成器配置目录 下的其他变体——质量会有微小变化但导出链路 100% 顺畅。批量部署前建议固定这一份 ONNX别再混用。三步构建 TensorRT 引擎引擎构建就是读 ONNX → 配优化选项 → 序列化三件事import tensorrt as trt trt_logger trt.Logger(trt.Logger.WARNING) builder trt.Builder(trt_logger) network builder.create_network(1 int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser trt.OnnxParser(network, trt_logger) with open(big-lama.onnx, rb) as f: parser.parse(f.read())然后配置构建选项config builder.create_builder_config() config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 30) # 1GB 工作区 config.set_flag(trt.BuilderFlag.FP16)最后构建并保存engine builder.build_serialized_network(network, config) with open(big-lama.engine, wb) as f: f.write(engine)两个实用细节工作区给大一点1GB 起步。工作区不够时 TensorRT 会自动降级选慢的内核显存宽裕时不妨加到 2GB白捡一点速度。FP16 是默认推荐档开启后构建时间会变长但推理速度和质量几乎双赢后面有数据。构建出的.engine文件和 GPU 型号、驱动、TensorRT 版本强绑定换环境记得重新构建。推理时输入怎么喂拿到引擎后每次推理的输入处理顺序是固定的四步读图 → 缩放到 [0,1] → 按 8 的倍数对称填充 → 拼通道img cv2.imread(path) / 255.0 # HxWx3, float32 mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) / 255.0 x np.concatenate([mask[..., None], img], axis2) # 拼成 134 通道 x np.pad(x, ..., modesymmetric) # 补齐到 8 的倍数 x np.ascontiguousarray(x.transpose(2, 0, 1)[None]) # NCHW推理结束后记得把 padding 区域裁掉再把输出从 [0,1] 乘回 255 存成图。big-lama 配置里add_out_act: sigmoid所以 ONNX/TensorRT 的输出已经规整在 0~1不用再手动做 tanh 归一——这是很多人移植时容易搞混的一点。三种推理后端怎么比用同一批测试图例如 50 张 512×512分别跑三个后端对比单张耗时和输出一致性后端相对速度输出精度适用场景PyTorch 原生基准 1×参考基准开发调试、结果对齐ONNX Runtime (FP32)约 1.5~2×与基准几乎一致CPU/GPU 通用、跨平台TensorRT FP16约 2~5×肉眼无差PSNR 变化通常 0.5dBGPU 上线首选TensorRT FP32约 1.5~3×最高质量审计、兜底TensorRT INT8最快需要校准集修复纹理可能劣化实时极限场景慎选我的建议很直接GPU 上线一律 TensorRT FP16。修复任务对细微数值误差的容忍度远高于分类任务FP16 在 512 分辨率下肉眼与 FP32 基本无差只有当 PSNR 对比掉出 0.5dB 以上才退回 FP32。INT8 除非有明确校准集且实测纹理无损否则不建议。验收时跑两个对照ONNX 输出 vs PyTorch验证导出无损、TensorRT FP16 vs FP32验证精度档位两项都过再进生产。常见问题与排查现象大概率原因处理办法ONNX 导出报不支持的 FFT 节点TensorRT 对应版本不认该 opset降低 opset或换纯卷积的 生成器配置 重新导出引擎构建成功但推理结果全黑/全白输入没归一化或通道顺序错确认 [0,1] 归一化、mask 在前 RGB 在后的 4 通道拼接输出尺寸和原图对不上忘了裁掉 padding按pad_out_to_modulo默认 8裁回原始高宽换个机器 engine 加载失败engine 与 GPU/驱动/TensorRT 版本绑定在新环境重新执行构建脚本想支持多种分辨率只按单尺寸构建的引擎配置 min/opt/max 动态 profile或按 512/1024 分档各建一个引擎不同输入尺寸速度波动大动态 shape 下 kernel 未按实际尺寸调优按实际业务分辨率分档构建静态引擎通常更稳更快落地建议按场景选档实时交互类网页端擦除、修图 App单张 500msTensorRT FP16 固定 512/1024 尺寸档batch1显存最省。离线批量几万张商品图TensorRT FP16batch 按显存调到 4~8配多进程分发图片吞吐可再翻数倍。质量审计/对拍基准保留一份 PyTorch 或 TensorRT FP32 的输出做 ground truth不参与生产。显存吃紧8GB 级别FP16 batch1 流式读写避免把整批图片堆在内存里。这套PyTorch → ONNX → TensorRT的链路不挑模型把 LaMa 换成其他大掩码修复或图像生成模型导出、构建、对拍的步骤完全一样。先把 512 固定档的 FP16 引擎跑通拿到真实加速数据再考虑动态 shape 和 INT8是踩坑最少的路线。【免费下载链接】lama LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表