一、为什么要实现自己的模型版本控制系统?
在实际的 AI 工程化落地中,模型管理是一个经常被低估却又极其关键的环节。很多团队初期只是把模型文件往磁盘上一扔,用文件名带个日期来区分版本——"model_v2_final_real_final.pth" 这种命名,想必做算法的同学都不陌生。
但随着模型迭代次数增加、实验变多、多人协作介入,这种粗放的方式会带来一系列问题:
- 版本混乱:无法区分哪个是线上稳定版、哪个是实验废弃版。一个模型文件夹里可能有十几个命名各异的文件,根本分不清谁是谁。
- 回滚困难:生产环境模型出问题时,找不到上一个可用的版本。紧急修复时只能翻 Git 历史、翻聊天记录去找"上次那个好用的版本"。
- 缺乏元数据:训练参数、数据集版本、评估指标等关键信息散落在各处。三个月后回来看,根本记不清这个模型是用什么学习率、什么数据切分方式训的。
- 难以复现:六个月后回来看,完全不知道这个模型的训练数据来自哪个版本的数据集、用了什么预处理逻辑、超参数是什么。
- 协作困难:多人同时实验时,A 同学的训练结果覆盖了 B 同学的实验目录,或者模型文件被误删除。
业界的解决方案如 MLflow Model Registry、DVC、Weights & Biases 等虽然功能强大,但对于很多中小团队来说,或者在某些对部署环境有严格要求的场景下(如内网部署、离线环境、合规要求),一个轻量、可控的自实现版本控制系统反而更实用。它不需要安装额外的服务,不需要网络连接,只需要一个目录和几条 Python 代码,就能完成模型版本管理的核心功能。
本文将手把手带你从零实现一套轻量级 AI 模型版本控制系统,涵盖:
整套系统仅依赖 Python 标准库 + JSON,零外部框架依赖,可直接嵌入到现有项目中。
二、系统设计概览
2.1 核心概念
我们的模型版本控制系统包含以下几个核心组件:
| ModelRegistry | 注册中心,管理所有模型 | 类似 Docker Registry |
| ModelVersion | 单个版本的数据结构 | 类似 Git commit |
| ModelStore | 底层存储,负责序列化 | 类似 Git 对象存储 |
| ModelManager | 高层接口,加载/推理 | 类似模型服务层 |
2.2 版本号设计
我们采用业界标准的语义化版本号(Semantic Versioning):
MAJOR.MINOR.PATCH[-PRERELEASE]
- MAJOR:不兼容的 API 变更(新架构、输入输出格式变化)。例如从 TF-IDF 升级到 BERT,输入格式从稀疏向量变成 token_ids,旧版本的推理脚本直接跑不通,这时候就需要递增 MAJOR 版本。
- MINOR:向下兼容的功能新增(加了新输出字段)。比如在分类结果中额外返回了置信度分数,但原有的 label 字段保持不变,旧代码不需要修改就可以继续使用。
- PATCH:向下兼容的问题修复(重训、调参)。比如换了更大的 batch size 重新训练,精度提升了但接口完全不变。
- PRERELEASE:alpha/beta/rc 等预发布标记。实验中的版本先标记为 -alpha,验证通过后再去掉标记正式发布。
每个版本还附带一个 Stage(阶段)标签:
- development:开发中,还在调参实验
- staging:验证通过,准生产
- production:线上稳定版
- archived:已归档,不再使用
Stage 的作用是让团队成员一目了然地知道每个版本当前所处的生命周期阶段,避免把还在实验中的版本误推到线上。
2.3 存储架构
model_registry/
├── metadata.json # 注册中心元数据
├── models/
│ └── {model_name}/
│ ├── metadata.json # 模型元数据
│ └── versions/
│ ├── v1.0.0/
│ │ ├── model.pkl # 序列化模型
│ │ ├── version.json # 版本元数据
│ │ ├── config.json # 训练配置
│ │ └── metrics.json # 评估指标
│ ├── v1.1.0/
│ │ └── …
│ └── v2.0.0-alpha/
│ └── …
这种目录结构非常直观:模型名 → 版本号 → 版本文件。无论是人眼查看还是程序遍历都非常方便。所有元数据都使用 JSON 格式,可读性好,也方便其他工具(如 Python 的 json 模块)直接读取。
三、为什么需要语义化版本号?
在进入代码实现之前,我们先思考一个问题:版本号为什么不能是简单的 1、2、3?
3.1 简单的数字版本号有什么问题?
假设你的项目中有三个版本:v1、v2、v3。如果同事问你:"线上跑的是哪个版本?"你只能说"v2"。但如果是一个语义化版本号 2.1.0,你可以告诉他:"2.1.0 是第二个大版本的第一次小改进,兼容之前的输入输出格式,直接可以替换。"
这就是语义化版本号的核心价值:版本号本身携带了变更信息。
- MAJOR 递增 → 不兼容变更,需要评估适配成本
- MINOR 递增 → 新增功能,但向后兼容
- PATCH 递增 → bug 修复,无缝替换
再看一个实际场景:你发现线上模型在某个样本上推理结果不对,需要回滚到上一个版本。如果版本号是 v2、v3、v4,你根本不知道 v2 和 v4 之间差了什么。但如果版本号是 1.0.0 → 1.1.0 → 2.0.0,你就知道回滚到 1.1.0 意味着不仅回到了上一个 PATCH 版本,而是回到了上一个 MAJOR 系列的最后一个 MINOR 版本——这完全可能是两种不同的模型架构。
3.2 回滚时的兼容性判断
语义化版本号还能帮我们做自动兼容性判断。假设我们要从版本 A 回滚到版本 B:
- 如果 A 是 2.0.0,B 是 2.1.0(MAJOR 相同)→ 接口兼容,直接替换
- 如果 A 是 2.0.0,B 是 1.3.0(MAJOR 不同)→ 接口不兼容,需要适配代码
- 如果 A 是 2.0.0-alpha,B 是 2.0.0(有 prerelease)→ 实验版本回稳,需要验证
这个判断逻辑可以在代码中自动实现(就是我们后面会看到的 is_compatible_with() 方法),实现自动化的版本兼容性检查。
四、基础数据结构与版本号工具函数
我们先从最基础的数据结构开始实现。
4.1 语义化版本号解析与比较
# semver.py
import re
from typing import Dict, Optional, Tuple
def parse_semver(version: str) -> Dict[str, any]:
"""解析语义化版本号字符串为结构化字典
Args:
version: 版本号字符串,如 "1.2.3" 或 "2.0.0-alpha"
Returns:
包含 major/minor/patch/prerelease 字段的字典
Raises:
ValueError: 版本号格式无效
"""
prerelease = None
if "-" in version:
version, prerelease = version.split("-", 1)
parts = version.split(".")
if len(parts) != 3:
raise ValueError(f"无效的语义化版本号: {version}")
try:
major, minor, patch = int(parts[0]), int(parts[1]), int(parts[2])
except ValueError:
raise ValueError(f"版本号必须包含数字: {version}")
return {
"major": major,
"minor": minor,
"patch": patch,
"prerelease": prerelease,
}
def compare_versions(v1: str, v2: str) -> int:
"""比较两个版本号的大小
Args:
v1, v2: 版本号字符串
Returns:
-1: v1 < v2
0: v1 == v2
1: v1 > v2
"""
p1 = parse_semver(v1)
p2 = parse_semver(v2)
for key in ["major", "minor", "patch"]:
if p1[key] < p2[key]:
return -1
elif p1[key] > p2[key]:
return 1
# 处理 prerelease(有 prerelease 的版本更"小")
if p1["prerelease"] and not p2["prerelease"]:
return -1
elif not p1["prerelease"] and p2["prerelease"]:
return 1
elif p1["prerelease"] and p2["prerelease"]:
if p1["prerelease"] < p2["prerelease"]:
return -1
elif p1["prerelease"] > p2["prerelease"]:
return 1
return 0
def bump_version(current: str, bump_type: str, prerelease: Optional[str] = None) -> str:
"""基于当前版本号自动递增
Args:
current: 当前版本号,如 "1.2.3"
bump_type: "major" / "minor" / "patch"
prerelease: 可选预发布标记
Returns:
新的版本号
"""
parts = parse_semver(current)
major, minor, patch = parts["major"], parts["minor"], parts["patch"]
if bump_type == "major":
major += 1
minor = 0
patch = 0
elif bump_type == "minor":
minor += 1
patch = 0
elif bump_type == "patch":
patch += 1
else:
raise ValueError(f"无效的 bump_type: {bump_type}(可选: major/minor/patch)")
version = f"{major}.{minor}.{patch}"
if prerelease:
version = f"{version}-{prerelease}"
return version
# ———- 单元测试 ———-
if __name__ == "__main__":
# 测试解析
assert parse_semver("1.2.3")["major"] == 1
assert parse_semver("2.0.0-alpha")["prerelease"] == "alpha"
assert parse_semver("0.1.0")["patch"] == 0
# 测试比较
assert compare_versions("1.0.0", "2.0.0") == -1
assert compare_versions("2.0.0", "1.0.0") == 1
assert compare_versions("1.0.0", "1.0.0") == 0
assert compare_versions("1.0.0", "1.0.1") == -1
assert compare_versions("1.0.0-alpha", "1.0.0") == -1
assert compare_versions("1.0.0-beta", "1.0.0-alpha") == 1
# 测试版本递增
assert bump_version("1.2.3", "major") == "2.0.0"
assert bump_version("1.2.3", "minor") == "1.3.0"
assert bump_version("1.2.3", "patch") == "1.2.4"
assert bump_version("1.0.0", "minor", "alpha") == "1.1.0-alpha"
assert bump_version("2.5.0", "patch") == "2.5.1"
print("✅ 所有语义化版本号单元测试通过")
4.2 模型版本数据结构
# model_version.py
import json
import time
import hashlib
from dataclasses import dataclass, field, asdict
from typing import Any, Optional, Dict, List
@dataclass
class ModelVersion:
"""模型版本的数据结构
一个 ModelVersion 对象完整地描述了一个模型版本的所有信息,
包括它的标识、状态、训练参数和评估结果。
"""
model_name: str # 模型名称
version: str # 语义化版本号,如 "1.0.0"
stage: str = "development" # 阶段标签:development/staging/production/archived
created_at: float = field(default_factory=time.time) # 创建时间戳(Unix 时间戳)
description: str = "" # 版本描述,记录这个版本改了什么的说明
tags: Dict[str, str] = field(default_factory=dict) # 自定义标签,如 {"dataset": "v2", "author": "zhang"}
parent_version: Optional[str] = None # 父版本号,记录版本继承链
checkpoint: Optional[str] = None # 检查点路径,用于训练中断后续训
# 模型元数据字段
framework: str = "" # 深度学习框架,如 "pytorch" / "tensorflow" / "sklearn"
model_type: str = "" # 模型架构类型,如 "bert-base" / "resnet50"
input_schema: Dict[str, Any] = field(default_factory=dict) # 输入格式描述
output_schema: Dict[str, Any] = field(default_factory=dict) # 输出格式描述
model_hash: str = "" # 模型文件的 SHA256 哈希,用于完整性校验
def to_dict(self) -> Dict[str, Any]:
"""将 ModelVersion 序列化为字典"""
return asdict(self)
@classmethod
def from_dict(cls, data: Dict[str, Any]) -> "ModelVersion":
"""从字典反序列化还原 ModelVersion 对象"""
return cls(**data)
def compute_model_hash(self, model_bytes: bytes) -> str:
"""计算模型文件的 SHA256 哈希值
哈希值用于后续验证模型文件是否被篡改或损坏。
"""
return hashlib.sha256(model_bytes).hexdigest()
def is_compatible_with(self, other: "ModelVersion") -> bool:
"""检查两个版本是否兼容
兼容判断标准:MAJOR 版本号相同。
如果 MAJOR 相同,说明接口和输入输出格式没有破坏性变更。
Returns:
True 表示可以直接用 other 替换 self,
False 表示需要检查接口适配
"""
self_major = self.version.split(".")[0]
other_major = other.version.split(".")[0]
return self_major == other_major
4.3 数据类为什么不是 JSON 可序列化的?
细心的读者可能会发现,我们同时提供了 to_dict() 和 from_dict() 方法。这是因为 Python 的 @dataclass 虽然方便,但生成的 asdict() 转换后的 datetime 对象、bytes 等类型并不能被 json.dumps() 直接序列化。通过自定义的序列化/反序列化方法,我们可以确保:
五、模型注册中心实现
模型注册中心是整个系统的核心,负责管理所有已注册的模型及其版本。它提供四种核心操作:模型注册、版本注册、版本阶段管理和模型加载。
5.1 注册中心类
# model_registry.py
import json
import os
import pickle
import shutil
import hashlib
import tempfile
from typing import Any, Optional, Dict, List, Tuple
from datetime import datetime
from model_version import ModelVersion
from semver import compare_versions, bump_version, parse_semver
class ModelRegistry:
"""轻量级模型注册中心
这是一个基于文件系统的模型注册中心实现,使用目录结构和 JSON 文件
来存储模型和版本的元数据。所有模型通过 pickle 序列化保存。
Args:
registry_dir: 注册中心在磁盘上的根目录
"""
def __init__(self, registry_dir: str = "./model_registry"):
self.registry_dir = registry_dir
self._ensure_dirs()
self._load_metadata()
def _ensure_dirs(self):
"""确保注册中心的目录结构存在"""
os.makedirs(os.path.join(self.registry_dir, "models"), exist_ok=True)
def _load_metadata(self):
"""加载或初始化注册中心的全局元数据"""
meta_path = os.path.join(self.registry_dir, "metadata.json")
if os.path.exists(meta_path):
with open(meta_path, "r") as f:
self.metadata = json.load(f)
else:
self.metadata = {
"registry_name": "AI Model Registry",
"created_at": datetime.now().isoformat(),
"model_count": 0,
"models": [],
}
self._save_metadata()
def _save_metadata(self):
"""保存注册中心的全局元数据到磁盘"""
meta_path = os.path.join(self.registry_dir, "metadata.json")
with open(meta_path, "w") as f:
json.dump(self.metadata, f, indent=2, ensure_ascii=False)
def _get_model_dir(self, model_name: str) -> str:
"""获取模型在磁盘上的目录路径"""
return os.path.join(self.registry_dir, "models", model_name)
def _get_version_dir(self, model_name: str, version: str) -> str:
"""获取指定版本在磁盘上的目录路径"""
return os.path.join(self._get_model_dir(model_name), "versions", version)
# ———- 模型管理 ———-
def list_models(self) -> List[Dict[str, Any]]:
"""列出注册中心中所有已注册的模型
Returns:
模型信息列表,每项包含 name 和 description
"""
return self.metadata.get("models", [])
def register_model(self, model_name: str, description: str = "") -> bool:
"""在注册中心注册一个新的模型
注册一个新模型时,会创建对应的目录结构并初始化元数据文件。
如果模型名称已存在,则不会重复注册。
Args:
model_name: 模型名称(在当前注册中心中必须唯一)
description: 模型的文字描述
Returns:
True 表示注册成功,False 表示模型已存在
"""
# 检查是否已存在同名模型
for m in self.metadata["models"]:
if m["name"] == model_name:
return False
# 创建模型目录
model_dir = self._get_model_dir(model_name)
os.makedirs(model_dir, exist_ok=True)
# 初始化模型元数据
model_meta = {
"name": model_name,
"description": description,
"created_at": datetime.now().isoformat(),
"latest_version": None,
"production_version": None,
"version_count": 0,
}
with open(os.path.join(model_dir, "metadata.json"), "w") as f:
json.dump(model_meta, f, indent=2, ensure_ascii=False)
# 更新注册中心全局元数据
self.metadata["models"].append({
"name": model_name,
"description": description,
})
self.metadata["model_count"] = len(self.metadata["models"])
self._save_metadata()
return True
def get_model_metadata(self, model_name: str) -> Optional[Dict]:
"""获取模型的元数据信息
Args:
model_name: 模型名称
Returns:
模型元数字典,如果模型不存在则返回 None
"""
meta_path = os.path.join(self._get_model_dir(model_name), "metadata.json")
if not os.path.exists(meta_path):
return None
with open(meta_path, "r") as f:
return json.load(f)
def _update_model_metadata(self, model_name: str, updates: Dict):
"""更新模型元数据的指定字段"""
meta = self.get_model_metadata(model_name)
if meta is None:
raise ValueError(f"模型 '{model_name}' 不存在")
meta.update(updates)
meta_path = os.path.join(self._get_model_dir(model_name), "metadata.json")
with open(meta_path, "w") as f:
json.dump(meta, f, indent=2, ensure_ascii=False)
# ———- 版本管理 ———-
def register_version(
self,
model_name: str,
model_obj: Any,
version: Optional[str] = None,
bump_type: Optional[str] = None,
description: str = "",
framework: str = "",
model_type: str = "",
input_schema: Optional[Dict] = None,
output_schema: Optional[Dict] = None,
tags: Optional[Dict[str, str]] = None,
train_config: Optional[Dict] = None,
metrics: Optional[Dict] = None,
) -> str:
"""注册一个新版本
这是最核心的方法之一。它会:
1. 确定版本号(手动指定或自动递增)
2. 序列化模型到磁盘
3. 计算模型文件哈希
4. 保存版本元数据、训练配置和评估指标
5. 更新模型元数据中的最新版本记录
Args:
model_name: 模型名称
model_obj: 模型对象(必须是可 pickle 序列化的)
version: 手动指定版本号(和 bump_type 二选一)
bump_type: 自动递增类型 "major"/"minor"/"patch"
description: 版本描述
framework: 使用的深度学习框架
model_type: 模型架构类型
input_schema: 输入格式说明(如字段名、类型、维度)
output_schema: 输出格式说明
tags: 自定义标签(如数据版本、训练时长)
train_config: 训练超参数配置
metrics: 评估指标(如 accuracy、f1、latency)
Returns:
新注册的版本号字符串
Raises:
ValueError: 模型不存在、版本号冲突或参数错误
"""
# 获取模型元数据
meta = self.get_model_metadata(model_name)
if meta is None:
raise ValueError(f"模型 '{model_name}' 尚未注册,请先调用 register_model()")
# 确定版本号
if version and bump_type:
raise ValueError("version 和 bump_type 不能同时指定,请选择一个")
if version:
# 验证版本号格式(至少是 x.y.z 格式)
try:
parse_semver(version)
except ValueError:
pass # 不是标准语义化版本号但允许使用
# 检查版本号是否已存在
if self.has_version(model_name, version):
raise ValueError(f"版本 '{version}' 已存在")
elif bump_type:
if meta["latest_version"]:
version = bump_version(meta["latest_version"], bump_type)
else:
version = "1.0.0"
else:
# 默认行为:递增 patch 版本
if meta["latest_version"]:
version = bump_version(meta["latest_version"], "patch")
else:
version = "1.0.0"
# 创建版本目录
version_dir = self._get_version_dir(model_name, version)
os.makedirs(version_dir, exist_ok=True)
# 序列化模型对象
model_path = os.path.join(version_dir, "model.pkl")
with open(model_path, "wb") as f:
pickle.dump(model_obj, f)
# 计算模型文件哈希
with open(model_path, "rb") as f:
model_bytes = f.read()
model_hash = hashlib.sha256(model_bytes).hexdigest()
# 构建 ModelVersion 对象
mv = ModelVersion(
model_name=model_name,
version=version,
stage="development",
description=description,
tags=tags or {},
parent_version=meta["latest_version"],
framework=framework,
model_type=model_type,
input_schema=input_schema or {},
output_schema=output_schema or {},
model_hash=model_hash,
)
# 保存版本元数据
version_meta = mv.to_dict()
version_meta_path = os.path.join(version_dir, "version.json")
with open(version_meta_path, "w") as f:
json.dump(version_meta, f, indent=2, ensure_ascii=False)
# 保存训练配置(如果有)
if train_config:
with open(os.path.join(version_dir, "config.json"), "w") as f:
json.dump(train_config, f, indent=2, ensure_ascii=False)
# 保存评估指标(如果有)
if metrics:
with open(os.path.join(version_dir, "metrics.json"), "w") as f:
json.dump(metrics, f, indent=2, ensure_ascii=False)
# 更新模型元数据
self._update_model_metadata(model_name, {
"latest_version": version,
"version_count": (meta.get("version_count", 0) or 0) + 1,
})
return version
def list_versions(self, model_name: str) -> List[str]:
"""列出模型的所有版本(按版本号降序排列,最新的在前)"""
versions_dir = os.path.join(self._get_model_dir(model_name), "versions")
if not os.path.exists(versions_dir):
return []
versions = []
for v in os.listdir(versions_dir):
vdir = os.path.join(versions_dir, v)
if os.path.isdir(vdir) and os.path.exists(os.path.join(vdir, "version.json")):
versions.append(v)
return sorted(versions, key=lambda v: compare_versions(v, "0.0.0"), reverse=True)
def has_version(self, model_name: str, version: str) -> bool:
"""检查指定版本号是否已存在"""
return os.path.exists(
os.path.join(self._get_version_dir(model_name, version), "version.json")
)
def get_version_metadata(self, model_name: str, version: str) -> Optional[Dict]:
"""获取指定版本的元数据"""
meta_path = os.path.join(self._get_version_dir(model_name, version), "version.json")
if not os.path.exists(meta_path):
return None
with open(meta_path, "r") as f:
return json.load(f)
def get_version_metrics(self, model_name: str, version: str) -> Optional[Dict]:
"""获取指定版本的评估指标"""
metrics_path = os.path.join(self._get_version_dir(model_name, version), "metrics.json")
if not os.path.exists(metrics_path):
return None
with open(metrics_path, "r") as f:
return json.load(f)
def get_version_config(self, model_name: str, version: str) -> Optional[Dict]:
"""获取指定版本的训练配置"""
config_path = os.path.join(self._get_version_dir(model_name, version), "config.json")
if not os.path.exists(config_path):
return None
with open(config_path, "r") as f:
return json.load(f)
# ———- 版本阶段管理 ———-
def promote_version(self, model_name: str, version: str, stage: str):
"""提升版本的阶段状态
阶段流转通常遵循:development → staging → production → archived
Args:
model_name: 模型名称
version: 版本号
stage: 目标阶段
Raises:
ValueError: 版本不存在或阶段值无效
"""
meta = self.get_version_metadata(model_name, version)
if meta is None:
raise ValueError(f"版本 '{version}' 不存在")
if stage not in ["development", "staging", "production", "archived"]:
raise ValueError(f"无效的阶段值: {stage}(可选: development/staging/production/archived)")
# 更新版本元数据中的 stage
meta["stage"] = stage
meta_path = os.path.join(self._get_version_dir(model_name, version), "version.json")
with open(meta_path, "w") as f:
json.dump(meta, f, indent=2, ensure_ascii=False)
# 如果提升到 production,同步更新模型元数据中 production_version 字段
if stage == "production":
self._update_model_metadata(model_name, {"production_version": version})
def get_production_version(self, model_name: str) -> Optional[str]:
"""获取当前正在生产的版本号"""
meta = self.get_model_metadata(model_name)
return meta.get("production_version") if meta else None
# ———- 模型加载 ———-
def load_model(self, model_name: str, version: Optional[str] = None) -> Tuple[Any, Dict]:
"""加载模型
如果不指定版本号,优先加载 production 版本;如果没有 production 版本,则加载最新版本。
Args:
model_name: 模型名称
version: 版本号(可选)
Returns:
(模型对象, 版本元数据) 的元组
"""
if version is None:
meta = self.get_model_metadata(model_name)
if meta is None:
raise ValueError(f"模型 '{model_name}' 不存在")
version = meta.get("production_version") or meta.get("latest_version")
if version is None:
raise ValueError(f"模型 '{model_name}' 没有任何可用版本")
# 加载序列化的模型文件
model_path = os.path.join(self._get_version_dir(model_name, version), "model.pkl")
if not os.path.exists(model_path):
raise ValueError(f"模型文件不存在: {model_path}")
with open(model_path, "rb") as f:
model = pickle.load(f)
# 加载版本元数据
version_meta = self.get_version_metadata(model_name, version)
return model, version_meta
# ———- 版本删除 ———-
def delete_version(self, model_name: str, version: str):
"""删除指定版本
注意:这是一个危险操作,会彻底删除版本目录及其所有文件。
建议在删除前做好备份。
"""
version_dir = self._get_version_dir(model_name, version)
if not os.path.exists(version_dir):
raise ValueError(f"版本 '{version}' 不存在")
shutil.rmtree(version_dir)
# 如果删除的是 production 版本,清除标记
meta = self.get_model_metadata(model_name)
if meta and meta.get("production_version") == version:
self._update_model_metadata(model_name, {"production_version": None})
print(f"🗑️ 已删除版本 {model_name}:{version}")
# ———- 版本回滚 ———-
def rollback(self, model_name: str, target_version: str) -> str:
"""回滚到指定版本
回滚的本质是将目标版本提升到 production 阶段。
如果目标版本和当前 production 版本 MAJOR 不同,会发出兼容性警告。
Args:
model_name: 模型名称
target_version: 目标版本号
Returns:
回滚后的 production 版本号
"""
if not self.has_version(model_name, target_version):
raise ValueError(f"目标版本 '{target_version}' 不存在")
# 检查兼容性
current = self.get_production_version(model_name)
if current:
current_meta = self.get_version_metadata(model_name, current)
target_meta = self.get_version_metadata(model_name, target_version)
if current_meta and target_meta:
current_major = current_meta["version"].split(".")[0]
target_major = target_meta["version"].split(".")[0]
if current_major != target_major:
print(f"⚠️ 警告:MAJOR 版本不同 ({current} → {target_version}),请确认接口兼容性")
# 提升目标版本到 production
self.promote_version(model_name, target_version, "production")
print(f"🔄 回滚完成: {current or '(首次上线)'} → {target_version}")
return target_version
def rollback_to_previous(self, model_name: str) -> str:
"""回滚到上一个 production 版本
遍历所有版本,找到最近一个也曾标记为 production 的版本。
"""
versions = self.list_versions(model_name)
current_prod = self.get_production_version(model_name)
for v in versions:
if v == current_prod:
continue
meta = self.get_version_metadata(model_name, v)
if meta and meta.get("stage") == "production":
return self.rollback(model_name, v)
raise ValueError(f"模型 '{model_name}' 没有可回滚的历史版本")
# ———- 版本对比 ———-
def compare_versions(self, model_name: str, v1: str, v2: str) -> Dict:
"""对比两个版本的评估指标"""
metrics1 = self.get_version_metrics(model_name, v1) or {}
metrics2 = self.get_version_metrics(model_name, v2) or {}
diff = {}
all_keys = set(metrics1.keys()) | set(metrics2.keys())
for key in sorted(all_keys):
val1 = metrics1.get(key)
val2 = metrics2.get(key)
if val1 is not None and val2 is not None and isinstance(val1, (int, float)):
diff[key] = {
v1: val1,
v2: val2,
"diff": round(val2 – val1, 4),
"diff_pct": f"{round((val2 – val1) / abs(val1) * 100, 2)}%" if val1 != 0 else "N/A",
}
else:
diff[key] = {v1: val1, v2: val2}
return diff
# ———- 统计信息 ———-
def get_statistics(self, model_name: Optional[str] = None) -> Dict:
"""获取注册中心或指定模型的统计信息"""
if model_name:
meta = self.get_model_metadata(model_name)
versions = self.list_versions(model_name)
return {
"model_name": model_name,
"total_versions": len(versions),
"latest_version": meta.get("latest_version"),
"production_version": meta.get("production_version"),
"created_at": meta.get("created_at"),
}
models = self.list_models()
total_versions = 0
model_stats = []
for m in models:
name = m["name"]
versions = self.list_versions(name)
total_versions += len(versions)
meta = self.get_model_metadata(name)
model_stats.append({
"name": name,
"versions": len(versions),
"latest": meta.get("latest_version"),
"production": meta.get("production_version"),
})
return {
"total_models": len(models),
"total_versions": total_versions,
"registry_dir": self.registry_dir,
"models": model_stats,
}
5.2 完整工作流程图
训练完成后调用 register_version()
│
▼
确定版本号 ──→ 手动指定 / 自动递增
│
▼
pickle 序列化模型对象
│
▼
计算 SHA256 模型哈希
│
▼
保存到版本目录 ├── model.pkl
├── version.json
├── config.json (可选)
└── metrics.json (可选)
│
▼
更新模型元数据中的 latest_version
│
▼
可以调用 promote_version() 上线
│
▼
推理服务调用 load_model() 加载
5.3 完整功能测试
# test_registry.py
import sys
import os
import tempfile
import shutil
# 确保能导入
sys.path.insert(0, os.path.dirname(__file__))
from model_registry import ModelRegistry
def test_full_workflow():
"""模拟一个完整的模型版本管理生命周期"""
with tempfile.TemporaryDirectory() as tmpdir:
registry = ModelRegistry(registry_dir=os.path.join(tmpdir, "test_registry"))
# ====== 第一步:注册模型 ======
assert registry.register_model("text_classifier", "文本分类模型")
assert not registry.register_model("text_classifier", "重复注册应该失败")
print("✅ 1. 模型注册成功")
# ====== 第二步:注册首个版本 v1.0.0 ======
model_v1 = {"type": "tf-idf", "params": {"max_features": 5000}}
v1 = registry.register_version(
"text_classifier", model_v1,
version="1.0.0",
description="初始版本,基于 TF-IDF 特征的线性分类器",
framework="sklearn",
model_type="logistic_regression",
input_schema={"text": "str"},
output_schema={"label": "str", "score": "float"},
train_config={
"learning_rate": 0.01,
"max_iter": 100,
"solver": "lbfgs",
},
metrics={"accuracy": 0.85, "f1": 0.84, "precision": 0.86, "recall": 0.83},
)
assert v1 == "1.0.0"
print("✅ 2. v1.0.0 注册完成")
# ====== 第三步:注册小改进版本 v1.1.0 ======
model_v2 = {"type": "tf-idf", "params": {"max_features": 10000, "ngram_range": [1, 3]}}
v2 = registry.register_version(
"text_classifier", model_v2,
bump_type="minor",
description="增加 n-gram 特征,提升特征表达能力",
metrics={"accuracy": 0.87, "f1": 0.86, "precision": 0.88},
)
assert v2 == "1.1.0"
print("✅ 3. v1.1.0 注册完成(minor 升级)")
# ====== 第四步:注册重大重构版本 v2.0.0 ======
model_v3 = {"type": "transformer", "params": {"model_name": "bert-base", "hidden_dim": 768}}
v3 = registry.register_version(
"text_classifier", model_v3,
bump_type="major",
description="升级为 Transformer 架构,引入预训练模型",
framework="pytorch",
input_schema={"text": "str", "max_length": "int"},
output_schema={"label": "str", "score": "float", "embeddings": "List[float]"},
train_config={
"learning_rate": 2e-5,
"batch_size": 16,
"epochs": 3,
"warmup_steps": 500,
},
metrics={"accuracy": 0.93, "f1": 0.92, "precision": 0.94, "recall": 0.91},
)
assert v3 == "2.0.0"
print("✅ 4. v2.0.0 注册完成(major 升级到 Transformer)")
# ====== 第五步:将 v1.0.0 部署到生产环境 ======
registry.promote_version("text_classifier", "1.0.0", "production")
assert registry.get_production_version("text_classifier") == "1.0.0"
print("✅ 5. v1.0.0 提升到 production")
# ====== 第六步:加载 production 模型验证 ======
model, meta = registry.load_model("text_classifier")
assert model["type"] == "tf-idf"
assert meta["version"] == "1.0.0"
print("✅ 6. 加载 production 模型成功")
# ====== 第七步:升级到 v2.0.0 ======
registry.promote_version("text_classifier", "2.0.0", "production")
model, meta = registry.load_model("text_classifier")
assert model["type"] == "transformer"
print("✅ 7. 升级到 v2.0.0 production")
# ====== 第八步:回滚到 v1.1.0 ======
registry.rollback("text_classifier", "1.1.0")
model, meta = registry.load_model("text_classifier")
assert model["type"] == "tf-idf"
assert model["params"]["ngram_range"] == [1, 3]
print("✅ 8. 回滚到 v1.1.0")
# ====== 第九步:对比版本指标 ======
diff = registry.compare_versions("text_classifier", "1.0.0", "2.0.0")
print(f"✅ 9. 版本对比: accuracy = {diff['accuracy']}")
# ====== 第十步:统计信息 ======
stats = registry.get_statistics("text_classifier")
assert stats["total_versions"] == 3
print(f"✅ 10. 统计: {stats['total_versions']} 个版本")
print("\\n🎉 全部测试通过!")
if __name__ == "__main__":
test_full_workflow()
六、高级功能:A/B 测试版本管理
在生产环境中,我们经常需要同时运行多个版本的模型进行对比测试。A/B 测试允许我们将部分流量导向新版本的模型,收集线上真实性能数据后,再决定是否全量升级。
6.1 A/B 测试管理器
# ab_testing.py
import random
import hashlib
from typing import Any, Optional, Dict, List, Callable
class ABTestManager:
"""模型 A/B 测试管理器
管理多个实验,每个实验包含一个控制组版本和一个实验组版本,
并按流量比例分配请求。
"""
def __init__(self, registry: 'ModelRegistry', model_name: str):
self.registry = registry
self.model_name = model_name
self.experiments: Dict[str, 'ABExperiment'] = {}
def create_experiment(
self,
name: str,
control_version: str,
treatment_version: str,
traffic_split: float = 0.5,
) -> 'ABExperiment':
"""创建一个 A/B 实验
Args:
name: 实验名称,需要唯一
control_version: 控制组版本号(通常为当前 production 版本)
treatment_version: 实验组版本号(要测试的新版本)
traffic_split: 实验组流量比例,0.1 = 10% 流量,最大 0.9
Returns:
创建的 ABExperiment 实例
"""
if not self.registry.has_version(self.model_name, control_version):
raise ValueError(f"控制组版本 '{control_version}' 不存在")
if not self.registry.has_version(self.model_name, treatment_version):
raise ValueError(f"实验组版本 '{treatment_version}' 不存在")
if not 0 < traffic_split < 1:
raise ValueError(f"流量比例必须在 0~1 之间: {traffic_split}")
experiment = ABExperiment(
registry=self.registry,
model_name=self.model_name,
name=name,
control_version=control_version,
treatment_version=treatment_version,
traffic_split=traffic_split,
)
self.experiments[name] = experiment
return experiment
def get_experiment(self, name: str) -> Optional['ABExperiment']:
return self.experiments.get(name)
def list_experiments(self) -> List[str]:
return list(self.experiments.keys())
class ABExperiment:
"""单个 A/B 实验"""
def __init__(
self,
registry: 'ModelRegistry',
model_name: str,
name: str,
control_version: str,
treatment_version: str,
traffic_split: float = 0.5,
):
self.registry = registry
self.model_name = model_name
self.name = name
self.control_version = control_version
self.treatment_version = treatment_version
self.traffic_split = traffic_split
self.results: Dict[str, List[Dict]] = {
"control": [],
"treatment": [],
}
# 预加载两个版本的模型
self.control_model, self.control_meta = registry.load_model(
model_name, control_version
)
self.treatment_model, self.treatment_meta = registry.load_model(
model_name, treatment_version
)
def get_model_for_request(self, request_id: str) -> tuple:
"""根据请求 ID 决定使用哪个模型
使用一致性哈希算法,确保同一请求 ID 始终进入同一组,
避免用户在刷新时看到不一致的结果。
Args:
request_id: 请求唯一标识(可以是用户 ID 或请求 ID)
Returns:
(模型对象, 版本号, 分组名称)
"""
# 一致性哈希:确保同一 request_id 始终映射到同一组
hash_val = int(hashlib.md5(request_id.encode()).hexdigest(), 16) % 10000
if hash_val / 10000 < self.traffic_split:
return self.treatment_model, self.treatment_version, "treatment"
else:
return self.control_model, self.control_version, "control"
def record_result(
self,
request_id: str,
group: str,
metrics: Dict[str, float],
):
"""记录一次推理的结果"""
self.results[group].append({
"request_id": request_id,
"metrics": metrics,
})
def get_summary(self) -> Dict:
"""汇总实验结果,计算平均指标"""
control_results = self.results["control"]
treatment_results = self.results["treatment"]
def avg_metrics(results: List[Dict]) -> Dict:
if not results:
return {}
keys = results[0]["metrics"].keys()
avg = {}
for key in keys:
values = [r["metrics"][key] for r in results]
avg[f"avg_{key}"] = round(sum(values) / len(values), 4)
return avg
return {
"experiment": self.name,
"control_version": self.control_version,
"treatment_version": self.treatment_version,
"traffic_split": self.traffic_split,
"control_samples": len(control_results),
"treatment_samples": len(treatment_results),
"control_metrics": avg_metrics(control_results),
"treatment_metrics": avg_metrics(treatment_results),
"completed": len(treatment_results) > 0,
}
# ———- A/B 测试演示 ———-
def demo_ab_test():
"""演示 A/B 测试流程"""
import tempfile
from model_registry import ModelRegistry
with tempfile.TemporaryDirectory() as tmpdir:
registry = ModelRegistry(registry_dir=os.path.join(tmpdir, "ab_demo"))
registry.register_model("sentiment_model", "情感分析模型")
# 注册两个版本
registry.register_version(
"sentiment_model",
{"type": "lstm", "params": {"hidden_size": 128}},
version="1.0.0",
metrics={"accuracy": 0.88},
)
registry.register_version(
"sentiment_model",
{"type": "transformer", "params": {"hidden_size": 256}},
version="2.0.0",
metrics={"accuracy": 0.92},
)
registry.promote_version("sentiment_model", "1.0.0", "production")
# 创建 A/B 测试:30% 流量到 Transformer 版本
ab = ABTestManager(registry, "sentiment_model")
exp = ab.create_experiment(
name="transformer_vs_lstm",
control_version="1.0.0",
treatment_version="2.0.0",
traffic_split=0.3,
)
# 模拟 1000 次线上请求
import random as rnd
for i in range(1000):
req_id = f"user_{i % 200:04d}_req_{i:04d}" # 200 个不同用户
model, version, group = exp.get_model_for_request(req_id)
# 模拟推理耗时和精度
latency = rnd.uniform(5, 50) if version == "1.0.0" else rnd.uniform(20, 80)
accuracy = 0.88 + rnd.gauss(0, 0.02) if version == "1.0.0" else 0.92 + rnd.gauss(0, 0.01)
exp.record_result(req_id, group, {
"latency_ms": round(latency, 2),
"accuracy": round(min(accuracy, 1.0), 4),
})
# 输出结果
summary = exp.get_summary()
print(f"📊 A/B 测试结果: {summary['experiment']}")
print(f" 控制组 ({summary['control_version']}): {summary['control_samples']} 次请求")
print(f" 实验组 ({summary['treatment_version']}): {summary['treatment_samples']} 次请求")
print(f" 控制组指标: {summary['control_metrics']}")
print(f" 实验组指标: {summary['treatment_metrics']}")
# 分析结论
c_acc = summary['control_metrics'].get('avg_accuracy', 0)
t_acc = summary['treatment_metrics'].get('avg_accuracy', 0)
if t_acc > c_acc:
print(f"\\n🎯 结论:实验组精度提升 {round((t_acc – c_acc) * 100, 2)}%,建议全量发布")
else:
print(f"\\n📌 结论:实验组未显著优于控制组")
if __name__ == "__main__":
demo_ab_test()
七、命令行工具集成
为了方便在终端中快速操作,我们提供一个完整的 CLI 工具:
# cli.py
import sys
import os
import json
import argparse
# 确保能导入
sys.path.insert(0, os.path.dirname(__file__))
from model_registry import ModelRegistry
def main():
parser = argparse.ArgumentParser(
description="AI 模型版本管理工具 – 命令行接口",
formatter_class=argparse.RawDescriptionHelpFormatter,
)
parser.add_argument("–registry-dir", default="./model_registry",
help="注册中心目录路径(默认: ./model_registry)")
subparsers = parser.add_subparsers(dest="command", help="可用命令")
# list-models: 列出所有已注册的模型
subparsers.add_parser("list-models", help="列出所有已注册的模型")
# register-model: 注册新模型
parser_register = subparsers.add_parser("register-model", help="注册一个新模型")
parser_register.add_argument("name", help="模型名称")
parser_register.add_argument("–desc", default="", help="模型描述")
# list-versions: 列出模型的版本
parser_versions = subparsers.add_parser("list-versions", help="列出指定模型的所有版本")
parser_versions.add_argument("model", help="模型名称")
# promote: 提升版本阶段
parser_promote = subparsers.add_parser("promote", help="提升版本的阶段状态")
parser_promote.add_argument("model", help="模型名称")
parser_promote.add_argument("version", help="版本号")
parser_promote.add_argument("stage",
choices=["development", "staging", "production", "archived"],
help="目标阶段")
# rollback: 回滚版本
parser_rollback = subparsers.add_parser("rollback", help="回滚到指定版本")
parser_rollback.add_argument("model", help="模型名称")
parser_rollback.add_argument("version", help="目标版本号")
# compare: 对比版本
parser_compare = subparsers.add_parser("compare", help="对比两个版本的评估指标")
parser_compare.add_argument("model", help="模型名称")
parser_compare.add_argument("v1", help="第一个版本号")
parser_compare.add_argument("v2", help="第二个版本号")
# stats: 统计信息
parser_stats = subparsers.add_parser("stats", help="显示注册中心统计信息")
parser_stats.add_argument("–model", default=None, help="模型名称(可选,指定则只显示该模型信息)")
args = parser.parse_args()
registry = ModelRegistry(registry_dir=args.registry_dir)
if args.command == "list-models":
models = registry.list_models()
if not models:
print("📭 暂无注册的模型")
else:
print(f"📦 已注册模型 ({len(models)}):")
for m in models:
meta = registry.get_model_metadata(m["name"])
prod = meta.get("production_version", "无") if meta else "无"
print(f" ─ {m['name']}: {meta.get('version_count', 0)} 个版本 (production: {prod})")
elif args.command == "register-model":
if registry.register_model(args.name, args.desc):
print(f"✅ 模型 '{args.name}' 注册成功")
else:
print(f"⚠️ 模型 '{args.name}' 已存在")
elif args.command == "list-versions":
versions = registry.list_versions(args.model)
if not versions:
print(f"📭 模型 '{args.model}' 没有版本记录")
else:
print(f"📋 模型 '{args.model}' 的版本 ({len(versions)}):")
for v in versions:
meta = registry.get_version_metadata(args.model, v)
stage = meta.get("stage", "?") if meta else "?"
desc = meta.get("description", "") if meta else ""
marker = " ⭐" if stage == "production" else ""
print(f" {marker} v{v} [{stage}] {desc}")
elif args.command == "promote":
try:
registry.promote_version(args.model, args.version, args.stage)
print(f"✅ 版本 {args.model}:v{args.version} → [{args.stage}]")
except ValueError as e:
print(f"❌ {e}")
elif args.command == "rollback":
try:
registry.rollback(args.model, args.version)
except ValueError as e:
print(f"❌ {e}")
elif args.command == "compare":
diff = registry.compare_versions(args.model, args.v1, args.v2)
if diff:
print(f"📊 版本对比: {args.v1} vs {args.v2}")
for key, vals in diff.items():
print(f" {key}: {vals}")
else:
print("没有可对比的指标数据")
elif args.command == "stats":
stats = registry.get_statistics(args.model)
print(json.dumps(stats, indent=2, ensure_ascii=False))
else:
parser.print_help()
if __name__ == "__main__":
main()
7.1 CLI 使用示例
# 列出所有模型
python cli.py list-models
# 注册新模型
python cli.py register-model text_classifier –desc "文本分类模型"
# 列出模型版本
python cli.py list-versions text_classifier
# 将 v2.0.0 提升到生产
python cli.py promote text_classifier 2.0.0 production
# 回滚到 v1.1.0
python cli.py rollback text_classifier 1.1.0
# 对比版本指标
python cli.py compare text_classifier 1.0.0 2.0.0
# 查看统计信息
python cli.py stats –model text_classifier
八、最佳实践与注意事项
8.1 目录结构建议
在实际生产环境中,建议将注册中心目录与模型训练项目分离:
project/
├── src/
│ ├── training/ # 训练代码
│ ├── inference/ # 推理服务
│ └── registry/ # 版本管理代码
│ ├── model_registry.py
│ ├── model_version.py
│ └── ab_testing.py
├── model_registry/ # 版本数据目录(可 Git LFS 追踪)
│ ├── metadata.json
│ └── models/
└── tests/
└── test_registry.py
8.2 版本管理规范
实验版本:2.0.0-alpha、2.0.0-beta
阶段流转规范 development → staging → production → archived 每个阶段对应不同级别的测试验证,staging 阶段的版本至少应该通过自动评测集的测试才能升到 production。
回滚流程
8.3 生产环境注意事项
序列化方式:pickle 在不同 Python 版本间可能有兼容性问题。如果是跨环境部署,建议使用更稳定的序列化方式,如 ONNX 格式或自定义序列化协议。
文件锁:在多进程场景下(如同时训练和推理),register_version() 和 load_model() 可能并发访问同一文件。可以通过文件锁(fcntl.flock)来保护。
磁盘空间:每保存一个版本都会完整复制模型文件。对于大模型(如数十 GB),建议只保存模型文件的路径或对象存储的地址,而不是直接序列化。
自动备份:建议定期对 model_registry/ 目录进行备份,或者使用 Git LFS 追踪版本数据。
8.4 可扩展方向
本文实现的版本控制系统是一个基础框架,你可以根据实际需求扩展:
- 远程存储支持:将 model.pkl 上传到 S3/MinIO 等对象存储,本地只保留元数据
- 模型对比 UI:基于 Gradio 或 Streamlit 实现可视化版本对比界面
- 自动管线集成:训练完成后在 training pipeline 中自动调用 register_version()
- Web 管理后台:基于 FastAPI 提供 RESTful API,前端用 Vue/React
- 模型签名与验证:使用数字签名防止模型文件被篡改
- 自动化版本清理:保留最近 N 个版本的自动清理策略
九、总结
本文从零实现了一套完整的 AI 模型版本控制系统,涵盖:
整套系统仅依赖 Python 标准库,零外部依赖,可以无缝集成到任何 Python 项目中。无论是个人实验、小团队协作,还是对部署环境有严格要求的离线场景,这套系统都能帮你建立规范的模型版本管理流程。
在实际工程中,版本控制不是可选项——它是 AI 工程化的基石,决定了团队能否高效迭代、快速复盘、稳定上线。当你的第二个模型跑出结果的那一刻,版本管理就不再是"以后再说"的事了。
📚 延伸阅读
如果你对 DeepSeek 的实战用法感兴趣,推荐阅读我的另一篇文章:
👉 DeepSeek 实战指南:提示词工程、API 集成与效率提升全攻略
这篇文章系统地拆解了 DeepSeek 的提示词工程技巧、API 封装方法以及日常效率提升场景,全文代码可直接运行,适合已经上手 DeepSeek 但希望更高效使用的开发者。
本文是"手写 AI 系统"系列文章之一。该系列从零实现 AI 系统中的关键组件,涵盖 RAG、Agent、Function Calling、MCP 等核心技术,帮助你深入理解底层原理,构建属于自己的 AI 工具。


![[特殊字符]DeepSeek‑Harness(DSH)小白保姆教程-171主机测评](https://www.171host.com/wp-content/uploads/2026/08/20260816085112-6a817a009aabf-220x150.png)