如何为ONNX Runtime的TensorRT执行提供程序创建INT8校准表?
实现ONNX模型INT8校准表生成(适配Jetson + ONNX Runtime/TensorRT)
方法1:直接通过ONNX Runtime生成校准表
ONNX Runtime的TensorRT后端支持自动完成INT8校准流程,只要提供代表性校准数据并开启对应配置即可:
操作步骤
- 配置ONNX Runtime的TensorRT执行器参数,开启INT8模式并指定校准表保存路径
- 准备一批和模型输入形状、 dtype、归一化逻辑完全一致的代表性校准数据(建议100-500条样本,避免用随机数)
- 运行一次推理触发校准流程,校准表会自动保存到指定路径
代码示例
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中加载使用:
操作步骤
- 定义TensorRT INT8校准器类,实现校准数据的批量读取逻辑
- 用TensorRT加载ONNX模型,开启INT8模式并绑定校准器,构建引擎时自动生成校准表
- 在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
相关产品推荐
相关产品推荐

