feat(#41): 模型模板注册/加载/版本机制(多版本+阶段+提升+回滚+审计+持久化,PRD 5.3 ③模型模板注册加载版本机制) #108
@@ -0,0 +1,71 @@
|
|||||||
|
# iAOP-Core · 模型框架层(AI Model Framework)
|
||||||
|
|
||||||
|
对应 PRD 5.3「③ AI 模型框架」与 EPIC #5「内核平台化改造」。
|
||||||
|
|
||||||
|
本层把化工 AI 的「模型资产」统一注册、版本化、阶段化管理:模型 / 配方 /
|
||||||
|
估计器是资产,需要注册表统一管理——注册、加载、版本、阶段(灰度)、回滚、审计。
|
||||||
|
|
||||||
|
## 当前已交付
|
||||||
|
|
||||||
|
| 模块 | 对应 issue | PRD 5.3 模型 | 说明 |
|
||||||
|
|------|-----------|-------------|------|
|
||||||
|
| `template_registry` | #41 | ③ 模型模板注册/加载/版本机制 | 多版本 + 阶段(dev/staging/prod) + 提升 + 回滚 + 审计 + 持久化 |
|
||||||
|
|
||||||
|
## 模板注册表(`template_registry.py`)
|
||||||
|
|
||||||
|
模板化之后,模型 / 配方是「资产」,需要注册表统一管理:
|
||||||
|
- **注册**:登记一个模型模板(含主干、超参、特征列、指标、版本);
|
||||||
|
- **加载**:按 name + version/别名/stage 取出;
|
||||||
|
- **版本**:同模板多版本共存,可回滚、可审计;
|
||||||
|
- **阶段**:版本带 stage 标签(dev/staging/prod),灰度发布可控。
|
||||||
|
|
||||||
|
### 核心组件
|
||||||
|
|
||||||
|
- **`ModelTemplate`**:模型模板数据对象(name/version/backbone/hyperparams/
|
||||||
|
feature_columns/metrics/stage),不可变、可序列化往返。
|
||||||
|
- **`TemplateRegistry`**:注册表核心 API:
|
||||||
|
- `register`(校验完整性 + 同版本号拒重复,force 可覆盖)
|
||||||
|
- `get`(按 version / stage / 默认最新加载)
|
||||||
|
- `promote`(dev→staging→prod 逐级提升)
|
||||||
|
- `rollback`(stage 指针回退,保留历史可审计)
|
||||||
|
- `set_stage`(直接设 stage,紧急回滚)
|
||||||
|
- `list_versions` / `list_by_stage` / `stage_pointer` / `history`
|
||||||
|
- `save` / `load`(JSON 持久化,重启恢复)
|
||||||
|
- **`Stage`**:阶段枚举(DEV / STAGING / PROD),PRD 灰度三段制。
|
||||||
|
- **`is_valid_version`**:语义化版本号校验(v1 / 1.0.0 / v1.2-rc1)。
|
||||||
|
|
||||||
|
### 快速开始
|
||||||
|
|
||||||
|
```python
|
||||||
|
from template_registry import ModelTemplate, Stage, TemplateRegistry
|
||||||
|
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(
|
||||||
|
name="ti-quality", version="v1.0", backbone="quality_forecast",
|
||||||
|
feature_columns=("furnace_temp", "cl2_flow"),
|
||||||
|
metrics={"accuracy": 0.91}, stage="dev"))
|
||||||
|
reg.register(ModelTemplate(name="ti-quality", version="v1.1",
|
||||||
|
backbone="quality_forecast", metrics={"accuracy": 0.94}))
|
||||||
|
|
||||||
|
reg.promote("ti-quality", "v1.1") # dev -> staging
|
||||||
|
reg.promote("ti-quality", "v1.1") # staging -> prod
|
||||||
|
reg.rollback("ti-quality", Stage.PROD, "v1.0") # 紧急回滚
|
||||||
|
|
||||||
|
serving = reg.get("ti-quality", stage=Stage.PROD) # 按阶段加载
|
||||||
|
reg.save("registry.json") # 持久化审计
|
||||||
|
```
|
||||||
|
|
||||||
|
## 测试
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd core/model-framework
|
||||||
|
python -m unittest discover -s tests -v
|
||||||
|
python _sanity_check.py
|
||||||
|
```
|
||||||
|
|
||||||
|
## 与规划模块的关系
|
||||||
|
|
||||||
|
接口风格对齐 #34(Model Recipe)、#36(quality_forecast)、#38
|
||||||
|
(cross_process_optimizer)、#40(pipeline)。本模块是 #40 `ModelRegistry`
|
||||||
|
的完整版(多阶段 + 回滚 + 持久化),#40 的 `RegisterStep` 未来可直接对接本
|
||||||
|
注册表,业务侧零改动。
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""iAOP-Core · 模型框架层(AI Model Framework)。
|
||||||
|
|
||||||
|
对应 PRD 5.3「③ AI 模型框架」与 EPIC #5(内核平台化改造)。
|
||||||
|
|
||||||
|
当前已交付(自包含,不依赖未合并分支):
|
||||||
|
- ``template_registry``:模型模板注册 / 加载 / 版本机制(多版本 + 阶段 dev/
|
||||||
|
staging/prod + 提升 + 回滚 + 审计 + JSON 持久化),issue #41。
|
||||||
|
|
||||||
|
规划(待相关 PR 合入后无缝对接,业务侧零改动):
|
||||||
|
- ``pipeline``(issue #40,PR #107)、``cross_process_optimizer``(issue #38,
|
||||||
|
PR #106)等可注册为本注册表的具名模板。
|
||||||
|
"""
|
||||||
|
from model_framework.template_registry import ( # noqa: F401
|
||||||
|
ALLOWED_BACKBONES,
|
||||||
|
ModelTemplate,
|
||||||
|
Stage,
|
||||||
|
TemplateRegistry,
|
||||||
|
TemplateRegistryError,
|
||||||
|
is_valid_version,
|
||||||
|
next_stage,
|
||||||
|
)
|
||||||
@@ -0,0 +1,90 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""模型模板注册 / 加载 / 版本机制 sanity 检查(issue #41)。
|
||||||
|
|
||||||
|
验证 PRD 5.3 模板资产管理的端到端能力:
|
||||||
|
1. 注册多版本 → 阶段提升(dev→staging→prod)→ 回滚;
|
||||||
|
2. 按版本 / 阶段 / 默认加载;
|
||||||
|
3. JSON 持久化往返一致;
|
||||||
|
4. 审计日志完整。
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import tempfile
|
||||||
|
|
||||||
|
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
if HERE not in sys.path:
|
||||||
|
sys.path.insert(0, HERE)
|
||||||
|
|
||||||
|
from template_registry import ( # noqa: E402
|
||||||
|
ModelTemplate, Stage, TemplateRegistry, is_valid_version, next_stage)
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> int:
|
||||||
|
failures = []
|
||||||
|
|
||||||
|
# 版本号校验
|
||||||
|
for v in ["v1", "1.0.0", "v2.1-rc3"]:
|
||||||
|
if not is_valid_version(v):
|
||||||
|
failures.append(f"版本号 {v!r} 应合法")
|
||||||
|
|
||||||
|
# 注册 + 提升 + 回滚
|
||||||
|
try:
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(
|
||||||
|
name="ti-quality", version="v1.0", backbone="quality_forecast",
|
||||||
|
feature_columns=("furnace_temp", "cl2_flow"),
|
||||||
|
metrics={"accuracy": 0.91}, stage="dev"))
|
||||||
|
reg.register(ModelTemplate(
|
||||||
|
name="ti-quality", version="v1.1", backbone="quality_forecast",
|
||||||
|
feature_columns=("furnace_temp", "cl2_flow", "impurity_fe"),
|
||||||
|
metrics={"accuracy": 0.94}, stage="dev"))
|
||||||
|
|
||||||
|
assert reg.list_versions("ti-quality") == ["v1.0", "v1.1"]
|
||||||
|
|
||||||
|
# v1.1 逐级提升到 prod
|
||||||
|
reg.promote("ti-quality", "v1.1") # dev->staging
|
||||||
|
reg.promote("ti-quality", "v1.1") # staging->prod
|
||||||
|
assert reg.stage_pointer("ti-quality", Stage.PROD) == "v1.1"
|
||||||
|
|
||||||
|
# 紧急回滚 prod 到 v1.0
|
||||||
|
reg.rollback("ti-quality", Stage.PROD, "v1.0")
|
||||||
|
assert reg.stage_pointer("ti-quality", Stage.PROD) == "v1.0"
|
||||||
|
assert "v1.1" in reg.list_versions("ti-quality") # 历史保留
|
||||||
|
|
||||||
|
# 按阶段加载
|
||||||
|
serving = reg.get("ti-quality", stage=Stage.PROD)
|
||||||
|
assert serving.version == "v1.0", f"prod 应为 v1.0,实为 {serving.version}"
|
||||||
|
|
||||||
|
print(f"[ti-quality] 版本={reg.list_versions('ti-quality')} "
|
||||||
|
f"prod={reg.stage_pointer('ti-quality', Stage.PROD)}")
|
||||||
|
print(f" v1.0 metrics={reg.get('ti-quality','v1.0').metrics}")
|
||||||
|
print(f" v1.1 metrics={reg.get('ti-quality','v1.1').metrics}")
|
||||||
|
print(f" 审计日志 {len(reg.history('ti-quality'))} 条")
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
failures.append(f"注册/提升/回滚流程失败:{exc}")
|
||||||
|
|
||||||
|
# 持久化往返
|
||||||
|
try:
|
||||||
|
with tempfile.NamedTemporaryFile(
|
||||||
|
mode="w", suffix=".json", delete=False, encoding="utf-8") as fh:
|
||||||
|
path = fh.name
|
||||||
|
reg.save(path)
|
||||||
|
reg2 = TemplateRegistry.load(path)
|
||||||
|
assert reg2.list_versions("ti-quality") == ["v1.0", "v1.1"]
|
||||||
|
assert reg2.stage_pointer("ti-quality", Stage.PROD) == "v1.0"
|
||||||
|
print(f"[persist] save/load 往返一致,恢复 {len(reg2)} 个模板")
|
||||||
|
os.unlink(path)
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
failures.append(f"持久化往返失败:{exc}")
|
||||||
|
|
||||||
|
if failures:
|
||||||
|
print("\n失败项:")
|
||||||
|
for f in failures:
|
||||||
|
print(f" ✗ {f}")
|
||||||
|
return 1
|
||||||
|
print("\n✓ template_registry sanity check 通过")
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(main())
|
||||||
@@ -0,0 +1,395 @@
|
|||||||
|
# -*- 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
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""测试引导:把连字符目录 ``core/model-framework`` 加载为可导入包
|
||||||
|
``model_framework``,使测试可 ``from model_framework import ...``。
|
||||||
|
"""
|
||||||
|
import importlib.util
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
PKG_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
|
||||||
|
|
||||||
|
def _load_package(name: str, path: str) -> None:
|
||||||
|
if name in sys.modules:
|
||||||
|
return
|
||||||
|
init_py = os.path.join(path, "__init__.py")
|
||||||
|
spec = importlib.util.spec_from_file_location(
|
||||||
|
name, init_py, submodule_search_locations=[path])
|
||||||
|
module = importlib.util.module_from_spec(spec)
|
||||||
|
sys.modules[name] = module
|
||||||
|
spec.loader.exec_module(module)
|
||||||
|
|
||||||
|
|
||||||
|
_load_package("model_framework", PKG_DIR)
|
||||||
@@ -0,0 +1,235 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""模型模板注册 / 加载 / 版本机制单元测试(issue #41)。
|
||||||
|
|
||||||
|
覆盖:
|
||||||
|
- 版本号语义校验 ``is_valid_version``;
|
||||||
|
- Stage 枚举与 ``next_stage`` 阶段提升顺序;
|
||||||
|
- ModelTemplate 构造校验(name/version/backbone/stage)+ 序列化往返;
|
||||||
|
- TemplateRegistry:注册(拒重复 / force 覆盖)、加载(version/stage/默认)、
|
||||||
|
阶段提升 promote、回滚 rollback、set_stage、查询(list_*)、审计日志、
|
||||||
|
JSON 持久化 save/load 往返一致性。
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
if HERE not in sys.path:
|
||||||
|
sys.path.insert(0, HERE)
|
||||||
|
|
||||||
|
import _bootstrap # noqa: E402 加载 model_framework 包
|
||||||
|
|
||||||
|
from model_framework import ( # noqa: E402
|
||||||
|
ModelTemplate,
|
||||||
|
Stage,
|
||||||
|
TemplateRegistry,
|
||||||
|
TemplateRegistryError,
|
||||||
|
is_valid_version,
|
||||||
|
next_stage,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestVersionValidation(unittest.TestCase):
|
||||||
|
def test_valid_versions(self):
|
||||||
|
for v in ["v1", "1.0", "v1.2.3", "1.2.3", "v1.0-rc1", "v2.0.0+build5"]:
|
||||||
|
self.assertTrue(is_valid_version(v), f"应合法:{v}")
|
||||||
|
|
||||||
|
def test_invalid_versions(self):
|
||||||
|
for v in ["", "v", "abc", "v1.x", "1..2", None, "v 1"]:
|
||||||
|
self.assertFalse(is_valid_version(v), f"应非法:{v!r}")
|
||||||
|
|
||||||
|
|
||||||
|
class TestStage(unittest.TestCase):
|
||||||
|
def test_from_str(self):
|
||||||
|
self.assertEqual(Stage.from_str("dev"), Stage.DEV)
|
||||||
|
self.assertEqual(Stage.from_str("PROD"), Stage.PROD)
|
||||||
|
|
||||||
|
def test_from_str_invalid(self):
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
Stage.from_str("qa")
|
||||||
|
|
||||||
|
def test_next_stage(self):
|
||||||
|
self.assertEqual(next_stage(Stage.DEV), Stage.STAGING)
|
||||||
|
self.assertEqual(next_stage(Stage.STAGING), Stage.PROD)
|
||||||
|
self.assertIsNone(next_stage(Stage.PROD))
|
||||||
|
|
||||||
|
|
||||||
|
class TestModelTemplate(unittest.TestCase):
|
||||||
|
def test_construct_minimal(self):
|
||||||
|
t = ModelTemplate(name="m", version="v1")
|
||||||
|
self.assertEqual(t.backbone, "generic")
|
||||||
|
self.assertEqual(t.stage, Stage.DEV)
|
||||||
|
|
||||||
|
def test_rejects_empty_name(self):
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
ModelTemplate(name="", version="v1")
|
||||||
|
|
||||||
|
def test_rejects_bad_version(self):
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
ModelTemplate(name="m", version="abc")
|
||||||
|
|
||||||
|
def test_rejects_bad_backbone(self):
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
ModelTemplate(name="m", version="v1", backbone="magic")
|
||||||
|
|
||||||
|
def test_accepts_known_backbones(self):
|
||||||
|
for b in ("quality_forecast", "anomaly_detection",
|
||||||
|
"cross_process_opt", "recipe_opt", "generic"):
|
||||||
|
ModelTemplate(name="m", version="v1", backbone=b)
|
||||||
|
|
||||||
|
def test_roundtrip(self):
|
||||||
|
t = ModelTemplate(
|
||||||
|
name="qa-model", version="v1.2.0", backbone="quality_forecast",
|
||||||
|
hyperparams={"lr": 0.1}, feature_columns=("a", "b"),
|
||||||
|
target_column="y", metrics={"accuracy": 0.93},
|
||||||
|
stage="prod", description="d", extra={"k": "v"})
|
||||||
|
t2 = ModelTemplate.from_dict(json.loads(json.dumps(t.to_dict())))
|
||||||
|
self.assertEqual(t, t2)
|
||||||
|
self.assertEqual(t2.stage, Stage.PROD)
|
||||||
|
self.assertEqual(t2.metrics["accuracy"], 0.93)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRegistryRegisterLoad(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.reg = TemplateRegistry()
|
||||||
|
self.t1 = ModelTemplate(name="m", version="v1", backbone="quality_forecast")
|
||||||
|
self.t2 = ModelTemplate(name="m", version="v2", backbone="quality_forecast")
|
||||||
|
|
||||||
|
def test_register_and_get_by_version(self):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
self.assertEqual(self.reg.get("m", "v1").version, "v1")
|
||||||
|
|
||||||
|
def test_register_duplicate_rejected(self):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
|
||||||
|
def test_register_force_overwrites(self):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
t1_updated = ModelTemplate(
|
||||||
|
name="m", version="v1", description="updated")
|
||||||
|
self.reg.register(t1_updated, force=True)
|
||||||
|
self.assertEqual(self.reg.get("m", "v1").description, "updated")
|
||||||
|
|
||||||
|
def test_get_missing_name(self):
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
self.reg.get("nope")
|
||||||
|
|
||||||
|
def test_get_missing_version(self):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
self.reg.get("m", "v99")
|
||||||
|
|
||||||
|
def test_get_default_latest(self):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
time.sleep(0.01)
|
||||||
|
self.reg.register(self.t2)
|
||||||
|
self.assertEqual(self.reg.get("m").version, "v2")
|
||||||
|
|
||||||
|
def test_get_by_stage_pointer(self):
|
||||||
|
self.reg.register(self.t1)
|
||||||
|
# 新注册默认进 dev
|
||||||
|
self.assertEqual(self.reg.get("m", stage=Stage.DEV).version, "v1")
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
self.reg.get("m", stage=Stage.PROD)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPromoteRollback(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.reg = TemplateRegistry()
|
||||||
|
self.reg.register(ModelTemplate(name="m", version="v1"))
|
||||||
|
self.reg.register(ModelTemplate(name="m", version="v2"))
|
||||||
|
|
||||||
|
def test_promote_chain(self):
|
||||||
|
self.reg.promote("m", "v1") # dev -> staging
|
||||||
|
self.assertEqual(self.reg.stage_pointer("m", Stage.STAGING), "v1")
|
||||||
|
self.reg.promote("m", "v1") # staging -> prod
|
||||||
|
self.assertEqual(self.reg.stage_pointer("m", Stage.PROD), "v1")
|
||||||
|
|
||||||
|
def test_promote_prod_raises(self):
|
||||||
|
self.reg.promote("m", "v1")
|
||||||
|
self.reg.promote("m", "v1") # 到 prod
|
||||||
|
with self.assertRaises(TemplateRegistryError):
|
||||||
|
self.reg.promote("m", "v1") # prod 无法继续
|
||||||
|
|
||||||
|
def test_rollback_stage_pointer(self):
|
||||||
|
self.reg.promote("m", "v2") # v2 -> staging
|
||||||
|
self.reg.promote("m", "v2") # v2 -> prod
|
||||||
|
# 回滚 prod 到 v1
|
||||||
|
self.reg.rollback("m", Stage.PROD, "v1")
|
||||||
|
self.assertEqual(self.reg.stage_pointer("m", Stage.PROD), "v1")
|
||||||
|
# v2 版本本身仍在(可审计)
|
||||||
|
self.assertIn("v2", self.reg.list_versions("m"))
|
||||||
|
|
||||||
|
def test_set_stage_direct(self):
|
||||||
|
self.reg.set_stage("m", "v1", Stage.PROD)
|
||||||
|
self.assertEqual(self.reg.get("m", "v1").stage, Stage.PROD)
|
||||||
|
self.assertEqual(self.reg.stage_pointer("m", Stage.PROD), "v1")
|
||||||
|
|
||||||
|
|
||||||
|
class TestQueries(unittest.TestCase):
|
||||||
|
def test_list_names_and_versions(self):
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(name="a", version="v1"))
|
||||||
|
reg.register(ModelTemplate(name="a", version="v2"))
|
||||||
|
reg.register(ModelTemplate(name="b", version="v1"))
|
||||||
|
self.assertEqual(reg.list_names(), ["a", "b"])
|
||||||
|
self.assertEqual(reg.list_versions("a"), ["v1", "v2"])
|
||||||
|
self.assertIn("a", reg)
|
||||||
|
self.assertNotIn("c", reg)
|
||||||
|
self.assertEqual(len(reg), 3)
|
||||||
|
|
||||||
|
def test_list_by_stage(self):
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(name="m", version="v1"))
|
||||||
|
reg.register(ModelTemplate(name="m", version="v2"))
|
||||||
|
reg.promote("m", "v2") # v2 -> staging
|
||||||
|
self.assertEqual(reg.list_by_stage("m", Stage.DEV), ["v1"])
|
||||||
|
self.assertEqual(reg.list_by_stage("m", Stage.STAGING), ["v2"])
|
||||||
|
|
||||||
|
|
||||||
|
class TestHistoryAndPersist(unittest.TestCase):
|
||||||
|
def test_history_logged(self):
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(name="m", version="v1"))
|
||||||
|
reg.promote("m", "v1")
|
||||||
|
h = reg.history("m")
|
||||||
|
actions = [e["action"] for e in h]
|
||||||
|
self.assertIn("register", actions)
|
||||||
|
self.assertIn("promote", actions)
|
||||||
|
|
||||||
|
def test_history_filter_by_name(self):
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(name="a", version="v1"))
|
||||||
|
reg.register(ModelTemplate(name="b", version="v1"))
|
||||||
|
self.assertEqual(len(reg.history("a")), 1)
|
||||||
|
self.assertEqual(len(reg.history("b")), 1)
|
||||||
|
|
||||||
|
def test_save_load_roundtrip(self):
|
||||||
|
reg = TemplateRegistry()
|
||||||
|
reg.register(ModelTemplate(
|
||||||
|
name="m", version="v1", backbone="quality_forecast",
|
||||||
|
metrics={"accuracy": 0.9}, stage="dev"))
|
||||||
|
reg.promote("m", "v1")
|
||||||
|
with tempfile.NamedTemporaryFile(
|
||||||
|
mode="w", suffix=".json", delete=False, encoding="utf-8") as fh:
|
||||||
|
path = fh.name
|
||||||
|
try:
|
||||||
|
reg.save(path)
|
||||||
|
reg2 = TemplateRegistry.load(path)
|
||||||
|
self.assertEqual(reg2.list_names(), ["m"])
|
||||||
|
self.assertEqual(reg2.get("m", "v1").backbone, "quality_forecast")
|
||||||
|
self.assertEqual(reg2.get("m", "v1").metrics["accuracy"], 0.9)
|
||||||
|
# 阶段指针恢复
|
||||||
|
self.assertEqual(reg2.stage_pointer("m", Stage.STAGING), "v1")
|
||||||
|
# 审计日志恢复
|
||||||
|
self.assertTrue(len(reg2.history("m")) >= 2)
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user