TensorRT8中ONNX模型GPU内存占用过高问题咨询
TensorRT版本升级后模型构建阶段GPU内存暴涨的原因与优化方案
问题场景
我使用以下代码构建TensorRT引擎:
import pycuda.driver as cuda import pycuda.autoinit import tensorrt as trt from cryptography.fernet import Fernet import zipfile import io import os TRT_LOGGER = trt.Logger(trt.Logger.WARNING) model_file = "models.onnx" key = "key" def load_model(fname): if fname.endswith(".uff"): return fname with open(fname, 'rb') as f: model = f.read() return model def load_compressed_model(fname, key, temp): fernet = Fernet(key) with open(fname, "rb") as f: content = fernet.decrypt(f.read()) with zipfile.ZipFile(io.BytesIO(content)) as zf: models = [n for n in zf.namelist()] assert len(models) == 1 temp.write(zf.read(models[0])) temp.flush() return load_model(temp.name) def build_engine(input_names, model): def gib(val): return val * 1 << 30 EXPLICIT_BATCH = 1 << (int)(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH) builder = trt.Builder(TRT_LOGGER) builder_config = builder.create_builder_config() builder_config.max_workspace_size = int(gib(1) * 0.15) builder_config.flags = trt.BuilderFlag.FP16 builder.max_batch_size = 1 print(EXPLICIT_BATCH) network = builder.create_network(EXPLICIT_BATCH) parser = trt.OnnxParser(network, TRT_LOGGER) parser.parse(model) return builder.build_engine(network, builder_config) from tempfile import NamedTemporaryFile with NamedTemporaryFile(suffix=os.path.splitext(model_file)[-1]) as temp: model = load_model(model_file) if 'raw_models' in model_file \ else load_compressed_model(model_file, key, temp) engine = build_engine(model_file, model)
环境差异
旧环境(
nvcr.io/nvidia/tensorrt:20.09-py3):nvcc -V nvcc: NVIDIA (R) Cuda compiler driver Copyright (c) 2005-2020 NVIDIA Corporation Built on Wed_Jul_22_19:09:09_PDT_2020 Cuda compilation tools, release 11.0, V11.0.221 Build cuda_11.0_bu.TC445_37.28845127_0 env | grep "CUDNN" CUDNN_VERSION=8.0.4.12 >>> import tensorrt >>> tensorrt.__version__ '7.1.2.8'构建阶段GPU内存占用约400MB,运行正常。
新环境(
nvcr.io/nvidia/tensorrt:22.12-py3):nvcc -V nvcc: NVIDIA (R) Cuda compiler driver Copyright (c) 2005-2022 NVIDIA Corporation Built on Wed_Sep_21_10:33:58_PDT_2022 Cuda compilation tools, release 11.8, V11.8.89 Build cuda_11.8.r11.8/compiler.31833905_0 env | grep "CUDNN" CUDNN_VERSION=8.7.0.84 >>> import tensorrt >>> tensorrt.__version__ '8.5.1.7'执行
builder.build_engine(***)时GPU内存占用飙升至约6000MB,调整builder_config.max_workspace_size无效果,模型未变更,NVIDIA驱动版本为550.78。
可能的原因
- 优化策略默认变更:TensorRT 8.x版本在构建阶段默认启用了更激进的优化(如多层融合、精度校准临时内存分配),这些优化会增加构建时的GPU内存开销,而7.x版本默认优化程度较低。
- 显式批次模式处理调整:虽然代码中使用了
EXPLICIT_BATCH,但8.x版本对显式批次的网络构建流程进行了重构,中间张量存储、优化搜索过程中的内存占用显著提升。 - ONNX解析器行为变化:8.x的ONNX解析器对模型的处理更严格,可能生成更多中间节点,或在解析阶段预分配更多GPU内存用于后续优化步骤。
- Builder配置隐式变更:比如FP16优化逻辑在8.x中更复杂,或默认开启了动态形状相关的内存预留(即使未显式配置动态形状)。
优化方案
1. 采用新版本推荐的内存限制方式
TensorRT 8.x推荐使用set_memory_pool_limit替代旧的max_workspace_size,同时可以限制临时内存池大小:
def build_engine(input_names, model): def gib(val): return val * 1 << 30 EXPLICIT_BATCH = 1 << (int)(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH) builder = trt.Builder(TRT_LOGGER) builder_config = builder.create_builder_config() # 设置工作区和临时内存池限制 builder_config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, int(gib(1)*0.15)) builder_config.set_memory_pool_limit(trt.MemoryPoolType.TEMPORARY, int(gib(1)*0.5)) builder_config.flags = trt.BuilderFlag.FP16 builder.max_batch_size = 1 network = builder.create_network(EXPLICIT_BATCH) parser = trt.OnnxParser(network, TRT_LOGGER) parser.parse(model) return builder.build_engine(network, builder_config)
2. 关闭不必要的优化选项
尝试禁用部分激进优化,降低构建时内存占用:
# 禁用时序缓存,减少内存占用 builder_config.set_flag(trt.BuilderFlag.DISABLE_TIMING_CACHE) # 暂时关闭FP16验证内存变化,确认后再重新开启 # builder_config.clear_flag(trt.BuilderFlag.FP16)
3. 序列化引擎复用
提前构建引擎并序列化保存,后续直接加载序列化文件,避免重复构建的内存开销:
# 构建后序列化引擎 engine = build_engine(model_file, model) with open("model_engine.trt", "wb") as f: f.write(engine.serialize()) # 后续加载引擎无需重新构建 runtime = trt.Runtime(TRT_LOGGER) with open("model_engine.trt", "rb") as f: engine = runtime.deserialize_cuda_engine(f.read())
4. 固定输入形状
显式固定输入张量形状,避免新版本默认的动态形状内存预留:
network = builder.create_network(EXPLICIT_BATCH) parser = trt.OnnxParser(network, TRT_LOGGER) parser.parse(model) # 根据你的模型实际输入形状设置 input_tensor = network.get_input(0) input_tensor.shape = (1, 3, 224, 224)
5. 启用严格类型约束
开启严格类型约束,减少中间张量的类型转换内存开销:
builder_config.set_flag(trt.BuilderFlag.STRICT_TYPE_CONSTRAINTS)
6. 调整ONNX解析器选项
通过解析器选项简化模型处理,减少中间节点内存占用:
parser = trt.OnnxParser(network, TRT_LOGGER) # 启用原生InstanceNormalization实现,减少内存开销 parser.set_flag(trt.OnnxParserFlag.NATIVE_INSTANCENORM) parser.parse(model)
内容的提问来源于stack exchange,提问作者mingz
相关产品推荐
相关产品推荐

