Files

396 lines
16 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""模型模板注册 / 加载 / 版本机制(PRD 5.3 ③ 模型框架)。
对应 issue #41(父 EPIC #5「③ AI 模型框架 配置化重构」、PRD 5.3
「③ 模型模板注册 / 加载 / 版本机制」)。
PRD 5.3 的核心诉求
------------------
模板化之后,模型 / 配方 / 估计器是「资产」,需要**注册表**统一管理:
- **注册**:把一个模型模板(含主干、超参、特征列、指标、版本)登记入库;
- **加载**:按名字 + 版本(或别名)取出可用的模板实例;
- **版本**:同一模板多版本共存,可回滚、可审计;
- **阶段**:版本带 stage 标签(dev / staging / prod),灰度发布可控;
- **校验**:注册时校验模板完整性,防止脏数据进库。
本模块交付一个独立的 ``TemplateRegistry``,自包含、不依赖未合并分支,是
issue #40 ``ModelRegistry`` 的深化(完整版本 + 阶段 + 回滚 + 审计 + 持久化)。
本模块交付什么
--------------
1. **``ModelTemplate``**:模型模板数据对象(name / version / backbone /
hyperparams / feature_columns / metrics / stage / extra),不可变、可序列化。
2. **``TemplateRegistry``**:模板注册表核心 API:
- ``register``:注册一个版本(校验完整性 + 同版本号拒重复);
- ``get``:按 name + version/别名加载;
- ``promote``:把版本提升到下一 stage(dev→staging→prod);
- ``rollback``:把某 stage 回滚到指定版本;
- ``list_versions`` / ``list_by_stage`` / ``history``:查询;
- ``save`` / ``load``:JSON 持久化(文件 / 目录),便于审计与重启恢复。
3. **``Stage``**:阶段枚举(DEV / STAGING / PROD),PRD 灰度发布三段制。
4. **校验**:注册时校验 name/version/backbone 非空、版本号语义合法、
stage 合法,拒绝脏模板。
零外部强依赖
------------
纯 Python 实现,无 numpy/sklearn 依赖,CI 可加载与校验。
与 issue #34 / #36 / #38 / #40 的关系
--------------------------------------
接口风格对齐 #34 声明式数据对象、#36 / #38 ``Recipe``、#40 ``ModelRegistry``。
本模块是 #40 ``ModelRegistry`` 的完整版(多阶段 + 回滚 + 持久化),#40 的
``RegisterStep`` 未来可直接对接本注册表,业务侧零改动。
"""
from __future__ import annotations
import json
import os
import re
import time
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Dict, List, Optional, Tuple
__all__ = [
"ModelTemplate",
"TemplateRegistry",
"Stage",
"TemplateRegistryError",
"is_valid_version",
"next_stage",
]
class TemplateRegistryError(Exception):
"""模板注册表统一异常(模板非法 / 版本冲突 / 版本不存在 / 阶段非法)。"""
# ---------------------------------------------------------------------------
# 阶段(Stage):灰度发布三段制
# ---------------------------------------------------------------------------
class Stage(str, Enum):
"""模型版本阶段:DEV(开发)→ STAGING(预发)→ PROD(生产)。"""
DEV = "dev"
STAGING = "staging"
PROD = "prod"
@classmethod
def from_str(cls, s: str) -> "Stage":
try:
return cls(s.lower())
except ValueError:
raise TemplateRegistryError(
f"非法阶段 {s!r},允许:{[s.value for s in Stage]}")
# 阶段提升顺序:dev → staging → prod
_STAGE_ORDER = [Stage.DEV, Stage.STAGING, Stage.PROD]
def next_stage(stage: Stage) -> Optional[Stage]:
"""返回下一阶段;PROD 已是最高返回 None。"""
try:
idx = _STAGE_ORDER.index(stage)
except ValueError:
return None
if idx + 1 >= len(_STAGE_ORDER):
return None
return _STAGE_ORDER[idx + 1]
# ---------------------------------------------------------------------------
# 版本号语义校验(语义化版本 vMAJOR.MINOR.PATCH 或简单的 vN / N)
# ---------------------------------------------------------------------------
_VERSION_RE = re.compile(r"^v?\d+(\.\d+)*([\-+][0-9A-Za-z.\-]+)?$")
def is_valid_version(version: str) -> bool:
"""校验版本号是否合法(v1 / 1.0 / v1.2.3 / v1.0-rc1 等)。"""
if not isinstance(version, str) or not version.strip():
return False
return bool(_VERSION_RE.match(version.strip()))
# ---------------------------------------------------------------------------
# 模型模板数据对象
# ---------------------------------------------------------------------------
#: 允许的主干类型(对齐 PRD 5.3 四类模型模板 + 通用)
ALLOWED_BACKBONES = (
"quality_forecast", # ① 质量预测(issue #36)
"anomaly_detection", # 异常检测(issue #37)
"cross_process_opt", # ③ 跨工序寻优(issue #38)
"recipe_opt", # ② 配方优化
"generic", # 通用
)
@dataclass(frozen=True)
class ModelTemplate:
"""模型模板:描述一个可加载的模型资产(主干 + 超参 + 特征 + 指标 + 版本)。
不可变数据对象,``to_dict`` / ``from_dict`` 可序列化往返,便于持久化与
配置台展示。注册表以 (name, version) 为主键管理多个模板实例。
"""
name: str
version: str
backbone: str = "generic"
hyperparams: Dict[str, Any] = field(default_factory=dict)
feature_columns: Tuple[str, ...] = field(default_factory=tuple)
target_column: str = ""
metrics: Dict[str, float] = field(default_factory=dict)
stage: Stage = Stage.DEV
registered_at: float = field(default_factory=time.time)
description: str = ""
extra: Dict[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
if not self.name:
raise TemplateRegistryError("ModelTemplate 缺少 name")
if not is_valid_version(self.version):
raise TemplateRegistryError(
f"非法版本号 {self.version!r}(例:v1 / 1.0.0 / v1.2-rc1)")
if self.backbone not in ALLOWED_BACKBONES:
raise TemplateRegistryError(
f"非法主干 {self.backbone!r},允许:{ALLOWED_BACKBONES}")
if not isinstance(self.stage, Stage):
# 允许从字符串构造(dataclass frozen 用 object.__setattr__)
object.__setattr__(self, "stage", Stage.from_str(str(self.stage)))
@property
def key(self) -> Tuple[str, str]:
return (self.name, self.version)
def to_dict(self) -> Dict[str, Any]:
return {
"name": self.name,
"version": self.version,
"backbone": self.backbone,
"hyperparams": dict(self.hyperparams),
"feature_columns": list(self.feature_columns),
"target_column": self.target_column,
"metrics": dict(self.metrics),
"stage": self.stage.value,
"registered_at": self.registered_at,
"description": self.description,
"extra": dict(self.extra),
}
@classmethod
def from_dict(cls, d: Dict[str, Any]) -> "ModelTemplate":
try:
return cls(
name=d["name"],
version=d["version"],
backbone=d.get("backbone", "generic"),
hyperparams=dict(d.get("hyperparams", {})),
feature_columns=tuple(d.get("feature_columns", [])),
target_column=d.get("target_column", ""),
metrics={k: float(v) for k, v in d.get("metrics", {}).items()},
stage=Stage.from_str(d.get("stage", "dev")),
registered_at=float(d.get("registered_at", time.time())),
description=d.get("description", ""),
extra=dict(d.get("extra", {})),
)
except KeyError as exc: # pragma: no cover
raise TemplateRegistryError(f"模板缺少必填字段:{exc}") from exc
# ---------------------------------------------------------------------------
# 模板注册表
# ---------------------------------------------------------------------------
class TemplateRegistry:
"""模型模板注册表:注册 / 加载 / 版本 / 阶段提升 / 回滚 / 持久化。
* 以 ``name`` 维护多个 ``version``;
* 每个版本带 ``Stage``(dev/staging/prod),``promote`` 逐级提升;
* ``rollback`` 把某 stage 指针回退到指定版本(保留历史,可审计);
* ``save`` / ``load`` JSON 持久化(单文件或目录每模板一文件)。
"""
def __init__(self) -> None:
self._templates: Dict[str, Dict[str, ModelTemplate]] = {}
# 每个 name 的阶段指针:{stage: version}
self._stage_pointers: Dict[str, Dict[Stage, str]] = {}
# 审计日志:[(ts, action, name, version, detail)]
self._history: List[Tuple[float, str, str, str, str]] = []
# -- 注册 / 校验 --------------------------------------------------------
def register(self, template: ModelTemplate, *, force: bool = False) -> ModelTemplate:
"""注册一个模板版本。
* 同 (name, version) 默认拒绝重复(``force=True`` 可覆盖);
* 注册后自动成为该模板 dev 阶段指针(若该 stage 无指针);
* 记入审计日志。
"""
versions = self._templates.setdefault(template.name, {})
if template.version in versions and not force:
raise TemplateRegistryError(
f"{template.name}@{template.version} 已存在(force=True 可覆盖)")
versions[template.version] = template
pointers = self._stage_pointers.setdefault(template.name, {})
# 新注册默认进 dev;若 dev 无指针则指向它
if Stage.DEV not in pointers:
pointers[Stage.DEV] = template.version
self._log("register", template.name, template.version,
f"stage={template.stage.value}")
return template
# -- 加载 ---------------------------------------------------------------
def get(self, name: str,
version: Optional[str] = None,
stage: Optional[Stage] = None) -> ModelTemplate:
"""按 name + version 或 name + stage 加载模板。
* ``version`` 优先;其次 ``stage``(取该阶段指针);
* 都不提供则取该模板最新注册版本。
"""
versions = self._templates.get(name)
if not versions:
raise TemplateRegistryError(f"模板 {name!r} 未注册")
if version is not None:
if version not in versions:
raise TemplateRegistryError(
f"{name!r} 无版本 {version!r}(可用:{sorted(versions)})")
return versions[version]
if stage is not None:
ptr = self._stage_pointers.get(name, {}).get(stage)
if ptr is None:
raise TemplateRegistryError(
f"{name!r} 无 {stage.value} 阶段指针")
return versions[ptr]
# 默认:按注册时间最新;时间相同时按版本号字典序最新(确定性)
latest = max(versions.values(),
key=lambda t: (t.registered_at, t.version))
return latest
# -- 阶段提升 / 回滚 -----------------------------------------------------
def promote(self, name: str, version: str) -> ModelTemplate:
"""把指定版本提升到下一 stage(dev→staging→prod)。
提升后该版本成为新 stage 的指针版本。
"""
tpl = self.get(name, version)
nxt = next_stage(tpl.stage)
if nxt is None:
raise TemplateRegistryError(
f"{name}@{version} 已在 PROD,无法继续提升")
# 更新该模板版本的 stage(需重建不可变对象)
new_tpl = ModelTemplate(**{**tpl.to_dict(), "stage": nxt.value})
self._templates[name][version] = new_tpl
self._stage_pointers.setdefault(name, {})[nxt] = version
self._log("promote", name, version, f"{tpl.stage.value}->{nxt.value}")
return new_tpl
def rollback(self, name: str, stage: Stage, version: str) -> ModelTemplate:
"""把某 stage 的指针回退到指定版本(保留历史版本,可审计)。"""
tpl = self.get(name, version)
if not isinstance(stage, Stage):
stage = Stage.from_str(str(stage))
self._stage_pointers.setdefault(name, {})[stage] = version
self._log("rollback", name, version, f"stage={stage.value}")
return tpl
def set_stage(self, name: str, version: str, stage: Stage) -> ModelTemplate:
"""直接把某版本设到指定 stage(覆盖 promote 的逐级约束,用于紧急回滚)。"""
tpl = self.get(name, version)
if not isinstance(stage, Stage):
stage = Stage.from_str(str(stage))
new_tpl = ModelTemplate(**{**tpl.to_dict(), "stage": stage.value})
self._templates[name][version] = new_tpl
self._stage_pointers.setdefault(name, {})[stage] = version
self._log("set_stage", name, version, f"stage={stage.value}")
return new_tpl
# -- 查询 ---------------------------------------------------------------
def list_names(self) -> List[str]:
return sorted(self._templates.keys())
def list_versions(self, name: str) -> List[str]:
return sorted(self._templates.get(name, {}).keys())
def list_by_stage(self, name: str, stage: Stage) -> List[str]:
"""列出某模板处于指定 stage 的所有版本号。"""
if not isinstance(stage, Stage):
stage = Stage.from_str(str(stage))
return sorted(v for v, t in self._templates.get(name, {}).items()
if t.stage == stage)
def stage_pointer(self, name: str, stage: Stage) -> Optional[str]:
if not isinstance(stage, Stage):
stage = Stage.from_str(str(stage))
return self._stage_pointers.get(name, {}).get(stage)
def history(self, name: Optional[str] = None) -> List[Dict[str, Any]]:
"""返回审计日志(可按 name 过滤)。"""
out = []
for ts, action, n, ver, detail in self._history:
if name is not None and n != name:
continue
out.append({"time": ts, "action": action, "name": n,
"version": ver, "detail": detail})
return out
# -- 持久化 --------------------------------------------------------------
def save(self, path: str) -> None:
"""把整个注册表序列化到 JSON 文件(含模板 + 阶段指针 + 审计日志)。"""
data = {
"templates": {
name: {ver: t.to_dict() for ver, t in vers.items()}
for name, vers in self._templates.items()
},
"stage_pointers": {
name: {s.value: v for s, v in ptrs.items()}
for name, ptrs in self._stage_pointers.items()
},
"history": [
{"time": ts, "action": a, "name": n, "version": v, "detail": d}
for ts, a, n, v, d in self._history
],
}
with open(path, "w", encoding="utf-8") as fh:
json.dump(data, fh, ensure_ascii=False, indent=2)
@classmethod
def load(cls, path: str) -> "TemplateRegistry":
"""从 JSON 文件恢复注册表。"""
with open(path, "r", encoding="utf-8") as fh:
data = json.load(fh)
reg = cls()
for name, vers in data.get("templates", {}).items():
for ver, td in vers.items():
reg._templates.setdefault(name, {})[ver] = ModelTemplate.from_dict(td)
for name, ptrs in data.get("stage_pointers", {}).items():
for s_str, v in ptrs.items():
reg._stage_pointers.setdefault(name, {})[Stage.from_str(s_str)] = v
for h in data.get("history", []):
reg._history.append(
(h["time"], h["action"], h["name"], h["version"], h["detail"]))
return reg
# -- 内部 ----------------------------------------------------------------
def _log(self, action: str, name: str, version: str, detail: str) -> None:
self._history.append((time.time(), action, name, version, detail))
def __len__(self) -> int:
return sum(len(v) for v in self._templates.values())
def __contains__(self, name: str) -> bool:
return name in self._templates