使用TensorRT转换动态形状输入输出SRModel超分模型报错求助
问题解答
现有配置存在的问题
- 缺失显式批处理模式开关:TensorRT 7.0版本的动态形状功能依赖显式批处理模式,你当前未开启该模式,直接传入带-1的动态维度会触发维度合法性校验报错,也就是你最后自定义添加输入失败的直接原因。
- 输入维度配置错误:add_inputs方法中你对输入张量做了
shape=tuple(torch_input.shape)[1:]切片,丢弃了批处理维度,结合未开启显式批处理的配置,网络会默认使用隐式批处理逻辑,无法匹配你设置的带批处理维度的动态形状参数。 - 推理输出尺寸硬编码:推理代码中直接写死输出张量尺寸为(1,3,448,448),输入尺寸变化后,实际网络输出尺寸和硬编码尺寸不匹配,导致无法拿到正确结果,是B、C测试场景无输出的核心原因。
- 优化profile配置时机错误:你在调用ctx.add_inputs之后才设置profile的形状参数,torch2trt的转换逻辑需要在添加输入前完成profile配置,否则网络构建时会使用固定输入尺寸,动态形状不生效。
- 重复设置profile无效:同一个优化profile对同一个输入只能调用一次set_shape,你在C测试场景中两次调用set_shape会覆盖之前的配置,无法同时生效两组输入输出映射。
修复方案
- 开启显式批处理模式:创建网络实例时添加对应flag:
builder = trt.Builder(TRT_LOGGER) network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))
- 修正输入维度配置:add_inputs方法中保留完整输入维度,不要做切片:
trt_tensor = self.network.add_input( name=names[i], shape=tuple(torch_input.shape), dtype=torch_dtype_to_trt(torch_input.dtype), )
- 调整profile配置顺序:在调用ctx.add_inputs之前完成profile的shape设置,同时确保config的fp16等配置正确传入torch2trt转换逻辑。
- 修正输出张量创建逻辑:直接使用context返回的动态形状创建输出张量,不要硬编码尺寸:
shape = tuple(self.context.get_binding_shape(idx)) output = torch.empty(size=shape, dtype=dtype, device=device)
- 验证自定义算子兼容性:检查你自定义实现的shuffle_layer和interpolate算子,确保支持动态形状推理,添加维度校验逻辑避免形状计算错误。
内容的提问来源于stack exchange,提问作者Allen Kris
相关产品推荐
相关产品推荐

