一、部署流程概览

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
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐