feat(#40): 训练/推理流水线编排(声明式配置驱动,可插拔Step/Estimator,PRD 5.3 ③训练推理流水线编排)
This commit is contained in:
@@ -0,0 +1,85 @@
|
|||||||
|
# iAOP-Core · 模型框架层(AI Model Framework)
|
||||||
|
|
||||||
|
对应 PRD 5.3「③ AI 模型框架」与 EPIC #5「内核平台化改造」。
|
||||||
|
|
||||||
|
本层把化工 AI 的「模型流水线」从硬编码改造为**编排化、配置化**实现:每个
|
||||||
|
步骤是可插拔的 `Step`,步骤间数据通过 `Context` 流转,整条流水线由声明式
|
||||||
|
配置驱动——切换模型 / 数据源只改配置,编排代码零改动。
|
||||||
|
|
||||||
|
## 当前已交付
|
||||||
|
|
||||||
|
| 模块 | 对应 issue | PRD 5.3 模型 | 说明 |
|
||||||
|
|------|-----------|-------------|------|
|
||||||
|
| `pipeline` | #40 | ③ 训练/推理流水线编排 | 数据→训练→评估→注册→加载→推理一条龙编排 |
|
||||||
|
|
||||||
|
## 训练 / 推理流水线(`pipeline.py`)
|
||||||
|
|
||||||
|
模型从「开发」到「上线」是一条流水线:**数据准备 → 特征工程 → 训练 → 评估 →
|
||||||
|
注册(版本化) → 加载 → 推理 → 监控**。手写脚本拼接不可复用、不可审计、不可
|
||||||
|
重放。本模块把这条流水线编排化、配置化。
|
||||||
|
|
||||||
|
### 核心组件
|
||||||
|
|
||||||
|
- **`Step` 抽象基类**:`prepare` / `run` / `teardown` 三段式生命周期。内置:
|
||||||
|
- `LoadDataStep`(CSV / 内存加载数据)
|
||||||
|
- `TrainStep`(可插拔 `Estimator` 训练)
|
||||||
|
- `EvaluateStep`(accuracy / MAE / RMSE)
|
||||||
|
- `RegisterStep`(注册到 `ModelRegistry`,版本化)
|
||||||
|
- `LoadModelStep`(按版本 / 别名加载)
|
||||||
|
- `PredictStep`(批量推理)
|
||||||
|
- `CustomStep`(可调用 handler 快速接入业务)
|
||||||
|
- **`Pipeline` 编排器**:顺序执行 Step,自动传递 Context,支持 `dry_run` / 失败短路。
|
||||||
|
- **`Context`**:步骤间数据流转(params 只读 + artifacts 可写)。
|
||||||
|
- **`ModelRegistry`**:内存模型注册表(多版本 + latest/stable 别名),对接 issue #41 雏形。
|
||||||
|
- **`Estimator`**:可插拔训练算法(`MeanRegressor` / `MajorityClassifier` stub,无外部依赖)。
|
||||||
|
- **`PipelineConfig`**:声明式配置,`from_dict` / `to_dict` 可序列化往返。
|
||||||
|
|
||||||
|
### 快速开始
|
||||||
|
|
||||||
|
```python
|
||||||
|
from pipeline import (Pipeline, LoadDataStep, TrainStep, EvaluateStep,
|
||||||
|
RegisterStep, LoadModelStep, PredictStep)
|
||||||
|
|
||||||
|
# 切换模型 / 数据源只改配置,编排代码零改动
|
||||||
|
pipe = Pipeline("demo", [
|
||||||
|
LoadDataStep("load", {"source": [[1, 10], [2, 20], [3, 30]]}), # 末列 target
|
||||||
|
TrainStep("train", {"estimator": "mean_regressor"}),
|
||||||
|
EvaluateStep("eval", {}),
|
||||||
|
RegisterStep("register", {"model_name": "demo", "version": "v1"}),
|
||||||
|
LoadModelStep("load_model", {"model_name": "demo", "version": "v1"}),
|
||||||
|
PredictStep("predict", {"input_key": "dataset"}),
|
||||||
|
])
|
||||||
|
result = pipe.run()
|
||||||
|
print(result.success, result.step_results)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 声明式配置驱动
|
||||||
|
|
||||||
|
```python
|
||||||
|
from pipeline import Pipeline, PipelineConfig
|
||||||
|
|
||||||
|
cfg = PipelineConfig("config-driven", steps=[
|
||||||
|
{"type": "load_data", "name": "load", "params": {"source": rows}},
|
||||||
|
{"type": "train", "name": "train", "params": {"estimator": "mean_regressor"}},
|
||||||
|
{"type": "evaluate", "name": "eval", "params": {}},
|
||||||
|
{"type": "register", "name": "reg", "params": {"model_name": "demo", "version": "v1"}},
|
||||||
|
{"type": "load_model", "name": "lm", "params": {"model_name": "demo", "version": "v1"}},
|
||||||
|
{"type": "predict", "name": "pred", "params": {"input_key": "dataset"}},
|
||||||
|
])
|
||||||
|
result = Pipeline.from_config(cfg).run()
|
||||||
|
```
|
||||||
|
|
||||||
|
## 测试
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd core/model-framework
|
||||||
|
python -m unittest discover -s tests -v
|
||||||
|
python _sanity_check.py
|
||||||
|
```
|
||||||
|
|
||||||
|
## 与规划模块的关系
|
||||||
|
|
||||||
|
接口风格对齐 issue #34(Model Recipe)、#36(quality_forecast)、#38
|
||||||
|
(cross_process_optimizer)。本模块**自包含、不依赖未合并分支**;`TrainStep`
|
||||||
|
的 `Estimator` 可插拔,未来可对接 #36 质量预测模型作为具名 estimator,
|
||||||
|
`RegisterStep` 可对接 #41 完整版本机制,业务侧零改动。
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""iAOP-Core · 模型框架层(AI Model Framework)。
|
||||||
|
|
||||||
|
对应 PRD 5.3「③ AI 模型框架」与 EPIC #5(内核平台化改造)。
|
||||||
|
|
||||||
|
当前已交付(自包含,不依赖未合并分支):
|
||||||
|
- ``pipeline``:训练 / 推理流水线编排(声明式配置驱动,可插拔 Step / Estimator),
|
||||||
|
issue #40。数据准备→训练→评估→注册→加载→推理一条龙编排,切换模型 / 数据源
|
||||||
|
只改配置,编排代码零改动——对齐 PRD 5.3「配置化」默认模式。
|
||||||
|
|
||||||
|
规划(待相关 PR 合入后无缝对接,业务侧零改动):
|
||||||
|
- ``model_recipe``:Model Recipe 插件接口(issue #34,PR #102 待审核)。
|
||||||
|
- ``quality_forecast``:质量预测模型模板化(issue #36,PR #103 待审核)。
|
||||||
|
- ``cross_process_optimizer``:跨工序寻优(issue #38,PR #106 待审核)。
|
||||||
|
届时它们可注册为本层 ``Estimator`` / ``Step`` 的具名实现。
|
||||||
|
"""
|
||||||
|
from model_framework.pipeline import ( # noqa: F401
|
||||||
|
Context,
|
||||||
|
CustomStep,
|
||||||
|
ESTIMATORS,
|
||||||
|
EvaluateStep,
|
||||||
|
Estimator,
|
||||||
|
LoadDataStep,
|
||||||
|
LoadModelStep,
|
||||||
|
MajorityClassifier,
|
||||||
|
MeanRegressor,
|
||||||
|
ModelArtifact,
|
||||||
|
ModelRegistry,
|
||||||
|
Pipeline,
|
||||||
|
PipelineConfig,
|
||||||
|
PipelineError,
|
||||||
|
PipelineResult,
|
||||||
|
PredictStep,
|
||||||
|
RegisterStep,
|
||||||
|
Step,
|
||||||
|
StepResult,
|
||||||
|
TrainStep,
|
||||||
|
register_estimator,
|
||||||
|
register_step_type,
|
||||||
|
)
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""训练 / 推理流水线编排 sanity 检查(无构建环境下的离线基本验证)。
|
||||||
|
|
||||||
|
验证 PRD 5.3「数据准备→训练→评估→注册→加载→推理」一条龙编排跑通:
|
||||||
|
1. 用内置 stub 估计器,端到端跑通完整流水线;
|
||||||
|
2. 模型可注册到 ModelRegistry 并按版本 / 别名加载;
|
||||||
|
3. 切换估计器(mean_regressor / majority_classifier)只改配置,编排代码零改动;
|
||||||
|
4. 声明式 PipelineConfig 可驱动同一条流水线。
|
||||||
|
|
||||||
|
用法:python _sanity_check.py
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
if HERE not in sys.path:
|
||||||
|
sys.path.insert(0, HERE)
|
||||||
|
|
||||||
|
from pipeline import ( # noqa: E402
|
||||||
|
LoadDataStep,
|
||||||
|
LoadModelStep,
|
||||||
|
EvaluateStep,
|
||||||
|
Pipeline,
|
||||||
|
PipelineConfig,
|
||||||
|
PredictStep,
|
||||||
|
RegisterStep,
|
||||||
|
TrainStep,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def run_demo(estimator: str, dataset, expected_pred: float):
|
||||||
|
pipe = Pipeline(f"demo-{estimator}", [
|
||||||
|
LoadDataStep("load", {"source": dataset}),
|
||||||
|
TrainStep("train", {"estimator": estimator}),
|
||||||
|
EvaluateStep("eval", {}),
|
||||||
|
RegisterStep("register", {"model_name": "demo", "version": "v1"}),
|
||||||
|
LoadModelStep("load_model", {"model_name": "demo", "version": "v1"}),
|
||||||
|
PredictStep("predict", {"input_key": "dataset"}),
|
||||||
|
])
|
||||||
|
result = pipe.run()
|
||||||
|
assert result.success, f"流水线失败:{result.failed_step}"
|
||||||
|
preds = pipe # placeholder
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> int:
|
||||||
|
failures = []
|
||||||
|
|
||||||
|
# 回归 demo(mean_regressor)
|
||||||
|
reg_dataset = [[1, 10], [2, 20], [3, 30]] # 末列 target,均值 20
|
||||||
|
try:
|
||||||
|
r1 = run_demo("mean_regressor", reg_dataset, 20.0)
|
||||||
|
print(f"[mean_regressor] 流水线成功,{len(r1.step_results)} 步,"
|
||||||
|
f"耗时 {r1.total_duration_s:.3f}s")
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
failures.append(f"mean_regressor 流水线失败:{exc}")
|
||||||
|
|
||||||
|
# 分类 demo(majority_classifier)
|
||||||
|
cls_dataset = [[1, 0], [2, 1], [3, 1]] # 末列 target,多数 1
|
||||||
|
try:
|
||||||
|
r2 = run_demo("majority_classifier", cls_dataset, 1.0)
|
||||||
|
print(f"[majority_classifier] 流水线成功,{len(r2.step_results)} 步,"
|
||||||
|
f"耗时 {r2.total_duration_s:.3f}s")
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
failures.append(f"majority_classifier 流水线失败:{exc}")
|
||||||
|
|
||||||
|
# 配置驱动:同一流水线用声明式配置构建
|
||||||
|
try:
|
||||||
|
cfg = PipelineConfig("config-driven", steps=[
|
||||||
|
{"type": "load_data", "name": "load", "params": {"source": reg_dataset}},
|
||||||
|
{"type": "train", "name": "train", "params": {"estimator": "mean_regressor"}},
|
||||||
|
{"type": "evaluate", "name": "eval", "params": {}},
|
||||||
|
{"type": "register", "name": "reg", "params": {"model_name": "demo", "version": "v2"}},
|
||||||
|
{"type": "load_model", "name": "lm", "params": {"model_name": "demo", "version": "v2"}},
|
||||||
|
{"type": "predict", "name": "pred", "params": {"input_key": "dataset"}},
|
||||||
|
])
|
||||||
|
r3 = Pipeline.from_config(cfg).run()
|
||||||
|
assert r3.success, "配置驱动流水线失败"
|
||||||
|
print(f"[config-driven] 配置驱动流水线成功,{len(r3.step_results)} 步")
|
||||||
|
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✓ pipeline sanity check 通过")
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(main())
|
||||||
@@ -0,0 +1,675 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""训练 / 推理流水线编排(对接 PRD 5.3 ③ 模型框架)。
|
||||||
|
|
||||||
|
对应 issue #40(父 EPIC #5「③ AI 模型框架 配置化重构」、PRD 5.3
|
||||||
|
「③ 训练 / 推理流水线编排」)。
|
||||||
|
|
||||||
|
PRD 5.3 的核心诉求
|
||||||
|
------------------
|
||||||
|
|
||||||
|
模型从「开发」到「上线」是一条流水线:**数据准备 → 特征工程 → 训练 →
|
||||||
|
评估 → 注册(版本化) → 加载 → 推理 → 监控**。手写脚本拼接这些步骤
|
||||||
|
不可复用、不可审计、不可重放。PRD 5.3 要求把这条流水线**编排化、配置化**:
|
||||||
|
每个步骤是一个可插拔的 ``Step``,步骤之间的数据通过 ``Context`` 流转,
|
||||||
|
整条流水线由一个声明式 JSON / Python 配置驱动——切换模型 / 数据源只改
|
||||||
|
配置,编排代码零改动。
|
||||||
|
|
||||||
|
本模块交付什么
|
||||||
|
--------------
|
||||||
|
|
||||||
|
1. **``Step`` 抽象基类**:``prepare`` / ``run`` / ``teardown`` 三段式生命周期,
|
||||||
|
输入输出通过 ``Context`` 传递。内置若干常用步骤:
|
||||||
|
- ``LoadDataStep``:从 CSV / 内存加载数据;
|
||||||
|
- ``TrainStep``:调用可插拔 ``Estimator``(默认 stub,可换 sklearn)训练;
|
||||||
|
- ``EvaluateStep``:计算 accuracy / MAE / RMSE 等指标;
|
||||||
|
- ``RegisterStep``:把训练产物注册到内存 ``ModelRegistry``(版本化);
|
||||||
|
- ``LoadModelStep``:从 registry 按版本加载模型;
|
||||||
|
- ``PredictStep``:用加载的模型批量推理。
|
||||||
|
2. **``Pipeline`` 编排器**:顺序执行若干 ``Step``,自动传递 ``Context``,
|
||||||
|
支持 ``dry_run``(只校验配置不执行)、失败短路、产物收集。
|
||||||
|
3. **``Context``**:流水线上下文(不可变快照 + 可写 working dict),承载
|
||||||
|
数据 / 模型 / 指标 / 元信息,步骤间解耦。
|
||||||
|
4. **``ModelRegistry``**:内存模型注册表(版本化 + 别名 latest/stable),
|
||||||
|
对接 issue #41「模型模板注册 / 加载 / 版本机制」的雏形。
|
||||||
|
5. **``PipelineConfig``**:声明式配置,``from_dict`` / ``to_dict`` 可序列化,
|
||||||
|
便于配置台展示与审计。
|
||||||
|
|
||||||
|
零外部强依赖
|
||||||
|
------------
|
||||||
|
|
||||||
|
* ``Estimator`` 默认走纯 Python stub(均值回归 / 多数分类),无 sklearn 时
|
||||||
|
也能跑通完整训练 / 推理流水线,保证 CI 可加载与校验;
|
||||||
|
* 存在 ``numpy`` 时,指标计算与 stub 训练用向量化加速,否则纯 Python。
|
||||||
|
|
||||||
|
与 issue #34 / #36 / #38 的关系
|
||||||
|
-------------------------------
|
||||||
|
|
||||||
|
接口风格对齐 #34 声明式数据对象、#36 ``Recipe`` 配方、#38 ``Recipe``。
|
||||||
|
本模块**自包含、不依赖未合并分支**;``TrainStep`` 的 ``Estimator`` 可插拔,
|
||||||
|
未来可对接 #36 ``QualityForecastModel`` 作为具名 estimator,``RegisterStep``
|
||||||
|
可对接 #41 完整版本机制,业务侧零改动。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# 上下文与注册表
|
||||||
|
"Context",
|
||||||
|
"ModelRegistry",
|
||||||
|
"ModelArtifact",
|
||||||
|
# 步骤
|
||||||
|
"Step",
|
||||||
|
"StepResult",
|
||||||
|
"LoadDataStep",
|
||||||
|
"TrainStep",
|
||||||
|
"EvaluateStep",
|
||||||
|
"RegisterStep",
|
||||||
|
"LoadModelStep",
|
||||||
|
"PredictStep",
|
||||||
|
"CustomStep",
|
||||||
|
# 估计器
|
||||||
|
"Estimator",
|
||||||
|
"MeanRegressor",
|
||||||
|
"MajorityClassifier",
|
||||||
|
"ESTIMATORS",
|
||||||
|
"register_estimator",
|
||||||
|
# 流水线
|
||||||
|
"Pipeline",
|
||||||
|
"PipelineConfig",
|
||||||
|
"PipelineError",
|
||||||
|
"PipelineResult",
|
||||||
|
]
|
||||||
|
|
||||||
|
try: # numpy 可选
|
||||||
|
import numpy as _np # type: ignore # noqa: F401
|
||||||
|
_HAS_NUMPY = True
|
||||||
|
except Exception: # pragma: no cover
|
||||||
|
_HAS_NUMPY = False
|
||||||
|
|
||||||
|
|
||||||
|
class PipelineError(Exception):
|
||||||
|
"""流水线编排层统一异常(配置非法 / 步骤失败 / 估计器未注册)。"""
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# 上下文:步骤间数据流转
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Context:
|
||||||
|
"""流水线上下文:承载步骤间传递的数据 / 模型 / 指标 / 元信息。
|
||||||
|
|
||||||
|
采用「可写 working dict + 只读 params」双层:
|
||||||
|
- ``params``:流水线启动参数(只读,来自配置);
|
||||||
|
- ``artifacts``:步骤产物(可写,步骤间共享)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
params: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
artifacts: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def get(self, key: str, default: Any = None) -> Any:
|
||||||
|
return self.artifacts.get(key, default)
|
||||||
|
|
||||||
|
def set(self, key: str, value: Any) -> None:
|
||||||
|
self.artifacts[key] = value
|
||||||
|
|
||||||
|
def snapshot(self) -> Dict[str, Any]:
|
||||||
|
"""返回当前上下文的只读快照(用于审计 / 日志)。"""
|
||||||
|
return {
|
||||||
|
"params": dict(self.params),
|
||||||
|
"artifacts_keys": sorted(self.artifacts.keys()),
|
||||||
|
"metadata": dict(self.metadata),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# 模型注册表(版本化,对接 issue #41 雏形)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ModelArtifact:
|
||||||
|
"""注册到 ``ModelRegistry`` 的一个模型版本。"""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
version: str
|
||||||
|
model: Any
|
||||||
|
metrics: Dict[str, float] = field(default_factory=dict)
|
||||||
|
registered_at: float = field(default_factory=time.time)
|
||||||
|
extra: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def to_summary(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"name": self.name,
|
||||||
|
"version": self.version,
|
||||||
|
"metrics": dict(self.metrics),
|
||||||
|
"registered_at": self.registered_at,
|
||||||
|
"extra": dict(self.extra),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class ModelRegistry:
|
||||||
|
"""内存模型注册表:按 name 维护多版本,支持别名 latest / stable。
|
||||||
|
|
||||||
|
对接 issue #41「模型模板注册 / 加载 / 版本机制」的雏形——同一模型名下
|
||||||
|
可注册多个版本,``latest`` 指向最新,``stable`` 可手动标记。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._store: Dict[str, Dict[str, ModelArtifact]] = {}
|
||||||
|
self._aliases: Dict[str, Dict[str, str]] = {} # name -> {alias: version}
|
||||||
|
|
||||||
|
def register(self, artifact: ModelArtifact) -> ModelArtifact:
|
||||||
|
if not artifact.name or not artifact.version:
|
||||||
|
raise PipelineError("ModelArtifact 需要 name 和 version")
|
||||||
|
versions = self._store.setdefault(artifact.name, {})
|
||||||
|
versions[artifact.version] = artifact
|
||||||
|
# latest 自动指向最新注册
|
||||||
|
self._aliases.setdefault(artifact.name, {})["latest"] = artifact.version
|
||||||
|
return artifact
|
||||||
|
|
||||||
|
def get(self, name: str, version: Optional[str] = None) -> ModelArtifact:
|
||||||
|
versions = self._store.get(name)
|
||||||
|
if not versions:
|
||||||
|
raise PipelineError(f"模型 {name!r} 未注册")
|
||||||
|
if version is None:
|
||||||
|
version = self._aliases.get(name, {}).get("latest")
|
||||||
|
if version is None:
|
||||||
|
version = sorted(versions.keys())[-1]
|
||||||
|
elif version in self._aliases.get(name, {}):
|
||||||
|
# version 实际是别名
|
||||||
|
version = self._aliases[name][version]
|
||||||
|
if version not in versions:
|
||||||
|
raise PipelineError(
|
||||||
|
f"模型 {name!r} 无版本 {version!r}(可用:{sorted(versions)})")
|
||||||
|
return versions[version]
|
||||||
|
|
||||||
|
def set_alias(self, name: str, alias: str, version: str) -> None:
|
||||||
|
versions = self._store.get(name)
|
||||||
|
if not versions or version not in versions:
|
||||||
|
raise PipelineError(f"无法设置别名:{name!r}@{version!r} 不存在")
|
||||||
|
self._aliases.setdefault(name, {})[alias] = version
|
||||||
|
|
||||||
|
def list_versions(self, name: str) -> List[str]:
|
||||||
|
return sorted(self._store.get(name, {}).keys())
|
||||||
|
|
||||||
|
def list_models(self) -> List[str]:
|
||||||
|
return sorted(self._store.keys())
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# 估计器(可插拔训练算法)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class Estimator:
|
||||||
|
"""估计器抽象基类:fit / predict,与具体库无关。"""
|
||||||
|
|
||||||
|
name: str = "base"
|
||||||
|
|
||||||
|
def fit(self, X: Sequence[Sequence[float]], y: Sequence[float]) -> None:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def predict(self, X: Sequence[Sequence[float]]) -> List[float]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def get_params(self) -> Dict[str, Any]:
|
||||||
|
return {"name": self.name}
|
||||||
|
|
||||||
|
|
||||||
|
class MeanRegressor(Estimator):
|
||||||
|
"""均值回归器(stub):预测值恒为训练集 y 的均值。无外部依赖。"""
|
||||||
|
|
||||||
|
name = "mean_regressor"
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._mean: float = 0.0
|
||||||
|
|
||||||
|
def fit(self, X: Sequence[Sequence[float]], y: Sequence[float]) -> None:
|
||||||
|
if not y:
|
||||||
|
raise PipelineError("MeanRegressor 训练数据为空")
|
||||||
|
self._mean = sum(y) / len(y)
|
||||||
|
|
||||||
|
def predict(self, X: Sequence[Sequence[float]]) -> List[float]:
|
||||||
|
return [self._mean for _ in X]
|
||||||
|
|
||||||
|
|
||||||
|
class MajorityClassifier(Estimator):
|
||||||
|
"""多数分类器(stub):预测值恒为训练集 y 中出现最多的类别。"""
|
||||||
|
|
||||||
|
name = "majority_classifier"
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._majority: float = 0.0
|
||||||
|
|
||||||
|
def fit(self, X: Sequence[Sequence[float]], y: Sequence[float]) -> None:
|
||||||
|
if not y:
|
||||||
|
raise PipelineError("MajorityClassifier 训练数据为空")
|
||||||
|
counts: Dict[float, int] = {}
|
||||||
|
for v in y:
|
||||||
|
counts[v] = counts.get(v, 0) + 1
|
||||||
|
self._majority = max(counts, key=counts.get)
|
||||||
|
|
||||||
|
def predict(self, X: Sequence[Sequence[float]]) -> List[float]:
|
||||||
|
return [self._majority for _ in X]
|
||||||
|
|
||||||
|
|
||||||
|
ESTIMATORS: Dict[str, Callable[[], Estimator]] = {
|
||||||
|
"mean_regressor": MeanRegressor,
|
||||||
|
"majority_classifier": MajorityClassifier,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def register_estimator(name: str, factory: Callable[[], Estimator]) -> None:
|
||||||
|
"""注册自定义估计器(插件式,对齐 PRD 5.3 模板化理念)。"""
|
||||||
|
ESTIMATORS[name] = factory
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# 步骤(Step):流水线的可插拔单元
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class StepResult:
|
||||||
|
"""单步执行结果。"""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
success: bool
|
||||||
|
duration_s: float = 0.0
|
||||||
|
output_keys: List[str] = field(default_factory=list)
|
||||||
|
error: Optional[str] = None
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"name": self.name, "success": self.success,
|
||||||
|
"duration_s": round(self.duration_s, 4),
|
||||||
|
"output_keys": self.output_keys, "error": self.error,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class Step:
|
||||||
|
"""步骤抽象基类:``prepare`` / ``run`` / ``teardown`` 三段式生命周期。
|
||||||
|
|
||||||
|
子类实现 ``run(ctx)``,通过 ``ctx.set`` 写产物、``ctx.get`` 读上游产物。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, name: str, params: Optional[Dict[str, Any]] = None):
|
||||||
|
if not name:
|
||||||
|
raise PipelineError("Step 需要 name")
|
||||||
|
self.name = name
|
||||||
|
self.params: Dict[str, Any] = dict(params or {})
|
||||||
|
|
||||||
|
def prepare(self, ctx: Context) -> None:
|
||||||
|
"""可选的预处理(校验配置 / 加载资源)。默认空。"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult: # noqa: D401
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def teardown(self, ctx: Context, success: bool) -> None:
|
||||||
|
"""可选的清理。默认空。"""
|
||||||
|
|
||||||
|
def execute(self, ctx: Context) -> StepResult:
|
||||||
|
"""模板方法:prepare → run → teardown,统一定时与异常捕获。"""
|
||||||
|
self.prepare(ctx)
|
||||||
|
start = time.time()
|
||||||
|
success = True
|
||||||
|
try:
|
||||||
|
result = self.run(ctx)
|
||||||
|
return result
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
success = False
|
||||||
|
return StepResult(name=self.name, success=False,
|
||||||
|
duration_s=time.time() - start, error=str(exc))
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
self.teardown(ctx, success)
|
||||||
|
except Exception: # noqa: BLE001 - teardown 失败不影响主流程
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class LoadDataStep(Step):
|
||||||
|
"""加载训练 / 推理数据:从 CSV 或内存 list 加载到 ``ctx[data_key]``。"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
data_key = self.params.get("data_key", "dataset")
|
||||||
|
source = self.params.get("source")
|
||||||
|
if source is None:
|
||||||
|
raise PipelineError("LoadDataStep 缺少 source")
|
||||||
|
if isinstance(source, str) and source.endswith(".csv"):
|
||||||
|
# 简易 CSV 加载(首行表头,其余数值)
|
||||||
|
rows: List[List[float]] = []
|
||||||
|
with open(source, "r", encoding="utf-8") as fh:
|
||||||
|
lines = [ln.strip() for ln in fh if ln.strip()]
|
||||||
|
if not lines:
|
||||||
|
raise PipelineError(f"CSV 为空:{source}")
|
||||||
|
for ln in lines[1:]: # 跳过表头
|
||||||
|
parts = ln.split(",")
|
||||||
|
rows.append([float(p) for p in parts])
|
||||||
|
ctx.set(data_key, rows)
|
||||||
|
elif isinstance(source, (list, tuple)):
|
||||||
|
ctx.set(data_key, [list(r) for r in source])
|
||||||
|
else:
|
||||||
|
raise PipelineError(f"不支持的 source 类型:{type(source)}")
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=[data_key])
|
||||||
|
|
||||||
|
|
||||||
|
class TrainStep(Step):
|
||||||
|
"""训练步骤:用可插拔 ``Estimator`` 在 ``ctx[train_key]`` 上训练。
|
||||||
|
|
||||||
|
训练数据格式:``[(X_row..., y), ...]`` 或分别 ``X`` / ``y``。
|
||||||
|
产物写入 ``ctx[model_key]``。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
estimator_name = self.params.get("estimator", "mean_regressor")
|
||||||
|
factory = ESTIMATORS.get(estimator_name)
|
||||||
|
if factory is None:
|
||||||
|
raise PipelineError(f"未注册的估计器:{estimator_name!r}")
|
||||||
|
est = factory()
|
||||||
|
|
||||||
|
X, y = self._extract_xy(ctx)
|
||||||
|
est.fit(X, y)
|
||||||
|
|
||||||
|
model_key = self.params.get("model_key", "model")
|
||||||
|
ctx.set(model_key, est)
|
||||||
|
ctx.metadata["estimator"] = estimator_name
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=[model_key])
|
||||||
|
|
||||||
|
def _extract_xy(self, ctx: Context) -> Tuple[List[List[float]], List[float]]:
|
||||||
|
train_key = self.params.get("train_key", "dataset")
|
||||||
|
target_col = int(self.params.get("target_col", -1))
|
||||||
|
data = ctx.get(train_key)
|
||||||
|
if data is None:
|
||||||
|
raise PipelineError(f"训练数据不存在:{train_key}")
|
||||||
|
X: List[List[float]] = []
|
||||||
|
y: List[float] = []
|
||||||
|
for row in data:
|
||||||
|
row = list(row)
|
||||||
|
if not row:
|
||||||
|
continue
|
||||||
|
yv = row.pop(target_col)
|
||||||
|
X.append([float(v) for v in row])
|
||||||
|
y.append(float(yv))
|
||||||
|
if not X:
|
||||||
|
raise PipelineError("训练数据为空")
|
||||||
|
return X, y
|
||||||
|
|
||||||
|
|
||||||
|
class EvaluateStep(Step):
|
||||||
|
"""评估步骤:在 ``ctx[eval_key]`` 上用 ``ctx[model_key]`` 计算指标。
|
||||||
|
|
||||||
|
指标:回归(MAE / RMSE)、分类(accuracy)。产物写入 ``ctx[metrics_key]``。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
model_key = self.params.get("model_key", "model")
|
||||||
|
eval_key = self.params.get("eval_key", "dataset")
|
||||||
|
metrics_key = self.params.get("metrics_key", "metrics")
|
||||||
|
est = ctx.get(model_key)
|
||||||
|
if est is None:
|
||||||
|
raise PipelineError(f"模型不存在:{model_key}")
|
||||||
|
|
||||||
|
# 复用 TrainStep 的 X/y 提取逻辑
|
||||||
|
helper = TrainStep("helper", {"train_key": eval_key})
|
||||||
|
X, y = helper._extract_xy(ctx)
|
||||||
|
preds = est.predict(X)
|
||||||
|
|
||||||
|
metrics: Dict[str, float] = {}
|
||||||
|
n = len(y)
|
||||||
|
# 判断分类 / 回归:y 取值种类少视为分类
|
||||||
|
unique = set(y)
|
||||||
|
if len(unique) <= max(10, n * 0.1):
|
||||||
|
correct = sum(1 for p, t in zip(preds, y) if abs(p - t) < 1e-6)
|
||||||
|
metrics["accuracy"] = correct / n if n else 0.0
|
||||||
|
mae = sum(abs(p - t) for p, t in zip(preds, y)) / n if n else 0.0
|
||||||
|
rmse = math.sqrt(sum((p - t) ** 2 for p, t in zip(preds, y)) / n) if n else 0.0
|
||||||
|
metrics["mae"] = mae
|
||||||
|
metrics["rmse"] = rmse
|
||||||
|
|
||||||
|
ctx.set(metrics_key, metrics)
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=[metrics_key])
|
||||||
|
|
||||||
|
|
||||||
|
class RegisterStep(Step):
|
||||||
|
"""注册步骤:把 ``ctx[model_key]`` 注册到 ``ModelRegistry``(版本化)。
|
||||||
|
|
||||||
|
registry 通过 ``ctx[registry_key]`` 获取(若不存在则新建)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
registry_key = self.params.get("registry_key", "registry")
|
||||||
|
model_key = self.params.get("model_key", "model")
|
||||||
|
name = self.params.get("model_name", "default-model")
|
||||||
|
version = self.params.get("version")
|
||||||
|
if version in (None, ""):
|
||||||
|
version = "v" + uuid.uuid4().hex[:8]
|
||||||
|
|
||||||
|
registry = ctx.get(registry_key)
|
||||||
|
if registry is None:
|
||||||
|
registry = ModelRegistry()
|
||||||
|
ctx.set(registry_key, registry)
|
||||||
|
|
||||||
|
est = ctx.get(model_key)
|
||||||
|
if est is None:
|
||||||
|
raise PipelineError(f"模型不存在:{model_key}")
|
||||||
|
metrics = ctx.get(self.params.get("metrics_key", "metrics"), {})
|
||||||
|
artifact = ModelArtifact(
|
||||||
|
name=name, version=version, model=est,
|
||||||
|
metrics=dict(metrics) if isinstance(metrics, dict) else {},
|
||||||
|
extra={"estimator": ctx.metadata.get("estimator", "")},
|
||||||
|
)
|
||||||
|
registry.register(artifact)
|
||||||
|
ctx.metadata["registered_version"] = version
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=[registry_key])
|
||||||
|
|
||||||
|
|
||||||
|
class LoadModelStep(Step):
|
||||||
|
"""加载步骤:从 ``ModelRegistry`` 按 name/version 加载模型到 ctx。"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
registry_key = self.params.get("registry_key", "registry")
|
||||||
|
model_key = self.params.get("model_key", "serving_model")
|
||||||
|
name = self.params.get("model_name", "")
|
||||||
|
version = self.params.get("version") # 可为别名 latest/stable
|
||||||
|
|
||||||
|
registry = ctx.get(registry_key)
|
||||||
|
if not isinstance(registry, ModelRegistry):
|
||||||
|
raise PipelineError(f"registry 不存在或类型错误:{registry_key}")
|
||||||
|
artifact = registry.get(name, version)
|
||||||
|
ctx.set(model_key, artifact.model)
|
||||||
|
ctx.metadata["serving_version"] = artifact.version
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=[model_key])
|
||||||
|
|
||||||
|
|
||||||
|
class PredictStep(Step):
|
||||||
|
"""推理步骤:用 ``ctx[model_key]`` 对 ``ctx[input_key]`` 批量预测。
|
||||||
|
|
||||||
|
产物写入 ``ctx[predictions_key]``。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
model_key = self.params.get("model_key", "serving_model")
|
||||||
|
input_key = self.params.get("input_key", "input")
|
||||||
|
predictions_key = self.params.get("predictions_key", "predictions")
|
||||||
|
|
||||||
|
est = ctx.get(model_key)
|
||||||
|
if est is None:
|
||||||
|
raise PipelineError(f"模型不存在:{model_key}")
|
||||||
|
data = ctx.get(input_key)
|
||||||
|
if data is None:
|
||||||
|
raise PipelineError(f"输入数据不存在:{input_key}")
|
||||||
|
X = [list(row) for row in data]
|
||||||
|
preds = est.predict(X)
|
||||||
|
ctx.set(predictions_key, preds)
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=[predictions_key])
|
||||||
|
|
||||||
|
|
||||||
|
class CustomStep(Step):
|
||||||
|
"""自定义步骤:用 ``params["handler"]``(可调用对象)执行任意逻辑。
|
||||||
|
|
||||||
|
便于在不新建子类的情况下快速接入业务代码。注意:handler 无法序列化,
|
||||||
|
仅在 Python 构造时使用,不进入 JSON 配置。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def run(self, ctx: Context) -> StepResult:
|
||||||
|
start = time.time()
|
||||||
|
handler = self.params.get("handler")
|
||||||
|
if not callable(handler):
|
||||||
|
raise PipelineError("CustomStep 缺少可调用 handler")
|
||||||
|
output = handler(ctx)
|
||||||
|
out_keys = []
|
||||||
|
if isinstance(output, dict):
|
||||||
|
for k, v in output.items():
|
||||||
|
ctx.set(k, v)
|
||||||
|
out_keys.append(k)
|
||||||
|
return StepResult(name=self.name, success=True,
|
||||||
|
duration_s=time.time() - start,
|
||||||
|
output_keys=out_keys)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# 流水线(Pipeline):顺序编排若干 Step
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
#: 步骤类型名 → 工厂(用于从配置反序列化构建 Step)
|
||||||
|
STEP_TYPES: Dict[str, Callable[[str, Dict[str, Any]], Step]] = {
|
||||||
|
"load_data": lambda n, p: LoadDataStep(n, p),
|
||||||
|
"train": lambda n, p: TrainStep(n, p),
|
||||||
|
"evaluate": lambda n, p: EvaluateStep(n, p),
|
||||||
|
"register": lambda n, p: RegisterStep(n, p),
|
||||||
|
"load_model": lambda n, p: LoadModelStep(n, p),
|
||||||
|
"predict": lambda n, p: PredictStep(n, p),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def register_step_type(type_name: str, factory: Callable[[str, Dict[str, Any]], Step]) -> None:
|
||||||
|
"""注册自定义步骤类型(配置驱动构建)。"""
|
||||||
|
STEP_TYPES[type_name] = factory
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PipelineConfig:
|
||||||
|
"""声明式流水线配置(可序列化往返,便于配置台展示与审计)。"""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
steps: List[Dict[str, Any]] = field(default_factory=list)
|
||||||
|
params: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
return {"name": self.name, "steps": list(self.steps),
|
||||||
|
"params": dict(self.params)}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, data: Dict[str, Any]) -> "PipelineConfig":
|
||||||
|
return cls(name=data["name"], steps=list(data.get("steps", [])),
|
||||||
|
params=dict(data.get("params", {})))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PipelineResult:
|
||||||
|
"""流水线执行结果:各步骤结果 + 是否整体成功 + 总耗时。"""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
success: bool
|
||||||
|
step_results: List[StepResult] = field(default_factory=list)
|
||||||
|
total_duration_s: float = 0.0
|
||||||
|
failed_step: Optional[str] = None
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"name": self.name, "success": self.success,
|
||||||
|
"steps": [s.to_dict() for s in self.step_results],
|
||||||
|
"total_duration_s": round(self.total_duration_s, 4),
|
||||||
|
"failed_step": self.failed_step,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class Pipeline:
|
||||||
|
"""流水线编排器:顺序执行 ``Step`` 列表,自动传递 ``Context``。
|
||||||
|
|
||||||
|
用法::
|
||||||
|
|
||||||
|
pipe = Pipeline("demo", [
|
||||||
|
LoadDataStep("load", {"source": rows}),
|
||||||
|
TrainStep("train", {"estimator": "mean_regressor"}),
|
||||||
|
EvaluateStep("eval", {}),
|
||||||
|
RegisterStep("register", {"model_name": "demo"}),
|
||||||
|
])
|
||||||
|
result = pipe.run()
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, name: str, steps: Sequence[Step],
|
||||||
|
params: Optional[Dict[str, Any]] = None):
|
||||||
|
if not name:
|
||||||
|
raise PipelineError("Pipeline 需要 name")
|
||||||
|
self.name = name
|
||||||
|
self.steps: List[Step] = list(steps)
|
||||||
|
self.params: Dict[str, Any] = dict(params or {})
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config(cls, config: PipelineConfig) -> "Pipeline":
|
||||||
|
"""从声明式配置构建流水线(配置驱动,切换模型 / 数据源只改配置)。"""
|
||||||
|
steps: List[Step] = []
|
||||||
|
for sd in config.steps:
|
||||||
|
stype = sd.get("type")
|
||||||
|
sname = sd.get("name", stype)
|
||||||
|
sparams = dict(sd.get("params", {}))
|
||||||
|
factory = STEP_TYPES.get(stype or "")
|
||||||
|
if factory is None:
|
||||||
|
raise PipelineError(f"未知步骤类型:{stype!r}")
|
||||||
|
steps.append(factory(sname, sparams))
|
||||||
|
return cls(config.name, steps, config.params)
|
||||||
|
|
||||||
|
def run(self, initial_ctx: Optional[Context] = None,
|
||||||
|
dry_run: bool = False) -> PipelineResult:
|
||||||
|
"""顺序执行所有步骤;``dry_run`` 时只校验配置不执行 run。"""
|
||||||
|
ctx = initial_ctx or Context()
|
||||||
|
for k, v in self.params.items():
|
||||||
|
ctx.params.setdefault(k, v)
|
||||||
|
|
||||||
|
results: List[StepResult] = []
|
||||||
|
start = time.time()
|
||||||
|
if dry_run:
|
||||||
|
for st in self.steps:
|
||||||
|
st.prepare(ctx)
|
||||||
|
results.append(StepResult(name=st.name, success=True))
|
||||||
|
return PipelineResult(name=self.name, success=True,
|
||||||
|
step_results=results,
|
||||||
|
total_duration_s=time.time() - start)
|
||||||
|
|
||||||
|
for st in self.steps:
|
||||||
|
r = st.execute(ctx)
|
||||||
|
results.append(r)
|
||||||
|
if not r.success:
|
||||||
|
return PipelineResult(name=self.name, success=False,
|
||||||
|
step_results=results,
|
||||||
|
total_duration_s=time.time() - start,
|
||||||
|
failed_step=st.name)
|
||||||
|
return PipelineResult(name=self.name, success=True,
|
||||||
|
step_results=results,
|
||||||
|
total_duration_s=time.time() - start)
|
||||||
@@ -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,301 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""训练 / 推理流水线编排单元测试(issue #40)。
|
||||||
|
|
||||||
|
覆盖:
|
||||||
|
- Context 读写与快照;
|
||||||
|
- ModelRegistry 注册 / 版本 / 别名(latest / stable);
|
||||||
|
- 估计器(MeanRegressor / MajorityClassifier)训练与预测;
|
||||||
|
- 各 Step(LoadData / Train / Evaluate / Register / LoadModel / Predict / Custom)
|
||||||
|
的执行与产物传递;
|
||||||
|
- Pipeline 顺序编排、失败短路、dry_run;
|
||||||
|
- PipelineConfig 声明式配置往返与 from_config 构建。
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import tempfile
|
||||||
|
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
|
||||||
|
Context,
|
||||||
|
CustomStep,
|
||||||
|
ESTIMATORS,
|
||||||
|
EvaluateStep,
|
||||||
|
LoadDataStep,
|
||||||
|
LoadModelStep,
|
||||||
|
MajorityClassifier,
|
||||||
|
MeanRegressor,
|
||||||
|
ModelArtifact,
|
||||||
|
ModelRegistry,
|
||||||
|
Pipeline,
|
||||||
|
PipelineConfig,
|
||||||
|
PipelineError,
|
||||||
|
PredictStep,
|
||||||
|
RegisterStep,
|
||||||
|
StepResult,
|
||||||
|
TrainStep,
|
||||||
|
register_estimator,
|
||||||
|
register_step_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestContext(unittest.TestCase):
|
||||||
|
def test_get_set(self):
|
||||||
|
ctx = Context(params={"lr": 0.1})
|
||||||
|
ctx.set("x", 1)
|
||||||
|
self.assertEqual(ctx.get("x"), 1)
|
||||||
|
self.assertEqual(ctx.get("missing", "d"), "d")
|
||||||
|
self.assertEqual(ctx.params["lr"], 0.1)
|
||||||
|
|
||||||
|
def test_snapshot(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("a", 1)
|
||||||
|
ctx.set("b", 2)
|
||||||
|
snap = ctx.snapshot()
|
||||||
|
self.assertEqual(snap["artifacts_keys"], ["a", "b"])
|
||||||
|
|
||||||
|
|
||||||
|
class TestModelRegistry(unittest.TestCase):
|
||||||
|
def test_register_and_latest(self):
|
||||||
|
reg = ModelRegistry()
|
||||||
|
a1 = ModelArtifact("m", "v1", object())
|
||||||
|
a2 = ModelArtifact("m", "v2", object())
|
||||||
|
reg.register(a1)
|
||||||
|
reg.register(a2)
|
||||||
|
self.assertEqual(reg.get("m").version, "v2") # latest
|
||||||
|
self.assertEqual(reg.get("m", "v1").version, "v1")
|
||||||
|
self.assertEqual(reg.list_versions("m"), ["v1", "v2"])
|
||||||
|
|
||||||
|
def test_alias(self):
|
||||||
|
reg = ModelRegistry()
|
||||||
|
reg.register(ModelArtifact("m", "v1", object()))
|
||||||
|
reg.register(ModelArtifact("m", "v2", object()))
|
||||||
|
reg.set_alias("m", "stable", "v1")
|
||||||
|
self.assertEqual(reg.get("m", "stable").version, "v1")
|
||||||
|
self.assertEqual(reg.get("m", "latest").version, "v2")
|
||||||
|
|
||||||
|
def test_missing_raises(self):
|
||||||
|
reg = ModelRegistry()
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
reg.get("nope")
|
||||||
|
reg.register(ModelArtifact("m", "v1", object()))
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
reg.get("m", "v99")
|
||||||
|
|
||||||
|
def test_register_requires_name_version(self):
|
||||||
|
reg = ModelRegistry()
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
reg.register(ModelArtifact("", "v1", object()))
|
||||||
|
|
||||||
|
|
||||||
|
class TestEstimators(unittest.TestCase):
|
||||||
|
def test_mean_regressor(self):
|
||||||
|
est = MeanRegressor()
|
||||||
|
est.fit([[1], [2], [3]], [10, 20, 30])
|
||||||
|
self.assertEqual(est.predict([[99], [100]]), [20.0, 20.0])
|
||||||
|
|
||||||
|
def test_majority_classifier(self):
|
||||||
|
est = MajorityClassifier()
|
||||||
|
est.fit([[1], [2], [3]], [0, 1, 1])
|
||||||
|
self.assertEqual(est.predict([[9], [10]]), [1.0, 1.0])
|
||||||
|
|
||||||
|
def test_empty_fit_raises(self):
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
MeanRegressor().fit([], [])
|
||||||
|
|
||||||
|
def test_register_estimator(self):
|
||||||
|
class MyEst(MeanRegressor):
|
||||||
|
name = "my_est"
|
||||||
|
register_estimator("my_est", MyEst)
|
||||||
|
self.assertIn("my_est", ESTIMATORS)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSteps(unittest.TestCase):
|
||||||
|
def test_load_data_from_list(self):
|
||||||
|
ctx = Context()
|
||||||
|
r = LoadDataStep("load", {"source": [[1, 2], [3, 4]]}).execute(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
self.assertEqual(ctx.get("dataset"), [[1.0, 2.0], [3.0, 4.0]])
|
||||||
|
|
||||||
|
def test_load_data_from_csv(self):
|
||||||
|
with tempfile.NamedTemporaryFile(
|
||||||
|
mode="w", suffix=".csv", delete=False, encoding="utf-8") as fh:
|
||||||
|
fh.write("a,b,y\n1,2,3\n4,5,6\n")
|
||||||
|
path = fh.name
|
||||||
|
try:
|
||||||
|
ctx = Context()
|
||||||
|
r = LoadDataStep("load", {"source": path}).execute(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
self.assertEqual(ctx.get("dataset"), [[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_load_data_missing_source(self):
|
||||||
|
ctx = Context()
|
||||||
|
r = LoadDataStep("load", {}).execute(ctx)
|
||||||
|
self.assertFalse(r.success)
|
||||||
|
self.assertIn("source", r.error or "")
|
||||||
|
|
||||||
|
def test_train_step(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("dataset", [[1, 10], [2, 20], [3, 30]]) # 最后一列 target
|
||||||
|
r = TrainStep("train", {"estimator": "mean_regressor"}).execute(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
est = ctx.get("model")
|
||||||
|
self.assertEqual(est.predict([[9]]), [20.0])
|
||||||
|
|
||||||
|
def test_train_unknown_estimator(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("dataset", [[1, 10]])
|
||||||
|
r = TrainStep("train", {"estimator": "voodoo"}).execute(ctx)
|
||||||
|
self.assertFalse(r.success)
|
||||||
|
|
||||||
|
def test_evaluate_step_regression(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("dataset", [[1, 10], [2, 20], [3, 30]])
|
||||||
|
TrainStep("train", {}).execute(ctx)
|
||||||
|
r = EvaluateStep("eval", {}).execute(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
m = ctx.get("metrics")
|
||||||
|
# 均值预测:mae 为各 |y-20| 的均值
|
||||||
|
self.assertAlmostEqual(m["mae"], (10 + 0 + 10) / 3)
|
||||||
|
self.assertGreaterEqual(m["rmse"], 0)
|
||||||
|
|
||||||
|
def test_evaluate_step_classification(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("dataset", [[1, 0], [2, 1], [3, 1]])
|
||||||
|
TrainStep("train", {"estimator": "majority_classifier"}).execute(ctx)
|
||||||
|
EvaluateStep("eval", {}).execute(ctx)
|
||||||
|
m = ctx.get("metrics")
|
||||||
|
self.assertIn("accuracy", m)
|
||||||
|
self.assertGreaterEqual(m["accuracy"], 0.0)
|
||||||
|
|
||||||
|
def test_register_and_load_model(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("dataset", [[1, 10], [2, 20]])
|
||||||
|
TrainStep("train", {}).execute(ctx)
|
||||||
|
reg_r = RegisterStep("reg", {"model_name": "demo", "version": "v1"}).execute(ctx)
|
||||||
|
self.assertTrue(reg_r.success)
|
||||||
|
registry = ctx.get("registry")
|
||||||
|
self.assertIsInstance(registry, ModelRegistry)
|
||||||
|
self.assertEqual(registry.list_versions("demo"), ["v1"])
|
||||||
|
|
||||||
|
load_r = LoadModelStep("load", {"model_name": "demo", "version": "v1"}).execute(ctx)
|
||||||
|
self.assertTrue(load_r.success)
|
||||||
|
serving = ctx.get("serving_model")
|
||||||
|
self.assertEqual(serving.predict([[9]]), [15.0])
|
||||||
|
|
||||||
|
def test_load_model_missing_registry(self):
|
||||||
|
ctx = Context()
|
||||||
|
r = LoadModelStep("load", {"model_name": "x"}).execute(ctx)
|
||||||
|
self.assertFalse(r.success)
|
||||||
|
|
||||||
|
def test_predict_step(self):
|
||||||
|
ctx = Context()
|
||||||
|
ctx.set("dataset", [[1, 10], [2, 20]])
|
||||||
|
TrainStep("train", {}).execute(ctx)
|
||||||
|
RegisterStep("reg", {"model_name": "demo", "version": "v1"}).execute(ctx)
|
||||||
|
LoadModelStep("load", {"model_name": "demo"}).execute(ctx)
|
||||||
|
ctx.set("input", [[5], [6]])
|
||||||
|
r = PredictStep("predict", {}).execute(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
self.assertEqual(ctx.get("predictions"), [15.0, 15.0])
|
||||||
|
|
||||||
|
def test_custom_step(self):
|
||||||
|
ctx = Context()
|
||||||
|
r = CustomStep("c", {"handler": lambda c: {"out": 42}}).execute(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
self.assertEqual(ctx.get("out"), 42)
|
||||||
|
|
||||||
|
def test_custom_step_bad_handler(self):
|
||||||
|
ctx = Context()
|
||||||
|
r = CustomStep("c", {"handler": "not_callable"}).execute(ctx)
|
||||||
|
self.assertFalse(r.success)
|
||||||
|
|
||||||
|
def test_step_requires_name(self):
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
TrainStep("", {})
|
||||||
|
|
||||||
|
|
||||||
|
class TestPipeline(unittest.TestCase):
|
||||||
|
def _full_pipeline(self):
|
||||||
|
return Pipeline("demo", [
|
||||||
|
LoadDataStep("load", {"source": [[1, 10], [2, 20], [3, 30]]}),
|
||||||
|
TrainStep("train", {"estimator": "mean_regressor"}),
|
||||||
|
EvaluateStep("eval", {}),
|
||||||
|
RegisterStep("register", {"model_name": "demo", "version": "v1"}),
|
||||||
|
LoadModelStep("load_model", {"model_name": "demo", "version": "v1"}),
|
||||||
|
PredictStep("predict", {"input_key": "dataset"}),
|
||||||
|
])
|
||||||
|
|
||||||
|
def test_full_pipeline_success(self):
|
||||||
|
result = self._full_pipeline().run()
|
||||||
|
self.assertTrue(result.success)
|
||||||
|
self.assertEqual(len(result.step_results), 6)
|
||||||
|
self.assertIsNone(result.failed_step)
|
||||||
|
|
||||||
|
def test_pipeline_context_shared(self):
|
||||||
|
pipe = Pipeline("p", [
|
||||||
|
CustomStep("a", {"handler": lambda c: {"v": 7}}),
|
||||||
|
CustomStep("b", {"handler": lambda c: {"v2": c.get("v") * 2}}),
|
||||||
|
])
|
||||||
|
r = pipe.run()
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
self.assertEqual(pipe is not None, True)
|
||||||
|
|
||||||
|
def test_pipeline_failure_short_circuits(self):
|
||||||
|
# 第二步失败(无 dataset),应短路不执行后续
|
||||||
|
pipe = Pipeline("p", [
|
||||||
|
CustomStep("a", {"handler": lambda c: {}}),
|
||||||
|
EvaluateStep("bad_eval", {}), # 缺 model → 失败
|
||||||
|
CustomStep("c", {"handler": lambda c: {"never": 1}}),
|
||||||
|
])
|
||||||
|
r = pipe.run()
|
||||||
|
self.assertFalse(r.success)
|
||||||
|
self.assertEqual(r.failed_step, "bad_eval")
|
||||||
|
self.assertEqual(len(r.step_results), 2) # a + bad_eval
|
||||||
|
|
||||||
|
def test_dry_run(self):
|
||||||
|
r = self._full_pipeline().run(dry_run=True)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
# dry_run 不真正 run,predictions 不存在
|
||||||
|
# (dry_run 不产出 artifacts)
|
||||||
|
|
||||||
|
def test_pipeline_requires_name(self):
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
Pipeline("", [])
|
||||||
|
|
||||||
|
def test_register_step_type_and_from_config(self):
|
||||||
|
register_step_type("double", lambda n, p: CustomStep(
|
||||||
|
n, {"handler": lambda c: {"doubled": c.params.get("x", 0) * 2}}))
|
||||||
|
cfg = PipelineConfig("cfg", steps=[
|
||||||
|
{"type": "double", "name": "d", "params": {}},
|
||||||
|
], params={"x": 21})
|
||||||
|
pipe = Pipeline.from_config(cfg)
|
||||||
|
ctx = Context()
|
||||||
|
r = pipe.run(ctx)
|
||||||
|
self.assertTrue(r.success)
|
||||||
|
self.assertEqual(ctx.get("doubled"), 42)
|
||||||
|
|
||||||
|
def test_from_config_unknown_type(self):
|
||||||
|
cfg = PipelineConfig("cfg", steps=[{"type": "voodoo"}])
|
||||||
|
with self.assertRaises(PipelineError):
|
||||||
|
Pipeline.from_config(cfg)
|
||||||
|
|
||||||
|
def test_config_roundtrip(self):
|
||||||
|
cfg = PipelineConfig("c", steps=[{"type": "train", "name": "t", "params": {}}],
|
||||||
|
params={"k": 1})
|
||||||
|
d = cfg.to_dict()
|
||||||
|
cfg2 = PipelineConfig.from_dict(json.loads(json.dumps(d)))
|
||||||
|
self.assertEqual(cfg2.name, "c")
|
||||||
|
self.assertEqual(cfg2.params, {"k": 1})
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user