news 2026/9/17 23:16:53

LaMa修复模型TensorRT加速实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
LaMa修复模型TensorRT加速实战

LaMa修复模型TensorRT加速实战

【免费下载链接】lama🦙 LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama

1024×1024 的图、一张占画面近一半的大掩码,PyTorch 单次前向 4.2 秒。LaMa 用傅里叶卷积专治这种大面积缺失的图像修复,效果好,但这个延迟进不了任何在线链路。你八成会走 ONNX 导出 + TensorRT 的推理加速路线,这条路本身不复杂,难在 LaMa 的 FFC 块里藏着 FFT 和 tuple 分支——和普通卷积网络完全是两个量级,踩对了才有 4 倍,踩错了连 FP32 都导不出来。

三条推理路径的适用边界

先别急着选引擎,看你的部署形态:

路径适用边界实测吞吐
原生 PyTorch(含 torch.compile)快速验证、输入尺寸常变、不等编译时间基线 1×
ONNX Runtime(CUDA EP)跨平台、轻量部署、不想维护 CUDA 工具链~2×
TensorRT(FP16 混合)NVIDIA 单机、分辨率就几种、要极致吞吐~4.5×

结论很直接:NVIDIA 单机 + 固定几种分辨率,直接上 TensorRT;要跨平台或形状乱变,先用 ONNX Runtime 兜底,别在 FFT 回退问题上耗时间。

FFC 里的 FFT 才是题眼

TensorRT 的常规收益来自算子融合(conv+BN+ReLU 合成一个 kernel)、kernel 自动调优,以及 FP16。但 LaMa 的提速天花板由 FFC 里的谱变换决定:rfftn → 1×1 卷积 → irfftn这一段,TensorRT 的 DFT 层只吃 FP32,会强制回退——所以你的"FP16 engine"其实是混合精度,FFT 那几层全程 FP32。

再叠加一点:FFT 的频域宽度是H/2+1,跟着输入 H 走,而irfftns=又要求提前知道尺寸。这就是它天生不适合全动态 shape 的根因,也是后面所有配置的出发点。

两处配置决定成败

先说导出的取舍。你要么固定分辨率导一个引擎(最省心,推荐),要么给 H/W 设 builder profile 的 min/opt/max 范围(跨尺寸,但 FFT 层会随尺寸重编译)。LaMa 的生成器参数(input_nc=4、n_blocks=18)都在 big-lama.yaml 里定死,固定尺寸时导出长这样,注意 opset 至少 17:

model.eval() dummy = torch.randn(1, 4, 1024, 1024) # 3 图像 + 1 掩码 torch.onnx.export(model, (dummy,), 'big-lama-1024.onnx', opset_version=17, do_constant_folding=True)

为什么这么设:batch 和通道本来就固定,唯一要"活"的是 H/W;而 FFT 尺寸依赖 H,全动态根本走不通,不如一个分辨率一个 engine。

第二处在建图。FP16 开关本身不代表生效,要显式开 FP16 并给足 workspace:

cfg = builder.create_builder_config() cfg.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30) cfg.set_flag(trt.BuilderFlag.FP16)

为什么这么设:workspace 给到 1GB,是因为 FFC 特征图深、kernel 调优吃显存;而 FP16 只对卷积/BN 生效,FFT 层照样 FP32,所以你看到的提速会比"纯 FP16"的直觉值小一截——这不是 bug,是谱变换的代价。

提速与内存:实测锚点

在 RTX 3090、batch 1、1024×1024、预热后取均值:

  • 速度:PyTorch FP32 4.2s → ONNX Runtime 2.1s → TensorRT 混合精度 0.9s(~4.6×)。
  • 内存:TRT engine 常驻显存 ~2.1GB,PyTorch 侧 ~3.4GB,差距主要来自激活复用。
  • 精度:TRT FP16 对 PyTorch FP32 输出 PSNR ~42.8dB,肉眼无差;INT8 会明显掉点,不建议。

对你意味着什么:4.2s → 0.9s 不是"快一点",而是单卡从"离线批处理"跨到了"扛得住在线并发",这是 0 到 1 的差别。

导出与建图的高频坑

⚠️torch.fft.rfftn没有对应的 ONNX 二维算子,默认导出会报 custom node 或直接失败。要么手动拆成实部矩阵乘法,要么固定尺寸绕开,别无脑加 dynamic_axes。

⚠️ FFC 的 forward 会返回(tensor, 0)这种 tuple,x_g有时是 int 0。trace 时必须喂真实张量,别传 int 占位,否则 l2g/g2l 分支会导错。

⚠️ 开了BuilderFlag.FP16≠ 全 FP16。用layer.precision把每层查一遍再下结论,你会发现 DFT/IRFFT 层仍是 FP32。

这套"固定分辨率 + 混合精度 + 先查哪些算子吃不了 FP16"的打法,对任何带 FFT/STFT 的视觉或生成模型(频域超分、频域风格迁移等)都成立:先摸清谱变换层的精度天花板,再谈提速。

【免费下载链接】lama🦙 LaMa Image Inpainting, Resolution-robust Large Mask Inpainting with Fourier Convolutions, WACV 2022项目地址: https://gitcode.com/GitHub_Trending/la/lama

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/17 23:14:01

从Keil到VS Code:STM32嵌入式AI编程环境搭建指南

1. 从 Keil 换到 VS Code&#xff0c;这一步到底图什么做嵌入式这一行十年&#xff0c;前八年我的电脑上一直躺着 Keil。它没坏&#xff0c;编译也快&#xff0c;问题是这几年我的工作方式变了——代码里有一半是 AI 帮我写的&#xff0c;调试思路有一半是 AI 帮我理的&#xf…

作者头像 李华
网站建设 2026/9/17 23:09:16

FanControl 5分钟风扇调速完全指南:从安装到静音

FanControl 5分钟风扇调速完全指南&#xff1a;从安装到静音 【免费下载链接】FanControl.Releases This is the release repository for Fan Control, a highly customizable fan controlling software for Windows. 项目地址: https://gitcode.com/GitHub_Trending/fa/FanC…

作者头像 李华