Appearance
第四章 模型部署与优化
课程导航
| 项目 | 内容 |
|---|---|
| 课时 | 2 小时 |
| 类型 | 理论 + 代码实战 |
| 前置知识 | PyTorch 模型训练、ONNX 基础概念 |
一、学习目标
- 掌握 PyTorch 模型导出为 ONNX 的完整流程
- 理解 TensorRT 的 FP16 / INT8 量化加速原理
- 掌握模型量化与剪枝的基本方法
- 了解端侧移动端部署框架(NCNN / TNN / MNN)
- 能搭建基于 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.table2.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:80002.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 |
四、课后练习
基础题
- 将训练好的 YOLOv8 模型导出为 ONNX 格式,并用 ONNX Runtime 进行推理,对比 PyTorch 和 ONNX Runtime 的输出一致性。
- 使用 FastAPI 封装一个图像分类 API,支持批量图片上传,返回 JSON 格式结果。
进阶题
- 将 ONNX 模型转为 TensorRT FP16 引擎,编写基准测试对比 PyTorch / ONNX Runtime / TensorRT 三者的延迟和吞吐量。
- 对一个分类模型执行训练后 INT8 量化,记录量化前后的模型大小和精度变化。
- 使用 Locust 对 FastAPI 推理服务进行压测,找到服务的最大 QPS 和最佳并发数。
思考题
- ONNX 导出时
dynamic_axes的作用是什么?为什么导出 Transformer 模型时动态轴更容易出错? - INT8 量化相比 FP16 能实现更高的压缩率,但为什么实际部署中 INT8 的使用率不如 FP16 广泛?
- 端侧部署(手机/嵌入式)和云端部署在技术选型上分别要考虑哪些关键因素?