轻量化部署突破:PETRv2-BEV模型ONNX转换与优化全流程

最近在搞自动驾驶相关的项目,用到了PETRv2这个模型。说实话,这模型效果确实不错,但部署起来真是让人头疼。特别是想把PyTorch模型转成ONNX格式,中间各种坑,自定义算子、动态轴设置、模型简化,每一步都可能出问题。

今天我就把自己踩过的坑和解决方案整理出来,手把手教你如何把PETRv2模型顺利转成ONNX格式,并且能在不同推理引擎上跑起来。整个过程我尽量用大白话讲清楚,就算你之前没怎么接触过模型转换,跟着做也能搞定。

1. 环境准备与模型理解

在开始转换之前,得先把环境搭好,还得搞清楚PETRv2模型的结构特点。这模型是旷视孙剑团队搞出来的,专门用于自动驾驶的3D感知,能同时做目标检测和BEV分割。

1.1 环境配置

首先得把需要的包都装上。我建议用conda新建个环境,这样干净,不容易出问题。

# 创建新的conda环境
conda create -n petr-onnx python=3.8
conda activate petr-onnx

# 安装PyTorch(根据你的CUDA版本选择)
pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 --extra-index-url https://download.pytorch.org/whl/cu113

# 安装ONNX相关包
pip install onnx==1.13.0
pip install onnxruntime==1.14.0
pip install onnxruntime-gpu==1.14.0  # 如果需要GPU推理

# 安装PETRv2依赖
pip install mmcv-full==1.7.0 -f https://download.openmmlab.com/mmcv/dist/cu113/torch1.12/index.html
pip install mmdet==2.28.0
pip install mmsegmentation==0.29.0

# 安装其他工具
pip install onnx-simplifier==0.4.8
pip install onnxoptimizer==0.3.10

1.2 模型结构要点

PETRv2有几个关键特点,转换的时候要特别注意:

  1. 3D位置编码:这是PETR系列的核心,把2D图像特征和3D位置信息结合起来
  2. 时序融合:能利用前后帧的信息,提升检测效果
  3. 多任务输出:同时输出3D检测框和BEV分割结果
  4. 自定义算子:模型里有些特殊的操作,转ONNX时需要特殊处理

我建议你先下载官方的预训练模型,然后加载看看结构。可以从GitHub上找到PETRv2的代码仓库,里面有详细的配置文件和模型权重。

2. 基础转换:从PyTorch到ONNX

这一步是最基础的,但也是最容易出问题的。我刚开始转的时候,各种错误报个不停,后来才发现是有些细节没注意到。

2.1 准备转换脚本

先写个简单的转换脚本,把模型加载进来,看看能不能正常推理。

import torch
import onnx
from mmdet.apis import init_detector

def load_petrv2_model(config_path, checkpoint_path):
    """加载PETRv2模型"""
    print("正在加载模型...")
    
    # 初始化模型
    model = init_detector(config_path, checkpoint_path, device='cuda:0')
    
    # 设置为评估模式
    model.eval()
    
    # 打印模型结构信息
    print(f"模型加载成功!")
    print(f"模型类型: {type(model)}")
    print(f"模型设备: {next(model.parameters()).device}")
    
    return model

def create_dummy_input():
    """创建虚拟输入数据"""
    # PETRv2通常需要6个摄像头的图像
    batch_size = 1
    num_cams = 6
    img_h, img_w = 256, 704  # 根据你的配置调整
    
    # 图像数据
    imgs = torch.randn(batch_size, num_cams, 3, img_h, img_w).cuda()
    
    # 相机参数(内参、外参)
    # 这里简化处理,实际使用时需要根据你的数据格式调整
    cam_intrinsics = torch.randn(batch_size, num_cams, 3, 3).cuda()
    cam_extrinsics = torch.randn(batch_size, num_cams, 4, 4).cuda()
    
    # 时间戳(用于时序融合)
    timestamp = torch.tensor([0.0]).cuda()
    
    return {
        'img': imgs,
        'cam_intrinsics': cam_intrinsics,
        'cam_extrinsics': cam_extrinsics,
        'timestamp': timestamp
    }

def test_model_inference(model, dummy_input):
    """测试模型推理"""
    print("\n测试模型推理...")
    
    with torch.no_grad():
        try:
            # 前向传播
            outputs = model(return_loss=False, rescale=True, **dummy_input)
            
            print("推理成功!")
            print(f"输出类型: {type(outputs)}")
            
            # 打印输出结构
            if isinstance(outputs, list):
                print(f"输出列表长度: {len(outputs)}")
                for i, out in enumerate(outputs):
                    print(f"  输出{i}: {type(out)}")
                    if hasattr(out, 'shape'):
                        print(f"    形状: {out.shape}")
            elif isinstance(outputs, dict):
                print(f"输出字典键: {list(outputs.keys())}")
                for k, v in outputs.items():
                    print(f"  {k}: {type(v)}")
                    if hasattr(v, 'shape'):
                        print(f"    形状: {v.shape}")
            
            return True
        except Exception as e:
            print(f"推理失败: {e}")
            return False

if __name__ == "__main__":
    # 配置文件路径
    config_path = "configs/petr/petrv2_r50_704x256.py"
    checkpoint_path = "checkpoints/petrv2_r50.pth"
    
    # 加载模型
    model = load_petrv2_model(config_path, checkpoint_path)
    
    # 创建虚拟输入
    dummy_input = create_dummy_input()
    
    # 测试推理
    test_model_inference(model, dummy_input)

2.2 第一次转换尝试

模型能正常推理后,就可以尝试转ONNX了。第一次转的时候,先别加太多复杂参数,看看基础转换能不能成功。

def export_to_onnx_basic(model, dummy_input, onnx_path):
    """基础ONNX导出"""
    print(f"\n开始导出ONNX到: {onnx_path}")
    
    # 定义前向函数
    def forward_func(img, cam_intrinsics, cam_extrinsics, timestamp):
        return model(
            return_loss=False,
            rescale=True,
            img=img,
            cam_intrinsics=cam_intrinsics,
            cam_extrinsics=cam_extrinsics,
            timestamp=timestamp
        )
    
    # 设置动态轴
    dynamic_axes = {
        'img': {0: 'batch_size', 2: 'height', 3: 'width'},
        'cam_intrinsics': {0: 'batch_size'},
        'cam_extrinsics': {0: 'batch_size'},
        'output': {0: 'batch_size'}
    }
    
    try:
        torch.onnx.export(
            model,
            (dummy_input['img'], 
             dummy_input['cam_intrinsics'],
             dummy_input['cam_extrinsics'],
             dummy_input['timestamp']),
            onnx_path,
            input_names=['img', 'cam_intrinsics', 'cam_extrinsics', 'timestamp'],
            output_names=['output'],
            dynamic_axes=dynamic_axes,
            opset_version=13,  # 使用较新的opset
            do_constant_folding=True,
            verbose=True
        )
        
        print("ONNX导出成功!")
        
        # 验证导出的ONNX模型
        onnx_model = onnx.load(onnx_path)
        onnx.checker.check_model(onnx_model)
        print("ONNX模型验证通过!")
        
        return True
        
    except Exception as e:
        print(f"ONNX导出失败: {e}")
        import traceback
        traceback.print_exc()
        return False

# 使用示例
onnx_path = "petrv2_basic.onnx"
export_to_onnx_basic(model, dummy_input, onnx_path)

3. 处理自定义算子

PETRv2里有些自定义算子,直接转ONNX会报错。我遇到的第一个大坑就是这个,报错信息看得一头雾水。后来才发现,得把这些特殊操作替换成ONNX支持的标准操作。

3.1 识别自定义算子

先看看模型里有哪些ONNX不支持的算子:

def analyze_model_operations(model, dummy_input):
    """分析模型中的操作"""
    print("\n分析模型操作...")
    
    # 使用torch.jit.trace获取计算图
    try:
        traced_model = torch.jit.trace(
            model,
            example_inputs=(dummy_input['img'], 
                           dummy_input['cam_intrinsics'],
                           dummy_input['cam_extrinsics'],
                           dummy_input['timestamp'])
        )
        
        # 获取计算图
        graph = traced_model.graph
        print("计算图获取成功")
        
        # 统计操作类型
        op_counter = {}
        for node in graph.nodes():
            op_name = node.kind()
            op_counter[op_name] = op_counter.get(op_name, 0) + 1
        
        print("\n操作统计:")
        for op, count in sorted(op_counter.items(), key=lambda x: x[1], reverse=True):
            print(f"  {op}: {count}")
            
    except Exception as e:
        print(f"分析失败: {e}")
        # 尝试其他方法
        print("\n尝试直接分析模型结构...")
        print(f"模型类: {model.__class__.__name__}")
        print(f"模块数量: {len(list(model.modules()))}")

3.2 常见自定义算子处理

PETRv2里常见的自定义算子包括:

  1. 3D位置编码生成:需要拆解成标准操作
  2. BEV特征生成:涉及特殊的坐标变换
  3. 多尺度特征融合:有特殊的采样方式

这里我提供一个处理3D位置编码的例子:

class PETRv2ONNXWrapper(torch.nn.Module):
    """包装PETRv2模型,处理自定义算子"""
    
    def __init__(self, original_model):
        super().__init__()
        self.model = original_model
        
    def generate_3d_coords(self, cam_intrinsics, cam_extrinsics, depth_num=64):
        """生成3D坐标(ONNX兼容版本)"""
        batch_size, num_cams = cam_intrinsics.shape[:2]
        
        # 生成视锥网格点
        # 原始实现可能用了一些特殊操作,这里用标准操作替换
        h, w = 256, 704  # 特征图大小
        u = torch.arange(w, device=cam_intrinsics.device).float()
        v = torch.arange(h, device=cam_intrinsics.device).float()
        
        # 创建网格
        grid_u, grid_v = torch.meshgrid(u, v, indexing='xy')
        grid_u = grid_u.reshape(1, 1, h, w)
        grid_v = grid_v.reshape(1, 1, h, w)
        
        # 扩展维度
        grid_u = grid_u.expand(batch_size, num_cams, h, w)
        grid_v = grid_v.expand(batch_size, num_cams, h, w)
        
        # 生成深度采样
        depth = torch.linspace(0.1, 60.0, depth_num, 
                              device=cam_intrinsics.device)
        depth = depth.view(1, 1, depth_num, 1, 1)
        depth = depth.expand(batch_size, num_cams, depth_num, h, w)
        
        # 这里简化处理,实际需要根据相机参数进行坐标变换
        # 原始代码可能用了mmcv的Camera2Lidar操作,需要拆解
        
        return {
            'grid_u': grid_u,
            'grid_v': grid_v,
            'depth': depth
        }
    
    def forward(self, img, cam_intrinsics, cam_extrinsics, timestamp):
        """前向传播"""
        # 生成3D坐标
        coords_3d = self.generate_3d_coords(cam_intrinsics, cam_extrinsics)
        
        # 调用原始模型,但使用我们处理过的中间结果
        # 注意:这里需要根据实际模型结构调整
        outputs = self.model(
            return_loss=False,
            rescale=True,
            img=img,
            cam_intrinsics=cam_intrinsics,
            cam_extrinsics=cam_extrinsics,
            timestamp=timestamp,
            # 可以传递额外的参数
            coords_3d=coords_3d
        )
        
        return outputs

# 使用包装器
wrapped_model = PETRv2ONNXWrapper(model)
wrapped_model.eval()

# 测试包装后的模型
test_model_inference(wrapped_model, dummy_input)

4. 动态轴设置与模型简化

模型能转成ONNX后,还得让它能在不同输入尺寸下工作,这就是动态轴的作用。另外,原始ONNX模型可能有很多冗余,需要简化。

4.1 精细化的动态轴设置

def export_to_onnx_advanced(model, dummy_input, onnx_path):
    """高级ONNX导出,支持动态轴"""
    print(f"\n开始高级ONNX导出到: {onnx_path}")
    
    # 更精细的动态轴设置
    dynamic_axes = {
        # 输入动态轴
        'img': {
            0: 'batch_size',
            1: 'num_cams',  # 摄像头数量
            3: 'height',    # 图像高度
            4: 'width'      # 图像宽度
        },
        'cam_intrinsics': {
            0: 'batch_size',
            1: 'num_cams'
        },
        'cam_extrinsics': {
            0: 'batch_size',
            1: 'num_cams'
        },
        'timestamp': {0: 'batch_size'},
        
        # 输出动态轴(根据实际输出调整)
        'bboxes': {0: 'num_bboxes'},
        'scores': {0: 'num_bboxes'},
        'labels': {0: 'num_bboxes'},
        'bev_seg': {2: 'bev_height', 3: 'bev_width'}
    }
    
    try:
        # 如果模型输出是字典或列表,需要特殊处理
        class ModelExporter(torch.nn.Module):
            def __init__(self, model):
                super().__init__()
                self.model = model
            
            def forward(self, img, cam_intrinsics, cam_extrinsics, timestamp):
                outputs = self.model(
                    return_loss=False,
                    rescale=True,
                    img=img,
                    cam_intrinsics=cam_intrinsics,
                    cam_extrinsics=cam_extrinsics,
                    timestamp=timestamp
                )
                
                # 将输出转换为ONNX友好的格式
                # 这里假设输出是字典,包含bboxes、scores、labels、bev_seg
                if isinstance(outputs, dict):
                    return (
                        outputs.get('bboxes', torch.tensor([])),
                        outputs.get('scores', torch.tensor([])),
                        outputs.get('labels', torch.tensor([])),
                        outputs.get('bev_seg', torch.tensor([]))
                    )
                else:
                    # 如果是其他格式,根据实际情况调整
                    return outputs
        
        # 创建导出器
        exporter = ModelExporter(model)
        exporter.eval()
        
        # 导出
        torch.onnx.export(
            exporter,
            (dummy_input['img'], 
             dummy_input['cam_intrinsics'],
             dummy_input['cam_extrinsics'],
             dummy_input['timestamp']),
            onnx_path,
            input_names=['img', 'cam_intrinsics', 'cam_extrinsics', 'timestamp'],
            output_names=['bboxes', 'scores', 'labels', 'bev_seg'],
            dynamic_axes=dynamic_axes,
            opset_version=13,
            do_constant_folding=True,
            keep_initializers_as_inputs=True,  # 重要:保持初始化为输入
            verbose=False
        )
        
        print("高级ONNX导出成功!")
        return True
        
    except Exception as e:
        print(f"高级ONNX导出失败: {e}")
        import traceback
        traceback.print_exc()
        return False

# 使用高级导出
onnx_advanced_path = "petrv2_advanced.onnx"
export_to_onnx_advanced(model, dummy_input, onnx_advanced_path)

4.2 模型简化与优化

导出的ONNX模型可能有很多冗余节点,用onnx-simplifier简化一下:

def simplify_onnx_model(input_onnx_path, output_onnx_path):
    """简化ONNX模型"""
    print(f"\n简化ONNX模型: {input_onnx_path} -> {output_onnx_path}")
    
    try:
        import onnxsim
        
        # 加载模型
        model = onnx.load(input_onnx_path)
        
        # 简化
        model_simp, check = onnxsim.simplify(
            model,
            input_shapes={
                'img': [1, 6, 3, 256, 704],
                'cam_intrinsics': [1, 6, 3, 3],
                'cam_extrinsics': [1, 6, 4, 4],
                'timestamp': [1]
            },
            dynamic_input_shape=True  # 保持动态输入
        )
        
        if check:
            # 保存简化后的模型
            onnx.save(model_simp, output_onnx_path)
            print(f"模型简化成功!保存到: {output_onnx_path}")
            
            # 验证简化后的模型
            onnx.checker.check_model(model_simp)
            print("简化模型验证通过!")
            
            # 统计简化效果
            orig_nodes = len(model.graph.node)
            simp_nodes = len(model_simp.graph.node)
            reduction = (orig_nodes - simp_nodes) / orig_nodes * 100
            
            print(f"节点数: {orig_nodes} -> {simp_nodes} (减少{reduction:.1f}%)")
            
            return True
        else:
            print("模型简化检查失败")
            return False
            
    except Exception as e:
        print(f"模型简化失败: {e}")
        return False

# 进一步优化
def optimize_onnx_model(input_onnx_path, output_onnx_path):
    """优化ONNX模型"""
    print(f"\n优化ONNX模型: {input_onnx_path} -> {output_onnx_path}")
    
    try:
        import onnxoptimizer
        
        # 加载模型
        model = onnx.load(input_onnx_path)
        
        # 定义优化passes
        passes = [
            'eliminate_deadend',
            'eliminate_identity',
            'eliminate_nop_dropout',
            'eliminate_nop_cast',
            'eliminate_nop_monotone_argmax',
            'eliminate_unused_initializer',
            'extract_constant_to_initializer',
            'fuse_add_bias_into_conv',
            'fuse_bn_into_conv',
            'fuse_consecutive_concats',
            'fuse_consecutive_reduce_unsqueeze',
            'fuse_consecutive_squeezes',
            'fuse_consecutive_transposes',
            'fuse_matmul_add_bias_into_gemm',
            'fuse_pad_into_conv',
            'fuse_transpose_into_gemm',
        ]
        
        # 优化
        optimized_model = onnxoptimizer.optimize(model, passes)
        
        # 保存优化后的模型
        onnx.save(optimized_model, output_onnx_path)
        print(f"模型优化成功!保存到: {output_onnx_path}")
        
        return True
        
    except Exception as e:
        print(f"模型优化失败: {e}")
        return False

# 执行简化和优化
onnx_simp_path = "petrv2_simplified.onnx"
onnx_opt_path = "petrv2_optimized.onnx"

if simplify_onnx_model(onnx_advanced_path, onnx_simp_path):
    optimize_onnx_model(onnx_simp_path, onnx_opt_path)

5. 推理引擎兼容性测试

模型转好了,还得看看在不同推理引擎上能不能跑起来。我测试了ONNX Runtime、TensorRT和OpenVINO,每个都有需要注意的地方。

5.1 ONNX Runtime测试

def test_onnxruntime(onnx_path, dummy_input):
    """测试ONNX Runtime推理"""
    print(f"\n测试ONNX Runtime: {onnx_path}")
    
    import onnxruntime as ort
    import numpy as np
    
    # 创建ONNX Runtime会话
    providers = ['CUDAExecutionProvider', 'CPUExecutionProvider']
    session = ort.InferenceSession(onnx_path, providers=providers)
    
    # 准备输入数据(转成numpy)
    input_feed = {
        'img': dummy_input['img'].cpu().numpy(),
        'cam_intrinsics': dummy_input['cam_intrinsics'].cpu().numpy(),
        'cam_extrinsics': dummy_input['cam_extrinsics'].cpu().numpy(),
        'timestamp': dummy_input['timestamp'].cpu().numpy()
    }
    
    # 推理
    try:
        print("开始ONNX Runtime推理...")
        
        # 预热
        for _ in range(5):
            outputs = session.run(None, input_feed)
        
        # 正式测试
        import time
        start_time = time.time()
        
        num_runs = 100
        for i in range(num_runs):
            outputs = session.run(None, input_feed)
            if i % 20 == 0:
                print(f"  已完成 {i+1}/{num_runs} 次推理")
        
        end_time = time.time()
        avg_time = (end_time - start_time) / num_runs * 1000  # 毫秒
        
        print(f"推理完成!平均耗时: {avg_time:.2f} ms")
        
        # 打印输出信息
        print(f"\n输出数量: {len(outputs)}")
        for i, output in enumerate(outputs):
            print(f"  输出{i}: 形状={output.shape}, 类型={output.dtype}")
        
        return True, avg_time
        
    except Exception as e:
        print(f"ONNX Runtime推理失败: {e}")
        return False, None

# 测试不同精度的模型
def test_different_precisions():
    """测试不同精度模型"""
    print("\n测试不同精度模型...")
    
    # 测试FP32
    print("\n1. 测试FP32模型:")
    success_fp32, time_fp32 = test_onnxruntime("petrv2_fp32.onnx", dummy_input)
    
    # 测试FP16(如果有的话)
    try:
        print("\n2. 测试FP16模型:")
        success_fp16, time_fp16 = test_onnxruntime("petrv2_fp16.onnx", dummy_input)
        
        if success_fp16 and success_fp32:
            speedup = time_fp32 / time_fp16
            print(f"FP16相对于FP32加速: {speedup:.2f}x")
    except:
        print("FP16模型测试跳过")
    
    return success_fp32

# 运行测试
test_different_precisions()

5.2 TensorRT测试

def test_tensorrt(onnx_path, dummy_input):
    """测试TensorRT推理"""
    print(f"\n测试TensorRT: {onnx_path}")
    
    try:
        import tensorrt as trt
        
        # 创建TensorRT日志记录器
        TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
        
        # 创建构建器
        builder = trt.Builder(TRT_LOGGER)
        
        # 创建网络定义
        network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))
        
        # 创建ONNX解析器
        parser = trt.OnnxParser(network, TRT_LOGGER)
        
        # 解析ONNX模型
        with open(onnx_path, 'rb') as model:
            if not parser.parse(model.read()):
                print("ONNX解析失败:")
                for error in range(parser.num_errors):
                    print(parser.get_error(error))
                return False, None
        
        print("ONNX解析成功!")
        
        # 配置构建器
        config = builder.create_builder_config()
        config.max_workspace_size = 1 << 30  # 1GB
        
        # 设置精度(可以根据需要调整)
        if builder.platform_has_fast_fp16:
            config.set_flag(trt.BuilderFlag.FP16)
            print("启用FP16精度")
        
        # 构建引擎
        print("正在构建TensorRT引擎...")
        engine = builder.build_engine(network, config)
        
        if engine is None:
            print("引擎构建失败")
            return False, None
        
        print(f"TensorRT引擎构建成功!")
        print(f"绑定数量: {engine.num_bindings}")
        
        # 这里可以继续实现推理逻辑
        # 由于代码较长,实际使用时需要根据TensorRT API实现
        
        return True, None
        
    except ImportError:
        print("TensorRT未安装,跳过测试")
        return False, None
    except Exception as e:
        print(f"TensorRT测试失败: {e}")
        return False, None

# 测试TensorRT
test_tensorrt(onnx_opt_path, dummy_input)

5.3 多引擎兼容性检查

def check_onnx_compatibility(onnx_path):
    """检查ONNX模型兼容性"""
    print(f"\n检查ONNX模型兼容性: {onnx_path}")
    
    import onnx
    from onnx import helper
    
    model = onnx.load(onnx_path)
    
    # 检查opset版本
    print(f"ONNX opset版本: {model.opset_import[0].version}")
    
    # 检查使用的算子
    op_counter = {}
    for node in model.graph.node:
        op_type = node.op_type
        op_counter[op_type] = op_counter.get(op_type, 0) + 1
    
    print("\n算子使用统计:")
    for op, count in sorted(op_counter.items(), key=lambda x: x[1], reverse=True)[:20]:
        print(f"  {op}: {count}")
    
    # 检查不支持的算子
    unsupported_ops = []
    common_ops = {
        'Conv', 'Relu', 'Add', 'Mul', 'Div', 'Sub', 'Pow', 'Sqrt',
        'Exp', 'Log', 'Sin', 'Cos', 'Tanh', 'Sigmoid', 'Softmax',
        'MatMul', 'Gemm', 'BatchNormalization', 'Dropout',
        'Reshape', 'Transpose', 'Concat', 'Split', 'Slice',
        'Gather', 'Unsqueeze', 'Squeeze', 'Flatten',
        'Constant', 'Shape', 'Cast', 'Where', 'Clip',
        'ReduceMean', 'ReduceSum', 'ReduceMax', 'ReduceMin',
        'ArgMax', 'ArgMin', 'TopK', 'NonZero',
        'Scatter', 'ScatterND', 'Tile', 'Pad',
        'ConvTranspose', 'MaxPool', 'AveragePool', 'GlobalAveragePool',
        'LSTM', 'GRU', 'RNN'
    }
    
    for op in op_counter.keys():
        if op not in common_ops:
            unsupported_ops.append(op)
    
    if unsupported_ops:
        print(f"\n  发现非常用算子: {unsupported_ops}")
        print("这些算子在某些推理引擎上可能不支持")
    else:
        print("\n 所有算子都是常用算子,兼容性良好")
    
    # 检查动态轴
    print("\n输入动态轴信息:")
    for input in model.graph.input:
        print(f"  {input.name}:")
        for dim in input.type.tensor_type.shape.dim:
            if dim.dim_param:
                print(f"    {dim.dim_param} (动态)")
            else:
                print(f"    {dim.dim_value} (固定)")
    
    return len(unsupported_ops) == 0

# 检查兼容性
check_onnx_compatibility(onnx_opt_path)

6. 实际部署建议

根据我的经验,这里给几个实际部署的建议:

6.1 性能优化建议

  1. 使用FP16精度:如果硬件支持,FP16能显著提升推理速度,内存占用也少一半
  2. 批处理优化:合理设置batch size,太小浪费计算资源,太大可能内存不够
  3. 使用TensorRT:如果部署在NVIDIA硬件上,TensorRT通常比ONNX Runtime更快
  4. 模型剪枝:如果对精度要求不是极致,可以适当剪枝,进一步提升速度

6.2 内存优化建议

  1. 分阶段加载:大模型可以分成多个部分加载,减少峰值内存
  2. 使用内存池:重复利用内存,避免频繁分配释放
  3. 梯度检查点:训练时用,用时间换空间

6.3 部署 checklist

这里我总结了一个部署检查清单,你可以对照着检查:

def deployment_checklist(onnx_path):
    """部署检查清单"""
    print("\n" + "="*50)
    print("部署检查清单")
    print("="*50)
    
    checks = []
    
    # 1. 模型文件检查
    try:
        import os
        file_size = os.path.getsize(onnx_path) / (1024**2)  # MB
        checks.append(("模型文件存在", True))
        checks.append((f"模型大小: {file_size:.1f}MB", file_size < 500))  # 假设500MB为合理大小
    except:
        checks.append(("模型文件存在", False))
    
    # 2. ONNX验证
    try:
        import onnx
        model = onnx.load(onnx_path)
        onnx.checker.check_model(model)
        checks.append(("ONNX模型验证", True))
    except:
        checks.append(("ONNX模型验证", False))
    
    # 3. 算子兼容性
    try:
        from onnx import helper
        opset_version = model.opset_import[0].version
        checks.append((f"ONNX opset版本: {opset_version}", opset_version >= 11))
    except:
        checks.append(("获取opset版本", False))
    
    # 4. 动态轴检查
    try:
        has_dynamic = False
        for input in model.graph.input:
            for dim in input.type.tensor_type.shape.dim:
                if dim.dim_param:
                    has_dynamic = True
                    break
        checks.append(("支持动态输入", has_dynamic))
    except:
        checks.append(("动态轴检查", False))
    
    # 5. 测试推理
    try:
        import onnxruntime as ort
        import numpy as np
        
        # 创建简单输入
        dummy_input = {
            'img': np.random.randn(1, 6, 3, 256, 704).astype(np.float32),
            'cam_intrinsics': np.random.randn(1, 6, 3, 3).astype(np.float32),
            'cam_extrinsics': np.random.randn(1, 6, 4, 4).astype(np.float32),
            'timestamp': np.array([0.0], dtype=np.float32)
        }
        
        providers = ['CPUExecutionProvider']
        session = ort.InferenceSession(onnx_path, providers=providers)
        
        # 单次推理
        outputs = session.run(None, dummy_input)
        checks.append(("ONNX Runtime推理", True))
        
    except Exception as e:
        checks.append((f"ONNX Runtime推理: {str(e)[:50]}...", False))
    
    # 打印检查结果
    print("\n检查结果:")
    all_pass = True
    for check_name, check_result in checks:
        status = " PASS" if check_result else " FAIL"
        print(f"  {status}: {check_name}")
        if not check_result:
            all_pass = False
    
    print("\n" + "="*50)
    if all_pass:
        print(" 所有检查通过,可以部署!")
    else:
        print("  部分检查未通过,请根据上述提示修复")
    print("="*50)
    
    return all_pass

# 运行检查清单
deployment_checklist(onnx_opt_path)

7. 总结

整个流程走下来,PETRv2转ONNX确实比普通模型复杂一些,主要是因为有自定义算子和复杂的结构。但按照上面这些步骤,一步步来,基本上都能搞定。

关键是要理解模型的结构,知道哪些地方需要特殊处理。动态轴设置也很重要,这决定了模型能不能适应不同的输入尺寸。最后一定要在各个推理引擎上测试,确保真的能用。

我把自己遇到的坑和解决方案都写出来了,你应该能少走很多弯路。实际做的时候,可能还会遇到一些具体问题,但大方向就是这些。多试试,多调调,总能解决的。


获取更多AI镜像

想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐