第16章 ML 平台架构

第16章 ML 平台架构

“一个成熟的 ML 平台不是某个工具,而是一套让数据科学家不需要关心基础设施的能力集合。”

16.1 为什么需要 ML 平台

在深入组件之前,先看一个常见场景:

一个数据科学家训练了一个推荐模型,在离线评估中表现优异。然后呢?

  • 模型代码在本地笔记本上,训练数据在 HDFS 上,特征管线在另一个团队的 Airflow 里
  • 要上线,需要找 SRE 配置服务,找平台组写推理容器,找数据组验证特征一致性
  • 两周后模型终于上线了,但线上指标和离线对不上——因为特征工程逻辑在搬迁过程中被改了一行代码

这个场景揭示了一个核心问题:从 notebook 到生产环境之间存在巨大的鸿沟。ML 平台的使命就是填平这道鸿沟。

传统软件 vs ML 系统的开发流程

graph LR
    subgraph 传统软件
        A1[需求] --> A2[编码] --> A3[测试] --> A4[部署] --> A5[监控]
    end
    subgraph ML系统
        B1[问题定义] --> B2[数据采集] --> B3[特征工程] --> B4[训练] --> B5[评估]
        B5 --> B6[验证] --> B7[部署] --> B8[监控]
        B8 -.->|数据漂移| B2
        B8 -.->|概念漂移| B4
    end

ML 系统的复杂性在于:数据会变、模型会退化、反馈环路是常态。一个好的 ML 平台必须原生处理这些特性,而不是事后补丁。

16.2 ML 平台的核心组件

一个功能完整的 ML 平台通常包含以下层次:

graph TB
    subgraph 接入层
        UI[Web UI / IDE 插件]
        API[REST API]
        CLI[CLI 工具]
    end
    
    subgraph 平台服务层
        ET[实验追踪]
        MR[模型注册中心]
        FM[特征存储]
        WF[工作流编排]
        SERV[模型服务]
        MON[监控告警]
    end
    
    subgraph 基础设施层
        K8S[Kubernetes 集群]
        GPU[GPU/NPU 资源池]
        STOR[对象存储 / NAS]
        DB[元数据数据库]
    end
    
    UI --> 平台服务层
    API --> 平台服务层
    CLI --> 平台服务层
    平台服务层 --> 基础设施层

组件 核心职责 常见选型
实验追踪 记录每次训练的参数、指标、代码版本 MLflow, W&B, Neptune
模型注册中心 模型版本管理、阶段流转 MLflow Registry, Vertex AI Model Registry
特征存储 离线/在线特征一致性与复用 Feast, Tecton, Vertex Feature Store
工作流编排 DAG 化的训练/数据管线 Airflow, Kubeflow Pipelines, Argo Workflows
模型服务 推理部署、版本切换、A/B 测试 KServe, BentoML, Triton Inference Server
监控告警 线上模型性能、数据质量、漂移检测 Evidently, Arize, Grafana + Prometheus
Tip

平台建设原则: 不要一开始就上全套。先解决最痛的瓶颈——通常是实验追踪和模型服务。特征存储和工作流编排可以后续迭代。

一个参考的 ML 平台技术栈

以下是一个经过生产验证的技术栈组合:

# platform-stack.yaml - ML 平台技术栈参考
platform:
  experiment_tracking:
    tool: mlflow
    backend: postgresql  # 元数据
    artifact_store: s3://ml-artifacts  # 模型文件
  
  model_registry:
    tool: mlflow
    stages: [staging, production, archived]
  
  feature_store:
    tool: feast
    offline_store: spark  # 离线训练
    online_store: redis   # 在线推理
  
  orchestration:
    tool: argo_workflows
    runtime: kubernetes
  
  serving:
    tool: kserve
    runtime: knative
    autoscaling: kpa  # Knative Pod Autoscaler
  
  monitoring:
    metrics: prometheus
    logs: loki
    traces: tempo
    drift_detection: evidently

16.3 实验追踪(Experiment Tracking)

实验追踪解决一个看似简单但极其折磨人的问题:“上周那个效果不错的模型,用了什么参数?”

核心概念

实验追踪系统记录每次训练运行的三个维度信息:

  1. 输入:代码版本、数据版本、超参数、环境配置
  2. 过程:训练曲线、中间指标、系统指标(GPU 利用率等)
  3. 输出:模型文件、评估报告、可视化图表

MLflow 实战示例

import mlflow
import mlflow.pytorch
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

# 设置实验
mlflow.set_experiment("llm-finetune-experiments")

with mlflow.start_run(run_name="llama3-lora-r64") as run:
    # 记录超参数
    params = {
        "model_name": "meta-llama/Llama-3-8B",
        "lora_r": 64,
        "lora_alpha": 128,
        "learning_rate": 2e-4,
        "batch_size": 16,
        "gpu_type": "A100-80GB",
        "gpu_count": 4,
        "training_data": "dataset_v3.2",
    }
    mlflow.log_params(params)
    
    # 训练逻辑(简化)
    model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3-8B")
    # ... 训练代码 ...
    
    # 记录指标
    for epoch in range(num_epochs):
        # train_loss = train_one_epoch(...)
        mlflow.log_metric("train_loss", train_loss, step=epoch)
        mlflow.log_metric("eval_loss", eval_loss, step=epoch)
        mlflow.log_metric("gpu_memory_gb", gpu_mem, step=epoch)
    
    # 记录模型
    mlflow.pytorch.log_model(model, "model", registered_model_name="llama3-8b")
    
    # 记录训练数据 hash 作为数据版本
    mlflow.set_tag("data_sha256", compute_dataset_hash("dataset_v3.2"))
    
    # 记录 Git commit
    mlflow.set_tag("git_commit", get_git_commit_hash())

实验追踪的常见陷阱

Warning

典型问题: 实验追踪系统里记录了上千个 run,但没有人在意数据版本。当模型效果无法复现时,才发现训练数据已经被覆盖更新了。

解决方案: 将数据版本(DVC/LakeFS hash)作为强制 tag,没有数据版本的 run 直接拒绝记录。

团队协作中的实验管理

当一个团队有 10+ 数据科学家时,实验管理需要约定:

# 团队实验命名规范示例
def get_standard_run_name(project, model_arch, key_param, version):
    """
    命名规范: {project}_{arch}_{key_param}_v{version}
    示例: recsys_din_emb256_v3
    """
    return f"{project}_{model_arch}_{key_param}_v{version}"

# 强制标签
REQUIRED_TAGS = {
    "owner": "负责人的姓名/工号",
    "stage": "dev|staging|prod-candidate",
    "dataset": "数据集 DVC hash",
    "objective": "优化目标说明",
    "baseline_run": "对照实验 run_id",
}

# 上线前检查
def validate_run_for_staging(run_id):
    run = mlflow.get_run(run_id)
    for tag in REQUIRED_TAGS:
        if tag not in run.data.tags:
            raise ValueError(f"Run {run_id} 缺少必填标签: {tag}")

16.4 模型注册中心与版本管理

模型注册中心是实验追踪到生产部署之间的桥梁。它要回答的问题是:“当前线上跑的是哪个版本的模型?上一个版本是什么?如何回滚?”

模型生命周期

stateDiagram-v2
    [*] --> Registered: 实验完成,注册模型
    Registered --> Staging: 通过离线评估
    Staging --> Production: 通过线上验证
    Staging --> Registered: 评估不通过
    Production --> Archived: 被新版本替代
    Production --> Production: 回滚
    Archived --> [*]

使用 MLflow Model Registry

import mlflow

client = mlflow.tracking.MlflowClient()

# 创建 Registered Model(如果不存在)
try:
    client.create_registered_model("llama3-8b-chat")
except Exception:
    pass  # 已存在

# 将某个 run 的模型注册为新版本
result = mlflow.register_model(
    model_uri="runs:/abc123def/model",
    name="llama3-8b-chat"
)
print(f"注册版本: {result.version}")

# 转换模型阶段
client.transition_model_version_stage(
    name="llama3-8b-chat",
    version=result.version,
    stage="Staging",
    archive_existing_versions=True  # 自动归档同阶段的旧版本
)

# 获取当前生产版本
prod_versions = client.get_latest_versions(
    "llama3-8b-chat", stages=["Production"]
)
for v in prod_versions:
    print(f"Production: v{v.version}, run_id={v.run_id}")

模型版本管理的最佳实践

Tip

版本命名约定: 不要只用自增整数。使用语义化版本号 MAJOR.MINOR.PATCH: - MAJOR:更换了基础模型架构(如 Llama-3 → Llama-4) - MINOR:重新训练或更换了训练数据 - PATCH:只调整了推理参数或后处理逻辑

16.5 工作流编排

工作流编排解决的是”把这些步骤按正确顺序串起来”的问题。在 ML 场景中,这些步骤通常包括数据预处理、特征计算、模型训练、评估、注册和部署。

三大编排工具对比

特性 Airflow Kubeflow Pipelines Argo Workflows
定位 通用数据管线 ML 专用 K8s 原生工作流
DAG 定义 Python Python / YAML YAML
K8s 原生 否(需要 operator)
ML 工具集成 丰富 MLflow/Katib 内置 需要自行集成
GPU 调度 通过 K8s 原生支持 原生支持
适合场景 数据团队主导 ML 团队主导 平台团队主导

Argo Workflows 示例:端到端训练管线

# training-pipeline.yaml
apiVersion: argoproj.io/v1alpha1
kind: Workflow
metadata:
  generateName: llm-finetune-
spec:
  entrypoint: main
  arguments:
    parameters:
      - name: model_name
        value: "meta-llama/Llama-3-8B"
      - name: dataset_version
        value: "v3.2"
      - name: lora_r
        value: "64"
  
  templates:
    - name: main
      dag:
        tasks:
          - name: prepare-data
            template: data-prep
            arguments:
              parameters:
                - name: version
                  value: "{{workflow.parameters.dataset_version}}"
          
          - name: compute-features
            dependencies: [prepare-data]
            template: feature-engine
          
          - name: train
            dependencies: [compute-features]
            template: train-model
            arguments:
              parameters:
                - name: model_name
                  value: "{{workflow.parameters.model_name}}"
                - name: lora_r
                  value: "{{workflow.parameters.lora_r}}"
                - name: data_path
                  value: "{{tasks.compute-features.outputs.parameters.output_path}}"
          
          - name: evaluate
            dependencies: [train]
            template: eval-model
            arguments:
              parameters:
                - name: model_path
                  value: "{{tasks.train.outputs.parameters.model_path}}"
          
          - name: register
            dependencies: [evaluate]
            template: register-model
            when: "{{tasks.evaluate.outputs.parameters.passed}} == true"
            arguments:
              parameters:
                - name: model_path
                  value: "{{tasks.train.outputs.parameters.model_path}}"
                - name: metrics
                  value: "{{tasks.evaluate.outputs.parameters.metrics}}"

    - name: data-prep
      inputs:
        parameters:
          - name: version
      container:
        image: registry.internal/data-prep:latest
        command: [python, -u, prepare.py]
        args: ["--version", "{{inputs.parameters.version}}"]
      outputs:
        parameters:
          - name: output_path
            valueFrom:
              path: /tmp/output_path.txt

    - name: train-model
      inputs:
        parameters:
          - name: model_name
          - name: lora_r
          - name: data_path
      container:
        image: registry.internal/training:latest
        resources:
          limits:
            nvidia.com/gpu: 4
            memory: 256Gi
          requests:
            nvidia.com/gpu: 4
            memory: 200Gi
        command: [python, -u, train.py]
        args:
          - --model
          - "{{inputs.parameters.model_name}}"
          - --data
          - "{{inputs.parameters.data_path}}"
          - --lora-r
          - "{{inputs.parameters.lora_r}}"
      outputs:
        parameters:
          - name: model_path
            valueFrom:
              path: /tmp/model_path.txt
Tip

实践经验: 工作流的每个步骤应该是幂等的。如果训练步骤失败,你应该能够直接重跑该步骤,而不需要从头开始准备数据。使用对象存储(S3/OSS)作为步骤间的数据传递中介,而不是本地磁盘。

16.6 CI/CD for ML

传统的 CI/CD 关心的是”代码变更是否破坏了功能”。ML 的 CI/CD 还需要关心”模型在数据变更后是否仍然有效”。

ML CI/CD Pipeline 的特殊之处

graph LR
    subgraph 传统CI/CD
        TC1[代码提交] --> TC2[单元测试] --> TC3[构建镜像] --> TC4[部署]
    end
    
    subgraph ML CI/CD
        MC1[代码/数据提交] --> MC2[代码测试]
        MC2 --> MC3[数据验证]
        MC3 --> MC4[训练]
        MC4 --> MC5[模型评估]
        MC5 --> MC6[公平性检查]
        MC6 --> MC7[构建推理镜像]
        MC7 --> MC8[金丝雀部署]
        MC8 --> MC9[线上指标验证]
        MC9 --> MC10[全量发布]
    end

GitHub Actions 实战:模型 CI/CD

# .github/workflows/model-ci-cd.yml
name: Model CI/CD

on:
  push:
    paths:
      - 'models/**'
      - 'data/**'
      - '.github/workflows/model-ci-cd.yml'
  schedule:
    # 每周一凌晨自动触发(处理数据漂移)
    - cron: '0 0 * * 1'

jobs:
  # 阶段 1: 代码检查
  lint-test:
    runs-on: ubuntu-latest
    steps:
      - uses: actions/checkout@v4
      - uses: actions/setup-python@v5
        with:
          python-version: "3.11"
      - run: pip install ruff pytest
      - run: ruff check models/
      - run: pytest tests/

  # 阶段 2: 数据验证
  data-validation:
    needs: lint-test
    runs-on: ubuntu-latest
    steps:
      - uses: actions/checkout@v4
      - name: 验证数据质量
        run: |
          python -m pipeline.validate_data \
            --input data/train_v${{ github.sha }} \
            --checks null_check,distribution_check,leakage_check \
            --baseline data/baseline_stats.json

  # 阶段 3: 模型训练
  train:
    needs: data-validation
    runs-on: self-hosted  # GPU runner
    steps:
      - uses: actions/checkout@v4
      - name: 训练
        run: |
          python -m pipeline.train \
            --config configs/training.yaml \
            --output models/${{ github.sha }}
      - name: 上传模型 artifact
        uses: actions/upload-artifact@v4
        with:
          name: model-artifact
          path: models/${{ github.sha }}/

  # 阶段 4: 评估与对比
  evaluate:
    needs: train
    runs-on: self-hosted
    steps:
      - uses: actions/checkout@v4
      - uses: actions/download-artifact@v4
        with:
          name: model-artifact
          path: ./model
      - name: 评估并对比基线
        run: |
          python -m pipeline.evaluate \
            --model ./model \
            --test-set data/test.jsonl \
            --baseline-metrics metrics/baseline.json \
            --output metrics/eval_${{ github.sha }}.json
      - name: 检查是否优于基线
        run: |
          python -m pipeline.compare \
            --current metrics/eval_${{ github.sha }}.json \
            --baseline metrics/baseline.json \
            --threshold 0.02  # 至少提升 2%

  # 阶段 5: 金丝雀部署
  canary:
    needs: evaluate
    if: github.ref == 'refs/heads/main'
    runs-on: ubuntu-latest
    steps:
      - name: 部署到金丝雀环境(5% 流量)
        run: |
          kubectl apply -f deploy/canary.yaml
          # 等待 30 分钟收集线上指标
          sleep 1800
          python -m pipeline.check_canary_metrics \
            --threshold latency_p99_200ms,error_rate_1pct

持续训练(Continuous Training)

除了代码触发的 CI/CD,ML 系统还需要数据触发的持续训练:

# ct_trigger.py - 持续训练触发器
import json
from datetime import datetime, timedelta

def should_retrain(model_id, metrics_history, data_quality_report):
    """
    多触发条件判断
    """
    reasons = []
    
    # 条件 1: 性能退化
    latest_metrics = metrics_history[-1]
    baseline_metrics = metrics_history[0]
    if latest_metrics["accuracy"] < baseline_metrics["accuracy"] - 0.02:
        reasons.append(f"准确率下降 {baseline_metrics['accuracy'] - latest_metrics['accuracy']:.1%}")
    
    # 条件 2: 数据漂移
    drift_score = data_quality_report.get("psi_score", 0)  # Population Stability Index
    if drift_score > 0.2:
        reasons.append(f"数据漂移 PSI={drift_score:.3f} 超过阈值 0.2")
    
    # 条件 3: 定期重训
    days_since_last_train = (
        datetime.now() - metrics_history[-1]["trained_at"]
    ).days
    if days_since_last_train > 7:
        reasons.append(f"距上次训练已 {days_since_last_train} 天")
    
    # 条件 4: 数据量积累
    new_data_ratio = data_quality_report.get("new_data_ratio", 0)
    if new_data_ratio > 0.15:
        reasons.append(f"新数据占比 {new_data_ratio:.1%}")
    
    return reasons

# 定时任务调用
reasons = should_retrain("llama3-8b-v2", history, quality_report)
if reasons:
    trigger_retraining_pipeline(model_id="llama3-8b-v2", reasons=reasons)
Warning

金丝雀部署的陷阱: 线上指标验证窗口太短。模型可能在简单样本上表现正常,但在长尾分布上严重退化。建议金丝雀阶段至少覆盖一个完整的数据周期(如一周的业务数据),并对不同用户群体做分桶评估。

16.7 小结

ML 平台架构的核心目标是降低从实验到生产的摩擦。本章覆盖了关键组件:

  • 实验追踪确保每次训练可复现、可比较
  • 模型注册中心让模型版本可追溯、可回滚
  • 工作流编排自动化端到端训练管线
  • CI/CD 将代码变更和数据变更统一纳入持续交付流程

平台建设是一个渐进的过程。不要试图一次性搭建所有组件——从最痛的瓶颈开始,迭代演进。记住:最好的平台是让数据科学家忘记平台存在的平台

延伸阅读

  • Designing Machine Learning Systems (Chip Huyen, 2022) — 第 6-9 章深入讨论 ML 平台设计
  • MLflow 官方文档 mlflow.org/docs — 最新的 Tracking/Registry/Serving 指南
  • Kubeflow 官方文档 kubeflow.org — K8s 原生 ML 平台
  • Hidden Technical Debt in Machine Learning Systems (Sculley et al., NIPS 2015) — 经典论文,解释为什么 ML 系统的维护成本远高于想象
  • Practitioners Guide to MLOps (Google Cloud, 2023) — MLOps 成熟度模型
  • Continuous Delivery for Machine Learning (Sato, Wider, Windheuser — martinfowler.com) — ML 环境下的 CD 实践