ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

torch2trt深度评测:PyTorch模型高效转换TensorRT的工程实践

torch2trt深度评测:PyTorch模型高效转换TensorRT的工程实践 torch2trt 这名字在 NVIDIA 开发者社区里不算新鲜但大多数讨论都停留在“能用”和“不能用”的层面很少有人把它当成一个需要深度评估的工程组件来看待。我这次不打算只跑个 demo 就下结论而是把 torch2trt 的源码翻了一遍梳理了它的核心架构和转换链路再结合几轮真实工程验证写一份偏企业尽调性质的评测报告。这篇内容适合正在做 PyTorch 模型推理加速、准备引入 TensorRT、或者被 ONNX 导出算子坑到想换方案的同学看完你基本能判断这个工具能不能接进自己的项目流程。1. 为什么还要折腾 torch2trtTensorRT 工具链的选型对照1.1 先搞清楚 TensorRT 转换生态里都有谁TensorRT 本质上是 NVIDIA 的深度学习推理优化器和运行时它并不直接读 PyTorch 的权重格式更不认识 torch.nn.Module。要把 PyTorch 模型跑在 TensorRT 上得先把它翻译成 TensorRT 能理解的中间表示。目前主流的路子有这几条一是 PyTorch 导出 ONNX再通过 trtexec 或者 onnx-tensorrt 把 ONNX 转成 TensorRT engine二是用 NVIDIA 后来的 torch_tensorrt 工具把 TorchScript 或动态图直接编译到 TensorRT第三就是本文要讲的主角 torch2trt它走的是“在 PyTorch 前向执行时同步构建 TensorRT 网络”的路线。这三条路线的差异不只是工具链长短的问题。ONNX 导出适合结构规整、算子常见的模型可一旦遇到自定义算子、控制流或者某些动态 shapeONNX 的算子集映射就会出幺蛾子。torch_tensorrt 虽然官方支持度高但它要求模型能被 TorchScript 完整追踪这同样会卡住一部分带非标准控制流的模型。torch2trt 的做法是直接把转换器注册到 PyTorch 算子上在模型 forward 的时候逐算子捕获并构建 TensorRT 层相当于绕开了 ONNX 这个中间商让不少在 ONNX 上碰壁的模型能跑通。我在实际项目里遇到过很典型的例子一个带着 F.grid_sample 的检测模型ONNX 导出后算子支持度很差trtexec 解析直接报不支持后来换成 torch2trt因为社区里已经有人给 grid_sample 写了 converter居然很顺利就转过去了。这种场景就是 torch2trt 存在价值的直观体现。1.2 torch2trt 独有的转换思路要理解 torch2trt得先抓住一个核心设计它不是做一个“模型翻译器”而是一个“图捕获器 层构建器”。在调用 torch2trt 的转换函数时它会用 PyTorch 的 tracing 机制跑一遍模型的前向这一步不是真的为了计算而是为了捕获计算图中每个算子被调用的时刻。每当一个注册过的算子被调用对应的 converter 就会被触发此时 torch2trt 会在 TensorRT 的 network 里创建一个对应类型的层并把输入 tensor 和输出 tensor 的映射关系记录下来。这个设计有一个很爽的副作用你不需要维护一份完整的模型结构描述也不需要手动遍历 nn.Module 的嵌套关系。torch2trt 天然支持 nn.Sequential、nn.ModuleList、函数调用、循环里反复调用的同一子模块因为它的单位是“算子调用”不是“模块层级”。这也是它和很多“每层正名硬编码”的转换工具完全不同的地方。当然这种设计也有代价转换结果高度依赖 tracing 过程中实际执行到的分支。比如模型里有 if 分支虽然 PyTorch 的 tracing 会把当前走到的那条路径记录下来但如果分支条件依赖于输入张量而非常量torch2trt 和其他 tracing 类工具一样会漏掉另一个分支。这个局限我必须提前点名否则后面到真实业务模型你会踩坑。1.3 企业选型时该关注什么从选型角度看评估一个转换工具不能只看“能不能转”还要看转换后的代码是否可控、性能是否对齐、遇到不支持的算子是否好扩展。torch2trt 在这几个维度上给出的答案比较明确。算子覆盖率方面torch2trt 对 cv 领域常见算子覆盖相对完整卷积、全连接、归一化、激活、池化、上采样、拼接、切片这大类都没有问题transformer 里常见的 einsum、softmax、layer_norm 也有对应实现。如果碰到没覆盖的算子torch2trt 的扩展机制非常轻量写一个函数用 tensorrt_converter 装饰器注册到对应 PyTorch 算子上即可这种扩展方式比维护一整个 ONNX 自定义算子节点要简单得多。性能对齐方面torch2trt 构建的层直接映射到 TensorRT C API 的对应层并不是走某种模拟层再自动优化所以理论上转换出来的 engine 和手写 TensorRT 网络没有本质区别。实际测试中只要没有落到性能很差的回退实现模型的卷积层、BN 层融合、内存复用等优化都会被 TensorRT 正常执行。可维护性方面要打个折扣。这个工具核心维护者几乎是一个人在跟进更新节奏不算快新算子支持也有滞后。如果你项目里大量使用 torch 内部新推出的算子可能就得做好自己写 converter 的准备。但换个角度看它的核心架构非常稳定从 2021 年到现在的版本装饰器注册机制和 ConversionContext 的核心数据结构没怎么大改新版本 PyTorch 只要兼容 API基本就能继续用。2. 环境准备驱动、CUDA、TensorRT、PyTorch 之间的版本账2.1 版本对应关系先把账算清楚安装这步看起来简单但版本不匹配能折腾你一个下午。torch2trt 本质上是一个 Python 包它通过 pybind11 调用 TensorRT 的 C API因此它和 TensorRT 版本之间的耦合度比普通 Python 库要高不少。我建议先把驱动、CUDA、PyTorch、TensorRT 的版本对应关系列成一张表作为基线。这里给一个实测稳定的组合Ubuntu 22.04 NVIDIA 驱动 535 系列 CUDA 12.2 PyTorch 2.1.x TensorRT 8.6.x torch2trt master 分支。如果用 CUDA 12.0 以下的老版本TensorRT 8.4 也够但 PyTorch 最好对应 1.13 左右太新的 PyTorch 和旧 TensorRT 在 API 上偶有摩擦。组件建议版本我实测稳定备选版本Ubuntu22.0420.04NVIDIA 驱动535.x525.x / 545.xCUDA12.211.8PyTorch2.1.x2.0.x / 1.13TensorRT8.6.x8.4.xtorch2trtmaster 分支最新 release先装驱动再装 CUDA Toolkit然后建 conda 环境装 PyTorch最后单独装 TensorRT顺序不要乱。在 conda 环境里直接 pip install tensorrt 会拉到 Python 版 TensorRT 包但 torch2trt 还需要 TensorRT 的 libnvinfer 动态库和头文件这两者不完全是一回事所以我更推荐下载 NVIDIA 官方的 TensorRT tar 包解压后把 lib 目录加进 LD_LIBRARY_PATH。2.2 Ubuntu 驱动与 CUDA 的报错坑热词里出现一堆驱动相关的报错这里集中展开一下。很多同学装好驱动后跑 nvidia-smi 直接报 “nvidia-smi has failed because it couldnt communicate with the nvidia driver”十有八九是新内核和驱动版本不匹配。Ubuntu 自动更新内核后NVIDIA 驱动模块没跟着重新编译就会出现这种通信失败。解决方法是重装一遍与当前内核匹配的驱动或者在更新内核后执行 dkms autoinstall。另一个高频报错是 “failed to load module glxserver_nvidia (module does not exist)”。这个一般出现在驱动卸载不干净、或者 nouveau 内核模块没屏蔽的情况下。安装驱动前得先把 nouveau 禁掉在 /etc/modprobe.d/blacklist-nouveau.conf 里写入 blacklist nouveau 和 options nouveau modeset0更新 initramfs 后重启再安装驱动就顺畅了。如果你还碰到 3D Vision 安装卡住这类问题多半是安装器在等待图形会话释放 GPU建议直接切到纯命令行模式安装。有一个判断驱动是否就绪的技巧装完驱动跑 nvidia-smi看到右上角和下方都显示 CUDA Version说明驱动通过此时再运行 nvcc -V如果输出和 nvidia-smi 里的 CUDA 版本不一致非常正常因为 nvcc 是 CUDA Toolkit 的编译器跟驱动自带的 CUDA runtime 版本允许不同。只要 nvidia-smi 正常驱动层面就没问题。2.3 conda 环境安装与验证环境建议用 conda 隔离。创建一个干净的 Python 3.9 环境然后装匹配 CUDA 版本的 PyTorch比如conda create -n trt python3.9 -y conda activate trt pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121接着安装 TensorRT。这里我强烈建议直接下载 NVIDIA 官网的 TensorRT 8.6 tar 包解压后把路径写进环境变量export TRT_RELEASE/path/to/TensorRT-8.6.x.x export LD_LIBRARY_PATH$TRT_RELEASE/lib:$LD_LIBRARY_PATH最后安装 torch2trtgit clone https://github.com/NVIDIA-AI-IOT/torch2trt.git cd torch2trt python setup.py install验证是否装好可以跑一个最小样例把 ResNet18 转一遍。但这里要注意最小样例能跑通不代表一切正常我建议额外做一个“输入输出一致性”检查让转换后的模型跑同一个输入和 PyTorch 原模型的结果做逐元素对比这一步能快速暴露环境层面的问题。等到后面实操部分我会给出完整的验证脚本。装好环境后如果你的 CUDA、TensorRT、PyTorch 三个组件版本匹配torch2trt 会直接正常工作如果出现 ImportError 或者 undefined symbol先检查 LD_LIBRARY_PATH 是否指向正确的 TensorRT lib 目录常见问题是系统里有多个 TensorRT 版本Python 进程加载了旧版动态库。3. 源码架构解剖torch2trt 到底怎么运作3.1 仓库结构与核心文件把源码克隆下来先看目录结构。torch2trt 的核心代码非常集中主要分三块converters 目录下是所有内置转换器的实现按算子类别拆分成 conv.py、matmul.py、normalization.py、activation.py 等文件core.py 定义了整个转换框架的骨架也就是 ConversionContext、TRTModule、convert 主流程和装饰器注册表还有 setup.py 和版本控制相关文件。如果只挑一个文件精读那一定是 core.py它承载了整个框架的心智模型。理解了 core.py 里的 ConversionContext 如何工作你就理解了 torch2trt 的全部设计精髓。converters 目录里的文件反倒是次要的因为每个转换器的逻辑都不复杂核心套路都一样拿 PyTorch 算子入参映射成 TensorRT 层的参数然后向 network 里加一层。3.2 ConversionContext整个转换过程的“现场”ConversionContext 是 torch2trt 在转换过程中维护的一个上下文对象。它的职责很集中记录当前 TensorRT network、当前输入张量的映射表、当前输出张量的映射表以及已经处理过的参数集合和计算过程中用到的中间张量。可以把 ConversionContext 理解成一个“施工许可证”加“施工图纸”的结合体。它握着 TensorRT builder 和 network 的引用这是向 engine 里填层的唯一通道。每一次 PyTorch 算子被调用触发对应的 converterconverter 第一件事就是通过 context 拿到当前算子的 PyTorch 输入张量去 input_map 里找到对应的 TensorRT ITensor如果没有就现场创建一个输入输出张量同样会注册到 output_map供后续算子继续查表。有些细节值得注意比如 context 里还维护了 method_dic 和 function_dic 两张注册表。深入看一下就会发现torch2trt 对“如何捕获一个算子”的处理分了两条线一条是捕获 torch.nn.functional 下的函数式算子一条是捕获 torch.Tensor 的方法调用。这两类调用在 Python 层面的触发点不同torch2trt 分别做了处理这也是它能同时覆盖 F.conv2d、x.view、x.reshape 等形态的原因。3.3 装饰器注册机制扩展一个算子有多简单torch2trt 最值得称道的设计之一就是装饰器注册机制。所有内置转换器都是通过 tensorrt_converter 这个装饰器注册到注册表里的。它的用法非常直白比如要支持某个 PyTorch 函数只需要tensorrt_converter(torch.nn.functional.relu) def convert_relu(ctx): input ctx.method_args[0] output ctx.method_return layer ctx.network.add_activation( trt_get_tensor(input, ctx), typetrt.ActivationType.RELU ) trt_set_tensor(output, layer.get_output(0), ctx)这里 tensorrt_converter 接收的是一个字符串这个字符串是算子的完整导入路径。torch2trt 在初始化时会把这个注册表构建成一个字典key 是算子全名value 是对应的转换函数。当 tracing 阶段捕获到某一个 PyTorch 算子调用时torch2trt 就根据当前算子所属的模块路径去这个字典里查询命中则执行对应转换逻辑。这种注册机制带来的工程价值很大。遇到不支持的算子你不需要去 fork 整个 torch2trt 仓库改源码只需要在自己的工程里写一个新的转换函数用同一个装饰器注册进去这个函数就能被 torch2trt 自动识别并用于后续转换。企业项目里维护一个自定义算子转换库完全不需要改动第三方包的任何文件对代码审计和依赖管理都友好得多。3.4 前向传播捕获流程拆开看每一步把整个转换流程串起来看实际是四步走。第一步torch2trt 调用 convert 函数把输入的 torch.nn.Module 实例、输入张量示例、网络配置传进去第二步创建 ConversionContext并在 context 里初始化 TensorRT network第三步使用 PyTorch 的 torch.jit.trace 机制对模块进行 tracing这一步是最关键的因为 trace 过程中每执行到一个算子都会触发上面说的注册表查询第四步所有算子处理完context 里的 network 构建完成此时调用 builder 生成 engine并把结果包装成一个 TRTModule 返回。第三步里有个隐蔽的细节trace 不是直接把输入跑一遍就完事它会在执行每个算子时把调用信息写入 trace 图。torch2trt 精妙的地方在于它把自己的钩子挂在了算子执行点而不是等着 trace 结束再去解析图结构。这意味着它拿到的不是“图结构”而是“算子调用顺序”这正好是构建 TensorRT 网络需要的拓扑顺序。还有一处需要注意torch2trt 对每个 PyTorch Parameter 都有缓存机制。因为同一个权重参数可能在多个地方被引用比如共享权重的两个卷积层。如果每次遇到这个参数都在 TensorRT 里新建一个常量层会造成权重重复存储浪费显存。torch2trt 在 context 里记录了一个 tensor_map专门维护 PyTorch 张量到 TensorRT ITensor 的映射遇到已经转换过的参数直接复用这个细节对多分支共享权重的模型影响很大。3.5 一个 converter 的实现细节以 Conv2d 为例与其空谈机制不如看一个具体转换器的实现。torch2trt/converters/conv.py 里定义了卷积转换逻辑核心代码大概长这样tensorrt_converter(torch.nn.Conv2d.forward) def convert_conv2d(ctx): module ctx.method_args[0] input ctx.method_args[1] output ctx.method_return input_trt trt_get_tensor(input, ctx) kernel module.weight.detach().cpu().numpy() kernel_trt ctx.network.add_constant(module.weight.shape, kernel) layer ctx.network.add_convolution( input_trt, num_output_mapsmodule.out_channels, kernel_shapemodule.kernel_size, kernelkernel_trt.get_output(0) ) layer.stride module.stride layer.padding module.padding if module.bias is not None: layer.bias module.bias.detach().cpu().numpy() trt_set_tensor(output, layer.get_output(0), ctx)梳理这段代码的逻辑先从 context 里拿到当前算子的模块实例、输入张量和输出张量再把 PyTorch 权重转成 NumPy 数组通过 add_constant 构建一个常量层然后调用 TensorRT 的 add_convolution 构建卷积层最后把该层的输出和 PyTorch 计算图中的输出张量做映射。整个逻辑简单、直接、没有任何花哨的包装。这个模式在所有 converter 里是高度一致的取输入、取参数、加层、映射输出。读通了这一个其他 converter 基本都能看懂。不过卷积这里有个细节值得提醒在 TensorRT 8.x 中 add_convolution 的 kernel 参数接收的是一个 ITensor而不是常见的 numpy 数组torch2trt 通过 add_constant 做了一个中间层来满足这个 API 要求。某些旧版本写法直接把 numpy 传进去也是可以的但新版 API 已经变了如果你自己写 converter一定要对着当前版本 TensorRT 的 Python API 检查参数类型。3.6 TRTModule转换产物怎么承载推理转换完成后返回的 TRTModule 本质上是一个包装过的 torch.nn.Module。它的不同之处在于 forward 方法里不再执行任何 PyTorch 算子而是把输入张量复制到 GPU通过 pybind 调用 TensorRT engine 执行推理。TRTModule 的初始化参数包括 TensorRT engine、输入输出绑定名称以及是否启用 context。每次 forward 时它创建一个统一的执行上下文把 PyTorch 输入张量按绑定名传给 engine执行完毕后把输出张量取出来转成 PyTorch Tensor 返回。这套机制保证了 TRTModule 在外部用起来和普通 nn.Module 几乎一样可以无缝塞进已有推理代码里。如果转换时指定了动态 batch 或动态 shapeTRTModule 会持有一些额外的 shape 信息并在每次推理时调用 set_binding_shape 来调整输入输出维度。这块稍复杂我在实操部分会再说明。4. ResNet18 实战从 PyTorch 模型到 TensorRT engine4.1 转换 API 参数详解实战前先把 torch2trt 的转换入口说透。核心方法是 torch2trt.torch2trt它的签名里最关键的几个参数是input输入张量示例可以是一个 torch.Tensor也可以是元组表示模型接收多个输入。这个参数既参与 tracing 计算图生成又决定了 engine 的输入 shape。max_batch_size静态 batch 上限。如果推理时 batch 不变或者固定保持默认 1 就行如果想支持动态 batch需要配合 opt_shape 参数。fp16_mode是否启用半精度转换。开启后 TensorRT 会在满足条件的层上自动替换为 FP16 精度对性能提升很直接但有精度下降风险。max_workspace_sizeTensorRT 构建时的最大可用显存单位字节。设小了可能优化不充分设大了可能超出显存导致构建失败。strict_type_constraints是否强制严格类型约束。一般不需要除非你有特殊的层必须保持 FP32。这里我给一个推荐配置思路先以 FP32、默认 workspace 跑通再开 FP16 做精度对比最后根据显存余量调大的 workspace让 TensorRT 有更多空间去做 kernel 自动调优。不要一上来就全参数拉满否则出问题很难定位是模型不支持还是参数设置不合理。4.2 一步转换与结果验证转换代码本身非常简洁PyTorch 里训练好的模型直接传进去就行。下面是一个完整的 ResNet18 转换和验证过程import torch import torchvision.models as models from torch2trt import torch2trt model models.resnet18(pretrainedTrue).cuda().eval() x torch.randn(1, 3, 224, 224).cuda() model_trt torch2trt(model, [x], fp16_modeFalse, max_workspace_size1 30) print(converted!) with torch.no_grad(): y_pytorch model(x) y_trt model_trt(x) max_abs_err float((y_pytorch - y_trt).abs().max()) print(max_abs_err:, max_abs_err)如果 max_abs_err 在 1e-4 量级甚至更小说明转换结果和 PyTorch 原模型输出基本一致。如果差距过大就要检查模型里是否有对精度特别敏感的层比如某些归一化、softmax以及 network 构建时是否报过警告。在实际工程里我还建议加一个固定的随机种子跑多次推理检查输出是否稳定。TensorRT engine 构建过程引入的优化不应该影响输出的确定性如果你发现同一输入不同批次推理结果有微小波动那大概率是模型里用了非确定性算子比如某些 Attention 实现的 dropout 在 eval 模式下仍然被启用这点需要到模型代码里排查而不是在转换工具上找问题。4.3 FP16 与工作空间调优把 fp16_mode 设为 True 再跑一次model_trt_fp16 torch2trt(model, [x], fp16_modeTrue, max_workspace_size1 30)FP16 带来的性能提升很直观尤其是卷积层和全连接层多的模型。但相应的代价是精度下降ResNet18 这种较为鲁棒的结构通常问题不大max_abs_err 可能从 1e-6 涨到 1e-3 左右对分类任务影响很小。如果是检测、分割任务FP16 对边界回归和分割边界可能有轻微影响需要你根据自己的精度要求判断。workspace 的调优逻辑是在显存允许范围内给 TensorRT 更大的构建空间它就能尝试更多 kernel 变体和融合策略。实际操作中我会从 1GB 开始逐步往上加直到构建时间明显变长或者显存不足报错为止。我测过一个语义分割模型workspace 从 512MB 提高到 2GB 后推理延迟提升了约 8%这个提升完全取决于模型结构得实测观察。另外如果你要部署到不同的 GPU 上强烈建议在目标 GPU 上重新构建 engine。TensorRT engine 和 GPU 架构和 TensorRT 版本强相关把 A100 上构建的 engine 拷贝到 4090 上跑大概率直接报错或者性能很差。4.4 推理性能对比与正确解读转换完成后最让人关心的就是性能。一个简单的 benchmark 脚本就可以测出差异import time def run_bench(model, x, n100): # 前几次预热让 CUDA context 和显存分配稳定下来 for _ in range(10): model(x) torch.cuda.synchronize() start time.time() for _ in range(n): model(x) torch.cuda.synchronize() return (time.time() - start) / n * 1000 print(PyTorch ms/iter: %.3f % run_bench(model, x)) print(TensorRT ms/iter: %.3f % run_bench(model_trt, x))我实测 ResNet18 在 RTX 3060 上PyTorch FP32 大约是 5.2msTensorRT FP32 大约 3.8msTensorRT FP16 大约 2.5ms。但这里必须提醒benchmark 的结果受很多因素影响输入 shape 大小、batch 大小、是否固定 memory 分配、CUDA 版本、TensorRT 版本。我见过很多人在自己的机器上测出和宣传完全不同的数字这不代表工具不行而是测试条件不同。解读性能时还有一个容易踩的坑小模型或者轻量模型在 PyTorch 里跑Python 侧的开销占比很高转成 TensorRT 后引擎执行时间很短反而数据拷贝和 Python 调用开销成了主要矛盾。如果你的推理管线里输入预处理、后处理都很重单纯优化模型推理这 1ms 可能只是杯水车薪企业级优化必须把整个推理链路放在一起做 profile不能只盯着 engine 本身的数字看。5. 常见问题与排查技巧实录5.1 高频报错速查表把实际操作中遇到的典型报错整理成一张表方便你对照排查。报错信息常见原因解决办法NotImplementedError: The converter for ... is not implemented算子未覆盖换用更常见的算子表达或者手写自定义 converterInvalidArgumentError: split can’t be applied on tensor...某层参数不兼容 TensorRT检查该算子的具体参数组合主动替换为等价算子Node ... did not match any convertertracing 时遇到了推理分支之外的算子确保转换时的输入能覆盖所有关键路径CUDA error: out of memoryworkspace 设置过大降低 max_workspace_size或者减小 batch 测试AssertionError: bindings for... are not consistent输入输出 shape 与 engine 不符检查 input 示例和实际推理 shape 是否一致engine 加载时报 architecture mismatchengine 不在当前 GPU 或版本上构建在部署机上重新构建 engineNotImplementedError 这个错误是 torch2trt 用户遇到最多的。报错信息里会直接给出算子路径比如 “The converter for torch.nn.functional.gaussian_blur is not implemented”。这里的处理思路不是干瞪眼而是先看 PyTorch 里这个算子是否能被等价替换。gaussian_blur 如果只是在推理前对输入做预处理完全可以在转 TensorRT 之前用 CPU 或者 OpenCV 处理掉没必要把它放进 engine 里如果确实在模型中间层那就走自定义 converter 方案。5.2 自定义 converter 扩展一个不支持的算子自定义 converter 是 torch2trt 最值得依赖的能力。以一个常见的 unsupported 算子为例比如模型里用到了某个自定义的加权求和算子它在 PyTorch 里长这样def weighted_sum(a, b, alpha): return a * alpha b * (1 - alpha)要给这个算子写转换器需要先明确它可以被拆成 TensorRT 的 elementwise 乘法和加法。转换器代码大致如下import tensorrt as trt from torch2trt import tensorrt_converter, trt_get_tensor, trt_set_tensor tensorrt_converter(__main__.weighted_sum) def convert_weighted_sum(ctx): a ctx.method_args[0] b ctx.method_args[1] alpha ctx.method_args[2] output ctx.method_return a_trt trt_get_tensor(a, ctx) b_trt trt_get_tensor(b, ctx) alpha_arr alpha.detach().cpu().numpy() alpha_trt ctx.network.add_constant((1,), alpha_arr).get_output(0) one_minus_alpha ctx.network.add_constant((1,), 1.0 - alpha_arr).get_output(0) a_scaled ctx.network.add_elementwise(a_trt, alpha_trt, trt.ElementWiseOperation.PROD).get_output(0) b_scaled ctx.network.add_elementwise(b_trt, one_minus_alpha, trt.ElementWiseOperation.PROD).get_output(0) out ctx.network.add_elementwise(a_scaled, b_scaled, trt.ElementWiseOperation.SUM).get_output(0) trt_set_tensor(output, out, ctx)写完这个函数后在转换之前 import 它即可torch2trt 会自动把这个装饰器注册到全局注册表。这个扩展模式是 torch2trt 能活这么久的核心原因——它把“支持新算子”的难度降到了“写一个普普通通的函数”而不是去改底层框架。写自定义 converter 时有几个细节要特别注意一是 ctx.method_args 的索引不要搞错尤其当目标算子是模块方法而非函数时第 0 个参数是模块实例本身二是 add_constant 创建常量时shape 要和参与运算的 TensorRT 张量兼容必要时要利用 broadcast 机制三是如果算子内部有多个输出每个输出都要有对应的 trt_set_tensor 注册否则后续算子拿到的是未映射的张量直接报错。5.3 动态 shape、序列化与部署注意事项torch2trt 原生对动态 shape 的支持比较有限它主要支持动态 batch而不太擅长 H、W 等维度任意变化。如果你的业务有动态分辨率需求需要传入多个不同 shape 的 input 示例或者借助 trt 的 optimization profile 配置。这个方案可行但配置复杂度和出问题的概率都会上升我的建议是如果业务允许优先把输入尺寸固定下来。很多推理框架都支持统一 resize 到固定大小这样能最大程度发挥 TensorRT 的优化能力。序列化方面转换得到的 TRTModule 可以通过 state_dict 的方式保存 engine 权重也可以直接用 trt 的 serialize 接口保存为 .engine 文件。我推荐用后者保存 engine 文件部署时直接反序列化加载省去每次重新构建的时间。但千万别忘了 engine 与硬件、TRT 版本绑定这个限制换了 GPU 型号或升级了 TensorRT 版本老 engine 基本作废。部署环节我不建议直接在产线环境里做转换而是把转换留到 CI/CD 或者离线构建阶段。典型做法是训练完成后跑一次 torch2trt 得到 engine保存文件部署端加载 engine 文件只负责推理。这样做的好处是部署机不需要安装 PyTorch也不依赖 torch2trt 环境镜像体积能小不少故障面也更小。写在最后的小提醒torch2trt 不是银弹选型时最重要一点是要评估你们模型里的算子是否在它的覆盖范围内最好拿真实模型和真实数据尽早验证。如果发现某个关键算子不支持先试试替换成等价操作不行再写自定义 converter实在兜不住才考虑换其他转换路线。我做了几次项目下来觉得 torch2trt 最舒服的地方是转换链路短代码可读性高出了问题愿意去翻源码的话半天时间基本能定位到根因。这套流程和源码阅读经验换个工具同样通用。
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表