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
+1
View File
@@ -0,0 +1 @@
# -*- coding: utf-8 -*-
+23
View File
@@ -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)
+301
View File
@@ -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()