feat(#34): Model Recipe 插件接口与样例协议(PRD 5.3 模型框架配置化/网络结构策略)

This commit is contained in:
2026-08-04 23:03:01 +08:00
parent 4e9d0e60c9
commit 9c07d19f63
6 changed files with 1144 additions and 0 deletions
+17
View File
@@ -0,0 +1,17 @@
# -*- coding: utf-8 -*-
"""测试引导:把 `core/model-framework` 以包名 `model_framework` 挂载到 sys.modules。
目录名 `model-framework` 含连字符,无法直接以包名 import;挂载后模块内
``from model_framework.model_recipe import ...`` 在 unittest 发现机制下可
正常解析(与 `core/data-bus/tests/_bootstrap.py` 一致)。
"""
import os
import sys
import types
MODEL_FRAMEWORK_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, MODEL_FRAMEWORK_DIR)
if "model_framework" not in sys.modules:
pkg = types.ModuleType("model_framework")
pkg.__path__ = [MODEL_FRAMEWORK_DIR]
sys.modules["model_framework"] = pkg
@@ -0,0 +1,303 @@
# -*- coding: utf-8 -*-
"""issue #34 Model Recipe 插件接口与样例协议 单元测试。
覆盖:
* 内置四类 Recipe 已注册、字段合法;
* build_model 跨主干(gbdt/dnn/lstm/gnn/stub)可构造、fit/predict 契约;
* 插件注册(register_recipe / register_backbone)零改码扩展;
* ModelRecipe 不可变 + to_dict/from_dict 往返;
* 超参包校验(recipe_id / 必需特征 / 主干可构造性);
* 样例协议:树脂 + Ti 两套 Recipe 同框架均跑通(EPIC #5 验收口径)。
"""
import os
import sys
# 引导:挂载 model_framework 包(目录含连字符)
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _bootstrap # noqa: F401,E402
import unittest
from model_framework.model_recipe import ( # noqa: E402
BACKBONES,
RECIPE_KINDS,
ModelRecipe,
RecipeError,
build_model,
get_recipe,
list_recipes,
load_sample_recipe,
register_backbone,
register_recipe,
validate_hyperparam_pack,
)
class TestBuiltinRecipes(unittest.TestCase):
"""内置四类 Recipe 注册与字段合法性。"""
def test_four_builtin_recipes_registered(self):
ids = {r["id"] for r in list_recipes()}
for rid in (
"quality_predict.default",
"process_optimize.default",
"anomaly_detect.default",
"cross_process.default",
):
self.assertIn(rid, ids, f"缺少内置 Recipe {rid}")
def test_each_builtin_kind_covered(self):
kinds = {get_recipe(rid).kind for rid in (
"quality_predict.default",
"process_optimize.default",
"anomaly_detect.default",
"cross_process.default",
)}
self.assertEqual(kinds, set(RECIPE_KINDS))
def test_backbone_registered(self):
for name in ("gbdt", "dnn", "lstm", "gnn", "stub"):
self.assertIn(name, BACKBONES, f"缺少内置主干 {name}")
class TestModelRecipeDataclass(unittest.TestCase):
"""ModelRecipe 不可变 + 序列化往返 + 校验。"""
def test_immutable(self):
r = get_recipe("quality_predict.default")
with self.assertRaises(Exception):
r.id = "x" # type: ignore[misc]
def test_to_from_dict_roundtrip(self):
r = get_recipe("quality_predict.default")
d = r.to_dict()
r2 = ModelRecipe.from_dict(d)
self.assertEqual(r2.to_dict(), d)
self.assertEqual(r2.id, r.id)
self.assertEqual(r2.backbone, r.backbone)
def test_invalid_kind_rejected(self):
with self.assertRaises(RecipeError):
ModelRecipe(id="x.bad", kind="bogus", backbone="gbdt")
def test_unregistered_backbone_rejected(self):
with self.assertRaises(RecipeError):
ModelRecipe(id="x.nobackbone", kind="quality_predict", backbone="no-such")
def test_merged_hyperparams_override_wins(self):
r = get_recipe("quality_predict.default")
base = r.default_hyperparams
merged = r.merged_hyperparams({"max_depth": 99})
self.assertEqual(merged["max_depth"], 99)
# 默认值未被污染
self.assertEqual(base["max_depth"], 6)
self.assertIn("eta", merged)
class TestBuildModel(unittest.TestCase):
"""build_model 跨主干构造 + fit/predict 契约。"""
def test_build_each_backbone(self):
for rid, bb in (
("quality_predict.default", "gbdt"),
("process_optimize.default", "gbdt"),
("anomaly_detect.default", "dnn"),
("cross_process.default", "gnn"),
):
m = build_model(rid)
self.assertEqual(m.backbone, bb)
self.assertFalse(m.fitted)
def test_fit_then_predict_returns_correct_length(self):
m = build_model("quality_predict.default")
X = [[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]
y = [1.0, 2.0, 3.0]
m.fit(X, y)
self.assertTrue(m.fitted)
pred = m.predict([[2.0, 3.0], [4.0, 5.0]])
self.assertEqual(len(pred), 2)
for v in pred:
self.assertIsInstance(v, float)
def test_predict_before_fit_fails_closed(self):
m = build_model("anomaly_detect.default")
with self.assertRaises(RecipeError):
m.predict([[1.0, 2.0]])
def test_unsupervised_fit_without_y(self):
# anomaly_detect 主干应允许无 y 拟合
m = build_model("anomaly_detect.default")
m.fit([[1.0, 2.0], [3.0, 4.0]])
self.assertTrue(m.fitted)
out = m.predict([[1.0, 2.0]])
self.assertEqual(len(out), 1)
def test_X_width_mismatch_rejected(self):
m = build_model("quality_predict.default")
with self.assertRaises(ValueError):
m.fit([[1.0, 2.0], [3.0]], [1.0, 2.0])
def test_Xy_length_mismatch_rejected(self):
m = build_model("quality_predict.default")
with self.assertRaises(ValueError):
m.fit([[1.0, 2.0], [3.0, 4.0]], [1.0])
def test_empty_X_rejected(self):
m = build_model("quality_predict.default")
with self.assertRaises(ValueError):
m.fit([], [])
def test_handle_to_dict(self):
m = build_model("quality_predict.default", {"max_depth": 7})
d = m.to_dict()
self.assertEqual(d["recipe_id"], "quality_predict.default")
self.assertEqual(d["backbone"], "gbdt")
self.assertEqual(d["hyperparams"]["max_depth"], 7)
self.assertFalse(d["fitted"])
def test_unknown_recipe_raises(self):
with self.assertRaises(RecipeError):
build_model("no.such.recipe")
class TestPluginRegistration(unittest.TestCase):
"""register_recipe / register_backbone 零改码扩展(PRD「新增结构走插件注册」)。"""
def test_register_custom_backbone_and_recipe(self):
seen = {}
def my_bb(hp):
class _Impl:
def iaop_fit(self, rows, y):
seen["fit_called"] = True
def iaop_predict(self, rows):
return [42.0 for _ in rows]
return _Impl()
register_backbone("my-gnn", my_bb)
self.assertIn("my-gnn", BACKBONES)
register_recipe(ModelRecipe(
id="cross_process.custom_gnn",
kind="cross_process",
backbone="my-gnn",
description="自研 GNN 主干,验证插件扩展",
))
m = build_model("cross_process.custom_gnn")
m.fit([[1.0, 2.0]], [1.0])
self.assertTrue(seen.get("fit_called"))
self.assertEqual(m.predict([[9.0, 9.0]]), [42.0])
def test_register_recipe_overwrites(self):
# 用独立的临时 recipe 验证"重复注册同 id 覆盖",不污染内置表
register_recipe(ModelRecipe(
id="quality_predict.temp",
kind="quality_predict",
backbone="gbdt",
description="第一版",
))
self.assertEqual(get_recipe("quality_predict.temp").description, "第一版")
register_recipe(ModelRecipe(
id="quality_predict.temp",
kind="quality_predict",
backbone="stub",
description="第二版覆盖",
))
self.assertEqual(get_recipe("quality_predict.temp").backbone, "stub")
self.assertEqual(get_recipe("quality_predict.temp").description, "第二版覆盖")
def test_register_invalid_backbone_name_rejected(self):
with self.assertRaises(RecipeError):
register_backbone("bad name!", lambda hp: None)
def test_register_non_callable_factory_rejected(self):
with self.assertRaises(RecipeError):
register_backbone("oops", "not callable") # type: ignore[arg-type]
def test_register_non_recipe_rejected(self):
with self.assertRaises(RecipeError):
register_recipe("not a recipe") # type: ignore[arg-type]
class TestHyperparamPackValidation(unittest.TestCase):
"""超参包校验(Recipe 视角)。"""
def test_valid_pack_no_issues(self):
pack = load_sample_recipe("ti")
self.assertEqual(validate_hyperparam_pack(pack), [])
def test_missing_required_field(self):
issues = validate_hyperparam_pack({"recipe_id": "quality_predict.default"})
msgs = " ".join(issues)
self.assertIn("model_id", msgs)
self.assertIn("features", msgs)
def test_unknown_recipe_id(self):
issues = validate_hyperparam_pack({
"model_id": "x", "recipe_id": "no.such", "features": [],
})
self.assertTrue(any("未注册" in i for i in issues))
def test_missing_required_feature(self):
# quality_predict.default 要求 'target' 特征
issues = validate_hyperparam_pack({
"model_id": "x",
"recipe_id": "quality_predict.default",
"features": [{"name": "only_a"}],
})
self.assertTrue(any("target" in i for i in issues))
class TestSampleRecipesAcceptance(unittest.TestCase):
"""EPIC #5 / PRD 5.3 验收口径:同框架加载树脂与 Ti 两套 Recipe 均跑通。"""
def test_both_samples_build_fit_predict(self):
for name in ("resin", "ti"):
pack = load_sample_recipe(name)
self.assertEqual(validate_hyperparam_pack(pack), [],
f"样例 {name} 校验未通过")
m = build_model(pack["recipe_id"], pack.get("hyperparams"))
# 构造与目标维度无关的训练样本(2 特征列)
X = [[float(i), float(i + 1)] for i in range(6)]
y = [float(i) for i in range(6)]
m.fit(X, y)
self.assertTrue(m.fitted)
pred = m.predict([[1.0, 2.0]])
self.assertEqual(len(pred), 1)
def test_samples_share_same_framework(self):
# 关键:两套样例用同一个 recipe_id(quality_predict.default),
# 仅超参不同——证明「切换模板仅改超参包,模型代码零改动」
r1 = load_sample_recipe("resin")
r2 = load_sample_recipe("ti")
self.assertEqual(r1["recipe_id"], r2["recipe_id"])
# 但超参不同(max_depth 4 vs 6)
self.assertNotEqual(
r1["hyperparams"]["max_depth"],
r2["hyperparams"]["max_depth"],
)
# 各自 build 得到不同超参的句柄
m1 = build_model(r1["recipe_id"], r1["hyperparams"])
m2 = build_model(r2["recipe_id"], r2["hyperparams"])
self.assertEqual(m1.hyperparams["max_depth"], 4)
self.assertEqual(m2.hyperparams["max_depth"], 6)
def test_load_unknown_sample_raises(self):
with self.assertRaises(RecipeError):
load_sample_recipe("bogus")
class TestBackboneFallback(unittest.TestCase):
"""主干在无第三方依赖时退化为 stub,接口契约不变。"""
def test_lstm_gnn_fallback_to_stub_contract(self):
# 无论是否有 torch,lstm/gnn 主干都应能构造并 fit/predict
for rid in ("cross_process.default",):
m = build_model(rid)
m.fit([[1.0, 2.0]], [1.0])
self.assertEqual(len(m.predict([[1.0, 2.0]])), 1)
if __name__ == "__main__":
unittest.main(verbosity=2)