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

Jetson Nano上TensorRT FP16引擎构建异常推理结果错误求助

TensorRT FP16模式输入输出仍为FP32的问题原因及修复

问题背景

在Jetson Nano上使用提供的代码基于ResNet50的ONNX模型构建TensorRT引擎,已配置FP16模式,但仅当输入输出均为FP32时能得到正确分类结果;使用FP16输入/输出时结果错误,调用engine.get_binding_dtype查看输入输出 dtype 为FP32,问题出在build_engine函数中。

核心原因分析

  1. ONNX模型输入输出 dtype 限制:你的ResNet50 ONNX模型本身的输入、输出节点数据类型是FP32,TensorRT默认会保留原模型的输入输出 dtype。开启BuilderFlag.FP16仅允许TensorRT在模型内部运算时使用FP16优化,不会自动修改输入输出节点的数据类型。
  2. 未显式指定输入输出 dtype:代码中仅开启了FP16优化开关,但没有手动将输入输出节点的数据类型设置为FP16,因此引擎的输入输出依然保持原模型的FP32类型。
  3. 旧引擎缓存干扰:如果之前已经用FP32模式构建过.engine文件,代码会直接加载缓存的旧引擎,不会重新构建FP16版本,导致输入输出 dtype 未更新。

修复后的build_engine函数

def build_engine(onnx_file_path):
    engine = None
    runtime = trt.Runtime(TRT_LOGGER)
    engine_file_path = os.path.join(os.getcwd(), ENGINE_FILE)
    
    # 注意:如果要重新构建FP16引擎,建议先删除旧的.engine文件
    # 可手动删除或添加代码判断是否需要强制重新构建
    if os.path.exists(engine_file_path) and os.path.isfile(engine_file_path):
        with open(engine_file_path,"rb") as fb:
            engine = fb.read()
            engine = runtime.deserialize_cuda_engine(engine)
    else:
        explicit_batch_flag = 1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
        builder = trt.Builder(TRT_LOGGER)
        network = builder.create_network(explicit_batch_flag)
        parser = trt.OnnxParser(network, TRT_LOGGER)
        
        with open(onnx_file_path, 'rb') as model:
            print('Beginning ONNX file parsing')
            parser.parse(model.read())
        print('Completed parsing of ONNX file')

        # 关键修改:显式设置输入输出节点为FP16
        # 获取输入tensor
        input_tensor = network.get_input(0)
        input_tensor.set_dtype(trt.DataType.HALF)
        # 获取输出tensor
        output_tensor = network.get_output(0)
        output_tensor.set_dtype(trt.DataType.HALF)

        config = builder.create_builder_config()
        config.max_workspace_size = 1 << 20
        if builder.platform_has_fast_fp16:
            print(builder.platform_has_fast_fp16)
            config.set_flag(trt.BuilderFlag.FP16)

        plan = builder.build_serialized_network(network, config)
        engine = runtime.deserialize_cuda_engine(plan)

        with open(engine_file_path,"wb")as f:
            f.write(plan)
 
    return engine

修复说明

  • 显式设置输入输出 dtype:在解析ONNX模型后,通过network.get_input(0)和network.get_output(0)获取输入输出tensor,调用set_dtype(trt.DataType.HALF)强制设置为FP16类型,确保引擎的输入输出与内部运算一致。
  • 清理旧引擎缓存:首次运行修复后的代码前,手动删除已存在的.engine文件,避免加载旧的FP32版本引擎。
  • FP16模式生效:开启BuilderFlag.FP16后,TensorRT会在内部使用FP16进行运算优化,同时输入输出也变为FP16,此时使用FP16输入输出即可得到正确的分类结果。

内容的提问来源于stack exchange,提问作者Blackat

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 17:15:37