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

如何为ONNX Runtime的TensorRT执行提供程序创建INT8校准表?

实现ONNX模型INT8校准表生成(适配Jetson + ONNX Runtime/TensorRT)

方法1:直接通过ONNX Runtime生成校准表

ONNX Runtime的TensorRT后端支持自动完成INT8校准流程,只要提供代表性校准数据并开启对应配置即可:

操作步骤

  1. 配置ONNX Runtime的TensorRT执行器参数,开启INT8模式并指定校准表保存路径
  2. 准备一批和模型输入形状、 dtype、归一化逻辑完全一致的代表性校准数据(建议100-500条样本,避免用随机数)
  3. 运行一次推理触发校准流程,校准表会自动保存到指定路径

代码示例

import onnxruntime as ort
import numpy as np

# 1. 配置ONNX Runtime会话选项
sess_options = ort.SessionOptions()
# 开启TensorRT后端INT8支持
sess_options.enable_trt_int8 = True
# 指定校准表保存路径
sess_options.trt_int8_calibration_table_name = "int8_calibration.table"

# 2. 加载模型并指定TensorRT执行器
providers = [
    ('TensorrtExecutionProvider', {
        'trt_int8_enable': True,
        'trt_int8_calibration_table_name': "int8_calibration.table",
        # 首次生成校准表设为False,后续复用设为True
        'trt_int8_use_native_calibration_table': False
    }),
    'CPUExecutionProvider'
]
session = ort.InferenceSession("your_model.onnx", sess_options=sess_options, providers=providers)

# 3. 准备校准数据(示例:假设模型输入为(1,3,224,224)的FP32张量)
# 注意:必须用真实业务场景数据,否则校准后精度会严重下降
calib_data = np.random.randn(100, 3, 224, 224).astype(np.float32)  # 100条校准样本

# 4. 运行校准推理
input_name = session.get_inputs()[0].name
for data in calib_data:
    # 补全batch维度,匹配模型输入形状
    session.run(None, {input_name: data[np.newaxis, ...]})

# 校准完成后,校准表会自动写入指定路径

方法2:通过TensorRT Python API生成校准表,再给ONNX Runtime使用

如果需要更精细控制校准过程(比如自定义batch size、校准逻辑),可以直接用TensorRT API生成校准表,再在ONNX Runtime中加载使用:

操作步骤

  1. 定义TensorRT INT8校准器类,实现校准数据的批量读取逻辑
  2. 用TensorRT加载ONNX模型,开启INT8模式并绑定校准器,构建引擎时自动生成校准表
  3. 在ONNX Runtime配置中指定加载已生成的校准表

代码示例

import tensorrt as trt
import numpy as np
import os

# 自定义INT8校准器类
class Int8Calibrator(trt.IInt8Calibrator):
    def __init__(self, calib_data, input_name, batch_size=1):
        trt.IInt8Calibrator.__init__(self)
        self.batch_size = batch_size
        self.input_name = input_name
        self.calib_data = calib_data
        self.current_idx = 0
        # 分配GPU输入缓存
        self.device_buf = trt.allocator().allocate(self.calib_data[0].nbytes * self.batch_size)

    def get_batch_size(self):
        return self.batch_size

    def get_batch(self, names):
        if self.current_idx + self.batch_size > len(self.calib_data):
            return None
        # 获取一批校准数据
        batch = self.calib_data[self.current_idx:self.current_idx+self.batch_size]
        self.current_idx += self.batch_size
        # 拷贝数据到GPU缓存
        trt.allocator().copy_to_device(self.device_buf, batch.tobytes())
        return [self.device_buf]

    def read_calibration_cache(self):
        return None  # 首次生成无需读取缓存

    def write_calibration_cache(self, cache):
        # 保存校准表到文件
        with open("int8_calibration.table", "wb") as f:
            f.write(cache)
        return

# 1. 准备校准数据(同方法1,用真实业务数据)
calib_data = np.random.randn(100, 3, 224, 224).astype(np.float32)

# 2. 用TensorRT生成校准表
TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
with trt.Builder(TRT_LOGGER) as builder, builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) as network, trt.OnnxParser(network, TRT_LOGGER) as parser:
    # 加载ONNX模型
    with open("your_model.onnx", "rb") as f:
        parser.parse(f.read())
    
    # 配置Builder参数
    config = builder.create_builder_config()
    config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30)  # 分配1GB工作空间
    builder.int8_mode = True
    # 绑定校准器
    input_name = network.get_input(0).name
    calibrator = Int8Calibrator(calib_data, input_name, batch_size=8)
    config.int8_calibrator = calibrator
    
    # 构建引擎(此过程会自动运行校准并保存校准表)
    builder.build_engine(network, config)

# 3. 在ONNX Runtime中复用校准表
import onnxruntime as ort
sess_options = ort.SessionOptions()
providers = [
    ('TensorrtExecutionProvider', {
        'trt_int8_enable': True,
        'trt_int8_calibration_table_name': "int8_calibration.table",
        # 开启使用已生成的校准表
        'trt_int8_use_native_calibration_table': True
    }),
    'CPUExecutionProvider'
]
session = ort.InferenceSession("your_model.onnx", sess_options=sess_options, providers=providers)
# 后续即可用INT8精度进行推理

关键注意事项

  • 校准数据质量:必须用和实际推理场景一致的数据,随机数或无关数据会导致模型精度严重下降
  • 环境适配:Jetson上需确保ONNX Runtime与TensorRT版本匹配(优先使用JetPack自带的TensorRT版本)
  • 输入一致性:校准数据的形状、数据类型、归一化方式必须和推理阶段完全一致
  • 校准表复用:生成一次校准表后,后续推理可直接加载,无需重复校准

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 15:19:54