训练好的模型只有部署到生产环境才能创造价值。本文覆盖模型格式转换、服务化架构、推理加速(量化、剪枝、蒸馏)与边缘端部署的完整链路。
1. 模型持久化与格式
1.1 PyTorch 模型保存
import torch
# 保存完整模型(不推荐,依赖环境)
torch.save(model, 'model.pt')
model = torch.load('model.pt')
# 保存状态字典(推荐,跨版本兼容)
torch.save(model.state_dict(), 'model_weights.pth')
# 加载
model = MyModel()
model.load_state_dict(torch.load('model_weights.pth'))
model.eval()
# 保存训练状态(含优化器、epoch)
checkpoint = {
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'loss': loss,
}
torch.save(checkpoint, 'checkpoint.tar')
1.2 Pickle vs Safetensors
# Safetensors:安全、快速、无代码执行风险
from safetensors.torch import save_file, load_file
save_file(model.state_dict(), 'model.safetensors')
state_dict = load_file('model.safetensors')
2. ONNX 跨平台部署
2.1 PyTorch → ONNX
ONNX (Open Neural Network Exchange) 是跨框架的中间表示格式。
import torch
# 导出 ONNX
model.eval()
dummy_input = torch.randn(1, 3, 224, 224).to('cuda')
torch.onnx.export(
model,
dummy_input,
'model.onnx',
export_params=True, # 存储训练好的参数
opset_version=17, # ONNX 算子集版本
do_constant_folding=True, # 常量折叠优化
input_names=['input'], # 输入名
output_names=['output'], # 输出名
dynamic_axes={ # 动态维度(batch size)
'input': {0: 'batch_size'},
'output': {0: 'batch_size'}
}
)
2.2 ONNX Runtime 推理
import onnxruntime as ort
import numpy as np
# 创建 InferenceSession
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider']
session = ort.InferenceSession('model.onnx', providers=providers)
# 查看输入输出信息
for inp in session.get_inputs():
print(f"Input: {inp.name}, Shape: {inp.shape}, Type: {inp.type}")
# 推理
input_name = session.get_inputs()[0].name
output_name = session.get_outputs()[0].name
x = np.random.randn(1, 3, 224, 224).astype(np.float32)
outputs = session.run([output_name], {input_name: x})
# 批量推理
x_batch = np.random.randn(32, 3, 224, 224).astype(np.float32)
outputs = session.run(None, {input_name: x_batch})
2.3 ONNX 优化
import onnx
from onnxruntime.tools import optimizer
# 加载并检查
model = onnx.load('model.onnx')
onnx.checker.check_model(model)
# 图优化
optimized_model = optimizer.optimize_model(
'model.onnx',
model_type='bert', # 或 'bert', 'gpt2', 'bart'
num_heads=12,
hidden_size=768,
use_gpu=True
)
optimized_model.convert_model_float32_to_float16()
optimized_model.save_model_to_file('model_optimized.onnx')
3. TorchServe 服务化
3.1 搭建模型服务
# model_handler.py
from ts.torch_handler.base_handler import BaseHandler
import torch
import json
import os
class CustomHandler(BaseHandler):
def initialize(self, context):
self.manifest = context.manifest
properties = context.system_properties
model_dir = properties.get("model_dir")
# 加载模型
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.model = torch.jit.load(os.path.join(model_dir, "model.pt"))
self.model.to(self.device)
self.model.eval()
# 加载预处理器
self.transform = ...
def preprocess(self, data):
images = []
for row in data:
image = row.get("data") or row.get("body")
image = Image.open(io.BytesIO(image))
image = self.transform(image)
images.append(image)
return torch.stack(images).to(self.device)
def inference(self, data):
with torch.no_grad():
return self.model(data)
def postprocess(self, inference_output):
probs = torch.softmax(inference_output, dim=1)
return json.dumps([{"class": p.argmax().item(),
"confidence": p.max().item()} for p in probs])
3.2 打包与部署
# 1. 创建 MAR 模型归档
torch-model-archiver --model-name my_model \
--version 1.0 \
--model-file model.py \
--serialized-file model.pth \
--handler custom_handler.py \
--export-path model_store
# 2. 启动 TorchServe
torchserve --start \
--model-store model_store \
--models my_model=my_model.mar \
--ts-config config.properties
# 3. 测试推理
curl -X POST http://localhost:8080/predictions/my_model \
-T sample.jpg
# 4. 管理 API
# 查看模型列表
curl http://localhost:8081/models
# 弹性伸缩 workers
curl -X PUT "http://localhost:8081/models/my_model?min_worker=2&max_worker=4"
3.3 配置优化
# config.properties
inference_address=http://0.0.0.0:8080
management_address=http://0.0.0.0:8081
metrics_address=http://0.0.0.0:8082
number_of_netty_threads=32
job_queue_size=1000
async_logging=true
# 默认 workers
default_workers_per_model=2
4. 模型量化
量化将 FP32 权重降低为 INT8/FP16,减少内存占用和计算量。
4.1 PyTorch 动态量化
import torch.quantization
# 动态量化(最简单,适用于 LSTM/Transformer)
model_int8 = torch.quantization.quantize_dynamic(
model,
{nn.Linear, nn.LSTM}, # 要量化的层
dtype=torch.qint8
)
# 对比大小
print(f"FP32 大小: {model.get_size()}")
print(f"INT8 大小: {model_int8.get_size()}")
4.2 静态量化(PTQ)
# 准备模型
model.qconfig = torch.quantization.get_default_qconfig('fbgemm')
torch.quantization.prepare(model, inplace=True)
# 用校准数据收集统计信息
with torch.no_grad():
for x, _ in calib_loader:
model(x)
# 转换
torch.quantization.convert(model, inplace=True)
4.3 量化感知训练 (QAT)
# 在训练中模拟量化效果,精度更高
model.qconfig = torch.quantization.get_default_qat_qconfig('fbgemm')
torch.quantization.prepare_qat(model, inplace=True)
# 正常训练几个 epoch
for epoch in range(3):
train(model, ...)
# 冻结 BN 统计并转换
model.apply(torch.quantization.disable_observer)
torch.quantization.convert(model, inplace=True)
4.4 ONNX 量化
from onnxruntime.quantization import quantize_dynamic, QuantType
quantize_dynamic(
model_input='model.onnx',
model_output='model_int8.onnx',
weight_type=QuantType.QInt8
)
| 量化方式 | 方法 | 精度损失 | 加速比 | 适用 |
|---|---|---|---|---|
| Post-Training | 校准后静态量化 | 中 | 2-4x | CNN、NLP |
| Quantization-Aware | 训练中模拟 | 低 | 2-4x | 对精度敏感 |
| Dynamic | 运行时动态 | 中高 | 1.5-2x | 快速上手 |
5. 知识蒸馏 (Knowledge Distillation)
让小模型(学生)学习大模型(教师)的输出分布。
class DistillationLoss(nn.Module):
def __init__(self, temperature=4.0, alpha=0.7):
super().__init__()
self.T = temperature
self.alpha = alpha
self.ce = nn.CrossEntropyLoss()
self.kl = nn.KLDivLoss(reduction='batchmean')
def forward(self, student_logits, teacher_logits, labels):
# Hard loss(学生自己学标签)
hard_loss = self.ce(student_logits, labels)
# Soft loss(学生学习教师"软标签")
student_soft = F.log_softmax(student_logits / self.T, dim=1)
teacher_soft = F.softmax(teacher_logits / self.T, dim=1)
soft_loss = self.kl(student_soft, teacher_soft) * (self.T ** 2)
return self.alpha * soft_loss + (1 - self.alpha) * hard_loss
# 训练循环
teacher.eval()
for epoch in range(epochs):
student.train()
for x, y in train_loader:
optimizer.zero_grad()
with torch.no_grad():
teacher_logits = teacher(x)
student_logits = student(x)
loss = distillation_loss(student_logits, teacher_logits, y)
loss.backward()
optimizer.step()
6. TensorRT 加速
6.1 ONNX → TensorRT
import tensorrt as trt
import pycuda.driver as cuda
import pycuda.autoinit
# 1. 创建 builder 和 logger
logger = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(logger)
network = builder.create_network(
1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
)
parser = trt.OnnxParser(network, logger)
# 2. 解析 ONNX
with open('model.onnx', 'rb') as f:
parser.parse(f.read())
# 3. 配置 builder
config = builder.create_builder_config()
config.max_workspace_size = 1 << 30 # 1GB
config.set_flag(trt.BuilderFlag.FP16)
# 4. 构建引擎
engine = builder.build_engine(network, config)
# 5. 序列化保存
with open('model.trt', 'wb') as f:
f.write(engine.serialize())
6.2 TensorRT 推理
# 反序列化引擎
with open('model.trt', 'rb') as f:
runtime = trt.Runtime(logger)
engine = runtime.deserialize_cuda_engine(f.read())
context = engine.create_execution_context()
# 分配 GPU 内存
h_input = cuda.pagelocked_empty(trt.volume(context.get_binding_shape(0)), dtype=np.float32)
h_output = cuda.pagelocked_empty(trt.volume(context.get_binding_shape(1)), dtype=np.float32)
d_input = cuda.mem_alloc(h_input.nbytes)
d_output = cuda.mem_alloc(h_output.nbytes)
stream = cuda.Stream()
# 推理
cuda.memcpy_htod_async(d_input, h_input, stream)
context.execute_async_v2(bindings=[int(d_input), int(d_output)], stream_handle=stream.handle)
cuda.memcpy_dtoh_async(h_output, d_output, stream)
stream.synchronize()
6.3 torch2trt
from torch2trt import torch2trt
model = model.eval().cuda()
x = torch.ones((1, 3, 224, 224)).cuda()
# 转换 FP16
model_trt = torch2trt(model, [x], fp16_mode=True)
# 推理
y_trt = model_trt(x)
7. 模型剪枝
移除不重要权重,减少参数量和计算量。
import torch.nn.utils.prune as prune
# 非结构化剪枝(随机)
prune.random_unstructured(module, name='weight', amount=0.3)
# 结构化剪枝(L1 范数)
prune.ln_structured(module, name='weight', amount=0.3, n=1, dim=0)
# 全局剪枝
parameters_to_prune = [
(model.conv1, 'weight'),
(model.conv2, 'weight'),
(model.fc1, 'weight'),
]
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=0.2
)
8. 边缘端部署
8.1 移动端:Core ML / TFLite
# PyTorch → ONNX → Core ML
import coremltools as ct
# 从 ONNX 转换
mlmodel = ct.converters.onnx.convert(
model='model.onnx',
minimum_ios_deployment_target='13'
)
mlmodel.save('model.mlmodel')
# PyTorch → TFLite (通过 ONNX 中转)
# torch → ONNX → TensorFlow → TFLite
8.2 树莓派 / Jetson
# Jetson Nano/Xavier 用 JetPack + TensorRT
# ARM 架构需交叉编译或使用 conda-forge
# 安装 PyTorch for ARM
wget https://nvidia.box.com/some_path/torch-2.0.0-cp38-cp38-linux_aarch64.whl
pip install torch-2.0.0-cp38-cp38-linux_aarch64.whl
9. 推理服务架构
9.1 Triton Inference Server
# 目录结构
model_repository/
└── my_model/
├── 1/
│ └── model.onnx
└── config.pbtxt
# config.pbtxt
name: "my_model"
platform: "onnxruntime_onnx"
max_batch_size: 32
input [
{
name: "input"
data_type: TYPE_FP32
dims: [3, 224, 224]
}
]
output [
{
name: "output"
data_type: TYPE_FP32
dims: [1000]
}
]
instance_group [
{
count: 2
kind: KIND_GPU
}
]
# 启动
docker run --gpus all -p 8000:8000 -p 8001:8001 -p 8002:8002 \
-v $(pwd)/model_repository:/models \
nvcr.io/nvidia/tritonserver:23.10-py3 \
tritonserver --model-repository=/models
9.2 性能基准测试
import time
# 预热
for _ in range(10):
_ = model(dummy_input)
torch.cuda.synchronize()
# 测试
start = time.time()
for _ in range(100):
output = model(batch_input)
torch.cuda.synchronize()
end = time.time()
throughput = (100 * batch_size) / (end - start)
latency_ms = (end - start) / 100 * 1000
print(f"吞吐: {throughput:.1f} samples/s, 延迟: {latency_ms:.2f} ms")
10. 监控与 A/B 测试
# 模型性能监控指标
metrics = {
'inference_latency_p99': np.percentile(latencies, 99),
'throughput': samples_per_second,
'gpu_utilization': gpustat_GPUUtilization(),
'memory_usage_mb': torch.cuda.memory_allocated() / 1e6,
'prediction_drift': ks_statistic(new_preds, baseline_preds),
'data_drift': ks_statistic(new_data, training_data)
}
# A/B 测试
class ModelRouter:
def __init__(self, model_a, model_b, split=0.5):
self.model_a = model_a
self.model_b = model_b
self.split = split
def predict(self, x, user_id):
if hash(user_id) % 100 < self.split * 100:
return self.model_a(x), 'A'
return self.model_b(x), 'B'
总结
| 部署方案 | 工具 | 适用场景 | 延迟 |
|---|---|---|---|
| 本地推理 | PyTorch/ONNX Runtime | 开发调试 | 取决于硬件 |
| GPU 加速 | TensorRT | 高吞吐 GPU 服务 | < 10ms |
| 服务化 | TorchServe/Triton | 多模型、高并发 | < 50ms |
| 移动端 | Core ML / TFLite | iOS/Android App | < 100ms |
| 边缘设备 | OpenVINO / TensorRT | Jetson/树莓派 | < 500ms |
| 优化手段 | 效果 | 难度 | 适用 |
|---|---|---|---|
| FP16 | 2x 加速,几乎无损 | 低 | 所有 GPU |
| INT8 量化 | 2-4x 加速,轻微损失 | 中 | 部署首选 |
| 知识蒸馏 | 小模型达到大模型 90%+ | 中 | 资源受限 |
| 剪枝 | 2-10x 压缩 | 中高 | 极端场景 |
| ONNX Runtime | 跨平台 2-3x | 低 | 通用部署 |
部署最佳实践:
- 基线:先用 ONNX Runtime 快速部署,验证业务逻辑
- 优化:针对硬件选 TensorRT(NVIDIA)或 OpenVINO(Intel)
- 量化:PTQ 先做,精度不够再 QAT
- 服务化:TorchServe 快速搭建,Triton 管理复杂多模型场景
- 监控:预测分布漂移、输入数据漂移、延迟与错误率
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。