Skip to content

第四章 模型部署与优化


课程导航

项目内容
课时2 小时
类型理论 + 代码实战
前置知识PyTorch 模型训练、ONNX 基础概念

一、学习目标

  1. 掌握 PyTorch 模型导出为 ONNX 的完整流程
  2. 理解 TensorRT 的 FP16 / INT8 量化加速原理
  3. 掌握模型量化与剪枝的基本方法
  4. 了解端侧移动端部署框架(NCNN / TNN / MNN)
  5. 能搭建基于 FastAPI + Triton 的推理服务并进行性能压测

二、核心知识点

2.1 ONNX 导出与转换

2.1.1 定义

ONNX(Open Neural Network Exchange)是开放的模型交换格式,让模型在不同框架间自由迁移。

PyTorch(.pth) ──[torch.onnx.export]──→ ONNX(.onnx) ──[onnxruntime/TensorRT]→ 推理

为什么要用 ONNX?

优势说明
框架互操作性PyTorch ↔ TensorFlow ↔ MXNet 互通
推理优化ONNX Runtime 提供算子融合、图优化
硬件适配可转为 TensorRT / OpenVINO / CoreML 等

2.1.2 代码示例

python
import torch
import torch.onnx
import onnx
import onnxruntime as ort
import numpy as np

# ---------- 1. 定义或加载模型 ----------
class SimpleModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = torch.nn.Conv2d(3, 16, 3)
        self.relu = torch.nn.ReLU()
        self.pool = torch.nn.AdaptiveAvgPool2d((1,1))
        self.fc = torch.nn.Linear(16, 10)

    def forward(self, x):
        x = self.relu(self.conv(x))
        x = self.pool(x).flatten(1)
        return self.fc(x)

model = SimpleModel()
model.eval()

# ---------- 2. 构造 dummy 输入 ----------
dummy_input = torch.randn(1, 3, 224, 224)

# ---------- 3. 导出 ONNX ----------
torch.onnx.export(
    model,
    dummy_input,
    "model.onnx",
    input_names=['input'],
    output_names=['output'],
    dynamic_axes={
        'input': {0: 'batch_size'},    # 动态 batch
        'output': {0: 'batch_size'},
    },
    opset_version=17,                   # 算子集版本
    do_constant_folding=True,           # 常量折叠优化
)

# ---------- 4. 验证 ONNX 模型 ----------
onnx_model = onnx.load("model.onnx")
onnx.checker.check_model(onnx_model)
print(f"ONNX 模型验证通过!输入节点: {[n.name for n in onnx_model.graph.input]}")

# ---------- 5. ONNX Runtime 推理 ----------
ort_session = ort.InferenceSession("model.onnx")

# 准备输入
input_data = np.random.randn(1, 3, 224, 224).astype(np.float32)
ort_inputs = {ort_session.get_inputs()[0].name: input_data}

# 推理
ort_outputs = ort_session.run(None, ort_inputs)
print(f"推理结果 shape: {ort_outputs[0].shape}")

2.1.3 常见问题与排查

python
# ---------- 动态 batch 下的尺寸对齐 ----------
torch.onnx.export(
    model, dummy_input, "model.onnx",
    dynamic_axes={
        'input': {0: 'batch_size', 2: 'height', 3: 'width'},
    }
    # 注意:如果有 reshape/flatten 操作,动态尺寸可能导致失败
    # 建议使用 torch.onnx.export(..., verbose=True) 查看算子图
)

# ---------- 自定义算子导出 ----------
# 使用 torch.onnx.register_custom_op_symbolic 注册自定义算子

2.2 TensorRT 推理加速

2.2.1 定义

NVIDIA TensorRT 是面向 GPU 的高性能推理优化器,通过层融合、精度校准、内存优化等手段大幅提升推理速度。

优化方式加速比(相对 PyTorch)说明
FP32 优化1.5× - 2×算子融合 + 图优化
FP16 量化2× - 4×半精度推理,精度损失极小
INT8 量化3× - 6×需校准数据集校准缩放因子

2.2.2 ONNX → TensorRT 转换与推理

python
import tensorrt as trt
import numpy as np
import pycuda.driver as cuda
import pycuda.autoinit

# ---------- 构建 TRT 引擎 ----------
TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(TRT_LOGGER)
network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))
parser = trt.OnnxParser(network, TRT_LOGGER)

# 解析 ONNX 模型
with open("model.onnx", "rb") as f:
    if not parser.parse(f.read()):
        for err in range(parser.num_errors):
            print(parser.get_error(err))

# ---------- 配置构建选项 ----------
config = builder.create_builder_config()
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30)  # 1GB

# FP16 模式
if builder.platform_has_fast_fp16:
    config.set_flag(trt.BuilderFlag.FP16)
    print("✅ 已启用 FP16")

# INT8 模式(需提供校准数据集)
# if builder.platform_has_fast_int8:
#     config.set_flag(trt.BuilderFlag.INT8)
#     config.int8_calibrator = MyCalibrator()

# 构建引擎
serialized_engine = builder.build_serialized_network(network, config)
with open("model_fp16.engine", "wb") as f:
    f.write(serialized_engine)

# ---------- TensorRT 推理 ----------
runtime = trt.Runtime(TRT_LOGGER)
engine = runtime.deserialize_cuda_engine(serialized_engine)
context = engine.create_execution_context()

# 分配 GPU 内存
input_buf = cuda.mem_alloc(1 * 3 * 224 * 224 * 4)  # float32
output_buf = cuda.mem_alloc(1 * 10 * 4)

# 执行推理
input_data = np.random.randn(1, 3, 224, 224).astype(np.float32)
cuda.memcpy_htod(input_buf, input_data)
context.execute_v2([int(input_buf), int(output_buf)])
output_data = np.empty(10, dtype=np.float32)
cuda.memcpy_dtoh(output_data, output_buf)
print(f"TRT 推理结果: {output_data[:5]}")

2.2.3 TensorRT 性能压测

python
import time
import numpy as np

def benchmark_trt(engine_path, n_warmup=10, n_bench=100):
    runtime = trt.Runtime(TRT_LOGGER)
    with open(engine_path, "rb") as f:
        engine = runtime.deserialize_cuda_engine(f.read())
    context = engine.create_execution_context()

    # 分配固定输入
    dummy = np.random.randn(1, 3, 224, 224).astype(np.float32)

    # Warm-up
    for _ in range(n_warmup):
        cuda.memcpy_htod(input_buf, dummy)
        context.execute_v2([int(input_buf), int(output_buf)])

    # Benchmark
    start = time.perf_counter()
    for _ in range(n_bench):
        cuda.memcpy_htod(input_buf, dummy)
        context.execute_v2([int(input_buf), int(output_buf)])
    end = time.perf_counter()

    avg_ms = (end - start) / n_bench * 1000
    fps = 1000 / avg_ms
    print(f"⏱ 平均延迟: {avg_ms:.2f} ms | FPS: {fps:.1f}")

2.3 模型量化与剪枝

2.3.1 定义

模型量化降低参数精度(FP32→INT8),减少模型体积和计算量。模型剪枝移除冗余参数/通道,减小模型规模。

方法说明优点缺点
训练后量化 (PTQ)训练完后直接量化权重无需重新训练大模型精度可能下降
量化感知训练 (QAT)模拟量化效果训练精度更高需要训练数据和时间
结构化剪枝移除整个通道或层直接减小模型尺寸可能需要微调
非结构化剪枝移除单个权重压缩率高硬件加速困难

2.3.2 PyTorch 训练后量化

python
import torch
import torch.quantization as quant

# ---------- 准备模型 ----------
model = SimpleModel()
model.load_state_dict(torch.load('model.pth'))
model.eval()

# ---------- 训练后动态量化(适用于 LSTM / Linear 等)----------
quantized_model = torch.quantization.quantize_dynamic(
    model,
    {torch.nn.Linear, torch.nn.LSTM},  # 只量化这些层
    dtype=torch.qint8
)
quantized_model.save('model_quantized.pt')
print(f"动态量化完成")

# ---------- 训练后静态量化(适用于 CNN)----------
# 1. 将模型设置为融合模式
model.qconfig = torch.quantization.get_default_qconfig('fbgemm')  # x86
# model.qconfig = torch.quantization.get_default_qconfig('qnnpack')  # ARM

# 2. 融合算子
model_fused = torch.quantization.fuse_modules(model, [['conv', 'relu']])

# 3. 准备校准
model_prepared = quant.prepare(model_fused)

# 4. 校准(用代表性数据)
def calibrate(model, calib_loader, n_samples=100):
    model.eval()
    with torch.no_grad():
        for i, (data, _) in enumerate(calib_loader):
            if i * data.size(0) >= n_samples:
                break
            model(data)

calibrate(model_prepared, calib_loader)

# 5. 转换
model_int8 = quant.convert(model_prepared)
torch.jit.save(torch.jit.script(model_int8), 'model_int8.pt')

# ---------- 对比模型大小 ----------
import os
print(f"原始模型: {os.path.getsize('model.pth') / 1e6:.2f} MB")
print(f"量化模型: {os.path.getsize('model_int8.pt') / 1e6:.2f} MB")

2.3.3 结构化剪枝(torch-pruning 示例)

python
# pip install torch-pruning
import torch_pruning as tp

model = SimpleModel()

# 1. 分析依赖图
DG = tp.DependencyGraph().build_dependency(model, torch.randn(1, 3, 224, 224))

# 2. 剪枝策略:按 L1 范数裁剪 Conv2d 的 50% 通道
strategy = tp.strategy.L1Strategy()
pruning_idxs = strategy(model.conv, amount=0.5)  # 剪掉 50% 的通道

# 3. 获取依赖组并剪枝
plan = DG.get_pruning_plan(model.conv, tp.prune_conv, idxs=pruning_idxs)
plan.exec()

print(f"剪枝前: {tp.utils.count_parameters(SimpleModel()):,} 参数")
print(f"剪枝后: {tp.utils.count_parameters(model):,} 参数")

2.4 端侧移动端部署

2.4.1 定义

移动端部署框架将模型转换为手机/嵌入式设备可运行的格式:

框架出品方特点
NCNN腾讯轻量、无依赖、ARM 优化、Vulkan 支持
TNN腾讯跨平台、ARM/GPU/OpenCL 支持
MNN阿里巴巴轻量(~5MB)、iOS/Android/Harmony 支持

2.4.2 ONNX → NCNN 转换

bash
# 1. 下载 ncnn 和 tools
# git clone https://github.com/Tencent/ncnn.git

# 2. ONNX → NCNN
./onnx2ncnn model.onnx model.param model.bin

# 3. 优化(FP16 存储)
./ncnnoptimize model.param model.bin model_opt.param model.bin 1

# 4. 量化(INT8)
./ncnn2table model_opt.param model_opt.bin calibration_set.txt model.table
./ncnn2int8 model_opt.param model_opt.bin model_int8.param model_int8.bin model.table

2.4.3 NCNN Android 集成(C++ 推理代码)

cpp
// Android JNI 推理代码片段
#include <ncnn/net.h>

extern "C" JNIEXPORT jfloatArray JNICALL
Java_com_example_app_NCNNEngine_nativeInfer(
    JNIEnv* env, jobject thiz, jbyteArray input_data) {

    // 1. 加载模型
    ncnn::Net net;
    net.load_param("model.param");
    net.load_model("model.bin");

    // 2. 构造输入
    ncnn::Mat in = ncnn::Mat::from_pixels(
        input_ptr, ncnn::Mat::PIXEL_BGR, 224, 224);

    // 3. 推理
    ncnn::Extractor ex = net.create_extractor();
    ex.input("input", in);
    ncnn::Mat out;
    ex.extract("output", out);

    // 4. 返回结果
    // ...
}

2.4.4 MNN Python 推理

python
# pip install MNN
import MNN

# 加载模型
interpreter = MNN.Interpreter("model.mnn")
session = interpreter.createSession()

# 获取输入输出
input_tensor = interpreter.getSessionInput(session, "input")
output_tensor = interpreter.getSessionOutput(session, "output")

# 设置输入
tmp = np.random.randn(1, 3, 224, 224).astype(np.float32)
tmp_input = MNN.Tensor((1, 3, 224, 224), MNN.Halide_Type_Float, tmp, MNN.Tensor_DimensionType_Caffe)
input_tensor.copyFrom(tmp_input)

# 推理
interpreter.runSession(session)

# 读取输出
output_data = np.zeros([1, 10], dtype=np.float32)
tmp_output = MNN.Tensor((1, 10), MNN.Halide_Type_Float, output_data, MNN.Tensor_DimensionType_Caffe)
output_tensor.copyToHostTensor(tmp_output)
print(output_data)

2.5 服务化部署与性能压测

2.5.1 定义

服务化部署将模型封装为 RESTful API 或 gRPC 服务,供业务系统调用。

方案特点适用场景
FastAPI + ONNX Runtime轻量、快速开发中小规模、快速上线
Triton Inference Server企业级、多模型、动态批处理高并发生产环境

2.5.2 FastAPI 推理服务

python
# ---------- app.py ----------
from fastapi import FastAPI, File, UploadFile
import numpy as np
import cv2
import onnxruntime as ort
import uvicorn

app = FastAPI(title="CV 推理服务")

# 加载模型
ort_session = ort.InferenceSession("model.onnx")
input_name = ort_session.get_inputs()[0].name

def preprocess(image_bytes):
    nparr = np.frombuffer(image_bytes, np.uint8)
    img = cv2.imdecode(nparr, cv2.IMREAD_COLOR)
    img = cv2.resize(img, (224, 224))
    img = img.astype(np.float32) / 255.0
    img = img.transpose(2, 0, 1)    # HWC → CHW
    img = np.expand_dims(img, axis=0)
    return img

@app.post("/predict")
async def predict(file: UploadFile = File(...)):
    image_bytes = await file.read()
    input_data = preprocess(image_bytes)
    outputs = ort_session.run(None, {input_name: input_data})
    pred_class = int(np.argmax(outputs[0][0]))
    confidence = float(np.max(outputs[0][0]))
    return {"class_id": pred_class, "confidence": confidence}

@app.get("/health")
async def health():
    return {"status": "ok"}

# 启动
if __name__ == "__main__":
    uvicorn.run(app, host="0.0.0.0", port=8000)

2.5.3 性能压测(locust)

python
# ---------- locustfile.py ----------
from locust import HttpUser, task, between
import io
import numpy as np
from PIL import Image

class CVInferenceUser(HttpUser):
    wait_time = between(0.1, 0.5)  # 请求间隔

    def on_start(self):
        # 生成测试图片
        img = Image.fromarray(
            np.random.randint(0, 255, (224, 224, 3), dtype=np.uint8))
        self.img_bytes = io.BytesIO()
        img.save(self.img_bytes, format='JPEG')
        self.img_bytes = self.img_bytes.getvalue()

    @task
    def predict(self):
        self.client.post("/predict",
            files={"file": ("test.jpg", self.img_bytes, "image/jpeg")})

    @task(1)
    def health(self):
        self.client.get("/health")

# 运行: locust -f locustfile.py --host=http://localhost:8000

2.5.4 NVIDIA Triton 部署

python
# ---------- 1. 模型仓库结构 ----------
# model_repository/
# └── yolo_model/
#     ├── 1/
#     │   └── model.onnx
#     └── config.pbtxt

# ---------- config.pbtxt ----------
# name: "yolo_model"
# platform: "onnxruntime_onnx"
# max_batch_size: 32
# input [
#   {
#     name: "images"
#     data_type: TYPE_FP32
#     dims: [3, 640, 640]
#   }
# ]
# output [
#   {
#     name: "output0"
#     data_type: TYPE_FP32
#     dims: [84, 8400]
#   }
# ]
# dynamic_batching {
#   preferred_batch_size: [1, 4, 8, 16, 32]
#   max_queue_delay_microseconds: 200
# }

# ---------- 启动 Triton ----------
# docker run --gpus all -p 8000:8000 -p 8001:8001 \
#   -v $(pwd)/model_repository:/models \
#   nvcr.io/nvidia/tritonserver:23.10-py3 \
#   tritonserver --model-repository=/models

# ---------- Python 客户端 ----------
import tritonclient.http as httpclient
import numpy as np

client = httpclient.InferenceServerClient(url="localhost:8000")

input_data = np.random.randn(1, 3, 640, 640).astype(np.float32)
inputs = [httpclient.InferInput("images", input_data.shape, "FP32")]
inputs[0].set_data_from_numpy(input_data)

outputs = [httpclient.InferRequestedOutput("output0")]
result = client.infer("yolo_model", inputs, outputs)
print(result.as_numpy("output0").shape)

三、小结

知识点掌握程度关键工具
ONNX 导出必须掌握torch.onnx.export → ONNX Runtime
TensorRT 加速理解流程FP16 / INT8 量化,算子融合
模型量化剪枝理解原理PTQ / QAT / 结构化剪枝
端侧部署了解生态NCNN / TNN / MNN
服务化部署必须掌握FastAPI + ONNX Runtime / Triton

四、课后练习

基础题

  1. 将训练好的 YOLOv8 模型导出为 ONNX 格式,并用 ONNX Runtime 进行推理,对比 PyTorch 和 ONNX Runtime 的输出一致性。
  2. 使用 FastAPI 封装一个图像分类 API,支持批量图片上传,返回 JSON 格式结果。

进阶题

  1. 将 ONNX 模型转为 TensorRT FP16 引擎,编写基准测试对比 PyTorch / ONNX Runtime / TensorRT 三者的延迟和吞吐量。
  2. 对一个分类模型执行训练后 INT8 量化,记录量化前后的模型大小和精度变化。
  3. 使用 Locust 对 FastAPI 推理服务进行压测,找到服务的最大 QPS 和最佳并发数。

思考题

  1. ONNX 导出时 dynamic_axes 的作用是什么?为什么导出 Transformer 模型时动态轴更容易出错?
  2. INT8 量化相比 FP16 能实现更高的压缩率,但为什么实际部署中 INT8 的使用率不如 FP16 广泛?
  3. 端侧部署(手机/嵌入式)和云端部署在技术选型上分别要考虑哪些关键因素?

Python 学习资料