Jetson Orin上FastAI2模型TensorRT优化后批量大小不匹配问题
修复TensorRT推理时批量维度超出Profile范围的问题
核心原因
你用detect_batch_size=8创建输入张量,但Torch2TRT默认未将batch维度设置为对应动态范围,或转换时的profile未正确覆盖batch=8的情况,导致生成的TensorRT engine仅支持batch size=1。
具体修复步骤
1. 转换模型时显式设置动态batch维度(推荐)
使用Torch2TRT转换时,通过input_shapes参数指定batch维度的动态范围,确保包含目标batch size(8)。示例代码:
from torch2trt import torch2trt model = learner.model model.cuda().eval() # 定义输入形状范围:batch维度覆盖1到8,其余维度固定 input_shapes = [ (torch.randn(1, 3, 200, 200).cuda(), torch.randn(8, 3, 200, 200).cuda()) ] # 转换并生成支持动态batch的TRT模型 trt_model = torch2trt( model, [torch.randn(8, 3, 200, 200).cuda()], input_shapes=input_shapes, fp16_mode=True # Jetson平台推荐启用FP16加速 ) # 保存优化后的模型 torch.save(trt_model.state_dict(), 'trt_model_batch8.pth')
input_shapes传入batch=1和batch=8的张量,让TensorRT生成覆盖该范围的profile,确保推理时batch=8被允许。
2. 转换时强制固定batch size为8(无需动态batch场景)
如果推理永远只用batch=8,可直接固定输入形状,生成仅支持该batch的engine:
trt_model = torch2trt( model, [torch.randn(8, 3, 200, 200).cuda()], fixed_batch_size=True, fp16_mode=True )
注意:启用fixed_batch_size后,无法切换其他batch大小。
3. 推理前确认输入张量形状
加载TRT模型后,确保DataLoader输出的张量形状为[8,3,200,200],避免被自动调整:
data = next(iter(dataloader)) data = data.cuda() print(f"Input shape: {data.shape}") # 检查是否为(8,3,200,200) output = trt_model(data)
若FastAI的learner包装了TRT模型,需确认learner未修改输入的batch维度。
4. 手动配置TensorRT Profile(进阶)
若上述方法无效,可手动创建TensorRT Profile指定batch维度范围:
import tensorrt as trt builder = trt.Builder(trt.Logger(trt.Logger.WARNING)) network = builder.create_network() config = builder.create_builder_config() # 创建profile,设置输入batch的min/opt/max范围 profile = builder.create_optimization_profile() input_tensor = network.add_input(name="input", dtype=trt.float32, shape=(1,3,200,200)) profile.set_shape("input", (1,3,200,200), (4,3,200,200), (8,3,200,200)) config.add_optimization_profile(profile) # 后续手动完成模型转换(适用于自定义TRT构建场景)
验证修复
加载模型后测试batch=8的输入:
from torch2trt import TRTModule trt_model = TRTModule() trt_model.load_state_dict(torch.load('trt_model_batch8.pth')) # 测试输入 test_input = torch.randn(8,3,200,200).cuda() output = trt_model(test_input) print(f"Output shape: {output.shape}") # 确认输出正常
内容的提问来源于stack exchange,提问作者PhilBot
相关产品推荐
相关产品推荐

