feat(#40): 训练/推理流水线编排(声明式配置驱动,可插拔Step/Estimator,PRD 5.3 ③训练推理流水线编排)

This commit is contained in:
2026-08-05 01:11:09 +08:00
parent 793dd0a3b8
commit 6a5ac40477
7 changed files with 1218 additions and 0 deletions
+675
View File
@@ -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)