目 录CONTENT

文章目录

模型版本管理:从 MLflow 到自建 Registry 的选型之路

PySuper
2025-06-07 / 0 评论 / 0 点赞 / 3 阅读 / 0 字
温馨提示:
所有牛逼的人都有一段苦逼的岁月。 但是你只要像SB一样去坚持,终将牛逼!!! ✊✊✊

前言

在机器学习项目中,模型是最核心的资产。一个成熟的 ML 系统往往需要管理数十个模型、数百个版本,如何高效地进行模型版本管理成为团队必须面对的挑战。

笔者在过去的两年里,经历了从简单文件管理,到 MLflow Model Registry,再到自建 Registry 的演进过程。本文将完整分享这一路上的思考、踩坑与收获。


一、为什么模型版本管理不同于代码版本管理

1.1 代码 vs 模型:本质差异

很多人会想当然地认为:代码可以用 Git 管理,模型为什么不行?这个问题值得我们深入探讨。

┌─────────────────────────────────────────────────────────────────┐
│                    代码版本 vs 模型版本                         │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│   代码 (Git)                     模型 (Model Registry)          │
│   ─────────────────              ──────────────────────         │
│                                                                 │
│   文件大小: KB~MB               文件大小: 几十MB~几百GB         │
│                                                                 │
│   变化类型:                       变化类型:                        │
│   ├─ 文本差异                     ├─ 参数值变化                   │
│   ├─ 可读可合并                   │  (二进制,无法diff)          │
│   └─ 增量存储                    ├─ 架构变化                    │
│                                   │  (层数、宽度变化)            │
│   版本数量: 1000+/repo           ├─ 训练数据变化                │
│                                   │  (影响模型但不内嵌)          │
│   克隆时间: 秒级                  └─ 训练配置变化                │
│                                   │  (超参数、seed)              │
│   存储成本: 低                   版本数量: 10~100/model         │
│                                   克隆时间: 分钟~小时           │
│                                   存储成本: 极高                │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

1.2 模型版本的独特挑战

# 模型版本管理的核心难点可视化
# challenges_visualization.py

"""
模型版本管理的独特挑战:

1. 非确定性训练
   - 相同代码 + 相同数据 ≠ 相同模型
   - GPU 浮点运算、随机 dropout、并行优化等导致结果差异
   
2. 隐式依赖链
   - 模型 ← 训练代码 ← 数据 ← 预处理代码
   - 任意一环变化都可能影响模型
   
3. 度量维度多样
   - 准确性、延迟、吞吐量、内存占用、能耗...
   - 不同业务场景优先级不同
   
4. 人工评审与自动化的平衡
   - 需要人工review才能上线
   - 但也需要自动化pipeline
"""

# 示例:非确定性训练导致的问题
TRAINING_SCENARIO = """
┌─────────────────────────────────────────────────────────────────┐
│                   非确定性训练问题                              │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│   Day 1: 工程师A 训练模型                                       │
│   ┌─────────────────────────────────────────────────────────┐   │
│   │ seed=42, acc=0.92, auc=0.95                             │   │
│   │ commit: abc123                                          │   │
│   │ data_version: dataset_v2                                │   │
│   └─────────────────────────────────────────────────────────┘   │
│                                                                 │
│   Day 5: 重新训练(代码数据未变)                                │
│   ┌─────────────────────────────────────────────────────────┐   │
│   │ seed=42, acc=0.93, auc=0.94  ⚠️ 指标变化了!           │   │
│   │ commit: abc123  (代码未变)                              │   │
│   │ data_version: dataset_v2  (数据未变)                    │   │
│   └─────────────────────────────────────────────────────────┘   │
│                                                                 │
│   可能原因:                                                     │
│   - PyTorch 升级了某个优化项                                    │
│   - CUDA 版本变化影响了计算顺序                                 │
│   - 底层库更新了随机数生成器                                    │
│                                                                 │
│   解决方案:                                                     │
│   - 环境完全固定 (conda/pip freeze)                            │
│   - 使用 deterministic 模式                                     │
│   - 模型 hash 需要包含环境和seed信息                           │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘
"""

print(TRAINING_SCENARIO)

1.3 版本依赖关系图

┌─────────────────────────────────────────────────────────────────┐
│                    模型版本依赖关系图                           │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│                     ┌─────────────────┐                        │
│                     │   Model v1.3    │                        │
│                     │  (Production)   │                        │
│                     └────────┬────────┘                        │
│                              │                                  │
│              ┌───────────────┼───────────────┐                  │
│              ▼               ▼               ▼                  │
│     ┌────────────┐   ┌────────────┐   ┌────────────┐            │
│     │ Eval Set C │   │ Train Code │   │ Hyperparams│            │
│     │  (最新)    │   │  (abc123)  │   │  (v2.1)    │            │
│     └─────┬──────┘   └─────┬──────┘   └─────┬──────┘           │
│           │                │                │                  │
│           └────────────────┼────────────────┘                  │
│                            ▼                                   │
│                     ┌─────────────────┐                        │
│                     │  Training Data  │                        │
│                     │  (dataset_v3)   │                        │
│                     └────────┬────────┘                        │
│                              │                                 │
│                     ┌────────┴────────┐                        │
│                     │  Preprocessing  │                        │
│                     │  Pipeline v1.2  │                        │
│                     └─────────────────┘                        │
│                                                                │
│   问题: 如何追踪这整条依赖链?                                      │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

二、MLflow Model Registry 详解

2.1 MLflow Registry 架构

┌─────────────────────────────────────────────────────────────────┐
│                  MLflow Model Registry 架构                    │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│   MLflow Tracking Server                                        │
│   ┌─────────────────────────────────────────────────────────┐   │
│   │                                                         │   │
│   │   ┌─────────┐  ┌─────────┐  ┌─────────┐               │   │
│   │   │Artifact │  │   API   │  │  Web UI │               │   │
│   │   │ Store   │  │ Server  │  │         │               │   │
│   │   │  (S3)   │  │         │  │         │               │   │
│   │   └────┬────┘  └────┬────┘  └────┬────┘               │   │
│   │        │            │            │                     │   │
│   │        └────────────┼────────────┘                     │   │
│   │                     │                                  │   │
│   │            ┌────────▼────────┐                         │   │
│   │            │  Metadata Store │                         │   │
│   │            │   (PostgreSQL)  │                         │   │
│   │            └─────────────────┘                         │   │
│   │                                                         │   │
│   └─────────────────────────────────────────────────────────┘   │
│                                                                 │
│   核心概念:                                                     │
│   - Registered Model: 模型的顶层容器                             │
│   - Model Version: 某个具体版本                                  │
│   - Stage: 生命周期阶段 (None, Staging, Production, Archived)   │
│   - Run: 训练实验的追踪记录                                      │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

2.2 MLflow 使用实战

# mlflow_model_registry.py
"""
MLflow Model Registry 完整使用示例
包含模型注册、版本管理、阶段转换、API调用等
"""

import os
import time
import hashlib
import mlflow
from mlflow.tracking import MlflowClient
from mlflow.entities import ViewType
from dataclasses import dataclass, field, asdict
from typing import Dict, List, Optional, Any
from datetime import datetime
from enum import Enum
import json


class ModelStage(Enum):
    """模型阶段"""
    NONE = "None"
    STAGING = "Staging"      # 预发布/测试
    PRODUCTION = "Production" # 生产环境
    ARCHIVED = "Archived"     # 归档


@dataclass
class ModelMetadata:
    """模型元数据"""
    # 基础信息
    model_name: str
    model_version: str
    
    # 模型标识
    model_hash: str = ""  # 模型文件的 SHA256
    signature_hash: str = ""  # 输入输出signature的hash
    
    # 血缘信息
    base_model: Optional[str] = None  # 基础模型(用于微调)
    training_code: str = ""  # 训练代码 commit hash
    training_data: str = ""  # 训练数据版本
    eval_data: str = ""  # 评测数据版本
    preprocessing_pipeline: str = ""  # 预处理pipeline版本
    
    # 训练配置
    hyperparameters: Dict[str, Any] = field(default_factory=dict)
    training_seed: int = 42
    environment: Dict[str, str] = field(default_factory=dict)  # pip freeze
    
    # 性能指标
    metrics: Dict[str, float] = field(default_factory=dict)
    
    # 业务指标
    business_metrics: Dict[str, float] = field(default_factory=dict)
    
    # 其他元数据
    description: str = ""
    tags: Dict[str, str] = field(default_factory=dict)
    created_by: str = ""
    created_at: str = ""
    
    def to_dict(self) -> Dict:
        """转换为字典"""
        data = asdict(self)
        return data
    
    def calculate_signature_hash(self) -> str:
        """计算signature hash"""
        signature_data = {
            "hyperparameters": self.hyperparameters,
            "training_seed": self.training_seed,
        }
        signature_str = json.dumps(signature_data, sort_keys=True)
        self.signature_hash = hashlib.sha256(signature_str.encode()).hexdigest()[:16]
        return self.signature_hash


class MLflowModelRegistry:
    """
    MLflow Model Registry 封装类
    提供更友好的模型版本管理接口
    """
    
    def __init__(
        self,
        tracking_uri: str = "http://localhost:5000",
        artifact_root: str = "s3://mlflow-artifacts/",
        registry_db: Optional[str] = None
    ):
        """
        初始化 MLflow Registry
        
        Args:
            tracking_uri: MLflow Tracking Server URI
            artifact_root: Artifacts 存储根目录
            registry_db: Registry 数据库连接(可选)
        """
        self.tracking_uri = tracking_uri
        self.artifact_root = artifact_root
        
        # 设置 MLflow
        mlflow.set_tracking_uri(tracking_uri)
        mlflow.set_experiment("model_registry")
        
        # 创建客户端
        self.client = MlflowClient(tracking_uri)
    
    def register_model(
        self,
        model_name: str,
        model_uri: str,
        metadata: ModelMetadata,
        await_registration_for: int = 300
    ) -> int:
        """
        注册新模型版本
        
        Args:
            model_name: 模型名称
            model_uri: 模型 artifacts URI
            metadata: 模型元数据
            await_registration_for: 等待注册的超时时间
        
        Returns:
            模型版本号
        """
        # 创建或获取已注册的模型
        try:
            self.client.create_registered_model(model_name)
            print(f"创建新注册模型: {model_name}")
        except mlflow.exceptions.MlflowException as e:
            if "already exists" not in str(e):
                raise
            print(f"模型已存在: {model_name}")
        
        # 准备 tags
        tags = {
            f"meta:{k}": str(v) for k, v in metadata.tags.items()
        }
        tags.update({
            "model_hash": metadata.model_hash,
            "created_by": metadata.created_by,
            "base_model": metadata.base_model or "",
            "training_data": metadata.training_data,
        })
        
        # 创建模型版本
        desc = metadata.description or ""
        if metadata.metrics:
            best_metric = max(metadata.metrics.items(), key=lambda x: x[1])
            desc += f"\nBest metric: {best_metric[0]}={best_metric[1]:.4f}"
        
        mv = self.client.create_model_version(
            name=model_name,
            source=model_uri,
            description=desc,
            tags=tags
        )
        
        version = mv.version
        
        # 记录参数和指标
        self._log_params_and_metrics(model_name, version, metadata)
        
        # 记录自定义元数据
        self._log_custom_metadata(model_name, version, metadata)
        
        return version
    
    def _log_params_and_metrics(
        self,
        model_name: str,
        version: int,
        metadata: ModelMetadata
    ):
        """记录参数和指标到 run"""
        # 找到关联的 run
        runs = self.client.search_runs(
            experiment_ids=["0"],  # 或者搜索特定的 experiment
            filter_string=f"tag.mlflow.runName = '{model_name}_v{version}'",
            max_results=1
        )
        
        if runs:
            run_id = runs[0].info.run_id
            
            # 记录参数
            with mlflow.start_run(run_id=run_id):
                for key, value in metadata.hyperparameters.items():
                    mlflow.log_param(key, value)
                
                # 记录指标
                for key, value in metadata.metrics.items():
                    mlflow.log_metric(key, value)
                
                # 记录业务指标
                for key, value in metadata.business_metrics.items():
                    mlflow.log_metric(f"business_{key}", value)
    
    def _log_custom_metadata(
        self,
        model_name: str,
        version: int,
        metadata: ModelMetadata
    ):
        """记录自定义元数据"""
        custom_metadata = {
            "training_code": metadata.training_code,
            "preprocessing_pipeline": metadata.preprocessing_pipeline,
            "environment": json.dumps(metadata.environment),
            "signature_hash": metadata.signature_hash,
            "created_at": metadata.created_at,
        }
        
        self.client.set_model_version_tag(
            model_name, version, "custom_metadata", json.dumps(custom_metadata)
        )
    
    def transition_stage(
        self,
        model_name: str,
        version: int,
        stage: ModelStage,
        archive_existing: bool = True
    ):
        """
        转换模型阶段
        
        Args:
            model_name: 模型名称
            version: 模型版本
            stage: 目标阶段
            archive_existing: 是否归档同阶段的现有版本
        """
        # 如果目标阶段是 Production,先归档现有的
        if stage == ModelStage.PRODUCTION and archive_existing:
            existing = self.get_latest_versions(model_name, stage.value)
            for v in existing:
                print(f"归档现有生产版本: {v}")
                self.client.transition_model_version_stage(
                    model_name, v, ModelStage.ARCHIVED.value
                )
        
        # 转换阶段
        self.client.transition_model_version_stage(
            model_name, version, stage.value
        )
        
        # 记录转换
        self.client.set_model_version_tag(
            model_name, version, "stage_transition_time",
            datetime.now().isoformat()
        )
        
        print(f"模型 {model_name}:v{version} 已转换到 {stage.value} 阶段")
    
    def get_latest_versions(
        self,
        model_name: str,
        stage: Optional[str] = None
    ) -> List[Dict]:
        """
        获取模型的最新版本
        
        Args:
            model_name: 模型名称
            stage: 筛选特定阶段
        
        Returns:
            版本列表
        """
        filter_str = f"name = '{model_name}'"
        if stage:
            filter_str += f" AND stage = '{stage}'"
        
        versions = self.client.search_model_versions(filter_str)
        
        return [
            {
                "version": v.version,
                "stage": v.current_stage,
                "description": v.description,
                "created": v.creation_timestamp,
                "last_updated": v.last_updated_timestamp,
                "tags": dict(v.tags),
                "run_id": v.run_id,
            }
            for v in versions
        ]
    
    def get_model_details(
        self,
        model_name: str,
        version: Optional[int] = None,
        stage: Optional[str] = None
    ) -> Dict:
        """
        获取模型详情
        
        Args:
            model_name: 模型名称
            version: 指定版本
            stage: 或指定阶段
        
        Returns:
            模型详情
        """
        if version:
            mv = self.client.get_model_version(model_name, version)
        elif stage:
            # 获取阶段对应的版本
            versions = self.get_latest_versions(model_name, stage)
            if not versions:
                raise ValueError(f"找不到 {stage} 阶段的模型")
            mv = self.client.get_model_version(model_name, versions[0]["version"])
        else:
            raise ValueError("必须指定 version 或 stage")
        
        return {
            "name": mv.name,
            "version": mv.version,
            "stage": mv.current_stage,
            "description": mv.description,
            "source": mv.source,
            "run_id": mv.run_id,
            "status": mv.status,
            "tags": dict(mv.tags),
            "created": mv.creation_timestamp,
            "last_updated": mv.last_updated_timestamp,
        }
    
    def download_model(
        self,
        model_name: str,
        version: int,
        local_path: str = "/tmp/mlflow_model"
    ) -> str:
        """
        下载模型到本地
        
        Args:
            model_name: 模型名称
            version: 版本号
            local_path: 本地路径
        
        Returns:
            下载后的模型路径
        """
        model_uri = f"models:/{model_name}/{version}"
        return mlflow.artifacts.download_artifacts(
            artifact_uri=model_uri,
            dst_path=local_path
        )
    
    def compare_versions(
        self,
        model_name: str,
        versions: List[int]
    ) -> Dict:
        """
        对比多个版本的性能
        
        Args:
            model_name: 模型名称
            versions: 版本列表
        
        Returns:
            对比结果
        """
        results = []
        
        for v in versions:
            details = self.get_model_details(model_name, v)
            
            # 获取关联 run 的指标
            run_id = details.get("run_id")
            metrics = {}
            if run_id:
                try:
                    run = self.client.get_run(run_id)
                    metrics = {
                        k: v for k, v in run.data.metrics.items()
                        if not k.startswith("business_")
                    }
                except Exception:
                    pass
            
            results.append({
                "version": v,
                "stage": details["stage"],
                "metrics": metrics,
                "tags": details.get("tags", {}),
            })
        
        return {"model_name": model_name, "versions": results}
    
    def delete_version(self, model_name: str, version: int):
        """删除模型版本"""
        self.client.delete_model_version(model_name, version)
        print(f"已删除 {model_name}:v{version}")
    
    def archive_old_versions(
        self,
        model_name: str,
        keep_last_n: int = 5,
        stages: List[str] = None
    ):
        """
        自动归档旧版本
        
        Args:
            model_name: 模型名称
            keep_last_n: 保留最近 N 个版本
            stages: 要处理的阶段列表
        """
        if stages is None:
            stages = [ModelStage.STAGING.value]
        
        versions = self.client.search_model_versions(
            f"name = '{model_name}'"
        )
        
        # 按创建时间排序
        versions = sorted(versions, key=lambda x: x.creation_timestamp)
        
        # 保留最近 N 个
        to_archive = versions[:-keep_last_n]
        
        for v in to_archive:
            if v.current_stage in stages:
                self.transition_stage(model_name, v.version, ModelStage.ARCHIVED)


# 使用示例
def mlflow_example():
    """MLflow 使用示例"""
    
    # 初始化 Registry
    registry = MLflowModelRegistry(
        tracking_uri="http://mlflow-server:5000",
        artifact_root="s3://my-mlflow-artifacts/"
    )
    
    # 准备元数据
    metadata = ModelMetadata(
        model_name="sentiment-classifier",
        model_version="v1.0.0",
        model_hash="a1b2c3d4e5f6",
        base_model="bert-base-chinese",
        training_code="abc123def456",
        training_data="sentiment_data_v2",
        eval_data="sentiment_eval_v1",
        preprocessing_pipeline="preprocess_v1.2",
        hyperparameters={
            "learning_rate": 2e-5,
            "batch_size": 32,
            "epochs": 3,
            "max_seq_length": 128,
        },
        training_seed=42,
        environment={
            "torch": "2.0.1",
            "transformers": "4.30.0",
        },
        metrics={
            "accuracy": 0.924,
            "f1": 0.918,
            "precision": 0.921,
            "recall": 0.915,
        },
        business_metrics={
            "latency_ms": 45.2,
            "throughput_qps": 1200,
        },
        description="基于 BERT 的中文情感分类模型",
        tags={
            "task": "classification",
            "language": "chinese",
            "domain": "sentiment",
        },
        created_by="zhangsan",
        created_at=datetime.now().isoformat(),
    )
    metadata.calculate_signature_hash()
    
    # 注册模型
    version = registry.register_model(
        model_name="sentiment-classifier",
        model_uri="runs:/abc123/model",
        metadata=metadata
    )
    print(f"注册成功,版本号: {version}")
    
    # 转换到预发布阶段
    registry.transition_stage("sentiment-classifier", version, ModelStage.STAGING)
    
    # 验证通过后,上线生产
    registry.transition_stage("sentiment-classifier", version, ModelStage.PRODUCTION)
    
    # 查看最新版本
    latest = registry.get_latest_versions("sentiment-classifier")
    print(f"最新版本: {latest}")
    
    # 对比版本
    comparison = registry.compare_versions("sentiment-classifier", [1, 2, 3])
    print(f"版本对比: {comparison}")


if __name__ == "__main__":
    mlflow_example()

2.3 MLflow 配置

# docker-compose.mlflow.yml
version: '3.8'

services:
  # MLflow Tracking Server
  mlflow:
    image: ghcr.io/mlflow/mlflow:v2.10.0
    container_name: mlflow-server
    restart: unless-stopped
    ports:
      - "5000:5000"
    environment:
      - MLFLOW_TRACKING_URI=postgresql://mlflow:mlflow@postgres:5432/mlflow
      - MLFLOW_ARTIFACT_ROOT=s3://mlflow-artifacts/
      - AWS_ACCESS_KEY_ID=${AWS_ACCESS_KEY_ID}
      - AWS_SECRET_ACCESS_KEY=${AWS_SECRET_ACCESS_KEY}
      - AWS_DEFAULT_REGION=us-east-1
    volumes:
      - mlflow_data:/mlflow
    depends_on:
      - postgres
    networks:
      - mlflow-net
    command: >
      mlflow server
      --host 0.0.0.0
      --port 5000
      --backend-store-uri postgresql://mlflow:mlflow@postgres:5432/mlflow
      --default-artifact-root s3://mlflow-artifacts/
      --workers 4

  # PostgreSQL for Metadata Store
  postgres:
    image: postgres:15-alpine
    container_name: mlflow-postgres
    restart: unless-stopped
    environment:
      - POSTGRES_USER=mlflow
      - POSTGRES_PASSWORD=${MLFLOW_DB_PASSWORD}
      - POSTGRES_DB=mlflow
    volumes:
      - postgres_data:/var/lib/postgresql/data
    networks:
      - mlflow-net
    command: >
      postgres
      -c max_connections=200
      -c shared_buffers=256MB
      -c effective_cache_size=1GB

  # MinIO for S3-compatible storage (development)
  minio:
    image: minio/minio:latest
    container_name: mlflow-minio
    restart: unless-stopped
    ports:
      - "9000:9000"
      - "9001:9001"
    environment:
      - MINIO_ROOT_USER=${MINIO_ACCESS_KEY}
      - MINIO_ROOT_PASSWORD=${MINIO_SECRET_KEY}
    volumes:
      - minio_data:/data
    networks:
      - mlflow-net
    command: server /data --console-address ":9001"

networks:
  mlflow-net:
    driver: bridge

volumes:
  mlflow_data:
  postgres_data:
  minio_data:

三、自建 Registry 的架构设计

3.1 什么时候需要自建 Registry

┌─────────────────────────────────────────────────────────────────┐
│                  选择自建 Registry 的信号                       │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│   信号1: MLflow 无法满足的需求                                  │
│   ├─ 需要细粒度的权限控制 (行级/列级)                           │
│   ├─ 需要与内部审批流程深度集成                                 │
│   ├─ 需要支持复杂的版本血缘关系                                │
│   └─ 需要跨云/混合云部署                                        │
│                                                                 │
│   信号2: 团队规模增长                                           │
│   ├─ 10+ 个模型                                                │
│   ├─ 100+ 个模型版本                                           │
│   ├─ 20+ 个开发者                                              │
│   └─ 多个团队协作                                              │
│                                                                 │
│   信号3: 合规要求                                               │
│   ├─ 需要完整的审计日志                                        │
│   ├─ 需要数据本地化存储                                         │
│   ├─ 需要模型可解释性报告                                      │
│   └─ 需要版本对比和回滚能力                                    │
│                                                                 │
│   信号4: 性能要求                                               │
│   ├─ 毫秒级的模型加载时间                                      │
│   ├─ 大规模模型版本搜索                                        │
│   └─ 高并发的模型元数据查询                                    │
│                                                                 │
│   成本考量:                                                    │
│   - MLflow: 2-4 周集成                                          │
│   - 自建 Registry: 3-6 个月开发                                 │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

3.2 自建 Registry 架构

┌─────────────────────────────────────────────────────────────────┐
│                    自建 Model Registry 架构                     │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│                         ┌─────────────────┐                      │
│                         │   API Gateway  │                      │
│                         │   (Kong/Nginx) │                      │
│                         └────────┬────────┘                      │
│                                  │                               │
│        ┌─────────────────────────┼─────────────────────────┐     │
│        │                         │                         │     │
│        ▼                         ▼                         ▼     │
│  ┌───────────┐           ┌───────────────┐         ┌─────────┐ │
│  │  Admin    │           │  Model CRUD   │         │ Version │ │
│  │  Portal   │           │    Service    │         │ Service │ │
│  └───────────┘           └───────┬───────┘         └────┬────┘ │
│                                  │                       │       │
│                                  ▼                       ▼       │
│                         ┌─────────────────────────────────┐     │
│                         │        Service Layer            │     │
│                         │  (Business Logic / Workflow)    │     │
│                         └────────────────┬────────────────┘     │
│                                          │                       │
│        ┌──────────────────────────────────┼────────────────┐      │
│        │                                  │                │      │
│        ▼                                  ▼                ▼      │
│  ┌───────────┐                    ┌───────────┐    ┌─────────┐  │
│  │ PostgreSQL│                    │   Redis   │    │  S3/MinIO│  │
│  │(Metadata) │                    │ (Cache)   │    │(Artifacts│  │
│  └───────────┘                    └───────────┘    └─────────┘  │
│                                                                 │
│  ┌───────────────────────────────────────────────────────────┐  │
│  │                    Artifact Storage                       │  │
│  │  ┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐        │  │
│  │  │ Model.1 │ │ Model.2 │ │ Model.3 │ │  ...    │        │  │
│  │  │  7B     │ │  13B    │ │  70B    │ │         │        │  │
│  │  └─────────┘ └─────────┘ └─────────┘ └─────────┘        │  │
│  └───────────────────────────────────────────────────────────┘  │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

3.3 数据库设计

-- schema.sql
-- Model Registry 数据库 Schema

-- 启用 UUID 扩展
CREATE EXTENSION IF NOT EXISTS "uuid-ossp";
CREATE EXTENSION IF NOT EXISTS "pg_trgm";  -- 用于模糊搜索

-- 1. 模型注册表 (顶层容器)
CREATE TABLE models (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    name VARCHAR(255) NOT NULL UNIQUE,
    description TEXT,
    model_type VARCHAR(50) NOT NULL,  -- 'classification', 'nlp', 'recommendation', etc.
    framework VARCHAR(50),  -- 'pytorch', 'tensorflow', 'onnx'
    owner_team VARCHAR(100),
    created_by VARCHAR(100) NOT NULL,
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    updated_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    
    CONSTRAINT unique_model_name UNIQUE (name)
);

CREATE INDEX idx_models_name ON models(name);
CREATE INDEX idx_models_type ON models(model_type);
CREATE INDEX idx_models_owner ON models(owner_team);

-- 2. 模型版本表
CREATE TABLE model_versions (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
    version VARCHAR(50) NOT NULL,
    version_number SERIAL,
    
    -- 存储信息
    artifact_uri TEXT NOT NULL,  -- S3 路径
    artifact_size BIGINT,  -- bytes
    artifact_checksum VARCHAR(64),  -- SHA256
    format VARCHAR(20) DEFAULT 'pytorch',  -- pytorch, tensorflow, onnx, etc.
    
    -- 血缘信息
    base_model_id UUID REFERENCES model_versions(id),  -- 基础模型 (微调场景)
    training_run_id VARCHAR(100),  -- 关联的训练 run
    training_code_commit VARCHAR(50),
    training_data_version VARCHAR(100),
    eval_data_version VARCHAR(100),
    preprocessing_pipeline_version VARCHAR(50),
    
    -- 训练配置
    hyperparameters JSONB DEFAULT '{}',
    training_seed INTEGER,
    environment JSONB DEFAULT '{}',  -- pip freeze
    
    -- 签名 (用于去重)
    signature_hash VARCHAR(64) UNIQUE,
    
    -- 生命周期
    stage VARCHAR(20) DEFAULT 'development' 
        CHECK (stage IN ('development', 'staging', 'production', 'archived')),
    status VARCHAR(20) DEFAULT 'pending' 
        CHECK (status IN ('pending', 'validated', 'rejected', 'deprecated')),
    
    -- 评审信息
    approved_by VARCHAR(100),
    approved_at TIMESTAMP WITH TIME ZONE,
    rejection_reason TEXT,
    
    -- 其他
    description TEXT,
    created_by VARCHAR(100) NOT NULL,
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    updated_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    
    CONSTRAINT unique_model_version UNIQUE (model_id, version)
);

CREATE INDEX idx_versions_model ON model_versions(model_id);
CREATE INDEX idx_versions_stage ON model_versions(stage);
CREATE INDEX idx_versions_status ON model_versions(status);
CREATE INDEX idx_versions_signature ON model_versions(signature_hash);
CREATE INDEX idx_versions_created ON model_versions(created_at DESC);

-- 3. 模型指标表
CREATE TABLE model_metrics (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    version_id UUID NOT NULL REFERENCES model_versions(id) ON DELETE CASCADE,
    
    -- 指标类型
    metric_type VARCHAR(50) NOT NULL,  -- 'accuracy', 'latency', 'custom'
    metric_name VARCHAR(100) NOT NULL,
    metric_value DOUBLE PRECISION NOT NULL,
    
    -- 上下文
    eval_dataset VARCHAR(100),
    eval_subset VARCHAR(50),  -- 'test', 'validation', 'holdout'
    
    -- 详细结果 (用于对比)
    detailed_results JSONB DEFAULT '{}',
    
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    
    CONSTRAINT unique_metric UNIQUE (version_id, metric_type, metric_name, eval_dataset)
);

CREATE INDEX idx_metrics_version ON model_metrics(version_id);
CREATE INDEX idx_metrics_type ON model_metrics(metric_type);

-- 4. 模型标签表
CREATE TABLE model_tags (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    version_id UUID NOT NULL REFERENCES model_versions(id) ON DELETE CASCADE,
    tag_key VARCHAR(100) NOT NULL,
    tag_value TEXT,
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    
    CONSTRAINT unique_tag UNIQUE (version_id, tag_key)
);

CREATE INDEX idx_tags_version ON model_tags(version_id);
CREATE INDEX idx_tags_key ON model_tags(tag_key);

-- 5. 版本阶段历史表 (审计用)
CREATE TABLE stage_transitions (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    version_id UUID NOT NULL REFERENCES model_versions(id) ON DELETE CASCADE,
    
    from_stage VARCHAR(20),
    to_stage VARCHAR(20) NOT NULL,
    
    triggered_by VARCHAR(100) NOT NULL,
    trigger_reason TEXT,
    approval_id UUID,  -- 关联审批记录
    
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
);

CREATE INDEX idx_transitions_version ON stage_transitions(version_id);
CREATE INDEX idx_transitions_time ON stage_transitions(created_at DESC);

-- 6. 模型对比记录表
CREATE TABLE comparison_reports (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    version_ids UUID[] NOT NULL,  -- 数组存储多个版本
    comparison_config JSONB DEFAULT '{}',
    comparison_results JSONB NOT NULL,
    winner_version_id UUID REFERENCES model_versions(id),
    
    created_by VARCHAR(100) NOT NULL,
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
);

CREATE INDEX idx_comparison_time ON comparison_reports(created_at DESC);

-- 7. Webhook 配置表
CREATE TABLE webhooks (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    name VARCHAR(100) NOT NULL,
    url TEXT NOT NULL,
    events VARCHAR[] NOT NULL,  -- 'version_created', 'stage_changed', etc.
    secret VARCHAR(255),  -- HMAC 签名密钥
    active BOOLEAN DEFAULT true,
    
    created_by VARCHAR(100) NOT NULL,
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
    updated_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
);

-- 8. 审批工作流表
CREATE TABLE approval_workflows (
    id UUID PRIMARY KEY DEFAULT uuid_generate_v4(),
    version_id UUID NOT NULL REFERENCES model_versions(id) ON DELETE CASCADE,
    workflow_type VARCHAR(50) NOT NULL,  -- 'auto', 'manual'
    
    -- 工作流配置
    required_approvals INTEGER DEFAULT 1,
    approvers VARCHAR[],  -- 指定的审批人列表
    
    -- 状态
    status VARCHAR(20) DEFAULT 'pending' 
        CHECK (status IN ('pending', 'approved', 'rejected', 'skipped')),
    current_step INTEGER DEFAULT 0,
    
    -- 结果
    decision TEXT,
    decided_by VARCHAR(100),
    decided_at TIMESTAMP WITH TIME ZONE,
    
    created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
);

CREATE INDEX idx_workflow_version ON approval_workflows(version_id);
CREATE INDEX idx_workflow_status ON approval_workflows(status);

-- 视图: 模型版本完整信息
CREATE VIEW v_model_version_full AS
SELECT 
    mv.*,
    m.name AS model_name,
    m.model_type,
    m.framework,
    m.owner_team,
    COALESCE(
        json_agg(json_build_object('key', mt.tag_key, 'value', mt.tag_value))
        FILTER (WHERE mt.id IS NOT NULL),
        '[]'
    ) AS tags,
    COALESCE(
        json_agg(json_build_object(
            'type', mm.metric_type,
            'name', mm.metric_name,
            'value', mm.metric_value
        )) FILTER (WHERE mm.id IS NOT NULL),
        '[]'
    ) AS metrics
FROM model_versions mv
JOIN models m ON mv.model_id = m.id
LEFT JOIN model_tags mt ON mv.id = mt.version_id
LEFT JOIN model_metrics mm ON mv.id = mm.version_id
GROUP BY mv.id, m.name, m.model_type, m.framework, m.owner_team;

3.4 自建 Registry API 设计

# model_registry_api.py
"""
自建 Model Registry API 实现
FastAPI + SQLAlchemy + Pydantic
"""

import os
import uuid
import hashlib
import json
from datetime import datetime
from typing import Dict, List, Optional, Any
from enum import Enum
from dataclasses import dataclass, field, asdict

from fastapi import FastAPI, HTTPException, Depends, BackgroundTasks, status
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field, field_validator
from sqlalchemy import create_engine, Column, String, Integer, DateTime, JSON, Text, BigInteger, Boolean, ARRAY, ForeignKey, Enum as SQLEnum
from sqlalchemy.dialects.postgresql import UUID, JSONB
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker, Session, relationship
from sqlalchemy.pool import QueuePool
import boto3
from botocore.exceptions import ClientError

# ============== 配置 ==============
DATABASE_URL = os.getenv("DATABASE_URL", "postgresql://registry:registry@localhost:5432/registry")
S3_ENDPOINT = os.getenv("S3_ENDPOINT", "http://localhost:9000")
S3_BUCKET = os.getenv("S3_BUCKET", "model-registry")
AWS_ACCESS_KEY = os.getenv("AWS_ACCESS_KEY_ID")
AWS_SECRET_KEY = os.getenv("AWS_SECRET_ACCESS_KEY")
REGION = os.getenv("AWS_DEFAULT_REGION", "us-east-1")

# ============== 数据库模型 ==============
Base = declarative_base()


class ModelDB(Base):
    """模型数据库模型"""
    __tablename__ = "models"
    
    id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
    name = Column(String(255), unique=True, nullable=False)
    description = Column(Text)
    model_type = Column(String(50), nullable=False)
    framework = Column(String(50))
    owner_team = Column(String(100))
    created_by = Column(String(100), nullable=False)
    created_at = Column(DateTime, default=datetime.utcnow)
    updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
    
    versions = relationship("ModelVersionDB", back_populates="model", cascade="all, delete-orphan")


class ModelVersionDB(Base):
    """模型版本数据库模型"""
    __tablename__ = "model_versions"
    
    id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
    model_id = Column(UUID(as_uuid=True), ForeignKey("models.id", ondelete="CASCADE"))
    version = Column(String(50), nullable=False)
    version_number = Column(Integer)
    
    # 存储
    artifact_uri = Column(Text, nullable=False)
    artifact_size = Column(BigInteger)
    artifact_checksum = Column(String(64))
    format = Column(String(20), default="pytorch")
    
    # 血缘
    base_model_id = Column(UUID(as_uuid=True))
    training_run_id = Column(String(100))
    training_code_commit = Column(String(50))
    training_data_version = Column(String(100))
    eval_data_version = Column(String(100))
    preprocessing_pipeline_version = Column(String(50))
    
    # 训练配置
    hyperparameters = Column(JSONB, default=dict)
    training_seed = Column(Integer)
    environment = Column(JSONB, default=dict)
    
    # 签名
    signature_hash = Column(String(64), unique=True)
    
    # 生命周期
    stage = Column(String(20), default="development")
    status = Column(String(20), default="pending")
    
    # 评审
    approved_by = Column(String(100))
    approved_at = Column(DateTime)
    rejection_reason = Column(Text)
    
    # 其他
    description = Column(Text)
    created_by = Column(String(100), nullable=False)
    created_at = Column(DateTime, default=datetime.utcnow)
    updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
    
    model = relationship("ModelDB", back_populates="versions")
    metrics = relationship("ModelMetricDB", back_populates="version", cascade="all, delete-orphan")
    tags = relationship("ModelTagDB", back_populates="version", cascade="all, delete-orphan")


class ModelMetricDB(Base):
    """模型指标数据库模型"""
    __tablename__ = "model_metrics"
    
    id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
    version_id = Column(UUID(as_uuid=True), ForeignKey("model_versions.id", ondelete="CASCADE"))
    
    metric_type = Column(String(50), nullable=False)
    metric_name = Column(String(100), nullable=False)
    metric_value = Column(BigInteger, nullable=False)
    
    eval_dataset = Column(String(100))
    eval_subset = Column(String(50))
    detailed_results = Column(JSONB, default=dict)
    
    created_at = Column(DateTime, default=datetime.utcnow)
    
    version = relationship("ModelVersionDB", back_populates="metrics")


class ModelTagDB(Base):
    """模型标签数据库模型"""
    __tablename__ = "model_tags"
    
    id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
    version_id = Column(UUID(as_uuid=True), ForeignKey("model_versions.id", ondelete="CASCADE"))
    tag_key = Column(String(100), nullable=False)
    tag_value = Column(Text)
    created_at = Column(DateTime, default=datetime.utcnow)
    
    version = relationship("ModelVersionDB", back_populates="tags")


class StageTransitionDB(Base):
    """阶段转换历史"""
    __tablename__ = "stage_transitions"
    
    id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
    version_id = Column(UUID(as_uuid=True), ForeignKey("model_versions.id", ondelete="CASCADE"))
    from_stage = Column(String(20))
    to_stage = Column(String(20), nullable=False)
    triggered_by = Column(String(100), nullable=False)
    trigger_reason = Column(Text)
    created_at = Column(DateTime, default=datetime.utcnow)


# ============== Pydantic 模型 ==============
class ModelStage(str, Enum):
    DEVELOPMENT = "development"
    STAGING = "staging"
    PRODUCTION = "production"
    ARCHIVED = "archived"


class ModelStatus(str, Enum):
    PENDING = "pending"
    VALIDATED = "validated"
    REJECTED = "rejected"
    DEPRECATED = "deprecated"


class ModelCreate(BaseModel):
    """创建模型请求"""
    name: str = Field(..., min_length=1, max_length=255)
    description: Optional[str] = None
    model_type: str = Field(..., min_length=1)
    framework: Optional[str] = None
    owner_team: Optional[str] = None


class ModelResponse(BaseModel):
    """模型响应"""
    id: str
    name: str
    description: Optional[str]
    model_type: str
    framework: Optional[str]
    owner_team: Optional[str]
    created_by: str
    created_at: datetime
    latest_version: Optional[str]
    production_version: Optional[str]
    
    class Config:
        from_attributes = True


class ModelVersionCreate(BaseModel):
    """创建模型版本请求"""
    version: str = Field(..., min_length=1)
    artifact_path: str  # S3 路径
    format: str = "pytorch"
    
    # 血缘
    base_model_version_id: Optional[str] = None
    training_run_id: Optional[str] = None
    training_code_commit: Optional[str] = None
    training_data_version: Optional[str] = None
    eval_data_version: Optional[str] = None
    preprocessing_pipeline_version: Optional[str] = None
    
    # 配置
    hyperparameters: Dict[str, Any] = Field(default_factory=dict)
    training_seed: Optional[int] = 42
    environment: Dict[str, str] = Field(default_factory=dict)
    
    # 指标
    metrics: Dict[str, float] = Field(default_factory=dict)
    metric_type: str = "accuracy"
    eval_dataset: Optional[str] = None
    
    # 标签
    tags: Dict[str, str] = Field(default_factory=dict)
    
    description: Optional[str] = None


class ModelVersionResponse(BaseModel):
    """模型版本响应"""
    id: str
    model_id: str
    model_name: str
    version: str
    artifact_uri: str
    artifact_size: Optional[int]
    artifact_checksum: Optional[str]
    format: str
    
    # 血缘
    base_model_id: Optional[str]
    training_code_commit: Optional[str]
    training_data_version: Optional[str]
    
    # 配置
    hyperparameters: Dict[str, Any]
    training_seed: Optional[int]
    environment: Dict[str, str]
    
    signature_hash: Optional[str]
    stage: str
    status: str
    
    metrics: List[Dict[str, Any]]
    tags: Dict[str, str]
    
    created_by: str
    created_at: datetime
    
    class Config:
        from_attributes = True


class StageTransitionRequest(BaseModel):
    """阶段转换请求"""
    target_stage: ModelStage
    trigger_reason: Optional[str] = None
    approval_id: Optional[str] = None


class VersionComparisonRequest(BaseModel):
    """版本对比请求"""
    version_ids: List[str] = Field(..., min_length=2, max_length=10)
    metrics_to_compare: List[str] = Field(default=["accuracy", "f1"])


# ============== API 实现 ==============
app = FastAPI(title="Model Registry API", version="1.0.0")

app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)

# 数据库连接
engine = create_engine(DATABASE_URL, poolclass=QueuePool, pool_size=10, max_overflow=20)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)


def get_db():
    """数据库会话依赖"""
    db = SessionLocal()
    try:
        yield db
    finally:
        db.close()


# S3 客户端
def get_s3_client():
    """获取 S3 客户端"""
    return boto3.client(
        "s3",
        endpoint_url=S3_ENDPOINT,
        aws_access_key_id=AWS_ACCESS_KEY,
        aws_secret_access_key=AWS_SECRET_KEY,
        region_name=REGION
    )


# ============== API 端点 ==============

@app.get("/health")
async def health_check():
    """健康检查"""
    return {"status": "healthy", "timestamp": datetime.utcnow().isoformat()}


@app.post("/models", response_model=ModelResponse)
async def create_model(model: ModelCreate, db: Session = Depends(get_db)):
    """创建新模型"""
    # 检查是否已存在
    existing = db.query(ModelDB).filter(ModelDB.name == model.name).first()
    if existing:
        raise HTTPException(status_code=400, detail=f"Model '{model.name}' already exists")
    
    db_model = ModelDB(
        name=model.name,
        description=model.description,
        model_type=model.model_type,
        framework=model.framework,
        owner_team=model.owner_team,
        created_by="system"  # TODO: 从认证获取
    )
    db.add(db_model)
    db.commit()
    db.refresh(db_model)
    
    return ModelResponse(
        id=str(db_model.id),
        name=db_model.name,
        description=db_model.description,
        model_type=db_model.model_type,
        framework=db_model.framework,
        owner_team=db_model.owner_team,
        created_by=db_model.created_by,
        created_at=db_model.created_at,
        latest_version=None,
        production_version=None
    )


@app.get("/models", response_model=List[ModelResponse])
async def list_models(
    model_type: Optional[str] = None,
    owner_team: Optional[str] = None,
    skip: int = 0,
    limit: int = 100,
    db: Session = Depends(get_db)
):
    """列出所有模型"""
    query = db.query(ModelDB)
    
    if model_type:
        query = query.filter(ModelDB.model_type == model_type)
    if owner_team:
        query = query.filter(ModelDB.owner_team == owner_team)
    
    models = query.offset(skip).limit(limit).all()
    
    results = []
    for m in models:
        # 获取最新版本
        latest = db.query(ModelVersionDB).filter(
            ModelVersionDB.model_id == m.id
        ).order_by(ModelVersionDB.created_at.desc()).first()
        
        # 获取生产版本
        prod = db.query(ModelVersionDB).filter(
            ModelVersionDB.model_id == m.id,
            ModelVersionDB.stage == "production"
        ).first()
        
        results.append(ModelResponse(
            id=str(m.id),
            name=m.name,
            description=m.description,
            model_type=m.model_type,
            framework=m.framework,
            owner_team=m.owner_team,
            created_by=m.created_by,
            created_at=m.created_at,
            latest_version=latest.version if latest else None,
            production_version=prod.version if prod else None
        ))
    
    return results


@app.get("/models/{model_id}", response_model=ModelResponse)
async def get_model(model_id: str, db: Session = Depends(get_db)):
    """获取模型详情"""
    model = db.query(ModelDB).filter(ModelDB.id == model_id).first()
    if not model:
        raise HTTPException(status_code=404, detail="Model not found")
    
    latest = db.query(ModelVersionDB).filter(
        ModelVersionDB.model_id == model.id
    ).order_by(ModelVersionDB.created_at.desc()).first()
    
    prod = db.query(ModelVersionDB).filter(
        ModelVersionDB.model_id == model.id,
        ModelVersionDB.stage == "production"
    ).first()
    
    return ModelResponse(
        id=str(model.id),
        name=model.name,
        description=model.description,
        model_type=model.model_type,
        framework=model.framework,
        owner_team=model.owner_team,
        created_by=model.created_by,
        created_at=model.created_at,
        latest_version=latest.version if latest else None,
        production_version=prod.version if prod else None
    )


@app.post("/models/{model_id}/versions", response_model=ModelVersionResponse)
async def create_version(
    model_id: str,
    version: ModelVersionCreate,
    background_tasks: BackgroundTasks,
    db: Session = Depends(get_db)
):
    """创建模型版本"""
    model = db.query(ModelDB).filter(ModelDB.id == model_id).first()
    if not model:
        raise HTTPException(status_code=404, detail="Model not found")
    
    # 计算版本号
    latest = db.query(ModelVersionDB).filter(
        ModelVersionDB.model_id == model.id
    ).order_by(ModelVersionDB.version_number.desc()).first()
    version_number = (latest.version_number or 0) + 1
    
    # 计算签名 hash (用于去重)
    signature_data = {
        "hyperparameters": version.hyperparameters,
        "training_seed": version.training_seed,
        "training_data_version": version.training_data_version,
    }
    signature_hash = hashlib.sha256(
        json.dumps(signature_data, sort_keys=True).encode()
    ).hexdigest()
    
    # 检查是否重复
    existing = db.query(ModelVersionDB).filter(
        ModelVersionDB.signature_hash == signature_hash
    ).first()
    if existing:
        raise HTTPException(
            status_code=409,
            detail=f"Duplicate model version with signature {signature_hash[:8]}"
        )
    
    # 计算 artifact checksum
    s3 = get_s3_client()
    artifact_key = f"{model.name}/{version.version}/{version.artifact_path}"
    full_uri = f"s3://{S3_BUCKET}/{artifact_key}"
    
    artifact_size = 0
    artifact_checksum = None
    try:
        response = s3.head_object(Bucket=S3_BUCKET, Key=artifact_key)
        artifact_size = response["ContentLength"]
        # ETag 可能不是 MD5,仅供参考
        artifact_checksum = response.get("ETag", "").strip('"')
    except ClientError:
        pass  # 可能文件不存在,忽略
    
    # 创建版本记录
    db_version = ModelVersionDB(
        model_id=model.id,
        version=version.version,
        version_number=version_number,
        artifact_uri=full_uri,
        artifact_size=artifact_size,
        artifact_checksum=artifact_checksum,
        format=version.format,
        base_model_id=version.base_model_version_id,
        training_run_id=version.training_run_id,
        training_code_commit=version.training_code_commit,
        training_data_version=version.training_data_version,
        eval_data_version=version.eval_data_version,
        preprocessing_pipeline_version=version.preprocessing_pipeline_version,
        hyperparameters=version.hyperparameters,
        training_seed=version.training_seed,
        environment=version.environment,
        signature_hash=signature_hash,
        stage="development",
        status="pending",
        description=version.description,
        created_by="system"
    )
    db.add(db_version)
    db.flush()
    
    # 添加指标
    for metric_name, metric_value in version.metrics.items():
        db_metric = ModelMetricDB(
            version_id=db_version.id,
            metric_type=version.metric_type,
            metric_name=metric_name,
            metric_value=int(metric_value * 10000),  # 存储为整数避免精度问题
            eval_dataset=version.eval_dataset
        )
        db.add(db_metric)
    
    # 添加标签
    for key, value in version.tags.items():
        db_tag = ModelTagDB(
            version_id=db_version.id,
            tag_key=key,
            tag_value=str(value)
        )
        db.add(db_tag)
    
    db.commit()
    db.refresh(db_version)
    
    return ModelVersionResponse(
        id=str(db_version.id),
        model_id=str(db_version.model_id),
        model_name=model.name,
        version=db_version.version,
        artifact_uri=db_version.artifact_uri,
        artifact_size=db_version.artifact_size,
        artifact_checksum=db_version.artifact_checksum,
        format=db_version.format,
        base_model_id=str(db_version.base_model_id) if db_version.base_model_id else None,
        training_code_commit=db_version.training_code_commit,
        training_data_version=db_version.training_data_version,
        hyperparameters=db_version.hyperparameters,
        training_seed=db_version.training_seed,
        environment=db_version.environment,
        signature_hash=db_version.signature_hash,
        stage=db_version.stage,
        status=db_version.status,
        metrics=[{"name": m.metric_name, "value": m.metric_value / 10000} 
                 for m in db_version.metrics],
        tags={t.tag_key: t.tag_value for t in db_version.tags},
        created_by=db_version.created_by,
        created_at=db_version.created_at
    )


@app.get("/models/{model_id}/versions", response_model=List[ModelVersionResponse])
async def list_versions(
    model_id: str,
    stage: Optional[str] = None,
    skip: int = 0,
    limit: int = 100,
    db: Session = Depends(get_db)
):
    """列出模型的版本"""
    query = db.query(ModelVersionDB).filter(ModelVersionDB.model_id == model_id)
    
    if stage:
        query = query.filter(ModelVersionDB.stage == stage)
    
    versions = query.order_by(ModelVersionDB.created_at.desc()).offset(skip).limit(limit).all()
    model = db.query(ModelDB).filter(ModelDB.id == model_id).first()
    
    return [
        ModelVersionResponse(
            id=str(v.id),
            model_id=str(v.model_id),
            model_name=model.name if model else "",
            version=v.version,
            artifact_uri=v.artifact_uri,
            artifact_size=v.artifact_size,
            artifact_checksum=v.artifact_checksum,
            format=v.format,
            base_model_id=str(v.base_model_id) if v.base_model_id else None,
            training_code_commit=v.training_code_commit,
            training_data_version=v.training_data_version,
            hyperparameters=v.hyperparameters,
            training_seed=v.training_seed,
            environment=v.environment,
            signature_hash=v.signature_hash,
            stage=v.stage,
            status=v.status,
            metrics=[{"name": m.metric_name, "value": m.metric_value / 10000} 
                     for m in v.metrics],
            tags={t.tag_key: t.tag_value for t in v.tags},
            created_by=v.created_by,
            created_at=v.created_at
        )
        for v in versions
    ]


@app.get("/versions/{version_id}", response_model=ModelVersionResponse)
async def get_version(version_id: str, db: Session = Depends(get_db)):
    """获取版本详情"""
    version = db.query(ModelVersionDB).filter(ModelVersionDB.id == version_id).first()
    if not version:
        raise HTTPException(status_code=404, detail="Version not found")
    
    model = db.query(ModelDB).filter(ModelDB.id == version.model_id).first()
    
    return ModelVersionResponse(
        id=str(version.id),
        model_id=str(version.model_id),
        model_name=model.name if model else "",
        version=version.version,
        artifact_uri=version.artifact_uri,
        artifact_size=version.artifact_size,
        artifact_checksum=version.artifact_checksum,
        format=version.format,
        base_model_id=str(version.base_model_id) if version.base_model_id else None,
        training_code_commit=version.training_code_commit,
        training_data_version=version.training_data_version,
        hyperparameters=version.hyperparameters,
        training_seed=version.training_seed,
        environment=version.environment,
        signature_hash=version.signature_hash,
        stage=version.stage,
        status=version.status,
        metrics=[{"name": m.metric_name, "value": m.metric_value / 10000} 
                 for m in version.metrics],
        tags={t.tag_key: t.tag_value for t in version.tags},
        created_by=version.created_by,
        created_at=version.created_at
    )


@app.post("/versions/{version_id}/transition")
async def transition_stage(
    version_id: str,
    request: StageTransitionRequest,
    db: Session = Depends(get_db)
):
    """转换模型阶段"""
    version = db.query(ModelVersionDB).filter(ModelVersionDB.id == version_id).first()
    if not version:
        raise HTTPException(status_code=404, detail="Version not found")
    
    # 验证转换规则
    current_stage = version.stage
    target_stage = request.target_stage.value
    
    # 业务规则验证
    if target_stage == "production":
        if version.status != "validated":
            raise HTTPException(
                status_code=400,
                detail="Only validated versions can transition to production"
            )
        
        # 如果有现有生产版本,需要先归档
        existing_prod = db.query(ModelVersionDB).filter(
            ModelVersionDB.model_id == version.model_id,
            ModelVersionDB.stage == "production",
            ModelVersionDB.id != version.id
        ).first()
        
        if existing_prod:
            # 创建转换历史
            transition = StageTransitionDB(
                version_id=existing_prod.id,
                from_stage="production",
                to_stage="archived",
                triggered_by="system",
                trigger_reason="Replaced by new production version"
            )
            db.add(transition)
            
            existing_prod.stage = "archived"
    
    # 创建转换历史
    transition = StageTransitionDB(
        version_id=version.id,
        from_stage=current_stage,
        to_stage=target_stage,
        triggered_by="system",
        trigger_reason=request.trigger_reason
    )
    db.add(transition)
    
    # 更新版本
    version.stage = target_stage
    version.updated_at = datetime.utcnow()
    
    db.commit()
    
    return {
        "message": f"Version transitioned from {current_stage} to {target_stage}",
        "version_id": str(version.id),
        "new_stage": target_stage
    }


@app.post("/versions/compare")
async def compare_versions(
    request: VersionComparisonRequest,
    db: Session = Depends(get_db)
):
    """对比多个版本"""
    versions = db.query(ModelVersionDB).filter(
        ModelVersionDB.id.in_(request.version_ids)
    ).all()
    
    if len(versions) != len(request.version_ids):
        found_ids = {str(v.id) for v in versions}
        missing = set(request.version_ids) - found_ids
        raise HTTPException(status_code=404, detail=f"Versions not found: {missing}")
    
    # 获取模型名称
    model_ids = {v.model_id for v in versions}
    models = db.query(ModelDB).filter(ModelDB.id.in_(model_ids)).all()
    model_names = {str(m.id): m.name for m in models}
    
    # 收集指标
    comparison = []
    for v in versions:
        metrics = {m.metric_name: m.metric_value / 10000 
                   for m in v.metrics 
                   if m.metric_name in request.metrics_to_compare}
        
        comparison.append({
            "version_id": str(v.id),
            "model_name": model_names.get(str(v.model_id), ""),
            "version": v.version,
            "stage": v.stage,
            "created_at": v.created_at.isoformat(),
            "metrics": metrics,
            "hyperparameters": v.hyperparameters
        })
    
    # 找出各指标最优版本
    winners = {}
    for metric in request.metrics_to_compare:
        best = max(comparison, key=lambda x: x["metrics"].get(metric, float("-inf")))
        if best["metrics"].get(metric):
            winners[metric] = {
                "version": best["version"],
                "value": best["metrics"][metric]
            }
    
    return {
        "versions": comparison,
        "winners": winners
    }


@app.get("/versions/{version_id}/download")
async def get_download_url(version_id: str, db: Session = Depends(get_db)):
    """获取模型下载 URL (预签名 URL)"""
    version = db.query(ModelVersionDB).filter(ModelVersionDB.id == version_id).first()
    if not version:
        raise HTTPException(status_code=404, detail="Version not found")
    
    s3 = get_s3_client()
    
    # 提取 S3 路径
    uri = version.artifact_uri.replace(f"s3://{S3_BUCKET}/", "")
    
    # 生成预签名 URL (有效期 1 小时)
    url = s3.generate_presigned_url(
        "get_object",
        Params={"Bucket": S3_BUCKET, "Key": uri},
        ExpiresIn=3600
    )
    
    return {
        "download_url": url,
        "expires_in": 3600,
        "artifact_size": version.artifact_size
    }


# 启动时创建表
@app.on_event("startup")
async def startup():
    Base.metadata.create_all(bind=engine)


# 使用示例
if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=8000)

四、模型版本生命周期

4.1 完整生命周期流程

┌─────────────────────────────────────────────────────────────────┐
│                    模型版本生命周期                              │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│   ┌─────────┐    ┌─────────┐    ┌───────────┐    ┌─────────┐  │
│   │ Develop │───►│  Stage  │───►│ Production│───►│ Archive │  │
│   │  开发   │    │  预发布 │    │   生产   │    │   归档  │  │
│   └─────────┘    └─────────┘    └───────────┘    └─────────┘  │
│        │              │               │               │        │
│        │              │               │               │        │
│        ▼              ▼               ▼               ▼        │
│   ┌─────────┐    ┌─────────┐    ┌───────────┐    ┌─────────┐  │
│   │ 自动验证 │    │ 人工评审 │    │ 灰度发布  │    │ 可恢复  │  │
│   │ 单元测试 │    │ 性能回归 │    │ 全量切换 │    │ 或删除  │  │
│   │ 格式检查 │    │ A/B测试  │    │ 监控告警 │    │         │  │
│   └─────────┘    └─────────┘    └───────────┘    └─────────┘  │
│                                                                 │
│   关键规则:                                                      │
│   1. Development → Staging: 自动验证通过                       │
│   2. Staging → Production: 人工审批                             │
│   3. Production → Archived: 新版本上线或人工触发               │
│   4. Archived 可恢复到任意阶段                                  │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

4.2 版本元数据设计

# metadata_design.py
"""
模型版本元数据设计
包含血缘追踪、性能指标、业务指标等
"""

from dataclasses import dataclass, field
from typing import Dict, List, Optional, Any
from datetime import datetime
from enum import Enum
import hashlib
import json


class ModelType(Enum):
    """模型类型"""
    CLASSIFICATION = "classification"
    REGRESSION = "regression"
    NLP = "nlp"
    RECOMMENDATION = "recommendation"
    GENERATION = "generation"
    MULTIMODAL = "multimodal"


class ModelFormat(Enum):
    """模型格式"""
    PYTORCH = "pytorch"
    TENSORFLOW = "tensorflow"
    ONNX = "onnx"
    JAX = "jax"
    TFLITE = "tflite"
    LLAMA_CPP = "llama.cpp"
    GGUF = "gguf"


@dataclass
class LineageInfo:
    """血缘信息"""
    base_model_id: Optional[str] = None  # 微调基础模型
    training_run_id: Optional[str] = None  # W&B/MLflow run ID
    training_code_commit: str = ""  # Git commit hash
    training_data_version: str = ""  # 数据集版本
    eval_data_version: str = ""  # 评测数据集版本
    preprocessing_pipeline_version: str = ""  # 预处理版本
    feature_pipeline_version: str = ""  # 特征工程版本
    
    # 数据集元数据
    training_samples: int = 0
    eval_samples: int = 0
    dataset_hash: str = ""  # 数据集内容的 hash
    
    def to_dict(self) -> Dict:
        return asdict(self)


@dataclass
class TrainingConfig:
    """训练配置"""
    hyperparameters: Dict[str, Any] = field(default_factory=dict)
    training_seed: int = 42
    optimizer: str = "adam"
    learning_rate: float = 1e-4
    batch_size: int = 32
    epochs: int = 10
    
    # 分布式训练配置
    num_gpus: int = 1
    gpu_type: str = ""
    training_time_hours: float = 0
    
    # 环境
    environment: Dict[str, str] = field(default_factory=dict)  # pip freeze
    cuda_version: str = ""
    python_version: str = ""
    
    def to_dict(self) -> Dict:
        return asdict(self)
    
    def calculate_signature(self) -> str:
        """计算配置签名(用于去重)"""
        sig_data = {
            "hyperparameters": self.hyperparameters,
            "training_seed": self.training_seed,
            "batch_size": self.batch_size,
            "epochs": self.epochs,
        }
        return hashlib.sha256(
            json.dumps(sig_data, sort_keys=True).encode()
        ).hexdigest()


@dataclass
class PerformanceMetrics:
    """性能指标"""
    # 准确性指标
    accuracy: Optional[float] = None
    precision: Optional[float] = None
    recall: Optional[float] = None
    f1_score: Optional[float] = None
    auc_roc: Optional[float] = None
    auc_pr: Optional[float] = None
    
    # 业务指标
    conversion_rate: Optional[float] = None
    ctr: Optional[float] = None
    retention: Optional[float] = None
    
    # 效率指标
    latency_p50_ms: Optional[float] = None
    latency_p95_ms: Optional[float] = None
    latency_p99_ms: Optional[float] = None
    throughput_qps: Optional[float] = None
    
    # 资源指标
    model_size_mb: Optional[float] = None
    memory_usage_mb: Optional[float] = None
    gpu_memory_mb: Optional[float] = None
    
    def to_dict(self) -> Dict:
        return {k: v for k, v in asdict(self).items() if v is not None}
    
    def get_summary(self) -> str:
        """获取指标摘要"""
        parts = []
        if self.accuracy is not None:
            parts.append(f"acc={self.accuracy:.4f}")
        if self.f1_score is not None:
            parts.append(f"f1={self.f1_score:.4f}")
        if self.latency_p95_ms is not None:
            parts.append(f"p95={self.latency_p95_ms:.1f}ms")
        return ", ".join(parts)


@dataclass
class ModelVersionMetadata:
    """完整的模型版本元数据"""
    # 基础信息
    model_name: str
    version: str
    
    # 模型信息
    model_type: ModelType
    model_format: ModelFormat
    architecture: str = ""  # 模型架构描述
    
    # 存储信息
    artifact_uri: str
    artifact_size_bytes: int = 0
    artifact_checksum: str = ""  # SHA256
    
    # 血缘
    lineage: LineageInfo = field(default_factory=LineageInfo)
    
    # 训练配置
    config: TrainingConfig = field(default_factory=TrainingConfig)
    config_signature: str = ""  # 配置签名
    
    # 性能指标
    metrics: PerformanceMetrics = field(default_factory=PerformanceMetrics)
    
    # 生命周期
    stage: str = "development"
    status: str = "pending"
    
    # 评审信息
    description: str = ""
    change_log: str = ""
    known_issues: List[str] = field(default_factory=list)
    recommendations: List[str] = field(default_factory=list)
    
    # 标签
    tags: Dict[str, str] = field(default_factory=dict)
    
    # 审计信息
    created_by: str = ""
    created_at: str = ""
    updated_by: str = ""
    updated_at: str = ""
    
    def __post_init__(self):
        """计算签名"""
        if not self.config_signature:
            self.config_signature = self.config.calculate_signature()
    
    def to_dict(self) -> Dict:
        """转换为字典"""
        return {
            "model_name": self.model_name,
            "version": self.version,
            "model_type": self.model_type.value if isinstance(self.model_type, Enum) else self.model_type,
            "model_format": self.model_format.value if isinstance(self.model_format, Enum) else self.model_format,
            "architecture": self.architecture,
            "artifact_uri": self.artifact_uri,
            "artifact_size_bytes": self.artifact_size_bytes,
            "artifact_checksum": self.artifact_checksum,
            "lineage": self.lineage.to_dict(),
            "config": self.config.to_dict(),
            "config_signature": self.config_signature,
            "metrics": self.metrics.to_dict(),
            "stage": self.stage,
            "status": self.status,
            "description": self.description,
            "change_log": self.change_log,
            "known_issues": self.known_issues,
            "recommendations": self.recommendations,
            "tags": self.tags,
            "created_by": self.created_by,
            "created_at": self.created_at,
            "updated_by": self.updated_by,
            "updated_at": self.updated_at,
        }
    
    def to_json(self, indent: int = 2) -> str:
        """转换为 JSON 字符串"""
        return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)
    
    @classmethod
    def from_dict(cls, data: Dict) -> "ModelVersionMetadata":
        """从字典创建"""
        # 处理枚举类型
        if isinstance(data.get("model_type"), str):
            data["model_type"] = ModelType(data["model_type"])
        if isinstance(data.get("model_format"), str):
            data["model_format"] = ModelFormat(data["model_format"])
        
        # 处理嵌套对象
        if "lineage" in data:
            data["lineage"] = LineageInfo(**data["lineage"])
        if "config" in data:
            data["config"] = TrainingConfig(**data["config"])
        if "metrics" in data:
            data["metrics"] = PerformanceMetrics(**data["metrics"])
        
        return cls(**data)


# 示例:创建完整的模型元数据
def create_example_metadata():
    """创建示例元数据"""
    metadata = ModelVersionMetadata(
        model_name="sentiment-classifier",
        version="v2.3.1",
        model_type=ModelType.NLP,
        model_format=ModelFormat.PYTORCH,
        architecture="bert-base-chinese",
        artifact_uri="s3://models/sentiment-classifier/v2.3.1/model.pt",
        artifact_size_bytes=420_000_000,  # 420MB
        artifact_checksum="a1b2c3d4e5f6...",
        lineage=LineageInfo(
            base_model_id="models/sentiment-classifier/v1.0.0",
            training_run_id="wandb/run-12345",
            training_code_commit="abc123def456",
            training_data_version="dataset-v3-20240115",
            eval_data_version="eval-set-v2",
            preprocessing_pipeline_version="preprocess-v1.2",
            training_samples=1_000_000,
            eval_samples=10_000,
            dataset_hash="sha256:xxxx...",
        ),
        config=TrainingConfig(
            hyperparameters={
                "learning_rate": 2e-5,
                "warmup_steps": 500,
                "weight_decay": 0.01,
                "max_seq_length": 128,
            },
            training_seed=42,
            optimizer="adamw",
            learning_rate=2e-5,
            batch_size=32,
            epochs=3,
            num_gpus=4,
            gpu_type="A100",
            training_time_hours=12.5,
            environment={
                "torch": "2.0.1",
                "transformers": "4.30.2",
                "cuda": "11.8",
            },
        ),
        metrics=PerformanceMetrics(
            accuracy=0.9245,
            precision=0.921,
            recall=0.918,
            f1_score=0.9195,
            auc_roc=0.97,
            latency_p50_ms=23,
            latency_p95_ms=45,
            latency_p99_ms=68,
            throughput_qps=1500,
            model_size_mb=420,
            memory_usage_mb=850,
            gpu_memory_mb=4096,
        ),
        stage="production",
        status="validated",
        description="基于 BERT 的中文情感分类模型 v2",
        change_log="""
        - v2.3.1: 优化推理速度,p95 延迟降低 15%
        - v2.3.0: 新增支持长文本分类
        - v2.2.0: 使用新数据集重新训练
        """,
        known_issues=[
            "在包含大量英文的混合文本上表现略差",
        ],
        recommendations=[
            "建议配合语言检测模型使用",
            "长文本建议分段处理",
        ],
        tags={
            "language": "chinese",
            "domain": "sentiment",
            "team": "nlp",
            "compliance": "gdpr-ready",
        },
        created_by="zhangsan",
        created_at=datetime.now().isoformat(),
    )
    
    return metadata


if __name__ == "__main__":
    metadata = create_example_metadata()
    print(metadata.to_json())

五、选型对比表

5.1 功能对比

功能项

MLflow Registry

自建 Registry

基础功能

模型注册

版本管理

阶段转换

元数据存储

高级功能

自定义元数据

⚠️ 有限

版本血缘追踪

审批工作流

Webhook 通知

A/B 测试集成

企业功能

细粒度权限控制

审计日志

⚠️ 基础

SSO/LDAP 集成

⚠️ 需企业版

高可用部署

⚠️ 需配置

数据本地化

生态集成

Git 集成

⚠️ 基础

CI/CD 集成

⚠️ 基础

监控告警

⚠️ 需额外配置

文档自动生成

5.2 成本对比

成本项

MLflow Registry

自建 Registry

初期投入

集成时间

2-4 周

3-6 个月

开发人力

1-2 人

3-5 人

学习成本

中等

较高

运维成本

基础设施

2-4 台机器

4-8 台机器

运维人力

0.5 人

1-2 人

监控告警

需自建

可内置

扩展成本

新功能开发

受限

灵活

定制化需求

困难

简单

第三方集成

有限

完全可控

5.3 推荐选择

┌─────────────────────────────────────────────────────────────────┐
│                       选型决策树                                │
├─────────────────────────────────────────────────────────────────┤
│                                                                 │
│   团队规模 < 5 人?                                             │
│   ├── 是 → MLflow Registry                                     │
│   │         ✓ 快速启动                                          │
│   │         ✓ 维护成本低                                       │
│   │         ✓ 功能足够                                          │
│   │                                                        │
│   └── 否 → 是否有复杂审批流程?                                │
│           ├── 否 → MLflow Registry + 定制开发                 │
│           │         ✓ 满足 80% 需求                            │
│           │         ✓ 按需扩展                                  │
│           │                                                        │
│           └── 是 → 自建 Registry                                │
│                   ✓ 完全可控                                    │
│                   ✓ 满足所有需求                                │
│                   ⚠️ 需要投入足够资源                          │
│                                                                 │
└─────────────────────────────────────────────────────────────────┘

总结

本文系统性地介绍了模型版本管理的演进之路,从为什么模型版本管理不同于代码版本管理,到 MLflow Model Registry 的使用,再到自建 Registry 的架构设计。

核心要点

  1. 理解差异:模型版本管理远比代码版本管理复杂,涉及非确定性、隐式依赖、多维度度量等挑战

  2. MLflow Registry 适用场景

  • 团队较小(< 10 人)

  • 需求相对标准

  • 快速验证阶段

  • 不想投入太多维护成本

  1. 自建 Registry 适用场景

  • 团队规模大、模型多

  • 有复杂的审批和合规要求

  • 需要深度集成内部系统

  • 有足够的开发和维护资源

  1. 关键设计点

  • 血缘追踪:连接模型、数据、代码

  • 元数据设计:支持查询、对比、审计

  • 生命周期:清晰的阶段流转

  • 签名机制:实现版本去重

建议

如果你正在从零开始,建议先从 MLflow Registry 起步,快速验证团队的工作流。当 MLflow 无法满足需求时,再考虑自建。

自建不是目的,解决问题才是目的。

希望本文对你有所帮助!如有问题,欢迎交流讨论。

0
  1. 支付宝打赏

    qrcode alipay
  2. 微信打赏

    qrcode weixin

评论区