feat(#34): Model Recipe 插件接口与样例协议(PRD 5.3 模型框架配置化/网络结构策略)
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user