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
第16章 ML 平台架构
第16章 ML 平台架构
“一个成熟的 ML 平台不是某个工具,而是一套让数据科学家不需要关心基础设施的能力集合。”
16.1 为什么需要 ML 平台
在深入组件之前,先看一个常见场景:
一个数据科学家训练了一个推荐模型,在离线评估中表现优异。然后呢?
- 模型代码在本地笔记本上,训练数据在 HDFS 上,特征管线在另一个团队的 Airflow 里
- 要上线,需要找 SRE 配置服务,找平台组写推理容器,找数据组验证特征一致性
- 两周后模型终于上线了,但线上指标和离线对不上——因为特征工程逻辑在搬迁过程中被改了一行代码
这个场景揭示了一个核心问题:从 notebook 到生产环境之间存在巨大的鸿沟。ML 平台的使命就是填平这道鸿沟。
传统软件 vs ML 系统的开发流程
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 |
平台建设原则: 不要一开始就上全套。先解决最痛的瓶颈——通常是实验追踪和模型服务。特征存储和工作流编排可以后续迭代。
一个参考的 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: evidently16.3 实验追踪(Experiment Tracking)
实验追踪解决一个看似简单但极其折磨人的问题:“上周那个效果不错的模型,用了什么参数?”
核心概念
实验追踪系统记录每次训练运行的三个维度信息:
- 输入:代码版本、数据版本、超参数、环境配置
- 过程:训练曲线、中间指标、系统指标(GPU 利用率等)
- 输出:模型文件、评估报告、可视化图表
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())实验追踪的常见陷阱
典型问题: 实验追踪系统里记录了上千个 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}")模型版本管理的最佳实践
版本命名约定: 不要只用自增整数。使用语义化版本号 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实践经验: 工作流的每个步骤应该是幂等的。如果训练步骤失败,你应该能够直接重跑该步骤,而不需要从头开始准备数据。使用对象存储(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)金丝雀部署的陷阱: 线上指标验证窗口太短。模型可能在简单样本上表现正常,但在长尾分布上严重退化。建议金丝雀阶段至少覆盖一个完整的数据周期(如一周的业务数据),并对不同用户群体做分桶评估。
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 实践