You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 06:25:11