CANN 网络模型部署实战:从训练到生产的完整流程
·
一、部署流程概览
1.1 端到端部署链路
┌──────────────────────────────────────────────────┐
│ CANN 模型部署完整链路 │
├──────────────────────────────────────────────────┤
│ │
│ 训练阶段 │
│ PyTorch / TF → ONNX → ATC → .om │
│ │
│ 部署阶段 │
│ .om → ACL 加载 → 推理 → 结果输出 │
│ │
│ 服务化阶段 │
│ Flask / Triton → HTTP / gRPC → 客户端 │
│ │
│ 监控阶段 │
│ Profiling → 指标采集 → 告警 → 自动化扩缩容 │
│ │
└──────────────────────────────────────────────────┘
1.2 部署方案选型
| 方案 | 适用场景 | 优点 | 缺点 |
|---|---|---|---|
| Flask HTTP | 小规模、低延迟 | 简单、快速 | 并发低 |
| Triton | 大规模、高吞吐 | 功能全、性能优 | 配置复杂 |
| gRPC | 微服务通信 | 高效、流式 | 需要 Protobuf |
| 自研框架 | 特殊需求 | 完全可控 | 开发成本高 |
二、模型导出与转换
2.1 PyTorch 模型导出
基础导出:
import torch
model = MyModel()
model.eval()
# 静态 shape 导出
torch.onnx.export(
model,
args=torch.randn(1, 3, 224, 224),
f="model.onnx",
input_names=["input"],
output_names=["output"],
opset_version=13
)
8.2 新增(动态 shape 导出):
import torch
model = MyModel()
model.eval()
# 动态 shape 导出(支持变长输入)
torch.onnx.export(
model,
args=torch.randn(1, 3, 224, 224),
f="model_dynamic.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={
"input": {0: "batch", 2: "height", 3: "width"},
"output": {0: "batch"}
},
opset_version=13
)
# 验证导出的模型
import onnx
model = onnx.load("model_dynamic.onnx")
onnx.checker.check_model(model)
print("ONNX model exported successfully")
2.2 ONNX 模型简化
复杂模型可能包含冗余节点,需要简化后再转换:
8.1 及之前:
# 手动调用 onnx-simplifier
python -m onnxsim model.onnx model_simplified.onnx
8.2 新增(ATC 内置简化):
# ATC 自动进行图简化
atc --model=model.onnx \
--framework=5 \
--output=model \
--input_shape="input:1,3,224,224" \
--soc_version=Ascend310 \
--graph_optimize_mode=GP_ENHANCE # 启用内置图优化
2.3 模型转换完整脚本
#!/bin/bash
# 模型转换完整脚本
MODEL_NAME="resnet50"
INPUT_MODEL="/workspace/${MODEL_NAME}.onnx"
OUTPUT_MODEL="/workspace/${MODEL_NAME}.om"
# 1. 检查模型文件
if [ ! -f "$INPUT_MODEL" ]; then
echo "Error: Input model not found: $INPUT_MODEL"
exit 1
fi
# 2. 执行转换
atc \
--model=${INPUT_MODEL} \
--framework=5 \
--output=${OUTPUT_MODEL} \
--input_shape="input:1,3,224,224" \
--input_format=NCHW \
--soc_version=Ascend310 \
--precision_mode=allow_fp32_to_fp16 \
--op_select_implmode=high_precision \
--output_type=FP16 \
--graph_optimize_mode=GP_ENHANCE \
--enable_single_stream=1 \
--log=INFO
# 3. 验证输出
if [ -f "$OUTPUT_MODEL" ]; then
echo "Model converted successfully: $OUTPUT_MODEL"
ls -lh ${OUTPUT_MODEL}
else
echo "Error: Model conversion failed"
exit 1
fi
三、Flask HTTP 服务化
3.1 基础 Flask 服务
8.1 及之前(基础实现):
from flask import Flask, request, jsonify
import numpy as np
import acl
app = Flask(__name__)
# ACL 初始化
ret = acl.init()
ret = acl.set_device(0)
# 加载模型
model_id, ret = acl.mdl.load_from_file("/path/to/model.om")
model_desc = acl.mdl.create_desc()
# 预处理
def preprocess(image_bytes):
image = np.frombuffer(image_bytes, dtype=np.uint8)
image = cv2.imdecode(image, cv2.IMREAD_COLOR)
image = cv2.resize(image, (224, 224))
image = image.astype(np.float32) / 255.0
image = (image - mean) / std
image = np.transpose(image, (2, 0, 1))
return image
@app.route('/predict', methods=['POST'])
def predict():
# 获取输入
file = request.files['image']
image_bytes = file.read()
# 预处理
input_data = preprocess(image_bytes)
# 拷贝到 NPU
input_npu = input_data.npu()
# 推理
output = acl.mdl.execute(model_id, input_dataset, output_dataset)
# 返回结果
result = postprocess(output)
return jsonify(result)
app.run(host="0.0.0.0", port=8080)
3.2 高并发 Flask 服务(8.2 新增)
from flask import Flask, request, jsonify
import asyncio
from threading import Lock
from concurrent.futures import ThreadPoolExecutor
import numpy as np
from ascend.acl import AclModel
app = Flask(__name__)
# 线程安全的模型加载
executor = ThreadPoolExecutor(max_workers=4)
model_lock = Lock()
class AclService:
def __init__(self, model_path):
self.model = AclModel(model_path)
self.warmed = False
def warmup(self):
if not self.warmed:
warmup_data = np.random.randn(1, 3, 224, 224).astype(np.float32)
for _ in range(5):
self.model.predict(warmup_data)
self.warmed = True
def predict(self, input_data):
with model_lock:
if not self.warmed:
self.warmup()
return self.model.predict(input_data)
service = AclService("/path/to/model.om")
# 异步路由
@app.route('/predict', methods=['POST'])
async def predict():
loop = asyncio.get_event_loop()
# 读取图片
file = request.files['image']
image_bytes = await loop.run_in_executor(executor, file.read)
# 预处理(同步执行)
input_data = await loop.run_in_executor(executor, preprocess, image_bytes)
# 推理
output = await loop.run_in_executor(executor, service.predict, input_data)
return jsonify({'result': output.tolist()})
@app.route('/health', methods=['GET'])
def health():
return jsonify({'status': 'ok'})
if __name__ == "__main__":
app.run(host="0.0.0.0", port=8080, threaded=True)
四、NVIDIA TensorRT 兼容部署
4.1 为什么需要 TensorRT 兼容
对于同时在 NVIDIA GPU 和昇腾 NPU 部署的场景,需要统一的模型格式转换流程:
# 统一接口定义
class UnifiedModel:
def __init__(self, model_path, backend="npu"):
self.backend = backend
if backend == "npu":
self.model = AclModel(model_path)
elif backend == "cuda":
import tensorrt as trt
# TensorRT 加载
else:
raise ValueError(f"Unknown backend: {backend}")
def predict(self, input_data):
if self.backend == "npu":
return self.model.predict(input_data)
elif self.backend == "cuda":
return self.model(input_data)
4.2 ONNX 到多平台转换
import subprocess
def convert_to_platform(model_path, platform, output_path):
"""转换模型到指定平台"""
if platform == "npu":
# 昇腾 NPU
cmd = [
"atc",
"--model", model_path,
"--framework", "5",
"--output", output_path,
"--input_shape", "input:1,3,224,224",
"--soc_version", "Ascend310"
]
elif platform == "cuda":
# NVIDIA TensorRT
cmd = [
"trtexec",
"--onnx", model_path,
"--saveEngine", output_path,
"--fp16"
]
else:
raise ValueError(f"Unknown platform: {platform}")
subprocess.run(cmd, check=True)
# 使用示例
convert_to_platform("model.onnx", "npu", "model_npu.om")
convert_to_platform("model.onnx", "cuda", "model_trt.engine")
五、生产级服务架构
5.1 模型版本管理
import hashlib
from datetime import datetime
class ModelVersion:
def __init__(self, version, model_path, metadata=None):
self.version = version
self.model_path = model_path
self.metadata = metadata or {}
self.load_time = None
self.model = None
def load(self):
"""加载模型"""
import time
start = time.time()
self.model = AclModel(self.model_path)
self.load_time = time.time() - start
return self
def get_info(self):
return {
"version": self.version,
"model_path": self.model_path,
"loaded": self.model is not None,
"load_time": self.load_time,
"metadata": self.metadata
}
class ModelRegistry:
def __init__(self):
self.versions = {}
self.current_version = None
def register(self, version, model_path, metadata=None):
model_version = ModelVersion(version, model_path, metadata)
self.versions[version] = model_version
return model_version
def load_version(self, version):
if version not in self.versions:
raise ValueError(f"Version {version} not registered")
self.versions[version].load()
return self.versions[version]
def switch_version(self, version):
"""平滑切换版本"""
old_version = self.current_version
if version not in self.versions:
raise ValueError(f"Version {version} not registered")
# 预加载新版本
self.versions[version].load()
# 切换
self.current_version = version
return old_version
def get_current(self):
return self.versions.get(self.current_version)
# 使用示例
registry = ModelRegistry()
registry.register("1.0.0", "/models/v1.0.0/model.om", {"accuracy": 0.95})
registry.register("1.1.0", "/models/v1.1.0/model.om", {"accuracy": 0.97})
registry.current_version = "1.0.0"
# 灰度发布:新版本切换 10% 流量
def switch_with_canary(registry, new_version, canary_ratio=0.1):
import random
if random.random() < canary_ratio:
registry.switch_version(new_version)
print(f"Switched to {new_version}")
else:
print(f"Stayed on {registry.current_version}")
5.2 灰度发布与 A/B 测试
import random
from dataclasses import dataclass
@dataclass
class Experiment:
name: str
model_version: str
traffic_ratio: float
class ABTester:
def __init__(self, registry):
self.registry = registry
self.experiments = {}
self.default_version = None
def add_experiment(self, name, model_version, traffic_ratio):
self.experiments[name] = Experiment(name, model_version, traffic_ratio)
def set_default(self, version):
self.default_version = version
def get_version(self, request):
"""根据请求决定使用哪个版本"""
# 检查是否有实验配置
experiment_name = request.headers.get('X-Experiment')
if experiment_name and experiment_name in self.experiments:
exp = self.experiments[experiment_name]
return exp.model_version
# 检查是否有指定版本
version_header = request.headers.get('X-Model-Version')
if version_header and version_header in self.registry.versions:
return version_header
# 默认版本(带流量分配)
return self._get_canary_version()
def _get_canary_version(self):
"""根据流量比例选择版本"""
rand = random.random()
cumulative = 0.0
for exp in self.experiments.values():
cumulative += exp.traffic_ratio
if rand < cumulative:
return exp.model_version
return self.default_version
# 使用示例
ab_tester = ABTester(registry)
ab_tester.set_default("1.0.0")
ab_tester.add_experiment("new_model", "1.1.0", traffic_ratio=0.1)
@app.route('/predict', methods=['POST'])
def predict():
version = ab_tester.get_version(request)
model = registry.versions[version].model
# 使用对应版本推理
result = model.predict(input_data)
return jsonify({"result": result, "version": version})
六、监控与运维
6.1 关键指标采集
import time
from collections import defaultdict
class MetricsCollector:
def __init__(self):
self.latencies = defaultdict(list)
self.counts = defaultdict(int)
self.errors = defaultdict(int)
def record_latency(self, endpoint, latency):
self.latencies[endpoint].append(latency)
def record_request(self, endpoint):
self.counts[endpoint] += 1
def record_error(self, endpoint):
self.errors[endpoint] += 1
def get_stats(self, endpoint):
latencies = self.latencies[endpoint]
if not latencies:
return {}
sorted_lat = sorted(latencies)
n = len(sorted_lat)
return {
"count": self.counts[endpoint],
"errors": self.errors[endpoint],
"avg_latency": sum(latencies) / n,
"p50_latency": sorted_lat[n // 2],
"p95_latency": sorted_lat[int(n * 0.95)],
"p99_latency": sorted_lat[int(n * 0.99)],
}
def export_prometheus(self):
"""导出 Prometheus 格式指标"""
lines = []
for endpoint in self.counts:
stats = self.get_stats(endpoint)
lines.append(f'request_count{{endpoint="{endpoint}"}} {stats["count"]}')
lines.append(f'request_latency_avg{{endpoint="{endpoint}"}} {stats["avg_latency"]:.3f}')
lines.append(f'request_latency_p95{{endpoint="{endpoint}"}} {stats["p95_latency"]:.3f}')
return "\n".join(lines)
collector = MetricsCollector()
@app.route('/metrics', methods=['GET'])
def metrics():
return collector.export_prometheus(), 200, {'Content-Type': 'text/plain'}
6.2 自动扩缩容
import time
from threading import Thread
class AutoScaler:
def __init__(self, min_replicas=1, max_replicas=10):
self.min_replicas = min_replicas
self.max_replicas = max_replicas
self.current_replicas = 1
self.metrics = MetricsCollector()
self.running = True
self.scaler_thread = Thread(target=self._scale_loop)
self.scaler_thread.start()
def _scale_loop(self):
while self.running:
stats = self.metrics.get_stats("/predict")
if not stats:
time.sleep(10)
continue
avg_latency = stats["avg_latency"]
p99_latency = stats["p99_latency"]
error_rate = stats["errors"] / max(stats["count"], 1)
# 扩容条件
if p99_latency > 1.0 or error_rate > 0.01:
new_replicas = min(self.current_replicas + 1, self.max_replicas)
if new_replicas != self.current_replicas:
print(f"Scaling up: {self.current_replicas} -> {new_replicas}")
self.current_replicas = new_replicas
# 缩容条件
elif avg_latency < 0.1 and error_rate < 0.001:
new_replicas = max(self.current_replicas - 1, self.min_replicas)
if new_replicas != self.current_replicas:
print(f"Scaling down: {self.current_replicas} -> {new_replicas}")
self.current_replicas = new_replicas
time.sleep(30) # 每 30 秒检查一次
def stop(self):
self.running = False
self.scaler_thread.join()
七、部署检查清单
| 检查项 | 说明 | 优先级 |
|---|---|---|
| 模型文件验证 | 检查 .om 文件完整性和权限 | 必须 |
| 预热推理 | 首次请求前执行 5-10 次预热 | 必须 |
| 监控指标接入 | 延迟、QPS、错误率接入监控 | 必须 |
| 告警配置 | 延迟阈值、错误率阈值告警 | 必须 |
| 版本回滚方案 | 新版本有问题时的回滚步骤 | 必须 |
| 流量切换策略 | 灰度发布比例和切换流程 | 推荐 |
| 资源预留 | CPU、内存、NPU 资源预留 | 推荐 |
| 日志规范 | 请求 ID、版本号、时间戳 | 推荐 |
八、常见问题
| 问题 | 原因 | 解决方案 |
|---|---|---|
| 模型加载失败 | .om 文件损坏或路径错误 | 检查文件完整性 |
| 推理超时 | batch size 过大或模型太复杂 | 减小 batch 或优化模型 |
| 预热无效 | 预热请求参数不正确 | 使用实际输入 shape 预热 |
| 版本切换失败 | 新版本加载失败 | 保留旧版本镜像,降级回滚 |
| 内存泄漏 | 模型未正确释放 | 使用 try-finally 确保释放 |
| 并发性能差 | GIL 限制或锁竞争 | 使用多进程或异步框架 |
相关仓库
- CANN - 昇腾异构计算架构 https://atomgit.com/cann
- runtime - 运行时 https://atomgit.com/cann/runtime
- torch-npu - PyTorch 适配 https://atomgit.com/cann/torch-npu
- cann-recipes-infer - 推理配方 https://atomgit.com/cann-recipes-infer
- cann-recipes-train - 训练配方 https://atomgit.com/cann-recipes-train
- cann-samples - 示例代码 https://atomgit.com/cann-samples
更多推荐



所有评论(0)