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

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。

可能的原因

  1. 优化策略默认变更:TensorRT 8.x版本在构建阶段默认启用了更激进的优化(如多层融合、精度校准临时内存分配),这些优化会增加构建时的GPU内存开销,而7.x版本默认优化程度较低。
  2. 显式批次模式处理调整:虽然代码中使用了EXPLICIT_BATCH,但8.x版本对显式批次的网络构建流程进行了重构,中间张量存储、优化搜索过程中的内存占用显著提升。
  3. ONNX解析器行为变化:8.x的ONNX解析器对模型的处理更严格,可能生成更多中间节点,或在解析阶段预分配更多GPU内存用于后续优化步骤。
  4. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 21:09:52